From d8341e71aec1cd80d44ead249898be6afe16a870 Mon Sep 17 00:00:00 2001 From: Eric Daniels Date: Mon, 1 Jul 2024 12:37:20 -0400 Subject: [PATCH] Cleanly close agent goroutines --- agent.go | 12 +++++++---- agent_handlers.go | 45 +++++++++++++++++++++++++++++++++++++++++- agent_handlers_test.go | 4 ++++ agent_test.go | 5 ++++- 4 files changed, 60 insertions(+), 6 deletions(-) diff --git a/agent.go b/agent.go index 761c07c..3412e10 100644 --- a/agent.go +++ b/agent.go @@ -220,9 +220,9 @@ func NewAgent(config *AgentConfig) (*Agent, error) { //nolint:gocognit userBindingRequestHandler: config.BindingRequestHandler, } - a.connectionStateNotifier = &handlerNotifier{connectionStateFunc: a.onConnectionStateChange} - a.candidateNotifier = &handlerNotifier{candidateFunc: a.onCandidate} - a.selectedCandidatePairNotifier = &handlerNotifier{candidatePairFunc: a.onSelectedCandidatePairChange} + a.connectionStateNotifier = &handlerNotifier{connectionStateFunc: a.onConnectionStateChange, done: make(chan struct{})} + a.candidateNotifier = &handlerNotifier{candidateFunc: a.onCandidate, done: make(chan struct{})} + a.selectedCandidatePairNotifier = &handlerNotifier{candidatePairFunc: a.onSelectedCandidatePairChange, done: make(chan struct{})} if a.net == nil { a.net, err = stdnet.NewNet() @@ -849,7 +849,11 @@ func (a *Agent) removeUfragFromMux() { // Close cleans up the Agent func (a *Agent) Close() error { - return a.loop.Close() + err := a.loop.Close() + a.connectionStateNotifier.Close() + a.candidateNotifier.Close() + a.selectedCandidatePairNotifier.Close() + return err } // Remove all candidates. This closes any listening sockets diff --git a/agent_handlers.go b/agent_handlers.go index bb0c8d3..7ebfedd 100644 --- a/agent_handlers.go +++ b/agent_handlers.go @@ -45,7 +45,8 @@ func (a *Agent) onConnectionStateChange(s ConnectionState) { type handlerNotifier struct { sync.Mutex - running bool + running bool + notifiers sync.WaitGroup connectionStates []ConnectionState connectionStateFunc func(ConnectionState) @@ -55,13 +56,38 @@ type handlerNotifier struct { selectedCandidatePairs []*CandidatePair candidatePairFunc func(*CandidatePair) + + // State for closing + done chan struct{} +} + +func (h *handlerNotifier) Close() { + h.Lock() + + select { + case <-h.done: + h.Unlock() + return + default: + } + close(h.done) + h.Unlock() + + h.notifiers.Wait() } func (h *handlerNotifier) EnqueueConnectionState(s ConnectionState) { h.Lock() defer h.Unlock() + select { + case <-h.done: + return + default: + } + notify := func() { + defer h.notifiers.Done() for { h.Lock() if len(h.connectionStates) == 0 { @@ -79,6 +105,7 @@ func (h *handlerNotifier) EnqueueConnectionState(s ConnectionState) { h.connectionStates = append(h.connectionStates, s) if !h.running { h.running = true + h.notifiers.Add(1) go notify() } } @@ -87,7 +114,14 @@ func (h *handlerNotifier) EnqueueCandidate(c Candidate) { h.Lock() defer h.Unlock() + select { + case <-h.done: + return + default: + } + notify := func() { + defer h.notifiers.Done() for { h.Lock() if len(h.candidates) == 0 { @@ -105,6 +139,7 @@ func (h *handlerNotifier) EnqueueCandidate(c Candidate) { h.candidates = append(h.candidates, c) if !h.running { h.running = true + h.notifiers.Add(1) go notify() } } @@ -113,7 +148,14 @@ func (h *handlerNotifier) EnqueueSelectedCandidatePair(p *CandidatePair) { h.Lock() defer h.Unlock() + select { + case <-h.done: + return + default: + } + notify := func() { + defer h.notifiers.Done() for { h.Lock() if len(h.selectedCandidatePairs) == 0 { @@ -131,6 +173,7 @@ func (h *handlerNotifier) EnqueueSelectedCandidatePair(p *CandidatePair) { h.selectedCandidatePairs = append(h.selectedCandidatePairs, p) if !h.running { h.running = true + h.notifiers.Add(1) go notify() } } diff --git a/agent_handlers_test.go b/agent_handlers_test.go index 35680ee..ce02733 100644 --- a/agent_handlers_test.go +++ b/agent_handlers_test.go @@ -19,6 +19,7 @@ func TestConnectionStateNotifier(t *testing.T) { connectionStateFunc: func(_ ConnectionState) { updates <- struct{}{} }, + done: make(chan struct{}), } // Enqueue all updates upfront to ensure that it // doesn't block @@ -38,6 +39,7 @@ func TestConnectionStateNotifier(t *testing.T) { close(done) }() <-done + c.Close() }) t.Run("TestUpdateOrdering", func(t *testing.T) { defer test.CheckRoutines(t)() @@ -46,6 +48,7 @@ func TestConnectionStateNotifier(t *testing.T) { connectionStateFunc: func(cs ConnectionState) { updates <- cs }, + done: make(chan struct{}), } done := make(chan struct{}) go func() { @@ -66,5 +69,6 @@ func TestConnectionStateNotifier(t *testing.T) { c.EnqueueConnectionState(ConnectionState(i)) } <-done + c.Close() }) } diff --git a/agent_test.go b/agent_test.go index 552d421..05dc19b 100644 --- a/agent_test.go +++ b/agent_test.go @@ -1379,11 +1379,12 @@ func TestCloseInConnectionStateCallback(t *testing.T) { isClosed := make(chan interface{}) isConnected := make(chan interface{}) + connectionStateConnectedSeen := make(chan interface{}) err = aAgent.OnConnectionStateChange(func(c ConnectionState) { switch c { case ConnectionStateConnected: <-isConnected - require.NoError(t, aAgent.Close()) + close(connectionStateConnectedSeen) case ConnectionStateClosed: close(isClosed) default: @@ -1393,6 +1394,8 @@ func TestCloseInConnectionStateCallback(t *testing.T) { connect(aAgent, bAgent) close(isConnected) + <-connectionStateConnectedSeen + require.NoError(t, aAgent.Close()) <-isClosed require.NoError(t, bAgent.Close())