diff --git a/message.go b/message.go new file mode 100644 --- /dev/null +++ b/message.go @@ -0,0 +1,6 @@ +package messagebroker + +type Message struct { + Topic string `json:"topic"` + Data []byte `json:"data"` +} diff --git a/peer.go b/server/peer.go rename from peer.go rename to server/peer.go --- a/peer.go +++ b/server/peer.go @@ -1,4 +1,4 @@ -package messagebroker +package server import ( "encoding/binary" diff --git a/server.go b/server/server.go rename from server.go rename to server/server.go --- a/server.go +++ b/server/server.go @@ -1,4 +1,4 @@ -package messagebroker +package server import ( "context" @@ -8,6 +8,8 @@ "log/slog" "net" "strings" "sync" + + "github.com/willdot/messagebroker" ) // Action represents the type of action that a peer requests to do @@ -19,11 +21,6 @@ Unsubscribe Action = 2 Publish Action = 3 ) -type Message struct { - Topic string `json:"topic"` - Data []byte `json:"data"` -} - type Server struct { addr string lis net.Listener @@ -32,7 +29,7 @@ mu sync.Mutex topics map[string]topic } -func NewServer(ctx context.Context, addr string) (*Server, error) { +func New(ctx context.Context, addr string) (*Server, error) { lis, err := net.Listen("tcp", addr) if err != nil { return nil, fmt.Errorf("failed to listen: %w", err) @@ -204,7 +201,7 @@ slog.Error("failed to read data from peer", "error", err, "peer", peer.addr()) return } - var msg Message + var msg messagebroker.Message err = json.Unmarshal(buf, &msg) if err != nil { _, _ = peer.Write([]byte("invalid message")) diff --git a/server_test.go b/server/server_test.go rename from server_test.go rename to server/server_test.go --- a/server_test.go +++ b/server/server_test.go @@ -1,4 +1,4 @@ -package messagebroker +package server import ( "context" @@ -11,10 +11,11 @@ "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "github.com/willdot/messagebroker" ) func createServer(t *testing.T) *Server { - srv, err := NewServer(context.Background(), ":3000") + srv, err := New(context.Background(), ":3000") require.NoError(t, err) t.Cleanup(func() { @@ -205,7 +206,7 @@ err = binary.Write(publisherConn, binary.BigEndian, Publish) require.NoError(t, err) // send a message - msg := Message{ + msg := messagebroker.Message{ Topic: "topic a", Data: []byte("hello world"), } @@ -247,7 +248,7 @@ require.NoError(t, err) messages := make([][]byte, 0, 10) for i := 0; i < 10; i++ { - msg := Message{ + msg := messagebroker.Message{ Topic: "topic a", Data: []byte(fmt.Sprintf("message %d", i)), } diff --git a/subscriber.go b/server/subscriber.go rename from subscriber.go rename to server/subscriber.go --- a/subscriber.go +++ b/server/subscriber.go @@ -1,4 +1,4 @@ -package messagebroker +package server import ( "encoding/binary" diff --git a/subscriber/subscriber.go b/subscriber/subscriber.go new file mode 100644 --- /dev/null +++ b/subscriber/subscriber.go @@ -0,0 +1,137 @@ +package subscriber + +import ( + "context" + "encoding/binary" + "encoding/json" + "fmt" + "log/slog" + "net" + "time" + + "github.com/willdot/messagebroker" + "github.com/willdot/messagebroker/server" +) + +type Subscriber struct { + conn net.Conn +} + +func New(addr string) (*Subscriber, error) { + conn, err := net.Dial("tcp", addr) + if err != nil { + return nil, fmt.Errorf("failed to dial: %w", err) + } + + return &Subscriber{ + conn: conn, + }, nil +} + +func (s *Subscriber) Close() error { + return s.conn.Close() +} + +func (s *Subscriber) SubscribeToTopics(topicNames []string) error { + err := binary.Write(s.conn, binary.BigEndian, server.Subscribe) + if err != nil { + return fmt.Errorf("failed to subscribe: %w", err) + } + + b, err := json.Marshal(topicNames) + if err != nil { + return fmt.Errorf("failed to marshal topic names: %w", err) + } + + err = binary.Write(s.conn, binary.BigEndian, uint32(len(b))) + if err != nil { + return fmt.Errorf("failed to write topic data length: %w", err) + } + + _, 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) + 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)) + } + + return nil +} + +type Consumer struct { + Msgs chan messagebroker.Message + Err error +} + +// TODO: maybe buffer the message channel up? +func (s *Subscriber) Consume(ctx context.Context) *Consumer { + consumer := &Consumer{ + Msgs: make(chan messagebroker.Message), + } + + go s.consume(ctx, consumer) + + return consumer +} + +func (s *Subscriber) consume(ctx context.Context, consumer *Consumer) { + defer close(consumer.Msgs) + for { + if ctx.Err() != nil { + return + } + + msg, err := s.readMessage() + if err != nil { + consumer.Err = err + return + } + + if msg != nil { + consumer.Msgs <- *msg + } + } +} + +func (s *Subscriber) readMessage() (*messagebroker.Message, error) { + err := s.conn.SetReadDeadline(time.Now().Add(time.Second)) + if err != nil { + return nil, err + } + + var dataLen uint64 + err = binary.Read(s.conn, binary.BigEndian, &dataLen) + if err != nil { + if neterr, ok := err.(net.Error); ok && neterr.Timeout() { + return nil, nil + } + return nil, err + } + + if dataLen <= 0 { + return nil, nil + } + + buf := make([]byte, dataLen) + _, err = s.conn.Read(buf) + if err != nil { + return nil, err + } + + var msg messagebroker.Message + err = json.Unmarshal(buf, &msg) + if err != nil { + slog.Error("failed to unmarshal message", "error", err) + return nil, nil + } + + return &msg, nil +} diff --git a/subscriber/subscriber_test.go b/subscriber/subscriber_test.go new file mode 100644 --- /dev/null +++ b/subscriber/subscriber_test.go @@ -0,0 +1,139 @@ +package subscriber_test + +import ( + "context" + "encoding/binary" + "encoding/json" + "fmt" + "net" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/willdot/messagebroker" + "github.com/willdot/messagebroker/server" + "github.com/willdot/messagebroker/subscriber" +) + +const ( + serverAddr = ":3000" +) + +func createServer(t *testing.T) { + server, err := server.New(context.Background(), serverAddr) + require.NoError(t, err) + + t.Cleanup(func() { + server.Shutdown() + }) +} + +func TestNew(t *testing.T) { + createServer(t) + + sub, err := subscriber.New(serverAddr) + require.NoError(t, err) + + t.Cleanup(func() { + sub.Close() + }) +} + +func TestNewInvalidServerAddr(t *testing.T) { + createServer(t) + + _, err := subscriber.New(":123456") + require.Error(t, err) +} + +func TestSubscribeToTopics(t *testing.T) { + createServer(t) + + sub, err := subscriber.New(serverAddr) + require.NoError(t, err) + + t.Cleanup(func() { + sub.Close() + }) + + topics := []string{"topic a", "topic b"} + + err = sub.SubscribeToTopics(topics) + require.NoError(t, err) +} + +func TestSubscribeConsumeFromSubscription(t *testing.T) { + createServer(t) + + sub, err := subscriber.New(serverAddr) + require.NoError(t, err) + + t.Cleanup(func() { + sub.Close() + }) + + topics := []string{"topic a", "topic b"} + + 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) + + var receivedMessages []messagebroker.Message + + consumerFinCh := make(chan struct{}) + go func() { + for msg := range consumer.Msgs { + receivedMessages = append(receivedMessages, msg) + } + + require.NoError(t, err) + consumerFinCh <- struct{}{} + }() + + publisherConn, err := net.Dial("tcp", "localhost:3000") + require.NoError(t, err) + + err = binary.Write(publisherConn, binary.BigEndian, server.Publish) + require.NoError(t, err) + + // send some messages + sentMessages := make([]messagebroker.Message, 0, 10) + for i := 0; i < 10; i++ { + msg := messagebroker.Message{ + Topic: "topic a", + Data: []byte(fmt.Sprintf("message %d", i)), + } + + sentMessages = append(sentMessages, msg) + + b, err := json.Marshal(msg) + require.NoError(t, err) + + err = binary.Write(publisherConn, binary.BigEndian, uint32(len(b))) + require.NoError(t, err) + n, err := publisherConn.Write(b) + require.NoError(t, err) + require.Equal(t, len(b), n) + } + + // 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): + t.Fatal("timed out waiting for consumer to read messages") + } + + assert.ElementsMatch(t, receivedMessages, sentMessages) +} diff --git a/topic.go b/server/topic.go rename from topic.go rename to server/topic.go --- a/topic.go +++ b/server/topic.go @@ -1,10 +1,12 @@ -package messagebroker +package server import ( "encoding/json" "log/slog" "net" "sync" + + "github.com/willdot/messagebroker" ) type topic struct { @@ -28,7 +30,7 @@ slog.Info("removing subscriber", "peer", addr) delete(t.subscriptions, addr) } -func (t *topic) sendMessageToSubscribers(msg Message) { +func (t *topic) sendMessageToSubscribers(msg messagebroker.Message) { t.mu.Lock() subscribers := t.subscriptions t.mu.Unlock()