diff --git a/internal/server/server.go b/internal/server/server.go index 7999c99..1621e66 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -376,7 +376,9 @@ func (s *Server) addSubsciberToTopic(topicName string, peer *Peer, startAt int) t = newTopic(topicName) } + t.mu.Lock() t.subscriptions[peer.Addr()] = newSubscriber(peer, t, s.ackDelay, s.ackTimeout, startAt) + t.mu.Unlock() s.topics[topicName] = t } @@ -396,25 +398,28 @@ func (s *Server) removeSubsciberFromTopic(topicName string, peer *Peer) { if !ok { return } - sub, ok := t.subscriptions[peer.Addr()] - if !ok { + + sub := t.findSubscription(peer.Addr()) + if sub == nil { return } + sub.unsubscribe() - delete(t.subscriptions, peer.Addr()) + t.removeSubscription(peer.Addr()) } func (s *Server) unsubscribePeerFromAllTopics(peer *Peer) { s.mu.Lock() defer s.mu.Unlock() - for _, topic := range s.topics { - sub, ok := topic.subscriptions[peer.Addr()] - if !ok { - continue + for _, t := range s.topics { + sub := t.findSubscription(peer.Addr()) + if sub == nil { + return } + sub.unsubscribe() - delete(topic.subscriptions, peer.Addr()) + t.removeSubscription(peer.Addr()) } } diff --git a/internal/server/server_test.go b/internal/server/server_test.go index 2d44f96..365d5ef 100644 --- a/internal/server/server_test.go +++ b/internal/server/server_test.go @@ -128,6 +128,8 @@ func TestSubscribeToTopics(t *testing.T) { _ = createConnectionAndSubscribe(t, []string{topicA, topicB}, Current, 0) + srv.mu.Lock() + defer srv.mu.Unlock() assert.Len(t, srv.topics, 2) assert.Len(t, srv.topics[topicA].subscriptions, 1) assert.Len(t, srv.topics[topicB].subscriptions, 1) @@ -139,9 +141,12 @@ func TestUnsubscribesFromTopic(t *testing.T) { conn := createConnectionAndSubscribe(t, []string{topicA, topicB, topicC}, Current, 0) assert.Len(t, srv.topics, 3) + + srv.mu.Lock() assert.Len(t, srv.topics[topicA].subscriptions, 1) assert.Len(t, srv.topics[topicB].subscriptions, 1) assert.Len(t, srv.topics[topicC].subscriptions, 1) + srv.mu.Unlock() topics := []string{topicA, topicB} @@ -156,9 +161,12 @@ func TestUnsubscribesFromTopic(t *testing.T) { assert.Equal(t, expectedRes, resp) assert.Len(t, srv.topics, 3) + + srv.mu.Lock() assert.Len(t, srv.topics[topicA].subscriptions, 0) assert.Len(t, srv.topics[topicB].subscriptions, 0) assert.Len(t, srv.topics[topicC].subscriptions, 1) + srv.mu.Unlock() } func TestSubscriberClosesWithoutUnsubscribing(t *testing.T) { @@ -167,8 +175,11 @@ func TestSubscriberClosesWithoutUnsubscribing(t *testing.T) { conn := createConnectionAndSubscribe(t, []string{topicA, topicB}, Current, 0) assert.Len(t, srv.topics, 2) + + srv.mu.Lock() assert.Len(t, srv.topics[topicA].subscriptions, 1) assert.Len(t, srv.topics[topicB].subscriptions, 1) + srv.mu.Unlock() // close the conn err := conn.Close() @@ -189,8 +200,11 @@ func TestSubscriberClosesWithoutUnsubscribing(t *testing.T) { time.Sleep(time.Millisecond * 100) assert.Len(t, srv.topics, 2) + + srv.mu.Lock() assert.Len(t, srv.topics[topicA].subscriptions, 0) assert.Len(t, srv.topics[topicB].subscriptions, 0) + srv.mu.Unlock() } func TestInvalidAction(t *testing.T) { diff --git a/internal/server/topic.go b/internal/server/topic.go index fa5cf5d..f1c11fd 100644 --- a/internal/server/topic.go +++ b/internal/server/topic.go @@ -46,3 +46,17 @@ func (t *topic) sendMessageToSubscribers(msg internal.Message) error { return nil } + +func (t *topic) findSubscription(addr net.Addr) *subscriber { + t.mu.Lock() + defer t.mu.Unlock() + + return t.subscriptions[addr] +} + +func (t *topic) removeSubscription(addr net.Addr) { + t.mu.Lock() + defer t.mu.Unlock() + + delete(t.subscriptions, addr) +}