diff --git a/server/peer/peer.go b/server/peer/peer.go index d9add96..052e31f 100644 --- a/server/peer/peer.go +++ b/server/peer/peer.go @@ -5,24 +5,30 @@ import ( "sync" ) +// Peer represents a remote connection to the server such as a publisher or subscriber type Peer struct { conn net.Conn connMu sync.Mutex } +// New returns a new peer func New(conn net.Conn) *Peer { return &Peer{ conn: conn, } } +// Addr returns the peers connections address func (p *Peer) Addr() net.Addr { return p.conn.RemoteAddr() } +// ConnOpp represents a set of actions on a connection that can be used synchrnously type ConnOpp func(conn net.Conn) error -func (p *Peer) ConnOperation(op ConnOpp) error { +// RunConnOperation will run the provided operation. It ensures that it is the only operation that is being +// run on the connection to ensure any other operations don't get mixed up. +func (p *Peer) RunConnOperation(op ConnOpp) error { p.connMu.Lock() defer p.connMu.Unlock() diff --git a/server/server.go b/server/server.go index c5f4893..a61708b 100644 --- a/server/server.go +++ b/server/server.go @@ -185,7 +185,7 @@ func (s *Server) subscribePeerToTopic(peer *peer.Peer) { return nil } - _ = peer.ConnOperation(op) + _ = peer.RunConnOperation(op) } func (s *Server) handleUnsubscribe(peer *peer.Peer) { @@ -224,7 +224,7 @@ func (s *Server) handleUnsubscribe(peer *peer.Peer) { return nil } - _ = peer.ConnOperation(op) + _ = peer.RunConnOperation(op) } type messageToSend struct { @@ -287,7 +287,7 @@ func (s *Server) handlePublish(peer *peer.Peer) { return nil } - _ = peer.ConnOperation(op) + _ = peer.RunConnOperation(op) if message == nil { continue @@ -377,7 +377,7 @@ func readAction(peer *peer.Peer, timeout time.Duration) (Action, error) { return nil } - err := peer.ConnOperation(op) + err := peer.RunConnOperation(op) if err != nil { return 0, fmt.Errorf("failed to read action from peer: %w", err) } @@ -391,7 +391,7 @@ func writeInvalidAction(peer *peer.Peer) { return nil } - _ = peer.ConnOperation(op) + _ = peer.RunConnOperation(op) } func dataLength(conn net.Conn) (uint32, error) { diff --git a/server/topic.go b/server/topic.go index a5352cd..3d901da 100644 --- a/server/topic.go +++ b/server/topic.go @@ -42,7 +42,7 @@ func (t *topic) sendMessageToSubscribers(msgData []byte) { t.mu.Unlock() for addr, subscriber := range subscribers { - err := subscriber.peer.ConnOperation(sendMessageOp(t.name, msgData)) + err := subscriber.peer.RunConnOperation(sendMessageOp(t.name, msgData)) if err != nil { slog.Error("failed to send to message", "error", err, "peer", addr) return