diff --git a/agent.go b/agent.go index 3412e10..761c07c 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, done: make(chan struct{})} - a.candidateNotifier = &handlerNotifier{candidateFunc: a.onCandidate, done: make(chan struct{})} - a.selectedCandidatePairNotifier = &handlerNotifier{candidatePairFunc: a.onSelectedCandidatePairChange, done: make(chan struct{})} + a.connectionStateNotifier = &handlerNotifier{connectionStateFunc: a.onConnectionStateChange} + a.candidateNotifier = &handlerNotifier{candidateFunc: a.onCandidate} + a.selectedCandidatePairNotifier = &handlerNotifier{candidatePairFunc: a.onSelectedCandidatePairChange} if a.net == nil { a.net, err = stdnet.NewNet() @@ -849,11 +849,7 @@ func (a *Agent) removeUfragFromMux() { // Close cleans up the Agent func (a *Agent) Close() error { - err := a.loop.Close() - a.connectionStateNotifier.Close() - a.candidateNotifier.Close() - a.selectedCandidatePairNotifier.Close() - return err + return a.loop.Close() } // Remove all candidates. This closes any listening sockets diff --git a/agent_handlers.go b/agent_handlers.go index 7ebfedd..bb0c8d3 100644 --- a/agent_handlers.go +++ b/agent_handlers.go @@ -45,8 +45,7 @@ func (a *Agent) onConnectionStateChange(s ConnectionState) { type handlerNotifier struct { sync.Mutex - running bool - notifiers sync.WaitGroup + running bool connectionStates []ConnectionState connectionStateFunc func(ConnectionState) @@ -56,38 +55,13 @@ 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 { @@ -105,7 +79,6 @@ func (h *handlerNotifier) EnqueueConnectionState(s ConnectionState) { h.connectionStates = append(h.connectionStates, s) if !h.running { h.running = true - h.notifiers.Add(1) go notify() } } @@ -114,14 +87,7 @@ 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 { @@ -139,7 +105,6 @@ func (h *handlerNotifier) EnqueueCandidate(c Candidate) { h.candidates = append(h.candidates, c) if !h.running { h.running = true - h.notifiers.Add(1) go notify() } } @@ -148,14 +113,7 @@ 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 { @@ -173,7 +131,6 @@ 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 ce02733..35680ee 100644 --- a/agent_handlers_test.go +++ b/agent_handlers_test.go @@ -19,7 +19,6 @@ 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 @@ -39,7 +38,6 @@ func TestConnectionStateNotifier(t *testing.T) { close(done) }() <-done - c.Close() }) t.Run("TestUpdateOrdering", func(t *testing.T) { defer test.CheckRoutines(t)() @@ -48,7 +46,6 @@ func TestConnectionStateNotifier(t *testing.T) { connectionStateFunc: func(cs ConnectionState) { updates <- cs }, - done: make(chan struct{}), } done := make(chan struct{}) go func() { @@ -69,6 +66,5 @@ func TestConnectionStateNotifier(t *testing.T) { c.EnqueueConnectionState(ConnectionState(i)) } <-done - c.Close() }) } diff --git a/agent_test.go b/agent_test.go index 05dc19b..552d421 100644 --- a/agent_test.go +++ b/agent_test.go @@ -1379,12 +1379,11 @@ 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 - close(connectionStateConnectedSeen) + require.NoError(t, aAgent.Close()) case ConnectionStateClosed: close(isClosed) default: @@ -1394,8 +1393,6 @@ func TestCloseInConnectionStateCallback(t *testing.T) { connect(aAgent, bAgent) close(isConnected) - <-connectionStateConnectedSeen - require.NoError(t, aAgent.Close()) <-isClosed require.NoError(t, bAgent.Close())