diff --git a/example/main.go b/example/main.go index e435bb4..87773a3 100644 --- a/example/main.go +++ b/example/main.go @@ -59,10 +59,7 @@ func sendMessages() { i := 0 for { i++ - msg := pubsub.Message{ - Topic: "topic a", - Data: []byte(fmt.Sprintf("message %d", i)), - } + msg := pubsub.NewMessage("topic a", []byte(fmt.Sprintf("message %d", i))) err = publisher.PublishMessage(msg) if err != nil { diff --git a/example/server/main.go b/example/server/main.go index 419b61c..3841b19 100644 --- a/example/server/main.go +++ b/example/server/main.go @@ -5,12 +5,13 @@ import ( "os" "os/signal" "syscall" + "time" "github.com/willdot/messagebroker/server" ) func main() { - srv, err := server.New(":3000") + srv, err := server.New(":3000", time.Second, time.Second*2) if err != nil { log.Fatal(err) } diff --git a/pubsub/message.go b/pubsub/message.go index 69d652c..3014e14 100644 --- a/pubsub/message.go +++ b/pubsub/message.go @@ -4,4 +4,20 @@ package pubsub type Message struct { Topic string `json:"topic"` Data []byte `json:"data"` + + ack chan bool +} + +// NewMessage creates a new message +func NewMessage(topic string, data []byte) *Message { + return &Message{ + Topic: topic, + Data: data, + ack: make(chan bool), + } +} + +// Ack will send the provided value of the ack to the server +func (m *Message) Ack(ack bool) { + m.ack <- ack } diff --git a/pubsub/publisher.go b/pubsub/publisher.go index 9f4b257..883baf3 100644 --- a/pubsub/publisher.go +++ b/pubsub/publisher.go @@ -39,7 +39,7 @@ func (p *Publisher) Close() error { } // Publish will publish the given message to the server -func (p *Publisher) PublishMessage(message Message) error { +func (p *Publisher) PublishMessage(message *Message) error { op := func(conn net.Conn) error { // send topic first topic := fmt.Sprintf("topic:%s", message.Topic) diff --git a/pubsub/subscriber.go b/pubsub/subscriber.go index d0d37fa..e01b87f 100644 --- a/pubsub/subscriber.go +++ b/pubsub/subscriber.go @@ -143,14 +143,14 @@ func (s *Subscriber) UnsubscribeToTopics(topicNames []string) error { // Consumer allows the consumption of messages. If during the consumer receiving messages from the // server an error occurs, it will be stored in Err type Consumer struct { - msgs chan Message + msgs chan *Message // TODO: better error handling? Maybe a channel of errors? Err error } // Messages returns a channel in which this consumer will put messages onto. It is safe to range over the channel since it will be closed once // the consumer has finished either due to an error or from being cancelled. -func (c *Consumer) Messages() <-chan Message { +func (c *Consumer) Messages() <-chan *Message { return c.msgs } @@ -158,7 +158,7 @@ func (c *Consumer) Messages() <-chan Message { // to read the messages func (s *Subscriber) Consume(ctx context.Context) *Consumer { consumer := &Consumer{ - msgs: make(chan Message), + msgs: make(chan *Message), } go s.consume(ctx, consumer) @@ -173,20 +173,16 @@ func (s *Subscriber) consume(ctx context.Context, consumer *Consumer) { return } - msg, err := s.readMessage() + err := s.readMessage(consumer.msgs) if err != nil { consumer.Err = err return } - - if msg != nil { - consumer.msgs <- *msg - } } } -func (s *Subscriber) readMessage() (*Message, error) { - var msg *Message +func (s *Subscriber) readMessage(msgChan chan *Message) error { + // var msg *Message op := func(conn net.Conn) error { err := s.conn.SetReadDeadline(time.Now().Add(time.Second)) if err != nil { @@ -225,25 +221,35 @@ func (s *Subscriber) readMessage() (*Message, error) { return err } - msg = &Message{ - Data: dataBuf, - Topic: string(topicBuf), + msg := NewMessage(string(topicBuf), dataBuf) + + msgChan <- msg + + ack := <-msg.ack + + ackMessage := server.Nack + if ack { + ackMessage = server.Ack + } + + err = binary.Write(s.conn, binary.BigEndian, ackMessage) + if err != nil { + return fmt.Errorf("failed to ack/nack message: %w", err) } return nil - } err := s.connOperation(op) if err != nil { var neterr net.Error if errors.As(err, &neterr) && neterr.Timeout() { - return nil, nil + return nil } - return nil, err + return err } - return msg, err + return err } func (s *Subscriber) connOperation(op connOpp) error { diff --git a/pubsub/subscriber_test.go b/pubsub/subscriber_test.go index ef0c840..0463aff 100644 --- a/pubsub/subscriber_test.go +++ b/pubsub/subscriber_test.go @@ -19,7 +19,7 @@ const ( ) func createServer(t *testing.T) { - server, err := server.New(serverAddr) + server, err := server.New(serverAddr, time.Millisecond*100, time.Millisecond*100) require.NoError(t, err) t.Cleanup(func() { @@ -105,14 +105,14 @@ func TestUnsubscribesFromTopic(t *testing.T) { consumer := sub.Consume(ctx) require.NoError(t, err) - var receivedMessages []Message + var receivedMessages []*Message consumerFinCh := make(chan struct{}) go func() { for msg := range consumer.Messages() { + msg.Ack(true) receivedMessages = append(receivedMessages, msg) } - require.NoError(t, err) consumerFinCh <- struct{}{} }() @@ -125,10 +125,7 @@ func TestUnsubscribesFromTopic(t *testing.T) { publisher.Close() }) - msg := Message{ - Topic: topicA, - Data: []byte("hello world"), - } + msg := NewMessage(topicA, []byte("hello world")) err = publisher.PublishMessage(msg) require.NoError(t, err) @@ -151,37 +148,17 @@ func TestUnsubscribesFromTopic(t *testing.T) { } func TestPublishAndSubscribe(t *testing.T) { - createServer(t) - - sub, err := NewSubscriber(serverAddr) - require.NoError(t, err) - - t.Cleanup(func() { - sub.Close() - }) - - topics := []string{topicA, topicB} - - err = sub.SubscribeToTopics(topics) - require.NoError(t, err) + consumer, cancel := setupConsumer(t) - ctx, cancel := context.WithCancel(context.Background()) - t.Cleanup(func() { - cancel() - }) - - consumer := sub.Consume(ctx) - require.NoError(t, err) - - var receivedMessages []Message + var receivedMessages []*Message consumerFinCh := make(chan struct{}) go func() { for msg := range consumer.Messages() { + msg.Ack(true) receivedMessages = append(receivedMessages, msg) } - require.NoError(t, err) consumerFinCh <- struct{}{} }() @@ -192,12 +169,9 @@ func TestPublishAndSubscribe(t *testing.T) { }) // send some messages - sentMessages := make([]Message, 0, 10) + sentMessages := make([]*Message, 0, 10) for i := 0; i < 10; i++ { - msg := Message{ - Topic: topicA, - Data: []byte(fmt.Sprintf("message %d", i)), - } + msg := NewMessage(topicA, []byte(fmt.Sprintf("message %d", i))) sentMessages = append(sentMessages, msg) @@ -212,9 +186,87 @@ func TestPublishAndSubscribe(t *testing.T) { select { case <-consumerFinCh: break - case <-time.After(time.Second): + case <-time.After(time.Second * 5): t.Fatal("timed out waiting for consumer to read messages") } + // THIS IS SO HACKY + for _, msg := range receivedMessages { + msg.ack = nil + } + + for _, msg := range sentMessages { + msg.ack = nil + } + assert.ElementsMatch(t, receivedMessages, sentMessages) } + +func TestPublishAndSubscribeNackMessage(t *testing.T) { + consumer, cancel := setupConsumer(t) + + var receivedMessages []*Message + + consumerFinCh := make(chan struct{}) + timesMsgWasReceived := 0 + go func() { + for msg := range consumer.Messages() { + msg.Ack(false) + timesMsgWasReceived++ + } + + consumerFinCh <- struct{}{} + }() + + publisher, err := NewPublisher("localhost:9999") + require.NoError(t, err) + t.Cleanup(func() { + publisher.Close() + }) + + // send a message + msg := NewMessage(topicA, []byte("hello world")) + + err = publisher.PublishMessage(msg) + require.NoError(t, err) + + // give the consumer some time to read the messages -- TODO: make better! + time.Sleep(time.Millisecond * 500) + cancel() + + select { + case <-consumerFinCh: + break + case <-time.After(time.Second * 5): + t.Fatal("timed out waiting for consumer to read messages") + } + + assert.Empty(t, receivedMessages) + assert.Equal(t, 5, timesMsgWasReceived) +} + +func setupConsumer(t *testing.T) (*Consumer, context.CancelFunc) { + createServer(t) + + sub, err := NewSubscriber(serverAddr) + require.NoError(t, err) + + t.Cleanup(func() { + sub.Close() + }) + + topics := []string{topicA, topicB} + + err = sub.SubscribeToTopics(topics) + require.NoError(t, err) + + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(func() { + cancel() + }) + + consumer := sub.Consume(ctx) + require.NoError(t, err) + + return consumer, cancel +} diff --git a/server/server.go b/server/server.go index 82f9d7d..fac9b8e 100644 --- a/server/server.go +++ b/server/server.go @@ -23,6 +23,8 @@ const ( Subscribe Action = 1 Unsubscribe Action = 2 Publish Action = 3 + Ack Action = 4 + Nack Action = 5 ) // Status represents the status of a request @@ -54,18 +56,23 @@ type Server struct { mu sync.Mutex topics map[string]*topic + + ackDelay time.Duration + ackTimeout time.Duration } // New creates and starts a new server -func New(Addr string) (*Server, error) { +func New(Addr string, ackDelay, ackTimeout time.Duration) (*Server, error) { lis, err := net.Listen("tcp", Addr) if err != nil { return nil, fmt.Errorf("failed to listen: %w", err) } srv := &Server{ - lis: lis, - topics: map[string]*topic{}, + lis: lis, + topics: map[string]*topic{}, + ackDelay: ackDelay, + ackTimeout: ackTimeout, } go srv.start() @@ -337,10 +344,7 @@ func (s *Server) addSubsciberToTopic(topicName string, peer *peer.Peer) { t = newTopic(topicName) } - t.subscriptions[peer.Addr()] = subscriber{ - peer: peer, - currentOffset: 0, - } + t.subscriptions[peer.Addr()] = newSubscriber(peer, topicName, s.ackDelay, s.ackTimeout) s.topics[topicName] = t } @@ -388,10 +392,14 @@ func readAction(peer *peer.Peer, timeout time.Duration) (Action, error) { var action Action op := func(conn net.Conn) error { if timeout > 0 { - err := conn.SetReadDeadline(time.Now().Add(timeout)) - if err != nil { + if err := conn.SetReadDeadline(time.Now().Add(timeout)); err != nil { slog.Error("failed to set connection read deadline", "error", err, "peer", peer.Addr()) } + defer func() { + if err := conn.SetReadDeadline(time.Time{}); err != nil { + slog.Error("failed to reset connection read deadline", "error", err, "peer", peer.Addr()) + } + }() } err := binary.Read(conn, binary.BigEndian, &action) diff --git a/server/server_test.go b/server/server_test.go index 0b2b02b..9660cc6 100644 --- a/server/server_test.go +++ b/server/server_test.go @@ -18,10 +18,13 @@ const ( topicC = "topic c" serverAddr = ":6666" + + ackDelay = time.Millisecond * 100 + ackTimeout = time.Millisecond * 100 ) func createServer(t *testing.T) *Server { - srv, err := New(serverAddr) + srv, err := New(serverAddr, ackDelay, ackTimeout) require.NoError(t, err) t.Cleanup(func() { @@ -35,7 +38,7 @@ func createServerWithExistingTopic(t *testing.T, topicName string) *Server { srv := createServer(t) srv.topics[topicName] = &topic{ name: topicName, - subscriptions: make(map[net.Addr]subscriber), + subscriptions: make(map[net.Addr]*subscriber), } return srv @@ -268,9 +271,13 @@ func TestSendsDataToTopicSubscribers(t *testing.T) { buf := make([]byte, dataLen) n, err := conn.Read(buf) require.NoError(t, err) + require.Equal(t, int(dataLen), n) assert.Equal(t, messageData, string(buf)) + + err = binary.Write(conn, binary.BigEndian, Ack) + require.NoError(t, err) } } @@ -314,6 +321,9 @@ func TestPublishMultipleTimes(t *testing.T) { require.Equal(t, int(dataLen), n) results = append(results, string(buf)) + + err = binary.Write(subscriberConn, binary.BigEndian, Ack) + require.NoError(t, err) } assert.ElementsMatch(t, results, messages) @@ -346,3 +356,208 @@ func TestPublishMultipleTimes(t *testing.T) { t.Fatal(fmt.Errorf("timed out waiting for subscriber to read messages")) } } + +func TestSendsDataToTopicSubscriberNacksThenAcks(t *testing.T) { + _ = createServer(t) + + subscriberConn := createConnectionAndSubscribe(t, []string{topicA, topicB}) + + publisherConn, err := net.Dial("tcp", fmt.Sprintf("localhost%s", serverAddr)) + require.NoError(t, err) + + err = binary.Write(publisherConn, binary.BigEndian, Publish) + require.NoError(t, err) + + topic := fmt.Sprintf("topic:%s", topicA) + messageData := "hello world" + + // send topic first + err = binary.Write(publisherConn, binary.BigEndian, uint32(len(topic))) + require.NoError(t, err) + _, err = publisherConn.Write([]byte(topic)) + require.NoError(t, err) + + // now send the data + err = binary.Write(publisherConn, binary.BigEndian, uint32(len(messageData))) + require.NoError(t, err) + n, err := publisherConn.Write([]byte(messageData)) + require.NoError(t, err) + require.Equal(t, len(messageData), n) + + // check the subsribers got the data + readMessage := func(conn net.Conn, ack Action) { + var topicLen uint64 + err = binary.Read(conn, binary.BigEndian, &topicLen) + require.NoError(t, err) + + topicBuf := make([]byte, topicLen) + _, err = conn.Read(topicBuf) + require.NoError(t, err) + assert.Equal(t, topicA, string(topicBuf)) + + var dataLen uint64 + err = binary.Read(conn, binary.BigEndian, &dataLen) + require.NoError(t, err) + + buf := make([]byte, dataLen) + n, err := conn.Read(buf) + require.NoError(t, err) + + require.Equal(t, int(dataLen), n) + + assert.Equal(t, messageData, string(buf)) + + err = binary.Write(conn, binary.BigEndian, ack) + require.NoError(t, err) + } + + // NACK the message and then ack it + readMessage(subscriberConn, Nack) + readMessage(subscriberConn, Ack) + // reading for another message should now timeout but give enough time for the ack delay to kick in + // should the second read of the message not have been ack'd properly + var topicLen uint64 + _ = subscriberConn.SetReadDeadline(time.Now().Add(ackDelay + time.Millisecond*100)) + err = binary.Read(subscriberConn, binary.BigEndian, &topicLen) + require.Error(t, err) +} + +func TestSendsDataToTopicSubscriberDoesntAckMessage(t *testing.T) { + _ = createServer(t) + + subscriberConn := createConnectionAndSubscribe(t, []string{topicA, topicB}) + + publisherConn, err := net.Dial("tcp", fmt.Sprintf("localhost%s", serverAddr)) + require.NoError(t, err) + + err = binary.Write(publisherConn, binary.BigEndian, Publish) + require.NoError(t, err) + + topic := fmt.Sprintf("topic:%s", topicA) + messageData := "hello world" + + // send topic first + err = binary.Write(publisherConn, binary.BigEndian, uint32(len(topic))) + require.NoError(t, err) + _, err = publisherConn.Write([]byte(topic)) + require.NoError(t, err) + + // now send the data + err = binary.Write(publisherConn, binary.BigEndian, uint32(len(messageData))) + require.NoError(t, err) + n, err := publisherConn.Write([]byte(messageData)) + require.NoError(t, err) + require.Equal(t, len(messageData), n) + + // check the subsribers got the data + readMessage := func(conn net.Conn, ack bool) { + var topicLen uint64 + err = binary.Read(conn, binary.BigEndian, &topicLen) + require.NoError(t, err) + + topicBuf := make([]byte, topicLen) + _, err = conn.Read(topicBuf) + require.NoError(t, err) + assert.Equal(t, topicA, string(topicBuf)) + + var dataLen uint64 + err = binary.Read(conn, binary.BigEndian, &dataLen) + require.NoError(t, err) + + buf := make([]byte, dataLen) + n, err := conn.Read(buf) + require.NoError(t, err) + + require.Equal(t, int(dataLen), n) + + assert.Equal(t, messageData, string(buf)) + + if ack { + err = binary.Write(conn, binary.BigEndian, Ack) + require.NoError(t, err) + return + } + } + + // don't send ack or nack and then ack on the second attempt + readMessage(subscriberConn, false) + readMessage(subscriberConn, true) + + // reading for another message should now timeout but give enough time for the ack delay to kick in + // should the second read of the message not have been ack'd properly + var topicLen uint64 + _ = subscriberConn.SetReadDeadline(time.Now().Add(ackDelay + time.Millisecond*100)) + err = binary.Read(subscriberConn, binary.BigEndian, &topicLen) + require.Error(t, err) +} + +func TestSendsDataToTopicSubscriberDeliveryCountTooHighWithNoAck(t *testing.T) { + _ = createServer(t) + + subscriberConn := createConnectionAndSubscribe(t, []string{topicA, topicB}) + + publisherConn, err := net.Dial("tcp", fmt.Sprintf("localhost%s", serverAddr)) + require.NoError(t, err) + + err = binary.Write(publisherConn, binary.BigEndian, Publish) + require.NoError(t, err) + + topic := fmt.Sprintf("topic:%s", topicA) + messageData := "hello world" + + // send topic first + err = binary.Write(publisherConn, binary.BigEndian, uint32(len(topic))) + require.NoError(t, err) + _, err = publisherConn.Write([]byte(topic)) + require.NoError(t, err) + + // now send the data + err = binary.Write(publisherConn, binary.BigEndian, uint32(len(messageData))) + require.NoError(t, err) + n, err := publisherConn.Write([]byte(messageData)) + require.NoError(t, err) + require.Equal(t, len(messageData), n) + + // check the subsribers got the data + readMessage := func(conn net.Conn, ack bool) { + var topicLen uint64 + err = binary.Read(conn, binary.BigEndian, &topicLen) + require.NoError(t, err) + + topicBuf := make([]byte, topicLen) + _, err = conn.Read(topicBuf) + require.NoError(t, err) + assert.Equal(t, topicA, string(topicBuf)) + + var dataLen uint64 + err = binary.Read(conn, binary.BigEndian, &dataLen) + require.NoError(t, err) + + buf := make([]byte, dataLen) + n, err := conn.Read(buf) + require.NoError(t, err) + + require.Equal(t, int(dataLen), n) + + assert.Equal(t, messageData, string(buf)) + + if ack { + err = binary.Write(conn, binary.BigEndian, Ack) + require.NoError(t, err) + return + } + } + + // nack the message 5 times + readMessage(subscriberConn, false) + readMessage(subscriberConn, false) + readMessage(subscriberConn, false) + readMessage(subscriberConn, false) + readMessage(subscriberConn, false) + + // reading for the message should now timeout as we have nack'd the message too many times + var topicLen uint64 + _ = subscriberConn.SetReadDeadline(time.Now().Add(ackDelay + time.Millisecond*100)) + err = binary.Read(subscriberConn, binary.BigEndian, &topicLen) + require.Error(t, err) +} diff --git a/server/subscriber.go b/server/subscriber.go new file mode 100644 index 0000000..a182070 --- /dev/null +++ b/server/subscriber.go @@ -0,0 +1,124 @@ +package server + +import ( + "encoding/binary" + "fmt" + "log/slog" + "net" + "time" + + "github.com/willdot/messagebroker/server/peer" +) + +type subscriber struct { + peer *peer.Peer + topic string + messages chan message + + ackDelay time.Duration + ackTimeout time.Duration +} + +type message struct { + data []byte + deliveryCount int +} + +func newMessage(data []byte) message { + return message{data: data, deliveryCount: 1} +} + +func newSubscriber(peer *peer.Peer, topic string, ackDelay, ackTimeout time.Duration) *subscriber { + s := &subscriber{ + peer: peer, + topic: topic, + messages: make(chan message), + ackDelay: ackDelay, + ackTimeout: ackTimeout, + } + + go s.sendMessages() + + return s +} + +func (s *subscriber) sendMessages() { + // TODO: should think about how to break out of this if the subsciber closes its connection etc + for msg := range s.messages { + ack, err := s.sendMessage(s.topic, msg) + if err != nil { + slog.Error("failed to send to message", "error", err, "peer", s.peer.Addr()) + } + + if ack { + continue + } + + if msg.deliveryCount >= 5 { + slog.Error("max delivery count for message. Dropping", "peer", s.peer.Addr()) + continue + } + + msg.deliveryCount++ + s.addMessage(msg, s.ackDelay) + } +} + +func (s *subscriber) addMessage(msg message, delay time.Duration) { + go func() { + time.Sleep(delay) + // TODO: should think about how to break out of this if the subsciber closes its connection etc + s.messages <- msg + }() +} + +func (s *subscriber) sendMessage(topic string, msg message) (bool, error) { + var ack bool + op := func(conn net.Conn) error { + topicLen := uint64(len(topic)) + err := binary.Write(conn, binary.BigEndian, topicLen) + if err != nil { + return fmt.Errorf("failed to send topic length: %w", err) + } + _, err = conn.Write([]byte(topic)) + if err != nil { + return fmt.Errorf("failed to send topic: %w", err) + } + + dataLen := uint64(len(msg.data)) + + err = binary.Write(conn, binary.BigEndian, dataLen) + if err != nil { + return fmt.Errorf("failed to send data length: %w", err) + } + + _, err = conn.Write(msg.data) + if err != nil { + return fmt.Errorf("failed to write to peer: %w", err) + } + + var ackRes Action + if err := conn.SetReadDeadline(time.Now().Add(s.ackTimeout)); err != nil { + slog.Error("failed to set connection read deadline", "error", err, "peer", s.peer.Addr()) + } + defer func() { + if err := conn.SetReadDeadline(time.Time{}); err != nil { + slog.Error("failed to reset connection read deadline", "error", err, "peer", s.peer.Addr()) + } + }() + err = binary.Read(conn, binary.BigEndian, &ackRes) + if err != nil { + return fmt.Errorf("failed to read ack from peer: %w", err) + } + + if ackRes == Ack { + ack = true + } + + return nil + } + + err := s.peer.RunConnOperation(op) + + return ack, err +} diff --git a/server/topic.go b/server/topic.go index a087425..723ed5c 100644 --- a/server/topic.go +++ b/server/topic.go @@ -1,30 +1,20 @@ package server import ( - "encoding/binary" - "fmt" - "log/slog" "net" "sync" - - "github.com/willdot/messagebroker/server/peer" ) type topic struct { name string - subscriptions map[net.Addr]subscriber + subscriptions map[net.Addr]*subscriber mu sync.Mutex } -type subscriber struct { - peer *peer.Peer - currentOffset int -} - func newTopic(name string) *topic { return &topic{ name: name, - subscriptions: make(map[net.Addr]subscriber), + subscriptions: make(map[net.Addr]*subscriber), } } @@ -33,51 +23,7 @@ func (t *topic) sendMessageToSubscribers(msgData []byte) { subscribers := t.subscriptions t.mu.Unlock() - 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 - } -} - -func sendMessageOp(topic string, data []byte) peer.ConnOpp { - return func(conn net.Conn) error { - topicLen := uint64(len(topic)) - err := binary.Write(conn, binary.BigEndian, topicLen) - if err != nil { - return fmt.Errorf("failed to send topic length: %w", err) - } - _, err = conn.Write([]byte(topic)) - if err != nil { - return fmt.Errorf("failed to send topic: %w", err) - } - - dataLen := uint64(len(data)) - - err = binary.Write(conn, binary.BigEndian, dataLen) - if err != nil { - return fmt.Errorf("failed to send data length: %w", err) - } - - _, err = conn.Write(data) - if err != nil { - return fmt.Errorf("failed to write to peer: %w", err) - } - return nil + subscriber.addMessage(newMessage(msgData), 0) } } -- 2.51.2 From 3aae3cde7f744f487719800e59ee5d072a63463c Mon Sep 17 00:00:00 2001 From: Will Date: Sat, 16 Dec 2023 21:39:36 +0000 Subject: [PATCH 2/3] better handling when unsubscribing --- server/server.go | 11 ++++++- server/subscriber.go | 68 +++++++++++++++++++++++++++----------------- 2 files changed, 52 insertions(+), 27 deletions(-) diff --git a/server/server.go b/server/server.go index fac9b8e..c5c7f2c 100644 --- a/server/server.go +++ b/server/server.go @@ -364,7 +364,11 @@ func (s *Server) removeSubsciberFromTopic(topicName string, peer *peer.Peer) { if !ok { return } - + sub, ok := t.subscriptions[peer.Addr()] + if !ok { + return + } + sub.unsubscribe() delete(t.subscriptions, peer.Addr()) } @@ -373,6 +377,11 @@ func (s *Server) unsubscribePeerFromAllTopics(peer *peer.Peer) { defer s.mu.Unlock() for _, topic := range s.topics { + sub, ok := topic.subscriptions[peer.Addr()] + if !ok { + continue + } + sub.unsubscribe() delete(topic.subscriptions, peer.Addr()) } } diff --git a/server/subscriber.go b/server/subscriber.go index a182070..344baea 100644 --- a/server/subscriber.go +++ b/server/subscriber.go @@ -11,9 +11,10 @@ import ( ) type subscriber struct { - peer *peer.Peer - topic string - messages chan message + peer *peer.Peer + topic string + messages chan message + unsubscribeCh chan struct{} ackDelay time.Duration ackTimeout time.Duration @@ -30,11 +31,12 @@ func newMessage(data []byte) message { func newSubscriber(peer *peer.Peer, topic string, ackDelay, ackTimeout time.Duration) *subscriber { s := &subscriber{ - peer: peer, - topic: topic, - messages: make(chan message), - ackDelay: ackDelay, - ackTimeout: ackTimeout, + peer: peer, + topic: topic, + messages: make(chan message), + ackDelay: ackDelay, + ackTimeout: ackTimeout, + unsubscribeCh: make(chan struct{}), } go s.sendMessages() @@ -43,32 +45,42 @@ func newSubscriber(peer *peer.Peer, topic string, ackDelay, ackTimeout time.Dura } func (s *subscriber) sendMessages() { - // TODO: should think about how to break out of this if the subsciber closes its connection etc - for msg := range s.messages { - ack, err := s.sendMessage(s.topic, msg) - if err != nil { - slog.Error("failed to send to message", "error", err, "peer", s.peer.Addr()) - } + for { + select { + case <-s.unsubscribeCh: + return + case msg := <-s.messages: + ack, err := s.sendMessage(s.topic, msg) + if err != nil { + slog.Error("failed to send to message", "error", err, "peer", s.peer.Addr()) + } - if ack { - continue - } + if ack { + continue + } - if msg.deliveryCount >= 5 { - slog.Error("max delivery count for message. Dropping", "peer", s.peer.Addr()) - continue - } + if msg.deliveryCount >= 5 { + slog.Error("max delivery count for message. Dropping", "peer", s.peer.Addr()) + continue + } - msg.deliveryCount++ - s.addMessage(msg, s.ackDelay) + msg.deliveryCount++ + s.addMessage(msg, s.ackDelay) + } } } func (s *subscriber) addMessage(msg message, delay time.Duration) { go func() { - time.Sleep(delay) - // TODO: should think about how to break out of this if the subsciber closes its connection etc - s.messages <- msg + timer := time.NewTimer(delay) + defer timer.Stop() + + select { + case <-s.unsubscribeCh: + return + case <-timer.C: + s.messages <- msg + } }() } @@ -122,3 +134,7 @@ func (s *subscriber) sendMessage(topic string, msg message) (bool, error) { return ack, err } + +func (s *subscriber) unsubscribe() { + close(s.unsubscribeCh) +} -- 2.51.2 From 091cfb46b56e8472ea502b53fdd48b8235efaea3 Mon Sep 17 00:00:00 2001 From: Will Date: Sat, 16 Dec 2023 21:52:58 +0000 Subject: [PATCH 3/3] Tidy up --- pubsub/subscriber.go | 13 +++++++++---- pubsub/subscriber_test.go | 5 +++-- 2 files changed, 12 insertions(+), 6 deletions(-) diff --git a/pubsub/subscriber.go b/pubsub/subscriber.go index e01b87f..b316a17 100644 --- a/pubsub/subscriber.go +++ b/pubsub/subscriber.go @@ -173,7 +173,7 @@ func (s *Subscriber) consume(ctx context.Context, consumer *Consumer) { return } - err := s.readMessage(consumer.msgs) + err := s.readMessage(ctx, consumer.msgs) if err != nil { consumer.Err = err return @@ -181,8 +181,7 @@ func (s *Subscriber) consume(ctx context.Context, consumer *Consumer) { } } -func (s *Subscriber) readMessage(msgChan chan *Message) error { - // var msg *Message +func (s *Subscriber) readMessage(ctx context.Context, msgChan chan *Message) error { op := func(conn net.Conn) error { err := s.conn.SetReadDeadline(time.Now().Add(time.Second)) if err != nil { @@ -225,7 +224,13 @@ func (s *Subscriber) readMessage(msgChan chan *Message) error { msgChan <- msg - ack := <-msg.ack + var ack bool + select { + case <-ctx.Done(): + return ctx.Err() + case ack = <-msg.ack: + } + //ack := <-msg.ack ackMessage := server.Nack if ack { diff --git a/pubsub/subscriber_test.go b/pubsub/subscriber_test.go index 0463aff..e4834a8 100644 --- a/pubsub/subscriber_test.go +++ b/pubsub/subscriber_test.go @@ -134,6 +134,7 @@ func TestUnsubscribesFromTopic(t *testing.T) { err = publisher.PublishMessage(msg) require.NoError(t, err) + time.Sleep(time.Second) cancel() select { @@ -180,7 +181,7 @@ func TestPublishAndSubscribe(t *testing.T) { } // give the consumer some time to read the messages -- TODO: make better! - time.Sleep(time.Millisecond * 500) + time.Sleep(time.Second) cancel() select { @@ -231,7 +232,7 @@ func TestPublishAndSubscribeNackMessage(t *testing.T) { require.NoError(t, err) // give the consumer some time to read the messages -- TODO: make better! - time.Sleep(time.Millisecond * 500) + time.Sleep(time.Second) cancel() select {