diff --git a/example/main.go b/example/main.go index 97e817e..d4dc09e 100644 --- a/example/main.go +++ b/example/main.go @@ -8,12 +8,15 @@ import ( "time" "github.com/willdot/messagebroker/pubsub" + "github.com/willdot/messagebroker/server" ) var consumeOnly *bool +var consumeFrom *int func main() { consumeOnly = flag.Bool("consume-only", false, "just consumes (doesn't start server and doesn't publish)") + consumeFrom = flag.Int("consume-from", -1, "index of message to start consuming from. If not set it will consume from the most recent") flag.Parse() if !*consumeOnly { @@ -28,8 +31,14 @@ func main() { defer func() { _ = sub.Close() }() + startAt := 0 + startAtType := server.Current + if *consumeFrom >= 0-1 { + startAtType = server.From + startAt = *consumeFrom + } - err = sub.SubscribeToTopics([]string{"topic a"}) + err = sub.SubscribeToTopics([]string{"topic a"}, startAtType, startAt) if err != nil { panic(err) } diff --git a/example/server/main.go b/example/server/main.go index 3841b19..4dece49 100644 --- a/example/server/main.go +++ b/example/server/main.go @@ -11,7 +11,7 @@ import ( ) func main() { - srv, err := server.New(":3000", time.Second, time.Second*2) + srv, err := server.New(":3000", time.Second, time.Second*2, server.NewMemoryStore()) if err != nil { log.Fatal(err) } diff --git a/pubsub/subscriber.go b/pubsub/subscriber.go index d146d49..16ab920 100644 --- a/pubsub/subscriber.go +++ b/pubsub/subscriber.go @@ -39,101 +39,120 @@ func (s *Subscriber) Close() error { } // SubscribeToTopics will subscribe to the provided topics -func (s *Subscriber) SubscribeToTopics(topicNames []string) error { +func (s *Subscriber) SubscribeToTopics(topicNames []string, startAtType server.StartAtType, startAtIndex int) error { op := func(conn net.Conn) error { - actionB := make([]byte, 2) - binary.BigEndian.PutUint16(actionB, server.Subscribed) - headers := actionB + return subscribeToTopics(conn, topicNames, startAtType, startAtIndex) + } - b, err := json.Marshal(topicNames) - if err != nil { - return fmt.Errorf("failed to marshal topic names: %w", err) - } + return s.connOperation(op) +} - topicNamesB := make([]byte, 4) - binary.BigEndian.PutUint32(topicNamesB, uint32(len(b))) - headers = append(headers, topicNamesB...) +// UnsubscribeToTopics will unsubscribe to the provided topics +func (s *Subscriber) UnsubscribeToTopics(topicNames []string) error { + op := func(conn net.Conn) error { + return unsubscribeToTopics(conn, topicNames) + } - _, err = conn.Write(append(headers, b...)) - if err != nil { - return fmt.Errorf("failed to subscribe to topics: %w", err) - } + return s.connOperation(op) +} - var resp server.Status - err = binary.Read(conn, binary.BigEndian, &resp) - if err != nil { - return fmt.Errorf("failed to read confirmation of subscription: %w", err) - } +func subscribeToTopics(conn net.Conn, topicNames []string, startAtType server.StartAtType, startAtIndex int) error { + actionB := make([]byte, 2) + binary.BigEndian.PutUint16(actionB, uint16(server.Subscribe)) + headers := actionB - if resp == server.Subscribed { - return nil - } + b, err := json.Marshal(topicNames) + if err != nil { + return fmt.Errorf("failed to marshal topic names: %w", err) + } - var dataLen uint32 - err = binary.Read(conn, binary.BigEndian, &dataLen) - if err != nil { - return fmt.Errorf("received status %s:", resp) - } + topicNamesB := make([]byte, 4) + binary.BigEndian.PutUint32(topicNamesB, uint32(len(b))) + headers = append(headers, topicNamesB...) + headers = append(headers, b...) - buf := make([]byte, dataLen) - _, err = conn.Read(buf) - if err != nil { - return fmt.Errorf("received status %s:", resp) - } + startAtTypeB := make([]byte, 2) + binary.BigEndian.PutUint16(startAtTypeB, uint16(startAtType)) + headers = append(headers, startAtTypeB...) - return fmt.Errorf("received status %s - %s", resp, buf) + if startAtType == server.From { + fromB := make([]byte, 2) + binary.BigEndian.PutUint16(fromB, uint16(startAtIndex)) + headers = append(headers, fromB...) } - return s.connOperation(op) -} + _, err = conn.Write(headers) + if err != nil { + return fmt.Errorf("failed to subscribe to topics: %w", err) + } -// UnsubscribeToTopics will unsubscribe to the provided topics -func (s *Subscriber) UnsubscribeToTopics(topicNames []string) error { - op := func(conn net.Conn) error { - actionB := make([]byte, 2) - binary.BigEndian.PutUint16(actionB, uint16(server.Unsubscribe)) - headers := actionB + var resp server.Status + err = binary.Read(conn, binary.BigEndian, &resp) + if err != nil { + return fmt.Errorf("failed to read confirmation of subscribe: %w", err) + } - b, err := json.Marshal(topicNames) - if err != nil { - return fmt.Errorf("failed to marshal topic names: %w", err) - } + if resp == server.Subscribed { + return nil + } - topicNamesB := make([]byte, 4) - binary.BigEndian.PutUint32(topicNamesB, uint32(len(b))) - headers = append(headers, topicNamesB...) + var dataLen uint32 + err = binary.Read(conn, binary.BigEndian, &dataLen) + if err != nil { + return fmt.Errorf("received status %s:", resp) + } - _, err = conn.Write(append(headers, b...)) - if err != nil { - return fmt.Errorf("failed to unsubscribe to topics: %w", err) - } + buf := make([]byte, dataLen) + _, err = conn.Read(buf) + if err != nil { + return fmt.Errorf("received status %s:", resp) + } - var resp server.Status - err = binary.Read(conn, binary.BigEndian, &resp) - if err != nil { - return fmt.Errorf("failed to read confirmation of unsubscription: %w", err) - } + return fmt.Errorf("received status %s - %s", resp, buf) +} - if resp == server.Unsubscribed { - return nil - } +func unsubscribeToTopics(conn net.Conn, topicNames []string) error { + actionB := make([]byte, 2) + binary.BigEndian.PutUint16(actionB, uint16(server.Unsubscribe)) + headers := actionB - var dataLen uint32 - err = binary.Read(conn, binary.BigEndian, &dataLen) - if err != nil { - return fmt.Errorf("received status %s:", resp) - } + b, err := json.Marshal(topicNames) + if err != nil { + return fmt.Errorf("failed to marshal topic names: %w", err) + } - buf := make([]byte, dataLen) - _, err = conn.Read(buf) - if err != nil { - return fmt.Errorf("received status %s:", resp) - } + topicNamesB := make([]byte, 4) + binary.BigEndian.PutUint32(topicNamesB, uint32(len(b))) + headers = append(headers, topicNamesB...) - return fmt.Errorf("received status %s - %s", resp, buf) + _, err = conn.Write(append(headers, b...)) + if err != nil { + return fmt.Errorf("failed to unsubscribe to topics: %w", err) } - return s.connOperation(op) + var resp server.Status + err = binary.Read(conn, binary.BigEndian, &resp) + if err != nil { + return fmt.Errorf("failed to read confirmation of unsubscribe: %w", err) + } + + if resp == server.Unsubscribed { + return nil + } + + var dataLen uint32 + err = binary.Read(conn, binary.BigEndian, &dataLen) + if err != nil { + return fmt.Errorf("received status %s:", resp) + } + + buf := make([]byte, dataLen) + _, err = conn.Read(buf) + if err != nil { + return fmt.Errorf("received status %s:", resp) + } + + return fmt.Errorf("received status %s - %s", resp, buf) } // Consumer allows the consumption of messages. If during the consumer receiving messages from the diff --git a/pubsub/subscriber_test.go b/pubsub/subscriber_test.go index e4834a8..a1c0b84 100644 --- a/pubsub/subscriber_test.go +++ b/pubsub/subscriber_test.go @@ -18,8 +18,19 @@ const ( topicB = "topic b" ) +type fakeStore struct { +} + +func (f *fakeStore) Write(msg server.MessageToSend) error { + return nil +} +func (f *fakeStore) ReadFrom(offset int, handleFunc func(msgs []server.MessageToSend)) error { + return nil +} + func createServer(t *testing.T) { - server, err := server.New(serverAddr, time.Millisecond*100, time.Millisecond*100) + fs := &fakeStore{} + server, err := server.New(serverAddr, time.Millisecond*100, time.Millisecond*100, fs) require.NoError(t, err) t.Cleanup(func() { @@ -75,7 +86,7 @@ func TestSubscribeToTopics(t *testing.T) { topics := []string{topicA, topicB} - err = sub.SubscribeToTopics(topics) + err = sub.SubscribeToTopics(topics, server.Current, 0) require.NoError(t, err) } @@ -91,7 +102,7 @@ func TestUnsubscribesFromTopic(t *testing.T) { topics := []string{topicA, topicB} - err = sub.SubscribeToTopics(topics) + err = sub.SubscribeToTopics(topics, server.Current, 0) require.NoError(t, err) err = sub.UnsubscribeToTopics([]string{topicA}) @@ -258,7 +269,7 @@ func setupConsumer(t *testing.T) (*Consumer, context.CancelFunc) { topics := []string{topicA, topicB} - err = sub.SubscribeToTopics(topics) + err = sub.SubscribeToTopics(topics, server.Current, 0) require.NoError(t, err) ctx, cancel := context.WithCancel(context.Background()) diff --git a/server/message_store.go b/server/message_store.go new file mode 100644 index 0000000..d936e1e --- /dev/null +++ b/server/message_store.go @@ -0,0 +1,47 @@ +package server + +import ( + "fmt" + "sync" +) + +type MemoryStore struct { + mu sync.Mutex + msgs map[int]MessageToSend + offset int +} + +func NewMemoryStore() *MemoryStore { + return &MemoryStore{ + msgs: make(map[int]MessageToSend), + } +} + +func (m *MemoryStore) Write(msg MessageToSend) error { + m.mu.Lock() + defer m.mu.Unlock() + + m.msgs[m.offset] = msg + + m.offset++ + + return nil +} + +func (m *MemoryStore) ReadFrom(offset int, handleFunc func(msgs []MessageToSend)) error { + if offset < 0 || offset > m.offset { + return fmt.Errorf("invalid offset provided") + } + + m.mu.Lock() + defer m.mu.Unlock() + + msgs := make([]MessageToSend, 0, len(m.msgs)) + for i := offset; i < len(m.msgs); i++ { + msgs = append(msgs, m.msgs[i]) + } + + handleFunc(msgs) + + return nil +} diff --git a/server/server.go b/server/server.go index 420b8a1..d9b4ed3 100644 --- a/server/server.go +++ b/server/server.go @@ -27,13 +27,30 @@ const ( Nack Action = 5 ) +func (a Action) String() string { + switch a { + case Subscribe: + return "subscribe" + case Unsubscribe: + return "unsubscribe" + case Publish: + return "publish" + case Ack: + return "ack" + case Nack: + return "nack" + } + + return "" +} + // Status represents the status of a request type Status uint16 const ( - Subscribed = 1 - Unsubscribed = 2 - Error = 3 + Subscribed Status = 1 + Unsubscribed Status = 2 + Error Status = 3 ) func (s Status) String() string { @@ -49,6 +66,20 @@ func (s Status) String() string { return "" } +// StartAtType represents where the subcriber wishes to start subscribing to a topic from +type StartAtType uint16 + +const ( + Begining StartAtType = 0 + Current StartAtType = 1 + From StartAtType = 2 +) + +type Store interface { + Write(msg MessageToSend) error + ReadFrom(offset int, handleFunc func(msgs []MessageToSend)) error +} + // Server accepts subscribe and publish connections and passes messages around type Server struct { Addr string @@ -59,20 +90,23 @@ type Server struct { ackDelay time.Duration ackTimeout time.Duration + + messageStore Store } // New creates and starts a new server -func New(Addr string, ackDelay, ackTimeout time.Duration) (*Server, error) { +func New(Addr string, ackDelay, ackTimeout time.Duration, messageStore Store) (*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{}, - ackDelay: ackDelay, - ackTimeout: ackTimeout, + lis: lis, + topics: map[string]*topic{}, + ackDelay: ackDelay, + ackTimeout: ackTimeout, + messageStore: messageStore, } go srv.start() @@ -191,6 +225,7 @@ func (s *Server) subscribePeerToTopic(peer *peer.Peer) { } var topics []string + fmt.Println(string(buf)) err = json.Unmarshal(buf, &topics) if err != nil { slog.Error("failed to unmarshal subscibers topic data", "error", err, "peer", peer.Addr()) @@ -198,7 +233,36 @@ func (s *Server) subscribePeerToTopic(peer *peer.Peer) { return nil } - s.subscribeToTopics(peer, topics) + var startAtType StartAtType + err = binary.Read(conn, binary.BigEndian, &startAtType) + if err != nil { + slog.Error(err.Error(), "peer", peer.Addr()) + writeStatus(Error, "invalid start at type provided", conn) + return nil + } + var startAt int + switch startAtType { + case From: + // read the from + var s uint16 + err = binary.Read(conn, binary.BigEndian, &s) + if err != nil { + slog.Error(err.Error(), "peer", peer.Addr()) + writeStatus(Error, "invalid start at value provided", conn) + return nil + } + startAt = int(s) + case Begining: + startAt = 0 + case Current: + startAt = -1 + default: + slog.Error("invalid start up type provided", "start up type", startAtType) + writeStatus(Error, "invalid start up type provided", conn) + return nil + } + + s.subscribeToTopics(peer, topics, startAt) writeStatus(Subscribed, "", conn) return nil @@ -247,7 +311,7 @@ func (s *Server) handleUnsubscribe(peer *peer.Peer) { _ = peer.RunConnOperation(op) } -type messageToSend struct { +type MessageToSend struct { topic string data []byte } @@ -255,7 +319,7 @@ type messageToSend struct { func (s *Server) handlePublish(peer *peer.Peer) { slog.Info("handling publisher", "peer", peer.Addr()) for { - var message *messageToSend + var message *MessageToSend op := func(conn net.Conn) error { dataLen, err := dataLength(conn) @@ -304,47 +368,48 @@ func (s *Server) handlePublish(peer *peer.Peer) { return nil } - message = &messageToSend{ + message = &MessageToSend{ topic: topicStr, data: dataBuf, } - return nil - } - _ = peer.RunConnOperation(op) + topic := s.getTopic(message.topic) + if topic == nil { + topic = newTopic(message.topic, s.messageStore) + s.topics[message.topic] = topic + } - if message == nil { - continue + err = topic.sendMessageToSubscribers(*message) + if err != nil { + slog.Error("failed to send message to subscribers", "error", err, "peer", peer.Addr()) + writeStatus(Error, "failed to send message to subscribers", conn) + return nil + } + + return nil } - // 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) - } - }() + _ = peer.RunConnOperation(op) } } -func (s *Server) subscribeToTopics(peer *peer.Peer, topics []string) { +func (s *Server) subscribeToTopics(peer *peer.Peer, topics []string, startAt int) { slog.Info("subscribing peer to topics", "topics", topics, "peer", peer.Addr()) for _, topic := range topics { - s.addSubsciberToTopic(topic, peer) + s.addSubsciberToTopic(topic, peer, startAt) } } -func (s *Server) addSubsciberToTopic(topicName string, peer *peer.Peer) { +func (s *Server) addSubsciberToTopic(topicName string, peer *peer.Peer, startAt int) { s.mu.Lock() defer s.mu.Unlock() t, ok := s.topics[topicName] if !ok { - t = newTopic(topicName) + t = newTopic(topicName, s.messageStore) } - t.subscriptions[peer.Addr()] = newSubscriber(peer, topicName, s.ackDelay, s.ackTimeout) + t.subscriptions[peer.Addr()] = newSubscriber(peer, topicName, s.ackDelay, s.ackTimeout, s.messageStore, startAt) s.topics[topicName] = t } diff --git a/server/server_test.go b/server/server_test.go index c12798b..3a836c2 100644 --- a/server/server_test.go +++ b/server/server_test.go @@ -24,7 +24,8 @@ const ( ) func createServer(t *testing.T) *Server { - srv, err := New(serverAddr, ackDelay, ackTimeout) + store := NewMemoryStore() + srv, err := New(serverAddr, ackDelay, ackTimeout, store) require.NoError(t, err) t.Cleanup(func() { @@ -44,11 +45,11 @@ func createServerWithExistingTopic(t *testing.T, topicName string) *Server { return srv } -func createConnectionAndSubscribe(t *testing.T, topics []string) net.Conn { +func createConnectionAndSubscribe(t *testing.T, topics []string, startAtType StartAtType, startAtIndex int) net.Conn { conn, err := net.Dial("tcp", fmt.Sprintf("localhost%s", serverAddr)) require.NoError(t, err) - subscribeOrUnsubscribetoTopics(t, conn, topics, Subscribe) + subscribeToTopics(t, conn, topics, startAtType, startAtIndex) expectedRes := Subscribed @@ -56,7 +57,7 @@ func createConnectionAndSubscribe(t *testing.T, topics []string) net.Conn { err = binary.Read(conn, binary.BigEndian, &resp) require.NoError(t, err) - assert.Equal(t, expectedRes, int(resp)) + assert.Equal(t, expectedRes, resp) return conn } @@ -76,9 +77,36 @@ func sendMessage(t *testing.T, conn net.Conn, topic string, message []byte) { require.NoError(t, err) } -func subscribeOrUnsubscribetoTopics(t *testing.T, conn net.Conn, topics []string, action Action) { +func subscribeToTopics(t *testing.T, conn net.Conn, topics []string, startAtType StartAtType, startAtIndex int) { actionB := make([]byte, 2) - binary.BigEndian.PutUint16(actionB, uint16(action)) + binary.BigEndian.PutUint16(actionB, uint16(Subscribe)) + headers := actionB + + b, err := json.Marshal(topics) + require.NoError(t, err) + + topicNamesB := make([]byte, 4) + binary.BigEndian.PutUint32(topicNamesB, uint32(len(b))) + headers = append(headers, topicNamesB...) + headers = append(headers, b...) + + startAtTypeB := make([]byte, 2) + binary.BigEndian.PutUint16(startAtTypeB, uint16(startAtType)) + headers = append(headers, startAtTypeB...) + + if startAtType == From { + fromB := make([]byte, 2) + binary.BigEndian.PutUint16(fromB, uint16(startAtIndex)) + headers = append(headers, fromB...) + } + + _, err = conn.Write(headers) + require.NoError(t, err) +} + +func unsubscribetoTopics(t *testing.T, conn net.Conn, topics []string) { + actionB := make([]byte, 2) + binary.BigEndian.PutUint16(actionB, uint16(Unsubscribe)) headers := actionB b, err := json.Marshal(topics) @@ -97,7 +125,7 @@ func TestSubscribeToTopics(t *testing.T) { // existing topic srv := createServerWithExistingTopic(t, topicA) - _ = createConnectionAndSubscribe(t, []string{topicA, topicB}) + _ = createConnectionAndSubscribe(t, []string{topicA, topicB}, Current, 0) assert.Len(t, srv.topics, 2) assert.Len(t, srv.topics[topicA].subscriptions, 1) @@ -107,7 +135,7 @@ func TestSubscribeToTopics(t *testing.T) { func TestUnsubscribesFromTopic(t *testing.T) { srv := createServerWithExistingTopic(t, topicA) - conn := createConnectionAndSubscribe(t, []string{topicA, topicB, topicC}) + conn := createConnectionAndSubscribe(t, []string{topicA, topicB, topicC}, Current, 0) assert.Len(t, srv.topics, 3) assert.Len(t, srv.topics[topicA].subscriptions, 1) @@ -116,7 +144,7 @@ func TestUnsubscribesFromTopic(t *testing.T) { topics := []string{topicA, topicB} - subscribeOrUnsubscribetoTopics(t, conn, topics, Unsubscribe) + unsubscribetoTopics(t, conn, topics) expectedRes := Unsubscribed @@ -124,7 +152,7 @@ func TestUnsubscribesFromTopic(t *testing.T) { err := binary.Read(conn, binary.BigEndian, &resp) require.NoError(t, err) - assert.Equal(t, expectedRes, int(resp)) + assert.Equal(t, expectedRes, resp) assert.Len(t, srv.topics, 3) assert.Len(t, srv.topics[topicA].subscriptions, 0) @@ -135,7 +163,7 @@ func TestUnsubscribesFromTopic(t *testing.T) { func TestSubscriberClosesWithoutUnsubscribing(t *testing.T) { srv := createServer(t) - conn := createConnectionAndSubscribe(t, []string{topicA, topicB}) + conn := createConnectionAndSubscribe(t, []string{topicA, topicB}, Current, 0) assert.Len(t, srv.topics, 2) assert.Len(t, srv.topics[topicA].subscriptions, 1) @@ -175,7 +203,7 @@ func TestInvalidAction(t *testing.T) { err = binary.Read(conn, binary.BigEndian, &resp) require.NoError(t, err) - assert.Equal(t, expectedRes, int(resp)) + assert.Equal(t, expectedRes, resp) expectedMessage := "unknown action" @@ -213,7 +241,7 @@ func TestInvalidTopicDataPublished(t *testing.T) { err = binary.Read(publisherConn, binary.BigEndian, &resp) require.NoError(t, err) - assert.Equal(t, expectedRes, int(resp)) + assert.Equal(t, expectedRes, resp) expectedMessage := "topic data does not contain 'topic:' prefix" @@ -234,7 +262,7 @@ func TestSendsDataToTopicSubscribers(t *testing.T) { subscribers := make([]net.Conn, 0, 10) for i := 0; i < 10; i++ { - subscriberConn := createConnectionAndSubscribe(t, []string{topicA, topicB}) + subscriberConn := createConnectionAndSubscribe(t, []string{topicA, topicB}, Current, 0) subscribers = append(subscribers, subscriberConn) } @@ -252,29 +280,8 @@ func TestSendsDataToTopicSubscribers(t *testing.T) { // check the subsribers got the data for _, conn := range subscribers { - 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) + msg := readMessage(t, conn) + assert.Equal(t, messageData, string(msg)) } } @@ -294,33 +301,13 @@ func TestPublishMultipleTimes(t *testing.T) { subscribeFinCh := make(chan struct{}) // create a subscriber that will read messages - subscriberConn := createConnectionAndSubscribe(t, []string{topicA, topicB}) + subscriberConn := createConnectionAndSubscribe(t, []string{topicA, topicB}, Current, 0) go func() { // check subscriber got all 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) - - topicBuf := make([]byte, topicLen) - _, err = subscriberConn.Read(topicBuf) - require.NoError(t, err) - assert.Equal(t, topicA, string(topicBuf)) - - var dataLen uint64 - err = binary.Read(subscriberConn, binary.BigEndian, &dataLen) - require.NoError(t, err) - - buf := make([]byte, dataLen) - n, err := subscriberConn.Read(buf) - require.NoError(t, err) - require.Equal(t, int(dataLen), n) - - results = append(results, string(buf)) - - err = binary.Write(subscriberConn, binary.BigEndian, Ack) - require.NoError(t, err) + msg := readMessage(t, subscriberConn) + results = append(results, string(msg)) } assert.ElementsMatch(t, results, messages) @@ -346,7 +333,7 @@ func TestPublishMultipleTimes(t *testing.T) { func TestSendsDataToTopicSubscriberNacksThenAcks(t *testing.T) { _ = createServer(t) - subscriberConn := createConnectionAndSubscribe(t, []string{topicA, topicB}) + subscriberConn := createConnectionAndSubscribe(t, []string{topicA, topicB}, Current, 0) publisherConn, err := net.Dial("tcp", fmt.Sprintf("localhost%s", serverAddr)) require.NoError(t, err) @@ -400,7 +387,7 @@ func TestSendsDataToTopicSubscriberNacksThenAcks(t *testing.T) { func TestSendsDataToTopicSubscriberDoesntAckMessage(t *testing.T) { _ = createServer(t) - subscriberConn := createConnectionAndSubscribe(t, []string{topicA, topicB}) + subscriberConn := createConnectionAndSubscribe(t, []string{topicA, topicB}, Current, 0) publisherConn, err := net.Dial("tcp", fmt.Sprintf("localhost%s", serverAddr)) require.NoError(t, err) @@ -458,7 +445,7 @@ func TestSendsDataToTopicSubscriberDoesntAckMessage(t *testing.T) { func TestSendsDataToTopicSubscriberDeliveryCountTooHighWithNoAck(t *testing.T) { _ = createServer(t) - subscriberConn := createConnectionAndSubscribe(t, []string{topicA, topicB}) + subscriberConn := createConnectionAndSubscribe(t, []string{topicA, topicB}, Current, 0) publisherConn, err := net.Dial("tcp", fmt.Sprintf("localhost%s", serverAddr)) require.NoError(t, err) @@ -514,3 +501,106 @@ func TestSendsDataToTopicSubscriberDeliveryCountTooHighWithNoAck(t *testing.T) { err = binary.Read(subscriberConn, binary.BigEndian, &topicLen) require.Error(t, err) } + +func TestSubscribeAndReplaysFromStart(t *testing.T) { + _ = createServer(t) + + 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) + + messages := make([]string, 0, 10) + for i := 0; i < 10; i++ { + messages = append(messages, fmt.Sprintf("message %d", i)) + } + + // send messages first + topic := fmt.Sprintf("topic:%s", topicA) + + // send multiple messages + for _, msg := range messages { + sendMessage(t, publisherConn, topic, []byte(msg)) + } + + subscriberConn := createConnectionAndSubscribe(t, []string{topicA}, From, 0) + results := make([]string, 0, len(messages)) + for i := 0; i < len(messages); i++ { + msg := readMessage(t, subscriberConn) + results = append(results, string(msg)) + } + assert.ElementsMatch(t, results, messages) +} + +func TestSubscribeAndReplaysFromIndex(t *testing.T) { + _ = createServer(t) + + 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) + + messages := make([]string, 0, 10) + for i := 0; i < 10; i++ { + messages = append(messages, fmt.Sprintf("message %d", i)) + } + + // send messages first + topic := fmt.Sprintf("topic:%s", topicA) + + // send multiple messages + for _, msg := range messages { + sendMessage(t, publisherConn, topic, []byte(msg)) + } + + subscriberConn := createConnectionAndSubscribe(t, []string{topicA}, From, 3) + + // now that the subscriber has subecribed send another message that should arrive after all the other messages were consumed + sendMessage(t, publisherConn, topic, []byte("hello there")) + + results := make([]string, 0, len(messages)) + for i := 0; i < len(messages)-3; i++ { + msg := readMessage(t, subscriberConn) + results = append(results, string(msg)) + } + require.Len(t, results, 7) + expMessages := make([]string, 0, 7) + for i, msg := range messages { + if i < 3 { + continue + } + expMessages = append(expMessages, msg) + } + assert.Equal(t, expMessages, results) + + // now check we can get the message that was sent after the subscription was created + msg := readMessage(t, subscriberConn) + assert.Equal(t, "hello there", string(msg)) +} + +func readMessage(t *testing.T, subscriberConn net.Conn) []byte { + var topicLen uint64 + err := binary.Read(subscriberConn, binary.BigEndian, &topicLen) + require.NoError(t, err) + + topicBuf := make([]byte, topicLen) + _, err = subscriberConn.Read(topicBuf) + require.NoError(t, err) + assert.Equal(t, topicA, string(topicBuf)) + + var dataLen uint64 + err = binary.Read(subscriberConn, binary.BigEndian, &dataLen) + require.NoError(t, err) + + buf := make([]byte, dataLen) + n, err := subscriberConn.Read(buf) + require.NoError(t, err) + require.Equal(t, int(dataLen), n) + + err = binary.Write(subscriberConn, binary.BigEndian, Ack) + require.NoError(t, err) + + return buf +} diff --git a/server/subscriber.go b/server/subscriber.go index 8fe8981..56360d1 100644 --- a/server/subscriber.go +++ b/server/subscriber.go @@ -29,18 +29,34 @@ func newMessage(data []byte) message { return message{data: data, deliveryCount: 1} } -func newSubscriber(peer *peer.Peer, topic string, ackDelay, ackTimeout time.Duration) *subscriber { +func newSubscriber(peer *peer.Peer, topic string, ackDelay, ackTimeout time.Duration, messageStore Store, startAt int) *subscriber { s := &subscriber{ peer: peer, topic: topic, messages: make(chan message), ackDelay: ackDelay, ackTimeout: ackTimeout, - unsubscribeCh: make(chan struct{}), + unsubscribeCh: make(chan struct{}, 1), } go s.sendMessages() + offset := startAt + + go func() { + // here we need to replay all messages from the store for the topic. + err := messageStore.ReadFrom(offset, func(msgs []MessageToSend) { + // go func() { + for _, msg := range msgs { + s.messages <- newMessage(msg.data) + } + // }() + }) + if err != nil { + slog.Error("failed to replay messages from offset", "error", err, "offset", offset) + } + }() + return s } @@ -79,7 +95,9 @@ func (s *subscriber) addMessage(msg message, delay time.Duration) { case <-s.unsubscribeCh: return case <-timer.C: + fmt.Printf("waiting to put message on queue: %s\n", msg.data) s.messages <- msg + fmt.Printf("put message on queue: %s\n", msg.data) } }() } diff --git a/server/topic.go b/server/topic.go index 723ed5c..c5d255a 100644 --- a/server/topic.go +++ b/server/topic.go @@ -1,6 +1,7 @@ package server import ( + "fmt" "net" "sync" ) @@ -9,21 +10,30 @@ type topic struct { name string subscriptions map[net.Addr]*subscriber mu sync.Mutex + messageStore Store } -func newTopic(name string) *topic { +func newTopic(name string, messageStore Store) *topic { return &topic{ name: name, subscriptions: make(map[net.Addr]*subscriber), + messageStore: messageStore, } } -func (t *topic) sendMessageToSubscribers(msgData []byte) { +func (t *topic) sendMessageToSubscribers(msg MessageToSend) error { + err := t.messageStore.Write(msg) + if err != nil { + return fmt.Errorf("failed to write message to store: %w", err) + } + t.mu.Lock() subscribers := t.subscriptions t.mu.Unlock() for _, subscriber := range subscribers { - subscriber.addMessage(newMessage(msgData), 0) + subscriber.addMessage(newMessage(msg.data), 0) } + + return nil }