From 0cbdc73538136566ebffe9316b83ea7966553299 Mon Sep 17 00:00:00 2001 From: Will Date: Fri, 8 Dec 2023 20:19:25 +0000 Subject: [PATCH] Refactor. Huge refactor to make conns synchronous --- .gitignore | 3 +- dockerfile.example-server | 20 +++ example/main.go | 18 +- example/server/main.go | 23 +++ go.mod | 6 +- go.sum | 4 + message.go => pubsub/message.go | 2 +- pubsub/publisher.go | 49 ++++-- pubsub/subscriber.go | 238 +++++++++++++++----------- pubsub/subscriber_test.go | 41 +++-- server/peer.go | 93 ++++------- server/server.go | 287 ++++++++++++++++++++++---------- server/server_test.go | 144 +++++++++------- server/subscriber.go | 26 --- server/topic.go | 50 ++++-- 15 files changed, 605 insertions(+), 399 deletions(-) create mode 100644 dockerfile.example-server create mode 100644 example/server/main.go rename message.go => pubsub/message.go (87%) delete mode 100644 server/subscriber.go diff --git a/.gitignore b/.gitignore index b7ff13a..ce82b5b 100644 --- a/.gitignore +++ b/.gitignore @@ -1 +1,2 @@ -.DS_STORE \ No newline at end of file +.DS_STORE +example/example diff --git a/dockerfile.example-server b/dockerfile.example-server new file mode 100644 index 0000000..8f91a5d --- /dev/null +++ b/dockerfile.example-server @@ -0,0 +1,20 @@ +FROM golang:latest as builder + +WORKDIR /app + +COPY go.mod go.sum ./ +COPY example/server/ ./ +RUN go mod download + +COPY . . + +RUN CGO_ENABLED=0 go build -o message-broker-server . + +FROM alpine:latest + +RUN apk --no-cache add ca-certificates + +WORKDIR /root/ +COPY --from=builder /app/message-broker-server . + +CMD ["./message-broker-server"] \ No newline at end of file diff --git a/example/main.go b/example/main.go index f023e8d..808ecac 100644 --- a/example/main.go +++ b/example/main.go @@ -2,22 +2,22 @@ package main import ( "context" + "flag" "fmt" "log/slog" - "github.com/willdot/messagebroker" "github.com/willdot/messagebroker/pubsub" - "github.com/willdot/messagebroker/server" ) +var consumeOnly *bool + func main() { - server, err := server.New(context.Background(), ":3000") - if err != nil { - panic(err) - } - defer server.Shutdown() + consumeOnly = flag.Bool("consume-only", false, "just consumes (doesn't start server and doesn't publish)") + flag.Parse() - go sendMessages() + if *consumeOnly == false { + go sendMessages() + } sub, err := pubsub.NewSubscriber(":3000") if err != nil { @@ -49,7 +49,7 @@ func sendMessages() { i := 0 for { i++ - msg := messagebroker.Message{ + msg := pubsub.Message{ Topic: "topic a", Data: []byte(fmt.Sprintf("message %d", i)), } diff --git a/example/server/main.go b/example/server/main.go new file mode 100644 index 0000000..827f72d --- /dev/null +++ b/example/server/main.go @@ -0,0 +1,23 @@ +package main + +import ( + "log" + "os" + "os/signal" + "syscall" + + "github.com/willdot/messagebroker/server" +) + +func main() { + srv, err := server.New(":3000") + if err != nil { + log.Fatal(err) + } + defer srv.Shutdown() + + signals := make(chan os.Signal, 1) + signal.Notify(signals, syscall.SIGTERM, syscall.SIGINT) + + <-signals +} diff --git a/go.mod b/go.mod index 8efa94d..0c1f350 100644 --- a/go.mod +++ b/go.mod @@ -2,7 +2,11 @@ module github.com/willdot/messagebroker go 1.21.0 -require github.com/stretchr/testify v1.8.4 +require ( + github.com/docker/distribution v2.8.3+incompatible + github.com/google/uuid v1.4.0 + github.com/stretchr/testify v1.8.4 +) require ( github.com/davecgh/go-spew v1.1.1 // indirect diff --git a/go.sum b/go.sum index fa4b6e6..42f8c32 100644 --- a/go.sum +++ b/go.sum @@ -1,5 +1,9 @@ github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/docker/distribution v2.8.3+incompatible h1:AtKxIZ36LoNK51+Z6RpzLpddBirtxJnzDrHLEKxTAYk= +github.com/docker/distribution v2.8.3+incompatible/go.mod h1:J2gT2udsDAN96Uj4KfcMRqY0/ypR+oyYUYmja8H+y+w= +github.com/google/uuid v1.4.0 h1:MtMxsa51/r9yyhkyLsVeVt0B+BGQZzpQiTQ4eHZ8bc4= +github.com/google/uuid v1.4.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/stretchr/testify v1.8.4 h1:CcVxjf3Q8PM0mHUKJCdn+eZZtm5yQwehR5yeSVQQcUk= diff --git a/message.go b/pubsub/message.go similarity index 87% rename from message.go rename to pubsub/message.go index b52fae9..69d652c 100644 --- a/message.go +++ b/pubsub/message.go @@ -1,4 +1,4 @@ -package messagebroker +package pubsub // Message represents a message that can be published or consumed type Message struct { diff --git a/pubsub/publisher.go b/pubsub/publisher.go index 87af684..9f4b257 100644 --- a/pubsub/publisher.go +++ b/pubsub/publisher.go @@ -2,17 +2,17 @@ package pubsub import ( "encoding/binary" - "encoding/json" "fmt" "net" + "sync" - "github.com/willdot/messagebroker" "github.com/willdot/messagebroker/server" ) // Publisher allows messages to be published to a server type Publisher struct { - conn net.Conn + conn net.Conn + connMu sync.Mutex } // NewPublisher connects to the server at the given address and registers as a publisher @@ -39,21 +39,38 @@ func (p *Publisher) Close() error { } // 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) - } +func (p *Publisher) PublishMessage(message Message) error { + op := func(conn net.Conn) error { + // send topic first + topic := fmt.Sprintf("topic:%s", message.Topic) + err := binary.Write(p.conn, binary.BigEndian, uint32(len(topic))) + if err != nil { + return fmt.Errorf("failed to write topic size to server") + } - 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([]byte(topic)) + if err != nil { + return fmt.Errorf("failed to write topic to server") + } - _, err = p.conn.Write(b) - if err != nil { - return fmt.Errorf("failed to publish data to server") + err = binary.Write(p.conn, binary.BigEndian, uint32(len(message.Data))) + if err != nil { + return fmt.Errorf("failed to write message size to server") + } + + _, err = p.conn.Write(message.Data) + if err != nil { + return fmt.Errorf("failed to publish data to server") + } + return nil } - return nil + return p.connOperation(op) +} + +func (p *Publisher) connOperation(op connOpp) error { + p.connMu.Lock() + defer p.connMu.Unlock() + + return op(p.conn) } diff --git a/pubsub/subscriber.go b/pubsub/subscriber.go index 297f11b..d0d37fa 100644 --- a/pubsub/subscriber.go +++ b/pubsub/subscriber.go @@ -4,18 +4,21 @@ import ( "context" "encoding/binary" "encoding/json" + "errors" "fmt" - "log/slog" "net" + "sync" "time" - "github.com/willdot/messagebroker" "github.com/willdot/messagebroker/server" ) +type connOpp func(conn net.Conn) error + // Subscriber allows subscriptions to a server and the consumption of messages type Subscriber struct { - conn net.Conn + conn net.Conn + connMu sync.Mutex } // NewSubscriber will connect to the server at the given address @@ -37,109 +40,117 @@ func (s *Subscriber) Close() error { // 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) - } + op := func(conn net.Conn) error { + err := binary.Write(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) - } + 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 = binary.Write(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) - } + _, err = conn.Write(b) + if err != nil { + return fmt.Errorf("failed to subscribe to topics: %w", err) + } - 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) - } + var resp server.Status + err = binary.Read(conn, binary.BigEndian, &resp) + if err != nil { + return fmt.Errorf("failed to read confirmation of subscription: %w", err) + } - if resp == server.Subscribed { - return nil - } + if resp == server.Subscribed { + return nil + } - var dataLen uint32 - err = binary.Read(s.conn, binary.BigEndian, &dataLen) - if err != nil { - return fmt.Errorf("received status %s:", resp) - } + var dataLen uint32 + err = binary.Read(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) + buf := make([]byte, dataLen) + _, err = conn.Read(buf) + if err != nil { + return fmt.Errorf("received status %s:", resp) + } + + return fmt.Errorf("received status %s - %s", resp, buf) } - return fmt.Errorf("received status %s - %s", resp, buf) + return s.connOperation(op) } // UnsubscribeToTopics will unsubscribe to the provided topics func (s *Subscriber) UnsubscribeToTopics(topicNames []string) error { - err := binary.Write(s.conn, binary.BigEndian, server.Unsubscribe) - if err != nil { - return fmt.Errorf("failed to unsubscribe: %w", err) - } + op := func(conn net.Conn) error { + err := binary.Write(conn, binary.BigEndian, server.Unsubscribe) + if err != nil { + return fmt.Errorf("failed to unsubscribe: %w", err) + } - b, err := json.Marshal(topicNames) - if err != nil { - return fmt.Errorf("failed to marshal topic names: %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 = binary.Write(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 unsubscribe to topics: %w", err) - } + _, err = conn.Write(b) + if err != nil { + return fmt.Errorf("failed to unsubscribe to topics: %w", err) + } - var resp server.Status - err = binary.Read(s.conn, binary.BigEndian, &resp) - if err != nil { - return fmt.Errorf("failed to read confirmation of unsubscription: %w", err) - } + var resp server.Status + err = binary.Read(conn, binary.BigEndian, &resp) + if err != nil { + return fmt.Errorf("failed to read confirmation of unsubscription: %w", err) + } - if resp == server.Unsubscribed { - return nil - } + if resp == server.Unsubscribed { + return nil + } - var dataLen uint32 - err = binary.Read(s.conn, binary.BigEndian, &dataLen) - if err != nil { - return fmt.Errorf("received status %s:", resp) - } + var dataLen uint32 + err = binary.Read(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) + buf := make([]byte, dataLen) + _, err = conn.Read(buf) + if err != nil { + return fmt.Errorf("received status %s:", resp) + } + + return fmt.Errorf("received status %s - %s", resp, buf) } - return fmt.Errorf("received status %s - %s", resp, buf) + return s.connOperation(op) } // Consumer allows the consumption of messages. 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 + msgs chan Message // TODO: better error handling? Maybe a channel of errors? Err error } // Messages returns a channel in which this consumer will put messages onto. It is safe to range over the channel since it will be closed once // the consumer has finished either due to an error or from being cancelled. -func (c *Consumer) Messages() <-chan messagebroker.Message { +func (c *Consumer) Messages() <-chan Message { return c.msgs } @@ -147,7 +158,7 @@ func (c *Consumer) Messages() <-chan messagebroker.Message { // to read the messages func (s *Subscriber) Consume(ctx context.Context) *Consumer { consumer := &Consumer{ - msgs: make(chan messagebroker.Message), + msgs: make(chan Message), } go s.consume(ctx, consumer) @@ -174,37 +185,70 @@ func (s *Subscriber) consume(ctx context.Context, consumer *Consumer) { } } -func (s *Subscriber) readMessage() (*messagebroker.Message, error) { - err := s.conn.SetReadDeadline(time.Now().Add(time.Second)) - if err != nil { - return nil, err - } +func (s *Subscriber) readMessage() (*Message, error) { + var msg *Message + op := func(conn net.Conn) error { + err := s.conn.SetReadDeadline(time.Now().Add(time.Second)) + if err != nil { + return 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 + var topicLen uint64 + err = binary.Read(s.conn, binary.BigEndian, &topicLen) + if err != nil { + // TODO: check if this is needed elsewhere. I'm not sure where the read deadline resets.... + if neterr, ok := err.(net.Error); ok && neterr.Timeout() { + return nil + } + return err + } + + topicBuf := make([]byte, topicLen) + _, err = s.conn.Read(topicBuf) + if err != nil { + return err + } + + var dataLen uint64 + err = binary.Read(s.conn, binary.BigEndian, &dataLen) + if err != nil { + return err + } + + if dataLen <= 0 { + return nil } - return nil, err - } - if dataLen <= 0 { - return nil, nil + dataBuf := make([]byte, dataLen) + _, err = s.conn.Read(dataBuf) + if err != nil { + return err + } + + msg = &Message{ + Data: dataBuf, + Topic: string(topicBuf), + } + + return nil + } - buf := make([]byte, dataLen) - _, err = s.conn.Read(buf) + err := s.connOperation(op) if err != nil { + var neterr net.Error + if errors.As(err, &neterr) && neterr.Timeout() { + return nil, 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, err +} + +func (s *Subscriber) connOperation(op connOpp) error { + s.connMu.Lock() + defer s.connMu.Unlock() - return &msg, nil + return op(s.conn) } diff --git a/pubsub/subscriber_test.go b/pubsub/subscriber_test.go index bbc92de..61212f3 100644 --- a/pubsub/subscriber_test.go +++ b/pubsub/subscriber_test.go @@ -8,17 +8,18 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" - "github.com/willdot/messagebroker" "github.com/willdot/messagebroker/server" ) const ( - serverAddr = ":3000" + serverAddr = ":9999" + topicA = "topic a" + topicB = "topic b" ) func createServer(t *testing.T) { - server, err := server.New(context.Background(), serverAddr) + server, err := server.New(serverAddr) require.NoError(t, err) t.Cleanup(func() { @@ -72,7 +73,7 @@ func TestSubscribeToTopics(t *testing.T) { sub.Close() }) - topics := []string{"topic a", "topic b"} + topics := []string{topicA, topicB} err = sub.SubscribeToTopics(topics) require.NoError(t, err) @@ -88,12 +89,12 @@ func TestUnsubscribesFromTopic(t *testing.T) { sub.Close() }) - topics := []string{"topic a", "topic b"} + topics := []string{topicA, topicB} err = sub.SubscribeToTopics(topics) require.NoError(t, err) - err = sub.UnsubscribeToTopics([]string{"topic a"}) + err = sub.UnsubscribeToTopics([]string{topicA}) require.NoError(t, err) ctx, cancel := context.WithCancel(context.Background()) @@ -104,7 +105,7 @@ func TestUnsubscribesFromTopic(t *testing.T) { consumer := sub.Consume(ctx) require.NoError(t, err) - var receivedMessages []messagebroker.Message + var receivedMessages []Message consumerFinCh := make(chan struct{}) go func() { for msg := range consumer.Messages() { @@ -118,30 +119,26 @@ func TestUnsubscribesFromTopic(t *testing.T) { // publish a message to both topics and check the subscriber only gets the message from the 1 topic // and not the unsubscribed topic - publisher, err := NewPublisher("localhost:3000") + publisher, err := NewPublisher("localhost:9999") require.NoError(t, err) t.Cleanup(func() { publisher.Close() }) - msg := messagebroker.Message{ - Topic: "topic a", + msg := Message{ + Topic: topicA, Data: []byte("hello world"), } err = publisher.PublishMessage(msg) require.NoError(t, err) - msg.Topic = "topic b" + msg.Topic = topicB err = publisher.PublishMessage(msg) require.NoError(t, err) cancel() - // give the consumer some time to read the messages -- TODO: make better! - time.Sleep(time.Millisecond * 500) - cancel() - select { case <-consumerFinCh: break @@ -150,7 +147,7 @@ func TestUnsubscribesFromTopic(t *testing.T) { } assert.Len(t, receivedMessages, 1) - assert.Equal(t, "topic b", receivedMessages[0].Topic) + assert.Equal(t, topicB, receivedMessages[0].Topic) } func TestPublishAndSubscribe(t *testing.T) { @@ -163,7 +160,7 @@ func TestPublishAndSubscribe(t *testing.T) { sub.Close() }) - topics := []string{"topic a", "topic b"} + topics := []string{topicA, topicB} err = sub.SubscribeToTopics(topics) require.NoError(t, err) @@ -176,7 +173,7 @@ func TestPublishAndSubscribe(t *testing.T) { consumer := sub.Consume(ctx) require.NoError(t, err) - var receivedMessages []messagebroker.Message + var receivedMessages []Message consumerFinCh := make(chan struct{}) go func() { @@ -188,17 +185,17 @@ func TestPublishAndSubscribe(t *testing.T) { consumerFinCh <- struct{}{} }() - publisher, err := NewPublisher("localhost:3000") + publisher, err := NewPublisher("localhost:9999") require.NoError(t, err) t.Cleanup(func() { publisher.Close() }) // send some messages - sentMessages := make([]messagebroker.Message, 0, 10) + sentMessages := make([]Message, 0, 10) for i := 0; i < 10; i++ { - msg := messagebroker.Message{ - Topic: "topic a", + msg := Message{ + Topic: topicA, Data: []byte(fmt.Sprintf("message %d", i)), } diff --git a/server/peer.go b/server/peer.go index 997c3af..b9a2e67 100644 --- a/server/peer.go +++ b/server/peer.go @@ -1,55 +1,12 @@ package server import ( - "encoding/binary" - "fmt" "log/slog" "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) - } + "sync" - return dataLen, nil -} + "github.com/google/uuid" +) // Status represents the status of a request type Status uint8 @@ -73,27 +30,33 @@ func (s Status) String() string { 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 - } +type peer struct { + conn net.Conn + connMu sync.Mutex + name string +} - if message == "" { - return +func newPeer(conn net.Conn) peer { + return peer{ + conn: conn, + name: uuid.New().String(), } +} - 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 - } +func (p *peer) addr() net.Addr { + return p.conn.RemoteAddr() +} - _, err = p.conn.Write(msgBytes) - if err != nil { - slog.Error("failed to write message to peers connection", "error", err, "peer", p.addr()) - return - } +type connOpp func(conn net.Conn) error + +func (p *peer) connOperation(op connOpp, from string) error { + slog.Info("operation running", "from", from, "peer", p.conn.RemoteAddr(), "name", p.name, "mu addr", &p.connMu) + + p.connMu.Lock() + err := op(p.conn) + p.connMu.Unlock() + + slog.Info("operation finished", "from", from, "peer", p.conn.RemoteAddr(), "name", p.name, "mu addr", &p.connMu) + + return err } diff --git a/server/server.go b/server/server.go index 9fedf5e..fc439f1 100644 --- a/server/server.go +++ b/server/server.go @@ -1,15 +1,15 @@ package server import ( - "context" + "encoding/binary" "encoding/json" "errors" "fmt" "log/slog" "net" + "strings" "sync" - - "github.com/willdot/messagebroker" + "time" ) // Action represents the type of action that a peer requests to do @@ -31,7 +31,7 @@ type Server struct { } // New creates and starts a new server -func New(ctx context.Context, addr string) (*Server, error) { +func New(addr string) (*Server, error) { lis, err := net.Listen("tcp", addr) if err != nil { return nil, fmt.Errorf("failed to listen: %w", err) @@ -42,7 +42,7 @@ func New(ctx context.Context, addr string) (*Server, error) { topics: map[string]topic{}, } - go srv.start(ctx) + go srv.start() return srv, nil } @@ -52,7 +52,7 @@ func (s *Server) Shutdown() error { return s.lis.Close() } -func (s *Server) start(ctx context.Context) { +func (s *Server) start() { for { conn, err := s.lis.Accept() if err != nil { @@ -70,7 +70,7 @@ func (s *Server) start(ctx context.Context) { func (s *Server) handleConn(conn net.Conn) { peer := newPeer(conn) - action, err := peer.readAction() + action, err := readAction(peer) if err != nil { slog.Error("failed to read action from peer", "error", err, "peer", peer.addr()) return @@ -85,19 +85,24 @@ func (s *Server) handleConn(conn net.Conn) { s.handlePublish(peer) default: slog.Error("unknown action", "action", action, "peer", peer.addr()) - peer.writeStatus(Error, "unknown action") + writeStatus(Error, "unknown action", peer.conn) } } func (s *Server) handleSubscribe(peer peer) { // subscribe the peer to the topic - s.subscribePeerToTopic(peer) + s.subscribePeerToTopic(&peer) // 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 := peer.readAction() + action, err := readAction(peer) if err != nil { + var neterr net.Error + if errors.As(err, &neterr) && neterr.Timeout() { + time.Sleep(time.Second) + continue + } // 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()) @@ -108,125 +113,177 @@ func (s *Server) handleSubscribe(peer peer) { switch action { case Subscribe: - s.subscribePeerToTopic(peer) + s.subscribePeerToTopic(&peer) case Unsubscribe: s.handleUnsubscribe(peer) default: slog.Error("unknown action for subscriber", "action", action, "peer", peer.addr()) - peer.writeStatus(Error, "unknown action") + writeStatus(Error, "unknown action", peer.conn) continue } } } -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(), "peer", peer.addr()) - peer.writeStatus(Error, "invalid data length of topics provided") - return - } - if dataLen == 0 { - peer.writeStatus(Error, "data length of topics is 0") - return - } - - 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.writeStatus(Error, "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, "peer", peer.addr()) - peer.writeStatus(Error, "invalid topic data provided") - return - } +func (s *Server) subscribePeerToTopic(peer *peer) { + op := func(conn net.Conn) error { + // get the topics the peer wishes to subscribe to + dataLen, err := dataLength(conn) + if err != nil { + slog.Error(err.Error(), "peer", peer.addr()) + writeStatus(Error, "invalid data length of topics provided", conn) + return nil + } + if dataLen == 0 { + writeStatus(Error, "data length of topics is 0", conn) + return nil + } - s.subscribeToTopics(peer, topics) - peer.writeStatus(Subscribed, "") -} + buf := make([]byte, dataLen) + _, err = conn.Read(buf) + if err != nil { + slog.Error("failed to read subscibers topic data", "error", err, "peer", peer.addr()) + writeStatus(Error, "failed to read topic data", conn) + return nil + } -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(), "peer", peer.addr()) - peer.writeStatus(Error, "invalid data length of topics provided") - return - } - if dataLen == 0 { - peer.writeStatus(Error, "data length of topics is 0") - return - } + var topics []string + err = json.Unmarshal(buf, &topics) + if err != nil { + slog.Error("failed to unmarshal subscibers topic data", "error", err, "peer", peer.addr()) + writeStatus(Error, "invalid topic data provided", conn) + return nil + } - 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.writeStatus(Error, "failed to read topic data") - return - } + s.subscribeToTopics(peer, topics) + writeStatus(Subscribed, "", conn) - 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.writeStatus(Error, "invalid topic data provided") - return + return nil } - s.unsubscribeToTopics(peer, topics) - peer.writeStatus(Unsubscribed, "") + _ = peer.connOperation(op, "subscribe peer to topic") } -func (s *Server) handlePublish(peer peer) { - for { - dataLen, err := peer.readDataLength() +func (s *Server) handleUnsubscribe(peer peer) { + op := func(conn net.Conn) error { + // get the topics the peer wishes to unsubscribe from + dataLen, err := dataLength(conn) if err != nil { slog.Error(err.Error(), "peer", peer.addr()) - peer.writeStatus(Error, "invalid data length of data provided") - return + writeStatus(Error, "invalid data length of topics provided", conn) + return nil } if dataLen == 0 { - continue + writeStatus(Error, "data length of topics is 0", conn) + return nil } buf := make([]byte, dataLen) - _, err = peer.Read(buf) + _, err = conn.Read(buf) if err != nil { - slog.Error("failed to read data from peer", "error", err, "peer", peer.addr()) - peer.writeStatus(Error, "failed to read data") - return + slog.Error("failed to read subscibers topic data", "error", err, "peer", peer.addr()) + writeStatus(Error, "failed to read topic data", conn) + return nil } - var msg messagebroker.Message - err = json.Unmarshal(buf, &msg) + var topics []string + err = json.Unmarshal(buf, &topics) if err != nil { - slog.Error("failed to unmarshal data to message", "error", err, "peer", peer.addr()) - peer.writeStatus(Error, "invalid message") + slog.Error("failed to unmarshal subscibers topic data", "error", err, "peer", peer.addr()) + writeStatus(Error, "invalid topic data provided", conn) + return nil + } + + s.unsubscribeToTopics(peer, topics) + writeStatus(Unsubscribed, "", conn) + + return nil + } + + _ = peer.connOperation(op, "handle unsubscribe") +} + +type messageToSend struct { + topic string + data []byte +} + +func (s *Server) handlePublish(peer peer) { + for { + var message *messageToSend + + op := func(conn net.Conn) error { + dataLen, err := dataLength(conn) + if err != nil { + slog.Error(err.Error(), "peer", peer.addr()) + writeStatus(Error, "invalid data length of data provided", conn) + return nil + } + if dataLen == 0 { + return nil + } + topicBuf := make([]byte, dataLen) + _, err = conn.Read(topicBuf) + if err != nil { + slog.Error("failed to read topic from peer", "error", err, "peer", peer.addr()) + writeStatus(Error, "failed to read topic", conn) + return nil + } + + topicStr := string(topicBuf) + if !strings.HasPrefix(topicStr, "topic:") { + slog.Error("topic data does not contain topic prefix", "peer", peer.addr()) + writeStatus(Error, "topic data does not contain 'topic:' prefix", conn) + return nil + } + topicStr = strings.TrimPrefix(topicStr, "topic:") + + dataLen, err = dataLength(conn) + if err != nil { + slog.Error(err.Error(), "peer", peer.addr()) + writeStatus(Error, "invalid data length of data provided", conn) + return nil + } + if dataLen == 0 { + return nil + } + + dataBuf := make([]byte, dataLen) + _, err = conn.Read(dataBuf) + if err != nil { + slog.Error("failed to read data from peer", "error", err, "peer", peer.addr()) + writeStatus(Error, "failed to read data", conn) + return nil + } + + message = &messageToSend{ + topic: topicStr, + data: dataBuf, + } + return nil + } + + _ = peer.connOperation(op, "handle publish") + + if message == nil { continue } + // TODO: this can be done in a go routine because once we've got the message from the publisher, the publisher + // doesn't need to wait for us to send the message to all peers - topic := s.getTopic(msg.Topic) + topic := s.getTopic(message.topic) if topic != nil { - topic.sendMessageToSubscribers(msg) + topic.sendMessageToSubscribers(message.data) } } } -func (s *Server) subscribeToTopics(peer peer, topics []string) { +func (s *Server) subscribeToTopics(peer *peer, topics []string) { for _, topic := range topics { s.addSubsciberToTopic(topic, peer) } } -func (s *Server) addSubsciberToTopic(topicName string, peer peer) { +func (s *Server) addSubsciberToTopic(topicName string, peer *peer) { s.mu.Lock() defer s.mu.Unlock() @@ -280,3 +337,57 @@ func (s *Server) getTopic(topicName string) *topic { return nil } + +func readAction(peer peer) (Action, error) { + var action Action + op := func(conn net.Conn) error { + conn.SetReadDeadline(time.Now().Add(time.Second)) + + err := binary.Read(conn, binary.BigEndian, &action) + if err != nil { + return err + } + return nil + } + + err := peer.connOperation(op, "read action") + if err != nil { + return 0, fmt.Errorf("failed to read action from peer: %w", err) + } + + return action, nil +} + +func dataLength(conn net.Conn) (uint32, error) { + var dataLen uint32 + err := binary.Read(conn, binary.BigEndian, &dataLen) + if err != nil { + return 0, err + } + return dataLen, nil +} + +func writeStatus(status Status, message string, conn net.Conn) { + err := binary.Write(conn, binary.BigEndian, status) + if err != nil { + slog.Error("failed to write status to peers connection", "error", err, "peer", conn.RemoteAddr()) + return + } + + if message == "" { + return + } + + msgBytes := []byte(message) + err = binary.Write(conn, binary.BigEndian, uint32(len(msgBytes))) + if err != nil { + slog.Error("failed to write message length to peers connection", "error", err, "peer", conn.RemoteAddr()) + return + } + + _, err = conn.Write(msgBytes) + if err != nil { + slog.Error("failed to write message to peers connection", "error", err, "peer", conn.RemoteAddr()) + return + } +} diff --git a/server/server_test.go b/server/server_test.go index 6c776d1..f75d5b8 100644 --- a/server/server_test.go +++ b/server/server_test.go @@ -1,7 +1,6 @@ package server import ( - "context" "encoding/binary" "encoding/json" "fmt" @@ -11,11 +10,18 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" - "github.com/willdot/messagebroker" +) + +const ( + topicA = "topic a" + topicB = "topic b" + topicC = "topic c" + + serverAddr = ":6666" ) func createServer(t *testing.T) *Server { - srv, err := New(context.Background(), ":3000") + srv, err := New(serverAddr) require.NoError(t, err) t.Cleanup(func() { @@ -36,7 +42,7 @@ func createServerWithExistingTopic(t *testing.T, topicName string) *Server { } func createConnectionAndSubscribe(t *testing.T, topics []string) net.Conn { - conn, err := net.Dial("tcp", "localhost:3000") + conn, err := net.Dial("tcp", fmt.Sprintf("localhost%s", serverAddr)) require.NoError(t, err) err = binary.Write(conn, binary.BigEndian, Subscribe) @@ -64,29 +70,29 @@ func createConnectionAndSubscribe(t *testing.T, topics []string) net.Conn { func TestSubscribeToTopics(t *testing.T) { // create a server with an existing topic so we can test subscribing to a new and // existing topic - srv := createServerWithExistingTopic(t, "topic a") + srv := createServerWithExistingTopic(t, topicA) - _ = createConnectionAndSubscribe(t, []string{"topic a", "topic b"}) + _ = createConnectionAndSubscribe(t, []string{topicA, topicB}) assert.Len(t, srv.topics, 2) - assert.Len(t, srv.topics["topic a"].subscriptions, 1) - assert.Len(t, srv.topics["topic b"].subscriptions, 1) + assert.Len(t, srv.topics[topicA].subscriptions, 1) + assert.Len(t, srv.topics[topicB].subscriptions, 1) } func TestUnsubscribesFromTopic(t *testing.T) { - srv := createServerWithExistingTopic(t, "topic a") + srv := createServerWithExistingTopic(t, topicA) - conn := createConnectionAndSubscribe(t, []string{"topic a", "topic b", "topic c"}) + conn := createConnectionAndSubscribe(t, []string{topicA, topicB, topicC}) assert.Len(t, srv.topics, 3) - assert.Len(t, srv.topics["topic a"].subscriptions, 1) - assert.Len(t, srv.topics["topic b"].subscriptions, 1) - assert.Len(t, srv.topics["topic c"].subscriptions, 1) + assert.Len(t, srv.topics[topicA].subscriptions, 1) + assert.Len(t, srv.topics[topicB].subscriptions, 1) + assert.Len(t, srv.topics[topicC].subscriptions, 1) err := binary.Write(conn, binary.BigEndian, Unsubscribe) require.NoError(t, err) - topics := []string{"topic a", "topic b"} + topics := []string{topicA, topicB} rawTopics, err := json.Marshal(topics) require.NoError(t, err) @@ -104,25 +110,25 @@ func TestUnsubscribesFromTopic(t *testing.T) { assert.Equal(t, expectedRes, int(resp)) assert.Len(t, srv.topics, 3) - assert.Len(t, srv.topics["topic a"].subscriptions, 0) - assert.Len(t, srv.topics["topic b"].subscriptions, 0) - assert.Len(t, srv.topics["topic c"].subscriptions, 1) + assert.Len(t, srv.topics[topicA].subscriptions, 0) + assert.Len(t, srv.topics[topicB].subscriptions, 0) + assert.Len(t, srv.topics[topicC].subscriptions, 1) } func TestSubscriberClosesWithoutUnsubscribing(t *testing.T) { srv := createServer(t) - conn := createConnectionAndSubscribe(t, []string{"topic a", "topic b"}) + conn := createConnectionAndSubscribe(t, []string{topicA, topicB}) assert.Len(t, srv.topics, 2) - assert.Len(t, srv.topics["topic a"].subscriptions, 1) - assert.Len(t, srv.topics["topic b"].subscriptions, 1) + assert.Len(t, srv.topics[topicA].subscriptions, 1) + assert.Len(t, srv.topics[topicB].subscriptions, 1) // close the conn err := conn.Close() require.NoError(t, err) - publisherConn, err := net.Dial("tcp", "localhost:3000") + publisherConn, err := net.Dial("tcp", fmt.Sprintf("localhost%s", serverAddr)) require.NoError(t, err) err = binary.Write(publisherConn, binary.BigEndian, Publish) @@ -137,14 +143,14 @@ func TestSubscriberClosesWithoutUnsubscribing(t *testing.T) { require.Equal(t, len(data), n) assert.Len(t, srv.topics, 2) - assert.Len(t, srv.topics["topic a"].subscriptions, 0) - assert.Len(t, srv.topics["topic b"].subscriptions, 0) + assert.Len(t, srv.topics[topicA].subscriptions, 0) + assert.Len(t, srv.topics[topicB].subscriptions, 0) } func TestInvalidAction(t *testing.T) { _ = createServer(t) - conn, err := net.Dial("tcp", "localhost:3000") + conn, err := net.Dial("tcp", fmt.Sprintf("localhost%s", serverAddr)) require.NoError(t, err) err = binary.Write(conn, binary.BigEndian, uint8(99)) @@ -170,24 +176,21 @@ func TestInvalidAction(t *testing.T) { assert.Equal(t, expectedMessage, string(buf)) } -func TestInvalidMessagePublished(t *testing.T) { +func TestInvalidTopicDataPublished(t *testing.T) { _ = createServer(t) - publisherConn, err := net.Dial("tcp", "localhost:3000") + publisherConn, err := net.Dial("tcp", fmt.Sprintf("localhost%s", serverAddr)) require.NoError(t, err) err = binary.Write(publisherConn, binary.BigEndian, Publish) require.NoError(t, err) - // send some data - data := []byte("this isn't wrapped in a message type") - - // send data length first - err = binary.Write(publisherConn, binary.BigEndian, uint32(len(data))) + // send topic + topic := topicA + err = binary.Write(publisherConn, binary.BigEndian, uint32(len(topic))) require.NoError(t, err) - n, err := publisherConn.Write(data) + _, err = publisherConn.Write([]byte(topic)) require.NoError(t, err) - require.Equal(t, len(data), n) expectedRes := Error @@ -196,7 +199,7 @@ func TestInvalidMessagePublished(t *testing.T) { assert.Equal(t, expectedRes, int(resp)) - expectedMessage := "invalid message" + expectedMessage := "topic data does not contain 'topic:' prefix" var dataLen uint32 err = binary.Read(publisherConn, binary.BigEndian, &dataLen) @@ -212,37 +215,45 @@ func TestInvalidMessagePublished(t *testing.T) { func TestSendsDataToTopicSubscribers(t *testing.T) { _ = createServer(t) - subscribers := make([]net.Conn, 0, 5) - for i := 0; i < 5; i++ { - subscriberConn := createConnectionAndSubscribe(t, []string{"topic a", "topic b"}) + subscribers := make([]net.Conn, 0, 1) + for i := 0; i < 1; i++ { + subscriberConn := createConnectionAndSubscribe(t, []string{topicA, topicB}) subscribers = append(subscribers, subscriberConn) } - publisherConn, err := net.Dial("tcp", "localhost:3000") + publisherConn, err := net.Dial("tcp", fmt.Sprintf("localhost%s", serverAddr)) require.NoError(t, err) err = binary.Write(publisherConn, binary.BigEndian, Publish) require.NoError(t, err) - // send a message - msg := messagebroker.Message{ - Topic: "topic a", - Data: []byte("hello world"), - } + topic := fmt.Sprintf("topic:%s", topicA) + messageData := "hello world" - rawMsg, err := json.Marshal(msg) + // send topic first + err = binary.Write(publisherConn, binary.BigEndian, uint32(len(topic))) + require.NoError(t, err) + _, err = publisherConn.Write([]byte(topic)) require.NoError(t, err) - // send data length first - err = binary.Write(publisherConn, binary.BigEndian, uint32(len(rawMsg))) + // now send the data + err = binary.Write(publisherConn, binary.BigEndian, uint32(len(messageData))) require.NoError(t, err) - n, err := publisherConn.Write(rawMsg) + n, err := publisherConn.Write([]byte(messageData)) require.NoError(t, err) - require.Equal(t, len(rawMsg), n) + require.Equal(t, len(messageData), n) // check the subsribers got the data for _, conn := range subscribers { + var topicLen uint64 + err = binary.Read(conn, binary.BigEndian, &topicLen) + require.NoError(t, err) + + topicBuf := make([]byte, topicLen) + _, err = conn.Read(topicBuf) + require.NoError(t, err) + assert.Equal(t, topicA, string(topicBuf)) var dataLen uint64 err = binary.Read(conn, binary.BigEndian, &dataLen) @@ -253,14 +264,14 @@ func TestSendsDataToTopicSubscribers(t *testing.T) { require.NoError(t, err) require.Equal(t, int(dataLen), n) - assert.Equal(t, rawMsg, buf) + assert.Equal(t, messageData, string(buf)) } } func TestPublishMultipleTimes(t *testing.T) { _ = createServer(t) - publisherConn, err := net.Dial("tcp", "localhost:3000") + publisherConn, err := net.Dial("tcp", fmt.Sprintf("localhost%s", serverAddr)) require.NoError(t, err) err = binary.Write(publisherConn, binary.BigEndian, Publish) @@ -268,23 +279,24 @@ func TestPublishMultipleTimes(t *testing.T) { messages := make([][]byte, 0, 10) for i := 0; i < 10; i++ { - msg := messagebroker.Message{ - Topic: "topic a", - Data: []byte(fmt.Sprintf("message %d", i)), - } - - rawMsg, err := json.Marshal(msg) - require.NoError(t, err) - - messages = append(messages, rawMsg) + messages = append(messages, []byte(fmt.Sprintf("message %d", i))) } subscribeFinCh := make(chan struct{}) // create a subscriber that will read messages - subscriberConn := createConnectionAndSubscribe(t, []string{"topic a", "topic b"}) + subscriberConn := createConnectionAndSubscribe(t, []string{topicA, topicB}) go func() { // check subscriber got all messages for _, msg := range messages { + var topicLen uint64 + err = binary.Read(subscriberConn, binary.BigEndian, &topicLen) + require.NoError(t, err) + + topicBuf := make([]byte, topicLen) + _, err = subscriberConn.Read(topicBuf) + require.NoError(t, err) + assert.Equal(t, topicA, string(topicBuf)) + var dataLen uint64 err = binary.Read(subscriberConn, binary.BigEndian, &dataLen) require.NoError(t, err) @@ -300,12 +312,20 @@ func TestPublishMultipleTimes(t *testing.T) { subscribeFinCh <- struct{}{} }() + topic := fmt.Sprintf("topic:%s", topicA) + // send multiple messages for _, msg := range messages { - // send data length first + // send topic first + err = binary.Write(publisherConn, binary.BigEndian, uint32(len(topic))) + require.NoError(t, err) + _, err = publisherConn.Write([]byte(topic)) + require.NoError(t, err) + + // now send the data err = binary.Write(publisherConn, binary.BigEndian, uint32(len(msg))) require.NoError(t, err) - n, err := publisherConn.Write(msg) + n, err := publisherConn.Write([]byte(msg)) require.NoError(t, err) require.Equal(t, len(msg), n) } diff --git a/server/subscriber.go b/server/subscriber.go deleted file mode 100644 index 8c55de3..0000000 --- a/server/subscriber.go +++ /dev/null @@ -1,26 +0,0 @@ -package server - -import ( - "encoding/binary" - "fmt" -) - -type subscriber struct { - peer peer - currentOffset int -} - -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 peer: %w", err) - } - return nil -} diff --git a/server/topic.go b/server/topic.go index 9ab22d7..3ea925c 100644 --- a/server/topic.go +++ b/server/topic.go @@ -1,12 +1,11 @@ package server import ( - "encoding/json" + "encoding/binary" + "fmt" "log/slog" "net" "sync" - - "github.com/willdot/messagebroker" ) type topic struct { @@ -15,6 +14,11 @@ type topic struct { mu sync.Mutex } +type subscriber struct { + peer *peer + currentOffset int +} + func newTopic(name string) topic { return topic{ name: name, @@ -30,21 +34,45 @@ func (t *topic) removeSubscriber(addr net.Addr) { delete(t.subscriptions, addr) } -func (t *topic) sendMessageToSubscribers(msg messagebroker.Message) { +func (t *topic) sendMessageToSubscribers(msgData []byte) { t.mu.Lock() 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(msgData) + //sendMessageOpFunc := sendMessageOp(t.name, msgData) + + err := subscriber.peer.connOperation(sendMessageOp(t.name, msgData), "send message to subscribers") if err != nil { slog.Error("failed to send to message", "error", err, "peer", addr) - continue + return + } + } +} + +func sendMessageOp(topic string, data []byte) connOpp { + return func(conn net.Conn) error { + topicLen := uint64(len(topic)) + err := binary.Write(conn, binary.BigEndian, topicLen) + if err != nil { + return fmt.Errorf("failed to send topic length: %w", err) + } + _, err = conn.Write([]byte(topic)) + if err != nil { + return fmt.Errorf("failed to send topic: %w", err) + } + + dataLen := uint64(len(data)) + + err = binary.Write(conn, binary.BigEndian, dataLen) + if err != nil { + return fmt.Errorf("failed to send data length: %w", err) + } + + _, err = conn.Write(data) + if err != nil { + return fmt.Errorf("failed to write to peer: %w", err) } + return nil } } -- 2.51.2