diff --git a/server/server.go b/server/server.go index 9f0ebea..bfb99a0 100644 --- a/server/server.go +++ b/server/server.go @@ -303,13 +303,15 @@ func (s *Server) handlePublish(peer *peer.Peer) { if message == nil { continue } - // TODO: this can be done in a go routine because once we've got the message from the publisher, the publisher - // doesn't need to wait for us to send the message to all peers - topic := s.getTopic(message.topic) - if topic != nil { - topic.sendMessageToSubscribers(message.data) - } + // sending messages to the subscribers can be done async because the publisher doesn't need to wait for + // subscribers to be sent the message + go func() { + topic := s.getTopic(message.topic) + if topic != nil { + topic.sendMessageToSubscribers(message.data) + } + }() } } diff --git a/server/server_test.go b/server/server_test.go index 04b0cde..0b2b02b 100644 --- a/server/server_test.go +++ b/server/server_test.go @@ -283,9 +283,9 @@ func TestPublishMultipleTimes(t *testing.T) { err = binary.Write(publisherConn, binary.BigEndian, Publish) require.NoError(t, err) - messages := make([][]byte, 0, 10) + messages := make([]string, 0, 10) for i := 0; i < 10; i++ { - messages = append(messages, []byte(fmt.Sprintf("message %d", i))) + messages = append(messages, fmt.Sprintf("message %d", i)) } subscribeFinCh := make(chan struct{}) @@ -293,7 +293,8 @@ func TestPublishMultipleTimes(t *testing.T) { subscriberConn := createConnectionAndSubscribe(t, []string{topicA, topicB}) go func() { // check subscriber got all messages - for _, msg := range messages { + results := make([]string, 0, len(messages)) + for i := 0; i < len(messages); i++ { var topicLen uint64 err = binary.Read(subscriberConn, binary.BigEndian, &topicLen) require.NoError(t, err) @@ -312,9 +313,11 @@ func TestPublishMultipleTimes(t *testing.T) { require.NoError(t, err) require.Equal(t, int(dataLen), n) - assert.Equal(t, msg, buf) + results = append(results, string(buf)) } + assert.ElementsMatch(t, results, messages) + subscribeFinCh <- struct{}{} }() diff --git a/server/topic.go b/server/topic.go index e6c5c11..a087425 100644 --- a/server/topic.go +++ b/server/topic.go @@ -33,12 +33,25 @@ func (t *topic) sendMessageToSubscribers(msgData []byte) { subscribers := t.subscriptions t.mu.Unlock() - for addr, subscriber := range subscribers { - err := subscriber.peer.RunConnOperation(sendMessageOp(t.name, msgData)) - if err != nil { - slog.Error("failed to send to message", "error", err, "peer", addr) - return - } + var wg sync.WaitGroup + + for _, subscriber := range subscribers { + wg.Add(1) + sub := subscriber + go func() { + defer wg.Done() + sendMessage(sub, t.name, msgData) + }() + } + + wg.Wait() +} + +func sendMessage(sub subscriber, topicName string, message []byte) { + err := sub.peer.RunConnOperation(sendMessageOp(topicName, message)) + if err != nil { + slog.Error("failed to send to message", "error", err, "peer", sub.peer.Addr()) + return } }