From b8557aede99f94ba280a2d37a9bb356fd77c6cb3 Mon Sep 17 00:00:00 2001 From: Will Andrews Date: Thu, 07 Dec 2023 18:58:26 +0000 Subject: [PATCH] Merge pull request #1 from willdot/subscriber PubSub code --- example/main.go | 63 +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++ message.go | 7 +++++++ server/peer.go | 2 +- pubsub/publisher.go | 59 +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++ pubsub/subscriber.go | 145 +++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++ pubsub/subscriber_test.go | 148 ++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++ server/server.go | 18 +++++++++--------- server/server_test.go | 11 ++++++----- server/subscriber.go | 6 +++--- server/topic.go | 12 +++++++----- 10 file(s) changed, 448 insertion(s)(+), 23 deletion(s)(-) diff --git a/example/main.go b/example/main.go new file mode 100644 --- /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 new file mode 100644 --- /dev/null +++ b/message.go @@ -0,0 +1,7 @@ +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/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/pubsub/publisher.go b/pubsub/publisher.go new file mode 100644 --- /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/pubsub/subscriber.go b/pubsub/subscriber.go new file mode 100644 --- /dev/null +++ b/pubsub/subscriber.go @@ -0,0 +1,145 @@ +package pubsub + +import ( + "context" + "encoding/binary" + "encoding/json" + "fmt" + "log/slog" + "net" + "time" + + "github.com/willdot/messagebroker" + "github.com/willdot/messagebroker/server" +) + +// Subscriber allows subscriptions to a server and the consumption of messages +type Subscriber struct { + conn net.Conn +} + +// 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) + } + + return &Subscriber{ + conn: conn, + }, 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 { + 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 +} + +// 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 + // TODO: better error handling? Maybe a channel of errors? + Err error +} + +// 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), + } + + 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/pubsub/subscriber_test.go b/pubsub/subscriber_test.go new file mode 100644 --- /dev/null +++ b/pubsub/subscriber_test.go @@ -0,0 +1,148 @@ +package pubsub + +import ( + "context" + "fmt" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/willdot/messagebroker" + + "github.com/willdot/messagebroker/server" +) + +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 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 := NewPublisher(serverAddr) + require.NoError(t, err) + + t.Cleanup(func() { + sub.Close() + }) +} + +func TestNewPublisherInvalidServerAddr(t *testing.T) { + createServer(t) + + _, err := NewPublisher(":123456") + require.Error(t, err) +} + +func TestSubscribeToTopics(t *testing.T) { + createServer(t) + + sub, err := NewSubscriber(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 TestPublishAndSubscribe(t *testing.T) { + createServer(t) + + sub, err := NewSubscriber(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{}{} + }() + + publisher, err := NewPublisher("localhost:3000") + require.NoError(t, err) + t.Cleanup(func() { + publisher.Close() + }) + + // 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) + + err = publisher.PublishMessage(msg) + require.NoError(t, err) + } + + // 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/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,7 @@ Unsubscribe Action = 2 Publish Action = 3 ) -type Message struct { - Topic string `json:"topic"` - Data []byte `json:"data"` -} - +// Server accepts subscribe and publish connections and passes messages around type Server struct { addr string lis net.Listener @@ -32,7 +30,8 @@ mu sync.Mutex topics map[string]topic } -func NewServer(ctx context.Context, addr string) (*Server, error) { +// 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 { return nil, fmt.Errorf("failed to listen: %w", err) @@ -48,6 +47,7 @@ return srv, nil } +// Shutdown will cleanly shutdown the server func (s *Server) Shutdown() error { return s.lis.Close() } @@ -204,7 +204,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")) @@ -234,7 +234,7 @@ if !ok { t = newTopic(topicName) } - t.subscriptions[peer.addr()] = Subscriber{ + t.subscriptions[peer.addr()] = subscriber{ peer: peer, currentOffset: 0, } 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() { @@ -28,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 @@ -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,16 +1,16 @@ -package messagebroker +package server import ( "encoding/binary" "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/topic.go b/server/topic.go rename from topic.go rename to server/topic.go --- a/topic.go +++ b/server/topic.go @@ -1,22 +1,24 @@ -package messagebroker +package server import ( "encoding/json" "log/slog" "net" "sync" + + "github.com/willdot/messagebroker" ) 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), } } @@ -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() @@ -39,7 +41,7 @@ slog.Error("failed to marshal message for subscribers", "error", err) } 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 -- tangled.sh