diff --git a/pubsub/subscriber.go b/pubsub/subscriber.go --- a/pubsub/subscriber.go +++ b/pubsub/subscriber.go @@ -56,18 +56,30 @@ _, err = s.conn.Write(b) if err != nil { return fmt.Errorf("failed to subscribe to topics: %w", err) } - buf := make([]byte, 512) - _, err = s.conn.Read(buf) + + var resp server.Status + err = binary.Read(s.conn, binary.BigEndian, &resp) if err != nil { return fmt.Errorf("failed to read confirmation of subscription: %w", err) } - // TODO: this is soooo hacky - need to have some sort of response code - if string(buf[:10]) != "subscribed" { - return fmt.Errorf("failed to subscribe: '%s'", string(buf)) + if resp == server.Subscribed { + return nil } - return nil + var dataLen uint32 + err = binary.Read(s.conn, binary.BigEndian, &dataLen) + if err != nil { + return fmt.Errorf("received status %s:", resp) + } + + buf := make([]byte, dataLen) + _, err = s.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. It is thread safe to range over the Msgs channel to consume. If during the consumer diff --git a/server/peer.go b/server/peer.go --- a/server/peer.go +++ b/server/peer.go @@ -3,6 +3,7 @@ import ( "encoding/binary" "fmt" + "log/slog" "net" ) @@ -49,3 +50,50 @@ } return dataLen, nil } + +// Status represents the status of a request +type Status uint8 + +const ( + Subscribed = 1 + Unsubscribed = 2 + Error = 3 +) + +func (s Status) String() string { + switch s { + case Subscribed: + return "subsribed" + case Unsubscribed: + return "unsubscribed" + case Error: + return "error" + } + + return "" +} + +func (p *peer) writeStatus(status Status, message string) { + err := binary.Write(p.conn, binary.BigEndian, status) + if err != nil { + slog.Error("failed to write status to peers connection", "error", err, "peer", p.addr()) + return + } + + if message == "" { + return + } + + msgBytes := []byte(message) + err = binary.Write(p.conn, binary.BigEndian, uint32(len(msgBytes))) + if err != nil { + slog.Error("failed to write message length to peers connection", "error", err, "peer", p.addr()) + return + } + + _, err = p.conn.Write(msgBytes) + if err != nil { + slog.Error("failed to write message to peers connection", "error", err, "peer", p.addr()) + return + } +} diff --git a/server/server.go b/server/server.go --- a/server/server.go +++ b/server/server.go @@ -84,7 +84,8 @@ case Publish: s.handlePublish(peer) default: slog.Error("unknown action", "action", action, "peer", peer.addr()) - _, _ = peer.Write([]byte("unknown action")) + peer.writeStatus(Error, "unknown action") + //_, _ = peer.Write([]byte("unknown action")) } } @@ -122,11 +123,13 @@ // get the topics the peer wishes to subscribe to dataLen, err := peer.readDataLength() if err != nil { slog.Error(err.Error(), "peer", peer.addr()) - _, _ = peer.Write([]byte("invalid data length of topics provided")) + peer.writeStatus(Error, "invalid data length of topics provided") + // _, _ = peer.Write([]byte("invalid data length of topics provided")) return } if dataLen == 0 { - _, _ = peer.Write([]byte("data length of topics is 0")) + peer.writeStatus(Error, "data length of topics is 0") + // _, _ = peer.Write([]byte("data length of topics is 0")) return } @@ -134,7 +137,8 @@ buf := make([]byte, dataLen) _, err = peer.Read(buf) if err != nil { slog.Error("failed to read subscibers topic data", "error", err, "peer", peer.addr()) - _, _ = peer.Write([]byte("failed to read topic data")) + peer.writeStatus(Error, "failed to read topic data") + //_, _ = peer.Write([]byte("failed to read topic data")) return } @@ -142,12 +146,14 @@ var topics []string err = json.Unmarshal(buf, &topics) if err != nil { slog.Error("failed to unmarshal subscibers topic data", "error", err, "peer", peer.addr()) - _, _ = peer.Write([]byte("invalid topic data provided")) + peer.writeStatus(Error, "invalid topic data provided") + //_, _ = peer.Write([]byte("invalid topic data provided")) return } s.subscribeToTopics(peer, topics) - _, _ = peer.Write([]byte("subscribed")) + //_, _ = peer.Write([]byte("subscribed")) + peer.writeStatus(Subscribed, "") } func (s *Server) handleUnsubscribe(peer peer) { @@ -155,11 +161,13 @@ // get the topics the peer wishes to unsubscribe from dataLen, err := peer.readDataLength() if err != nil { slog.Error(err.Error(), "peer", peer.addr()) - _, _ = peer.Write([]byte("invalid data length of topics provided")) + peer.writeStatus(Error, "invalid data length of topics provided") + //_, _ = peer.Write([]byte("invalid data length of topics provided")) return } if dataLen == 0 { - _, _ = peer.Write([]byte("data length of topics is 0")) + peer.writeStatus(Error, "data length of topics is 0") + //_, _ = peer.Write([]byte("data length of topics is 0")) return } @@ -167,7 +175,8 @@ buf := make([]byte, dataLen) _, err = peer.Read(buf) if err != nil { slog.Error("failed to read subscibers topic data", "error", err, "peer", peer.addr()) - _, _ = peer.Write([]byte("failed to read topic data")) + peer.writeStatus(Error, "failed to read topic data") + //_, _ = peer.Write([]byte("failed to read topic data")) return } @@ -175,13 +184,14 @@ var topics []string err = json.Unmarshal(buf, &topics) if err != nil { slog.Error("failed to unmarshal subscibers topic data", "error", err, "peer", peer.addr()) - _, _ = peer.Write([]byte("invalid topic data provided")) + peer.writeStatus(Error, "invalid topic data provided") + //_, _ = peer.Write([]byte("invalid topic data provided")) return } s.unsubscribeToTopics(peer, topics) - - _, _ = peer.Write([]byte("unsubscribed")) + peer.writeStatus(Unsubscribed, "") + //_, _ = peer.Write([]byte("unsubscribed")) } func (s *Server) handlePublish(peer peer) { @@ -189,7 +199,8 @@ for { dataLen, err := peer.readDataLength() if err != nil { slog.Error(err.Error(), "peer", peer.addr()) - _, _ = peer.Write([]byte("invalid data length of data provided")) + peer.writeStatus(Error, "invalid data length of data provided") + //_, _ = peer.Write([]byte("invalid data length of data provided")) return } if dataLen == 0 { @@ -199,15 +210,17 @@ 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()) + peer.writeStatus(Error, "failed to read data") + //_, _ = peer.Write([]byte("failed to read data")) return } var msg messagebroker.Message err = json.Unmarshal(buf, &msg) if err != nil { - _, _ = peer.Write([]byte("invalid message")) + peer.writeStatus(Error, "invalid message") + //_, _ = peer.Write([]byte("invalid message")) slog.Error("failed to unmarshal data to message", "error", err, "peer", peer.addr()) continue } diff --git a/server/server_test.go b/server/server_test.go --- a/server/server_test.go +++ b/server/server_test.go @@ -51,14 +51,12 @@ _, err = conn.Write(rawTopics) require.NoError(t, err) - expectedRes := "subscribed" + expectedRes := Subscribed - buf := make([]byte, len(expectedRes)) - n, err := conn.Read(buf) - require.NoError(t, err) - require.Equal(t, len(expectedRes), n) + var resp Status + err = binary.Read(conn, binary.BigEndian, &resp) - assert.Equal(t, expectedRes, string(buf)) + assert.Equal(t, expectedRes, int(resp)) return conn } @@ -98,14 +96,12 @@ _, err = conn.Write(rawTopics) require.NoError(t, err) - expectedRes := "unsubscribed" + expectedRes := Unsubscribed - buf := make([]byte, len(expectedRes)) - n, err := conn.Read(buf) - require.NoError(t, err) - require.Equal(t, len(expectedRes), n) + var resp Status + err = binary.Read(conn, binary.BigEndian, &resp) - assert.Equal(t, expectedRes, string(buf)) + assert.Equal(t, expectedRes, int(resp)) assert.Len(t, srv.topics, 3) assert.Len(t, srv.topics["topic a"].subscriptions, 0) @@ -154,14 +150,24 @@ err = binary.Write(conn, binary.BigEndian, uint8(99)) require.NoError(t, err) - expectedRes := "unknown action" + expectedRes := Error + + var resp Status + err = binary.Read(conn, binary.BigEndian, &resp) + + assert.Equal(t, expectedRes, int(resp)) - buf := make([]byte, len(expectedRes)) - n, err := conn.Read(buf) + expectedMessage := "unknown action" + + var dataLen uint32 + err = binary.Read(conn, binary.BigEndian, &dataLen) + assert.Equal(t, len(expectedMessage), int(dataLen)) + + buf := make([]byte, dataLen) + _, err = conn.Read(buf) require.NoError(t, err) - require.Equal(t, len(expectedRes), n) - assert.Equal(t, expectedRes, string(buf)) + assert.Equal(t, expectedMessage, string(buf)) } func TestInvalidMessagePublished(t *testing.T) { @@ -183,10 +189,24 @@ n, err := publisherConn.Write(data) require.NoError(t, err) require.Equal(t, len(data), n) - buf := make([]byte, 15) + expectedRes := Error + + var resp Status + err = binary.Read(publisherConn, binary.BigEndian, &resp) + + assert.Equal(t, expectedRes, int(resp)) + + expectedMessage := "invalid message" + + var dataLen uint32 + err = binary.Read(publisherConn, binary.BigEndian, &dataLen) + assert.Equal(t, len(expectedMessage), int(dataLen)) + + buf := make([]byte, dataLen) _, err = publisherConn.Read(buf) require.NoError(t, err) - assert.Equal(t, "invalid message", string(buf)) + + assert.Equal(t, expectedMessage, string(buf)) } func TestSendsDataToTopicSubscribers(t *testing.T) {