diff --git a/message.go b/message.go new file mode 100644 index 0000000..a653248 --- /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 similarity index 97% rename from peer.go rename to server/peer.go index 60ee7bb..3680f8f 100644 --- 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 similarity index 96% rename from server.go rename to server/server.go index 824789f..11a487a 100644 --- a/server.go +++ b/server/server.go @@ -1,4 +1,4 @@ -package messagebroker +package server import ( "context" @@ -8,6 +8,8 @@ import ( "net" "strings" "sync" + + "github.com/willdot/messagebroker" ) // Action represents the type of action that a peer requests to do @@ -19,11 +21,6 @@ const ( 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 @@ type Server struct { 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 @@ func (s *Server) handlePublish(peer peer) { 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 similarity index 97% rename from server_test.go rename to server/server_test.go index c3cd92c..0ed8628 100644 --- a/server_test.go +++ b/server/server_test.go @@ -1,4 +1,4 @@ -package messagebroker +package server import ( "context" @@ -11,10 +11,11 @@ import ( "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 @@ func TestSendsDataToTopicSubscribers(t *testing.T) { require.NoError(t, err) // send a message - msg := Message{ + msg := messagebroker.Message{ Topic: "topic a", Data: []byte("hello world"), } @@ -247,7 +248,7 @@ func TestPublishMultipleTimes(t *testing.T) { 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 similarity index 95% rename from subscriber.go rename to server/subscriber.go index 990769c..2c1ed0a 100644 --- a/subscriber.go +++ b/server/subscriber.go @@ -1,4 +1,4 @@ -package messagebroker +package server import ( "encoding/binary" diff --git a/topic.go b/server/topic.go similarity index 87% rename from topic.go rename to server/topic.go index 5b585d5..c09edf5 100644 --- 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 @@ func (t *topic) removeSubscriber(addr net.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() diff --git a/subscriber/subscriber.go b/subscriber/subscriber.go new file mode 100644 index 0000000..0321162 --- /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 index 0000000..b60080e --- /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) +} -- 2.51.2 From 5f39d2814a108aae93268dd714b2b63fbc4d5c67 Mon Sep 17 00:00:00 2001 From: Will Date: Thu, 7 Dec 2023 18:56:53 +0000 Subject: [PATCH 2/2] Implement publisher code and refactor / comments --- example/main.go | 63 +++++++++++++++++++++++ message.go | 1 + pubsub/publisher.go | 59 +++++++++++++++++++++ {subscriber => pubsub}/subscriber.go | 16 ++++-- {subscriber => pubsub}/subscriber_test.go | 55 +++++++++++--------- server/server.go | 5 +- server/server_test.go | 2 +- server/subscriber.go | 4 +- server/topic.go | 6 +-- 9 files changed, 177 insertions(+), 34 deletions(-) create mode 100644 example/main.go create mode 100644 pubsub/publisher.go rename {subscriber => pubsub}/subscriber.go (77%) rename {subscriber => pubsub}/subscriber_test.go (72%) diff --git a/example/main.go b/example/main.go new file mode 100644 index 0000000..6bebe1f --- /dev/null +++ b/example/main.go @@ -0,0 +1,63 @@ +package main + +import ( + "context" + "fmt" + "log/slog" + + "github.com/willdot/messagebroker" + "github.com/willdot/messagebroker/pubsub" + "github.com/willdot/messagebroker/server" +) + +func main() { + server, err := server.New(context.Background(), ":3000") + if err != nil { + panic(err) + } + defer server.Shutdown() + + go sendMessages() + + sub, err := pubsub.NewSubscriber(":3000") + if err != nil { + panic(err) + } + defer sub.Close() + + sub.SubscribeToTopics([]string{"topic a"}) + + consumer := sub.Consume(context.Background()) + if consumer.Err != nil { + panic(err) + } + + for msg := range consumer.Msgs { + slog.Info("received message", "message", string(msg.Data)) + } + +} + +func sendMessages() { + publisher, err := pubsub.NewPublisher("localhost:3000") + if err != nil { + panic(err) + } + defer publisher.Close() + + // send some messages + i := 0 + for { + i++ + msg := messagebroker.Message{ + Topic: "topic a", + Data: []byte(fmt.Sprintf("message %d", i)), + } + + err = publisher.PublishMessage(msg) + if err != nil { + slog.Error("failed to publish message", "error", err) + continue + } + } +} diff --git a/message.go b/message.go index a653248..b52fae9 100644 --- a/message.go +++ b/message.go @@ -1,5 +1,6 @@ package messagebroker +// Message represents a message that can be published or consumed type Message struct { Topic string `json:"topic"` Data []byte `json:"data"` diff --git a/pubsub/publisher.go b/pubsub/publisher.go new file mode 100644 index 0000000..87af684 --- /dev/null +++ b/pubsub/publisher.go @@ -0,0 +1,59 @@ +package pubsub + +import ( + "encoding/binary" + "encoding/json" + "fmt" + "net" + + "github.com/willdot/messagebroker" + "github.com/willdot/messagebroker/server" +) + +// Publisher allows messages to be published to a server +type Publisher struct { + conn net.Conn +} + +// NewPublisher connects to the server at the given address and registers as a publisher +func NewPublisher(addr string) (*Publisher, error) { + conn, err := net.Dial("tcp", addr) + if err != nil { + return nil, fmt.Errorf("failed to dial: %w", err) + } + + err = binary.Write(conn, binary.BigEndian, server.Publish) + if err != nil { + conn.Close() + return nil, fmt.Errorf("failed to register publish to server: %w", err) + } + + return &Publisher{ + conn: conn, + }, nil +} + +// Close cleanly shuts down the publisher +func (p *Publisher) Close() error { + return p.conn.Close() +} + +// Publish will publish the given message to the server +func (p *Publisher) PublishMessage(message messagebroker.Message) error { + b, err := json.Marshal(message) + if err != nil { + return fmt.Errorf("failed to marshal message: %w", err) + } + + err = binary.Write(p.conn, binary.BigEndian, uint32(len(b))) + if err != nil { + return fmt.Errorf("failed to write message size to server") + } + + _, err = p.conn.Write(b) + if err != nil { + return fmt.Errorf("failed to publish data to server") + } + + return nil +} diff --git a/subscriber/subscriber.go b/pubsub/subscriber.go similarity index 77% rename from subscriber/subscriber.go rename to pubsub/subscriber.go index 0321162..81c3253 100644 --- a/subscriber/subscriber.go +++ b/pubsub/subscriber.go @@ -1,4 +1,4 @@ -package subscriber +package pubsub import ( "context" @@ -13,11 +13,13 @@ import ( "github.com/willdot/messagebroker/server" ) +// Subscriber allows subscriptions to a server and the consumption of messages type Subscriber struct { conn net.Conn } -func New(addr string) (*Subscriber, error) { +// NewSubscriber will connect to the server at the given address +func NewSubscriber(addr string) (*Subscriber, error) { conn, err := net.Dial("tcp", addr) if err != nil { return nil, fmt.Errorf("failed to dial: %w", err) @@ -28,10 +30,12 @@ func New(addr string) (*Subscriber, error) { }, nil } +// Close cleanly shuts down the subscriber func (s *Subscriber) Close() error { return s.conn.Close() } +// SubscribeToTopics will subscribe to the provided topics func (s *Subscriber) SubscribeToTopics(topicNames []string) error { err := binary.Write(s.conn, binary.BigEndian, server.Subscribe) if err != nil { @@ -66,12 +70,16 @@ func (s *Subscriber) SubscribeToTopics(topicNames []string) error { return nil } +// Consumer allows the consumption of messages. It is thread safe to range over the Msgs channel to consume. If during the consumer +// receiving messages from the server an error occurs, it will be stored in Err type Consumer struct { Msgs chan messagebroker.Message - Err error + // TODO: better error handling? Maybe a channel of errors? + Err error } -// TODO: maybe buffer the message channel up? +// Consume will create a consumer and start it running in a go routine. You can then use the Msgs channel of the consumer +// to read the messages func (s *Subscriber) Consume(ctx context.Context) *Consumer { consumer := &Consumer{ Msgs: make(chan messagebroker.Message), diff --git a/subscriber/subscriber_test.go b/pubsub/subscriber_test.go similarity index 72% rename from subscriber/subscriber_test.go rename to pubsub/subscriber_test.go index b60080e..2f29430 100644 --- a/subscriber/subscriber_test.go +++ b/pubsub/subscriber_test.go @@ -1,19 +1,16 @@ -package subscriber_test +package pubsub 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 ( @@ -29,10 +26,28 @@ func createServer(t *testing.T) { }) } -func TestNew(t *testing.T) { +func TestNewSubscriber(t *testing.T) { + createServer(t) + + sub, err := NewSubscriber(serverAddr) + require.NoError(t, err) + + t.Cleanup(func() { + sub.Close() + }) +} + +func TestNewSubscriberInvalidServerAddr(t *testing.T) { + createServer(t) + + _, err := NewSubscriber(":123456") + require.Error(t, err) +} + +func TestNewPublisher(t *testing.T) { createServer(t) - sub, err := subscriber.New(serverAddr) + sub, err := NewPublisher(serverAddr) require.NoError(t, err) t.Cleanup(func() { @@ -40,17 +55,17 @@ func TestNew(t *testing.T) { }) } -func TestNewInvalidServerAddr(t *testing.T) { +func TestNewPublisherInvalidServerAddr(t *testing.T) { createServer(t) - _, err := subscriber.New(":123456") + _, err := NewPublisher(":123456") require.Error(t, err) } func TestSubscribeToTopics(t *testing.T) { createServer(t) - sub, err := subscriber.New(serverAddr) + sub, err := NewSubscriber(serverAddr) require.NoError(t, err) t.Cleanup(func() { @@ -63,10 +78,10 @@ func TestSubscribeToTopics(t *testing.T) { require.NoError(t, err) } -func TestSubscribeConsumeFromSubscription(t *testing.T) { +func TestPublishAndSubscribe(t *testing.T) { createServer(t) - sub, err := subscriber.New(serverAddr) + sub, err := NewSubscriber(serverAddr) require.NoError(t, err) t.Cleanup(func() { @@ -98,11 +113,11 @@ func TestSubscribeConsumeFromSubscription(t *testing.T) { consumerFinCh <- struct{}{} }() - publisherConn, err := net.Dial("tcp", "localhost:3000") - require.NoError(t, err) - - err = binary.Write(publisherConn, binary.BigEndian, server.Publish) + publisher, err := NewPublisher("localhost:3000") require.NoError(t, err) + t.Cleanup(func() { + publisher.Close() + }) // send some messages sentMessages := make([]messagebroker.Message, 0, 10) @@ -114,14 +129,8 @@ func TestSubscribeConsumeFromSubscription(t *testing.T) { 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) + err = publisher.PublishMessage(msg) require.NoError(t, err) - require.Equal(t, len(b), n) } // give the consumer some time to read the messages -- TODO: make better! diff --git a/server/server.go b/server/server.go index 11a487a..83ed020 100644 --- a/server/server.go +++ b/server/server.go @@ -21,6 +21,7 @@ const ( Publish Action = 3 ) +// Server accepts subscribe and publish connections and passes messages around type Server struct { addr string lis net.Listener @@ -29,6 +30,7 @@ type Server struct { topics map[string]topic } +// New creates and starts a new server func New(ctx context.Context, addr string) (*Server, error) { lis, err := net.Listen("tcp", addr) if err != nil { @@ -45,6 +47,7 @@ func New(ctx context.Context, addr string) (*Server, error) { return srv, nil } +// Shutdown will cleanly shutdown the server func (s *Server) Shutdown() error { return s.lis.Close() } @@ -231,7 +234,7 @@ func (s *Server) addSubsciberToTopic(topicName string, peer peer) { t = newTopic(topicName) } - t.subscriptions[peer.addr()] = Subscriber{ + t.subscriptions[peer.addr()] = subscriber{ peer: peer, currentOffset: 0, } diff --git a/server/server_test.go b/server/server_test.go index 0ed8628..5d00483 100644 --- a/server/server_test.go +++ b/server/server_test.go @@ -29,7 +29,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 diff --git a/server/subscriber.go b/server/subscriber.go index 2c1ed0a..8c55de3 100644 --- a/server/subscriber.go +++ b/server/subscriber.go @@ -5,12 +5,12 @@ import ( "fmt" ) -type Subscriber struct { +type subscriber struct { peer peer currentOffset int } -func (s *Subscriber) SendMessage(msg []byte) error { +func (s *subscriber) sendMessage(msg []byte) error { dataLen := uint64(len(msg)) err := binary.Write(&s.peer, binary.BigEndian, dataLen) diff --git a/server/topic.go b/server/topic.go index c09edf5..9ab22d7 100644 --- a/server/topic.go +++ b/server/topic.go @@ -11,14 +11,14 @@ import ( type topic struct { name string - subscriptions map[net.Addr]Subscriber + subscriptions map[net.Addr]subscriber mu sync.Mutex } func newTopic(name string) topic { return topic{ name: name, - subscriptions: make(map[net.Addr]Subscriber), + subscriptions: make(map[net.Addr]subscriber), } } @@ -41,7 +41,7 @@ func (t *topic) sendMessageToSubscribers(msg messagebroker.Message) { } for addr, subscriber := range subscribers { - err := subscriber.SendMessage(msgData) + err := subscriber.sendMessage(msgData) if err != nil { slog.Error("failed to send to message", "error", err, "peer", addr) continue