From 219f554cc92a945de512d6b9921c50a617db4151 Mon Sep 17 00:00:00 2001 From: Will Date: Wed, 6 Dec 2023 20:22:13 +0000 Subject: [PATCH] Some refactoring and publisher conn can now send multiple messages --- peer.go | 51 +++++++++++++ server.go | 192 ++++++++++++++++++++++--------------------------- server_test.go | 79 ++++++++++++++++++-- subscriber.go | 17 +++-- topic.go | 20 +++--- 5 files changed, 232 insertions(+), 127 deletions(-) create mode 100644 peer.go diff --git a/peer.go b/peer.go new file mode 100644 index 0000000..60ee7bb --- /dev/null +++ b/peer.go @@ -0,0 +1,51 @@ +package messagebroker + +import ( + "encoding/binary" + "fmt" + "net" +) + +type peer struct { + conn net.Conn +} + +func newPeer(conn net.Conn) peer { + return peer{ + conn: conn, + } +} + +// Read wraps the peers underlying connections Read function to satisfy io.Reader +func (p *peer) Read(b []byte) (n int, err error) { + return p.conn.Read(b) +} + +// Write wraps the peers underlying connections Write function to satisfy io.Writer +func (p *peer) Write(b []byte) (n int, err error) { + return p.conn.Write(b) +} + +func (p *peer) addr() net.Addr { + return p.conn.LocalAddr() +} + +func (p *peer) readAction() (Action, error) { + var action Action + err := binary.Read(p.conn, binary.BigEndian, &action) + if err != nil { + return 0, fmt.Errorf("failed to read action from peer: %w", err) + } + + return action, nil +} + +func (p *peer) readDataLength() (uint32, error) { + var dataLen uint32 + err := binary.Read(p.conn, binary.BigEndian, &dataLen) + if err != nil { + return 0, fmt.Errorf("failed to read data length from peer: %w", err) + } + + return dataLen, nil +} diff --git a/server.go b/server.go index ec85548..824789f 100644 --- a/server.go +++ b/server.go @@ -2,7 +2,6 @@ package messagebroker import ( "context" - "encoding/binary" "encoding/json" "fmt" "log/slog" @@ -11,7 +10,7 @@ import ( "sync" ) -// Action represents the type of action that a connection requests to do +// Action represents the type of action that a peer requests to do type Action uint8 const ( @@ -68,182 +67,165 @@ func (s *Server) start(ctx context.Context) { } } -func getActionFromConn(conn net.Conn) (Action, error) { - var action Action - err := binary.Read(conn, binary.BigEndian, &action) - if err != nil { - return 0, err - } - - return action, nil -} - -func getDataLengthFromConn(conn net.Conn) (uint32, error) { - var dataLen uint32 - err := binary.Read(conn, binary.BigEndian, &dataLen) - if err != nil { - return 0, fmt.Errorf("failed to read data length from conn: %w", err) - } - - return dataLen, nil -} - func (s *Server) handleConn(conn net.Conn) { - action, err := getActionFromConn(conn) + peer := newPeer(conn) + action, err := peer.readAction() if err != nil { - slog.Error("failed to read action from conn", "error", err, "conn", conn.LocalAddr()) + slog.Error("failed to read action from peer", "error", err, "peer", peer.addr()) return } switch action { case Subscribe: - s.handleSubscribingConn(conn) + s.handleSubscribe(peer) case Unsubscribe: - s.handleUnsubscribingConn(conn) + s.handleUnsubscribe(peer) case Publish: - s.handlePublisherConn(conn) + s.handlePublish(peer) default: - slog.Error("unknown action", "action", action, "conn", conn.LocalAddr()) - _, _ = conn.Write([]byte("unknown action")) + slog.Error("unknown action", "action", action, "peer", peer.addr()) + _, _ = peer.Write([]byte("unknown action")) } } -func (s *Server) handleSubscribingConn(conn net.Conn) { - // subscribe the connection to the topic - s.subscribeConnToTopic(conn) +func (s *Server) handleSubscribe(peer peer) { + // subscribe the peer to the topic + s.subscribePeerToTopic(peer) - // keep handling the connection, getting the action from the conection when it wishes to do something else. - // once the connection ends, it will be unsubscribed from all topics and returned + // keep handling the peers connection, getting the action from the peer when it wishes to do something else. + // once the peers connection ends, it will be unsubscribed from all topics and returned for { - action, err := getActionFromConn(conn) + action, err := peer.readAction() if err != nil { - // TODO: see if there's a way to check if the connection has been ended etc - slog.Error("failed to read action from subscriber", "error", err, "conn", conn.LocalAddr()) + // TODO: see if there's a way to check if the peers connection has been ended etc + slog.Error("failed to read action from subscriber", "error", err, "peer", peer.addr()) - s.unsubscribeConnectionFromAllTopics(conn.LocalAddr()) + s.unsubscribePeerFromAllTopics(peer) return } switch action { case Subscribe: - s.subscribeConnToTopic(conn) + s.subscribePeerToTopic(peer) case Unsubscribe: - s.handleUnsubscribingConn(conn) + s.handleUnsubscribe(peer) default: - slog.Error("unknown action for subscriber", "action", action, "conn", conn.LocalAddr()) + slog.Error("unknown action for subscriber", "action", action, "peer", peer.addr()) continue } } } -func (s *Server) subscribeConnToTopic(conn net.Conn) { - // get the topics the connection wishes to subscribe to - dataLen, err := getDataLengthFromConn(conn) +func (s *Server) subscribePeerToTopic(peer peer) { + // get the topics the peer wishes to subscribe to + dataLen, err := peer.readDataLength() if err != nil { - slog.Error(err.Error(), "conn", conn.LocalAddr()) - _, _ = conn.Write([]byte("invalid data length of topics provided")) + slog.Error(err.Error(), "peer", peer.addr()) + _, _ = peer.Write([]byte("invalid data length of topics provided")) return } if dataLen == 0 { - _, _ = conn.Write([]byte("data length of topics is 0")) + _, _ = peer.Write([]byte("data length of topics is 0")) return } buf := make([]byte, dataLen) - _, err = conn.Read(buf) + _, err = peer.Read(buf) if err != nil { - slog.Error("failed to read subscibers topic data", "error", err, "conn", conn.LocalAddr()) - _, _ = conn.Write([]byte("failed to read topic data")) + slog.Error("failed to read subscibers topic data", "error", err, "peer", peer.addr()) + _, _ = peer.Write([]byte("failed to read topic data")) return } var topics []string err = json.Unmarshal(buf, &topics) if err != nil { - slog.Error("failed to unmarshal subscibers topic data", "error", err, "conn", conn.LocalAddr()) - _, _ = conn.Write([]byte("invalid topic data provided")) + slog.Error("failed to unmarshal subscibers topic data", "error", err, "peer", peer.addr()) + _, _ = peer.Write([]byte("invalid topic data provided")) return } - s.subscribeToTopics(conn, topics) - _, _ = conn.Write([]byte("subscribed")) + s.subscribeToTopics(peer, topics) + _, _ = peer.Write([]byte("subscribed")) } -func (s *Server) handleUnsubscribingConn(conn net.Conn) { - // get the topics the connection wishes to unsubscribe from - dataLen, err := getDataLengthFromConn(conn) +func (s *Server) handleUnsubscribe(peer peer) { + // get the topics the peer wishes to unsubscribe from + dataLen, err := peer.readDataLength() if err != nil { - slog.Error(err.Error(), "conn", conn.LocalAddr()) - _, _ = conn.Write([]byte("invalid data length of topics provided")) + slog.Error(err.Error(), "peer", peer.addr()) + _, _ = peer.Write([]byte("invalid data length of topics provided")) return } if dataLen == 0 { - _, _ = conn.Write([]byte("data length of topics is 0")) + _, _ = peer.Write([]byte("data length of topics is 0")) return } buf := make([]byte, dataLen) - _, err = conn.Read(buf) + _, err = peer.Read(buf) if err != nil { - slog.Error("failed to read subscibers topic data", "error", err, "conn", conn.LocalAddr()) - _, _ = conn.Write([]byte("failed to read topic data")) + slog.Error("failed to read subscibers topic data", "error", err, "peer", peer.addr()) + _, _ = peer.Write([]byte("failed to read topic data")) return } var topics []string err = json.Unmarshal(buf, &topics) if err != nil { - slog.Error("failed to unmarshal subscibers topic data", "error", err, "conn", conn.LocalAddr()) - _, _ = conn.Write([]byte("invalid topic data provided")) + slog.Error("failed to unmarshal subscibers topic data", "error", err, "peer", peer.addr()) + _, _ = peer.Write([]byte("invalid topic data provided")) return } - s.unsubscribeToTopics(conn, topics) + s.unsubscribeToTopics(peer, topics) - _, _ = conn.Write([]byte("unsubscribed")) + _, _ = peer.Write([]byte("unsubscribed")) } -func (s *Server) handlePublisherConn(conn net.Conn) { - dataLen, err := getDataLengthFromConn(conn) - if err != nil { - slog.Error(err.Error(), "conn", conn.LocalAddr()) - _, _ = conn.Write([]byte("invalid data length of data provided")) - return - } - if dataLen == 0 { - return - } +func (s *Server) handlePublish(peer peer) { + for { + dataLen, err := peer.readDataLength() + if err != nil { + slog.Error(err.Error(), "peer", peer.addr()) + _, _ = peer.Write([]byte("invalid data length of data provided")) + return + } + if dataLen == 0 { + continue + } - buf := make([]byte, dataLen) - _, err = conn.Read(buf) - if err != nil { - _, _ = conn.Write([]byte("failed to read data")) - slog.Error("failed to read data from conn", "error", err, "conn", conn.LocalAddr()) - return - } + buf := make([]byte, dataLen) + _, err = peer.Read(buf) + if err != nil { + _, _ = peer.Write([]byte("failed to read data")) + slog.Error("failed to read data from peer", "error", err, "peer", peer.addr()) + return + } - var msg Message - err = json.Unmarshal(buf, &msg) - if err != nil { - _, _ = conn.Write([]byte("invalid message")) - slog.Error("failed to unmarshal data to message", "error", err, "conn", conn.LocalAddr()) - return - } + var msg Message + err = json.Unmarshal(buf, &msg) + if err != nil { + _, _ = peer.Write([]byte("invalid message")) + slog.Error("failed to unmarshal data to message", "error", err, "peer", peer.addr()) + continue + } - topic := s.getTopic(msg.Topic) - if topic != nil { - topic.sendMessageToSubscribers(msg) + topic := s.getTopic(msg.Topic) + if topic != nil { + topic.sendMessageToSubscribers(msg) + } } } -func (s *Server) subscribeToTopics(conn net.Conn, topics []string) { +func (s *Server) subscribeToTopics(peer peer, topics []string) { for _, topic := range topics { - s.addSubsciberToTopic(topic, conn) + s.addSubsciberToTopic(topic, peer) } } -func (s *Server) addSubsciberToTopic(topicName string, conn net.Conn) { +func (s *Server) addSubsciberToTopic(topicName string, peer peer) { s.mu.Lock() defer s.mu.Unlock() @@ -252,21 +234,21 @@ func (s *Server) addSubsciberToTopic(topicName string, conn net.Conn) { t = newTopic(topicName) } - t.subscriptions[conn.LocalAddr()] = Subscriber{ - conn: conn, + t.subscriptions[peer.addr()] = Subscriber{ + peer: peer, currentOffset: 0, } s.topics[topicName] = t } -func (s *Server) unsubscribeToTopics(conn net.Conn, topics []string) { +func (s *Server) unsubscribeToTopics(peer peer, topics []string) { for _, topic := range topics { - s.removeSubsciberFromTopic(topic, conn) + s.removeSubsciberFromTopic(topic, peer) } } -func (s *Server) removeSubsciberFromTopic(topicName string, conn net.Conn) { +func (s *Server) removeSubsciberFromTopic(topicName string, peer peer) { s.mu.Lock() defer s.mu.Unlock() @@ -275,15 +257,15 @@ func (s *Server) removeSubsciberFromTopic(topicName string, conn net.Conn) { return } - delete(t.subscriptions, conn.LocalAddr()) + delete(t.subscriptions, peer.addr()) } -func (s *Server) unsubscribeConnectionFromAllTopics(addr net.Addr) { +func (s *Server) unsubscribePeerFromAllTopics(peer peer) { s.mu.Lock() defer s.mu.Unlock() for _, topic := range s.topics { - delete(topic.subscriptions, addr) + delete(topic.subscriptions, peer.addr()) } } diff --git a/server_test.go b/server_test.go index 33d86a3..c3cd92c 100644 --- a/server_test.go +++ b/server_test.go @@ -4,8 +4,10 @@ import ( "context" "encoding/binary" "encoding/json" + "fmt" "net" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -202,11 +204,10 @@ func TestSendsDataToTopicSubscribers(t *testing.T) { err = binary.Write(publisherConn, binary.BigEndian, Publish) require.NoError(t, err) - // send some data - data := []byte("hello world") + // send a message msg := Message{ Topic: "topic a", - Data: data, + Data: []byte("hello world"), } rawMsg, err := json.Marshal(msg) @@ -221,11 +222,77 @@ func TestSendsDataToTopicSubscribers(t *testing.T) { // check the subsribers got the data for _, conn := range subscribers { - buf := make([]byte, len(data)) + + 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, len(data), n) + require.Equal(t, int(dataLen), n) + + assert.Equal(t, rawMsg, buf) + } +} + +func TestPublishMultipleTimes(t *testing.T) { + _ = createServer(t) + + publisherConn, err := net.Dial("tcp", "localhost:3000") + require.NoError(t, err) + + err = binary.Write(publisherConn, binary.BigEndian, Publish) + require.NoError(t, err) + + messages := make([][]byte, 0, 10) + for i := 0; i < 10; i++ { + msg := Message{ + Topic: "topic a", + Data: []byte(fmt.Sprintf("message %d", i)), + } + + rawMsg, err := json.Marshal(msg) + require.NoError(t, err) + + messages = append(messages, rawMsg) + } + + subscribeFinCh := make(chan struct{}) + // create a subscriber that will read messages + subscriberConn := createConnectionAndSubscribe(t, []string{"topic a", "topic b"}) + go func() { + // check subscriber got all messages + for _, msg := range messages { + 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) + + assert.Equal(t, msg, buf) + } + + subscribeFinCh <- struct{}{} + }() + + // send multiple messages + for _, msg := range messages { + // send data length first + err = binary.Write(publisherConn, binary.BigEndian, uint32(len(msg))) + require.NoError(t, err) + n, err := publisherConn.Write(msg) + require.NoError(t, err) + require.Equal(t, len(msg), n) + } - assert.Equal(t, data, buf) + select { + case <-subscribeFinCh: + break + case <-time.After(time.Second): + t.Fatal(fmt.Errorf("timed out waiting for subscriber to read messages")) } } diff --git a/subscriber.go b/subscriber.go index 9d36c0e..990769c 100644 --- a/subscriber.go +++ b/subscriber.go @@ -1,19 +1,26 @@ package messagebroker import ( + "encoding/binary" "fmt" - "net" ) type Subscriber struct { - conn net.Conn + peer peer currentOffset int } -func (s *Subscriber) SendMessage(data []byte) error { - _, err := s.conn.Write(data) +func (s *Subscriber) SendMessage(msg []byte) error { + dataLen := uint64(len(msg)) + + err := binary.Write(&s.peer, binary.BigEndian, dataLen) + if err != nil { + return fmt.Errorf("failed to send data length: %w", err) + } + + _, err = s.peer.Write(msg) if err != nil { - return fmt.Errorf("failed to write to connection: %w", err) + return fmt.Errorf("failed to write to peer: %w", err) } return nil } diff --git a/topic.go b/topic.go index 32e27a6..5b585d5 100644 --- a/topic.go +++ b/topic.go @@ -1,6 +1,7 @@ package messagebroker import ( + "encoding/json" "log/slog" "net" "sync" @@ -19,19 +20,11 @@ func newTopic(name string) topic { } } -func (t *topic) addSubscriber(conn net.Conn) { - t.mu.Lock() - defer t.mu.Unlock() - - slog.Info("adding subscriber", "conn", conn.LocalAddr()) - t.subscriptions[conn.LocalAddr()] = Subscriber{conn: conn} -} - func (t *topic) removeSubscriber(addr net.Addr) { t.mu.Lock() defer t.mu.Unlock() - slog.Info("removing subscriber", "conn", addr) + slog.Info("removing subscriber", "peer", addr) delete(t.subscriptions, addr) } @@ -40,10 +33,15 @@ func (t *topic) sendMessageToSubscribers(msg Message) { subscribers := t.subscriptions t.mu.Unlock() + msgData, err := json.Marshal(msg) + if err != nil { + slog.Error("failed to marshal message for subscribers", "error", err) + } + for addr, subscriber := range subscribers { - err := subscriber.SendMessage(msg.Data) + err := subscriber.SendMessage(msgData) if err != nil { - slog.Error("failed to send to message", "error", err, "conn", addr) + slog.Error("failed to send to message", "error", err, "peer", addr) continue } } -- 2.51.2