diff options
Diffstat (limited to 'hotline')
| -rw-r--r-- | hotline/client_conn.go | 81 | ||||
| -rw-r--r-- | hotline/client_conn_test.go | 220 | ||||
| -rw-r--r-- | hotline/server.go | 51 | ||||
| -rw-r--r-- | hotline/server_test.go | 52 |
4 files changed, 304 insertions, 100 deletions
diff --git a/hotline/client_conn.go b/hotline/client_conn.go index c17ef35..363376e 100644 --- a/hotline/client_conn.go +++ b/hotline/client_conn.go @@ -21,6 +21,10 @@ var clientConnSortFunc = func(a, b *ClientConn) int { ) } +// sendQueueDepth is the number of transactions that can be queued for delivery to a client before +// the client is considered too slow and is disconnected. +const sendQueueDepth = 64 + // ClientConn represents a client connected to a Server type ClientConn struct { Connection io.ReadWriteCloser @@ -43,6 +47,74 @@ type ClientConn struct { Logger *slog.Logger mu sync.RWMutex + + sendCh chan Transaction + sendInit sync.Once + sendMu sync.Mutex // guards sendClosed and close(sendCh) + sendClosed bool +} + +func (cc *ClientConn) initSendQueue() { + cc.sendInit.Do(func() { cc.sendCh = make(chan Transaction, sendQueueDepth) }) +} + +// Send enqueues t for delivery to this client by its writer goroutine, preserving enqueue order. +// It never blocks: if the queue is full, the client is considered too slow and its connection is +// closed, which unblocks the client's read loop and triggers the usual Disconnect cleanup. +func (cc *ClientConn) Send(t Transaction) { + cc.initSendQueue() + + cc.sendMu.Lock() + defer cc.sendMu.Unlock() + + if cc.sendClosed { + return + } + + select { + case cc.sendCh <- t: + default: + cc.sendClosed = true + close(cc.sendCh) + + if cc.Logger != nil { + cc.Logger.Warn("Send queue full; disconnecting slow client") + } + if cc.Connection != nil { + _ = cc.Connection.Close() + } + } +} + +// closeSendQueue idempotently closes the send queue, stopping the client's writer goroutine after +// it drains any remaining queued transactions. +func (cc *ClientConn) closeSendQueue() { + cc.initSendQueue() + + cc.sendMu.Lock() + defer cc.sendMu.Unlock() + + if !cc.sendClosed { + cc.sendClosed = true + close(cc.sendCh) + } +} + +// writeLoop is the single writer to cc.Connection. Serializing all writes through one goroutine +// prevents concurrent sends from interleaving bytes within the connection's transaction framing. +// It runs until the send queue is closed or a write fails. +func (cc *ClientConn) writeLoop() { + cc.initSendQueue() + + for t := range cc.sendCh { + if _, err := io.Copy(cc.Connection, &t); err != nil { + if cc.Logger != nil { + cc.Logger.Debug("error writing transaction to client", "err", err) + } + _ = cc.Connection.Close() + return + } + } } func (cc *ClientConn) TextDecoder() *encoding.Decoder { return cc.Server.TextDecoder } @@ -109,7 +181,7 @@ func (cftm *ClientFileTransferMgr) Delete(ftType FileTransferType, id FileTransf func (cc *ClientConn) SendAll(t [2]byte, fields ...Field) { for _, c := range cc.Server.ClientMgr.List() { - cc.Server.outbox <- NewTransaction(t, c.ID, fields...) + c.Send(NewTransaction(t, c.ID, fields...)) } } @@ -124,7 +196,7 @@ func (cc *ClientConn) handleTransaction(transaction Transaction) { } for _, t := range handler(cc, &transaction) { - cc.Server.outbox <- t + cc.Server.Send(t) } } @@ -169,12 +241,15 @@ func (cc *ClientConn) Authorize(access int) bool { // Disconnect notifies other clients that a client has disconnected and closes the connection. func (cc *ClientConn) Disconnect() { + // Remove the client from the manager first so no new transactions are routed to it. cc.Server.ClientMgr.Delete(cc.ID) for _, t := range cc.NotifyOthers(NewTransaction(TranNotifyDeleteUser, [2]byte{}, NewField(FieldUserID, cc.ID[:]))) { - cc.Server.outbox <- t + cc.Server.Send(t) } + cc.closeSendQueue() + if err := cc.Connection.Close(); err != nil { cc.Server.Logger.Debug("error closing client connection", "remoteAddr", cc.RemoteAddr) } diff --git a/hotline/client_conn_test.go b/hotline/client_conn_test.go index 94e38e9..5c5463d 100644 --- a/hotline/client_conn_test.go +++ b/hotline/client_conn_test.go @@ -1,7 +1,10 @@ package hotline import ( + "bufio" "bytes" + "fmt" + "sync" "testing" "github.com/stretchr/testify/assert" @@ -343,13 +346,11 @@ func TestClientConn_Disconnect(t *testing.T) { mockMgr.On("Delete", ClientID{0, 1}).Return() mockMgr.On("List").Return([]*ClientConn{}) - outbox := make(chan Transaction, 10) cc := &ClientConn{ ID: ClientID{0, 1}, Connection: &nopCloserRWC{Buffer: &bytes.Buffer{}}, Server: &Server{ ClientMgr: mockMgr, - outbox: outbox, Logger: NewTestLogger(), }, } @@ -357,39 +358,41 @@ func TestClientConn_Disconnect(t *testing.T) { cc.Disconnect() mockMgr.AssertCalled(t, "Delete", ClientID{0, 1}) - assert.Empty(t, outbox) // No other clients to notify }) t.Run("notifies other clients", func(t *testing.T) { + peer2 := &ClientConn{ID: ClientID{0, 2}} + peer3 := &ClientConn{ID: ClientID{0, 3}} + mockMgr := &MockClientMgr{} mockMgr.On("Delete", ClientID{0, 1}).Return() mockMgr.On("List").Return([]*ClientConn{ {ID: ClientID{0, 1}}, - {ID: ClientID{0, 2}}, - {ID: ClientID{0, 3}}, + peer2, + peer3, }) + mockMgr.On("Get", ClientID{0, 2}).Return(peer2) + mockMgr.On("Get", ClientID{0, 3}).Return(peer3) - outbox := make(chan Transaction, 10) cc := &ClientConn{ ID: ClientID{0, 1}, Connection: &nopCloserRWC{Buffer: &bytes.Buffer{}}, Server: &Server{ ClientMgr: mockMgr, - outbox: outbox, Logger: NewTestLogger(), }, } cc.Disconnect() - assert.Len(t, outbox, 2) + assert.Len(t, peer2.sendCh, 1) + assert.Len(t, peer3.sendCh, 1) mockMgr.AssertExpectations(t) }) } func TestClientConn_handleTransaction(t *testing.T) { t.Run("dispatches to registered handler", func(t *testing.T) { - outbox := make(chan Transaction, 10) mockMgr := &MockClientMgr{} cc := &ClientConn{ @@ -397,7 +400,6 @@ func TestClientConn_handleTransaction(t *testing.T) { Account: &Account{}, Logger: NewTestLogger(), Server: &Server{ - outbox: outbox, ClientMgr: mockMgr, handlers: map[TranType]HandlerFunc{ TranChatSend: func(cc *ClientConn, t *Transaction) []Transaction { @@ -406,15 +408,16 @@ func TestClientConn_handleTransaction(t *testing.T) { }, }, } + mockMgr.On("Get", ClientID{0, 1}).Return(cc) cc.handleTransaction(NewTransaction(TranChatSend, ClientID{0, 1})) - assert.Len(t, outbox, 1) + assert.Len(t, cc.sendCh, 1) assert.Equal(t, 0, cc.IdleTime) }) t.Run("keepalive does not reset idle time", func(t *testing.T) { - outbox := make(chan Transaction, 10) + mockMgr := &MockClientMgr{} cc := &ClientConn{ ID: ClientID{0, 1}, @@ -422,7 +425,7 @@ func TestClientConn_handleTransaction(t *testing.T) { IdleTime: 100, Logger: NewTestLogger(), Server: &Server{ - outbox: outbox, + ClientMgr: mockMgr, handlers: map[TranType]HandlerFunc{ TranKeepAlive: func(cc *ClientConn, t *Transaction) []Transaction { return []Transaction{cc.NewReply(t)} @@ -430,6 +433,7 @@ func TestClientConn_handleTransaction(t *testing.T) { }, }, } + mockMgr.On("Get", ClientID{0, 1}).Return(cc) cc.handleTransaction(NewTransaction(TranKeepAlive, ClientID{0, 1})) @@ -437,11 +441,9 @@ func TestClientConn_handleTransaction(t *testing.T) { }) t.Run("non-keepalive clears away flag", func(t *testing.T) { - outbox := make(chan Transaction, 10) + peer := &ClientConn{ID: ClientID{0, 1}} mockMgr := &MockClientMgr{} - mockMgr.On("List").Return([]*ClientConn{ - {ID: ClientID{0, 1}}, - }) + mockMgr.On("List").Return([]*ClientConn{peer}) cc := &ClientConn{ ID: ClientID{0, 1}, @@ -451,7 +453,6 @@ func TestClientConn_handleTransaction(t *testing.T) { IdleTime: 50, Logger: NewTestLogger(), Server: &Server{ - outbox: outbox, ClientMgr: mockMgr, handlers: map[TranType]HandlerFunc{ TranChatSend: func(cc *ClientConn, t *Transaction) []Transaction { @@ -467,40 +468,195 @@ func TestClientConn_handleTransaction(t *testing.T) { assert.Equal(t, 0, cc.IdleTime) assert.False(t, cc.Flags.IsSet(UserFlagAway)) // SendAll should have sent TranNotifyChangeUser - assert.Greater(t, len(outbox), 0) + assert.Greater(t, len(peer.sendCh), 0) }) } func TestClientConn_SendAll(t *testing.T) { - mockMgr := &MockClientMgr{} - mockMgr.On("List").Return([]*ClientConn{ + peers := []*ClientConn{ {ID: ClientID{0, 1}}, {ID: ClientID{0, 2}}, {ID: ClientID{0, 3}}, - }) + } + mockMgr := &MockClientMgr{} + mockMgr.On("List").Return(peers) - outbox := make(chan Transaction, 10) cc := &ClientConn{ ID: ClientID{0, 1}, Server: &Server{ ClientMgr: mockMgr, - outbox: outbox, }, } cc.SendAll(TranChatMsg, NewField(FieldData, []byte("hello"))) - assert.Len(t, outbox, 3) + for _, peer := range peers { + assert.Len(t, peer.sendCh, 1) - clientIDs := make(map[ClientID]bool) - for range 3 { - tran := <-outbox - clientIDs[tran.ClientID] = true + tran := <-peer.sendCh + assert.Equal(t, peer.ID, tran.ClientID) assert.Equal(t, TranChatMsg, tran.Type) } - assert.True(t, clientIDs[ClientID{0, 1}]) - assert.True(t, clientIDs[ClientID{0, 2}]) - assert.True(t, clientIDs[ClientID{0, 3}]) mockMgr.AssertExpectations(t) } + +// closeRecorderRWC records whether Close was called. +type closeRecorderRWC struct { + *bytes.Buffer + closed bool +} + +func (c *closeRecorderRWC) Close() error { + c.closed = true + return nil +} + +// TestClientConn_writeLoop_ordering verifies that transactions are written to the connection in +// the order they were enqueued. +func TestClientConn_writeLoop_ordering(t *testing.T) { + buf := &bytes.Buffer{} + cc := &ClientConn{ + ID: ClientID{0, 1}, + Connection: &nopCloserRWC{Buffer: buf}, + Logger: NewTestLogger(), + } + + done := make(chan struct{}) + go func() { + defer close(done) + cc.writeLoop() + }() + + const numTrans = 50 + for i := range numTrans { + cc.Send(NewTransaction(TranChatMsg, cc.ID, NewField(FieldData, fmt.Appendf(nil, "msg-%03d", i)))) + } + + // Closing the queue stops writeLoop after it drains the remaining transactions. + cc.closeSendQueue() + <-done + + scanner := bufio.NewScanner(bytes.NewReader(buf.Bytes())) + scanner.Split(transactionScanner) + + var count int + for scanner.Scan() { + var tran Transaction + _, err := tran.Write(scanner.Bytes()) + require.NoError(t, err, "transaction %d is malformed", count) + + assert.Equal(t, fmt.Sprintf("msg-%03d", count), string(tran.GetField(FieldData).Data)) + count++ + } + require.NoError(t, scanner.Err()) + assert.Equal(t, numTrans, count) +} + +// TestClientConn_writeLoop_noInterleaving verifies that concurrent senders cannot interleave bytes +// within the connection's transaction framing. +func TestClientConn_writeLoop_noInterleaving(t *testing.T) { + buf := &bytes.Buffer{} + cc := &ClientConn{ + ID: ClientID{0, 1}, + Connection: &nopCloserRWC{Buffer: buf}, + Logger: NewTestLogger(), + } + + done := make(chan struct{}) + go func() { + defer close(done) + cc.writeLoop() + }() + + // Total sends must stay within sendQueueDepth so the queue cannot overflow even if writeLoop + // has not started draining yet. + const senders, transPerSender = 4, 10 + + var wg sync.WaitGroup + for s := range senders { + wg.Add(1) + go func() { + defer wg.Done() + for i := range transPerSender { + cc.Send(NewTransaction(TranChatMsg, cc.ID, NewField(FieldData, fmt.Appendf(nil, "sender-%d-msg-%d", s, i)))) + } + }() + } + wg.Wait() + + cc.closeSendQueue() + <-done + + scanner := bufio.NewScanner(bytes.NewReader(buf.Bytes())) + scanner.Split(transactionScanner) + + var count int + for scanner.Scan() { + var tran Transaction + _, err := tran.Write(scanner.Bytes()) + require.NoError(t, err, "transaction %d is malformed", count) + assert.Equal(t, TranChatMsg, tran.Type) + count++ + } + require.NoError(t, scanner.Err()) + assert.Equal(t, senders*transPerSender, count) +} + +// TestClientConn_Send_slowClientDisconnected verifies that a client whose send queue overflows has +// its connection closed and that subsequent sends are dropped without panicking. +func TestClientConn_Send_slowClientDisconnected(t *testing.T) { + conn := &closeRecorderRWC{Buffer: &bytes.Buffer{}} + cc := &ClientConn{ + ID: ClientID{0, 1}, + Connection: conn, + Logger: NewTestLogger(), + } + + // No writeLoop is running, so the queue fills up after sendQueueDepth sends. + for i := range sendQueueDepth + 1 { + cc.Send(NewTransaction(TranChatMsg, cc.ID, NewField(FieldData, fmt.Appendf(nil, "msg-%d", i)))) + } + + assert.True(t, conn.closed, "connection should be closed when the send queue overflows") + + // Further sends must be silently dropped. + cc.Send(NewTransaction(TranChatMsg, cc.ID)) +} + +// TestClientConn_SendDisconnectRace exercises concurrent Send and Disconnect calls. Run with +// -race to detect data races and unsynchronized channel closes. +func TestClientConn_SendDisconnectRace(t *testing.T) { + mockMgr := &MockClientMgr{} + mockMgr.On("Delete", ClientID{0, 1}).Return() + mockMgr.On("List").Return([]*ClientConn{}) + + cc := &ClientConn{ + ID: ClientID{0, 1}, + Connection: &nopCloserRWC{Buffer: &bytes.Buffer{}}, + Logger: NewTestLogger(), + Server: &Server{ + ClientMgr: mockMgr, + Logger: NewTestLogger(), + }, + } + + var wg sync.WaitGroup + for range 4 { + wg.Add(1) + go func() { + defer wg.Done() + for i := range 100 { + cc.Send(NewTransaction(TranChatMsg, cc.ID, NewField(FieldData, fmt.Appendf(nil, "msg-%d", i)))) + } + }() + } + + wg.Add(1) + go func() { + defer wg.Done() + cc.Disconnect() + }() + + wg.Wait() +} diff --git a/hotline/server.go b/hotline/server.go index 6ac68d8..a892b2e 100644 --- a/hotline/server.go +++ b/hotline/server.go @@ -49,8 +49,6 @@ type Server struct { FS FileStore // Storage backend to use for File storage - outbox chan Transaction - Agreement io.ReadSeeker Banner []byte @@ -124,7 +122,6 @@ type ServerConfig struct { func NewServer(options ...Option) (*Server, error) { server := Server{ handlers: make(map[TranType]HandlerFunc), - outbox: make(chan Transaction), rateLimiters: make(map[string]*rate.Limiter), FS: &OSFileStore{}, ChatMgr: NewMemChatManager(), @@ -164,7 +161,6 @@ func (s *Server) CurrentStats() StatValues { func (s *Server) ListenAndServe(ctx context.Context) error { go s.registerWithTrackers(ctx) go s.keepaliveHandler(ctx) - go s.processOutbox() var wg sync.WaitGroup @@ -245,29 +241,11 @@ func (s *Server) ServeFileTransfersWithTLS(ctx context.Context, ln net.Listener) return s.ServeFileTransfers(ctx, tls.NewListener(ln, s.TLSConfig)) } -func (s *Server) sendTransaction(t Transaction) error { - client := s.ClientMgr.Get(t.ClientID) - - if client == nil { - return nil - } - - _, err := io.Copy(client.Connection, &t) - if err != nil { - return fmt.Errorf("failed to send transaction to client %v: %v", t.ClientID, err) - } - - return nil -} - -func (s *Server) processOutbox() { - for { - t := <-s.outbox - go func() { - if err := s.sendTransaction(t); err != nil { - s.Logger.Error("error sending transaction", "err", err) - } - }() +// Send routes t to the send queue of the client identified by t.ClientID. Transactions for +// clients that are no longer connected are dropped. +func (s *Server) Send(t Transaction) { + if c := s.ClientMgr.Get(t.ClientID); c != nil { + c.Send(t) } } @@ -532,7 +510,10 @@ func (s *Server) handleNewConnection(ctx context.Context, rwc io.ReadWriteCloser } c := s.NewClientConn(rwc, remoteAddr) - // Add the client to the list of connected clients + + // Start the client's writer goroutine: the single writer to the connection, which preserves + // transaction ordering and prevents interleaved writes. + go c.writeLoop() // TODO: refactor this into a connection manager interface, maybe? if s.Redis != nil { @@ -590,14 +571,14 @@ func (s *Server) handleNewConnection(ctx context.Context, rwc io.ReadWriteCloser c.Flags.Set(UserFlagAdmin, 1) } - s.outbox <- c.NewReply(&clientLogin, + c.Send(c.NewReply(&clientLogin, NewField(FieldVersion, []byte{0x00, 0xbe}), NewField(FieldCommunityBannerID, []byte{0, 0}), NewField(FieldServerName, []byte(s.Config.Name)), - ) + )) // Send user access privs so client UI knows how to behave - c.Server.outbox <- NewTransaction(TranUserAccess, c.ID, NewField(FieldUserAccess, c.Account.Access[:])) + c.Send(NewTransaction(TranUserAccess, c.ID, NewField(FieldUserAccess, c.Account.Access[:]))) // Accounts with AccessNoAgreement do not receive the server agreement on login. The behavior is different between // client versions. For 1.2.3 client, we do not send TranShowAgreement. For other client versions, we send @@ -605,13 +586,13 @@ func (s *Server) handleNewConnection(ctx context.Context, rwc io.ReadWriteCloser if c.Authorize(AccessNoAgreement) { // If client version is nil, then the client uses the 1.2.3 login behavior if c.Version != nil { - c.Server.outbox <- NewTransaction(TranShowAgreement, c.ID, NewField(FieldNoServerAgreement, []byte{1})) + c.Send(NewTransaction(TranShowAgreement, c.ID, NewField(FieldNoServerAgreement, []byte{1}))) } } else { _, _ = c.Server.Agreement.Seek(0, 0) data, _ := io.ReadAll(c.Server.Agreement) - c.Server.outbox <- NewTransaction(TranShowAgreement, c.ID, NewField(FieldData, data)) + c.Send(NewTransaction(TranShowAgreement, c.ID, NewField(FieldData, data))) } // If the client has provided a username as part of the login, we can infer that it is using the 1.2.3 login @@ -641,7 +622,7 @@ func (s *Server) handleNewConnection(ctx context.Context, rwc io.ReadWriteCloser NewField(FieldUserFlags, c.Flags[:]), ), ) { - c.Server.outbox <- t + c.Server.Send(t) } } @@ -776,7 +757,7 @@ func (s *Server) handleFileTransfer(ctx context.Context, rwc io.ReadWriter) erro func (s *Server) SendAll(t TranType, fields ...Field) { for _, c := range s.ClientMgr.List() { - s.outbox <- NewTransaction(t, c.ID, fields...) + c.Send(NewTransaction(t, c.ID, fields...)) } } diff --git a/hotline/server_test.go b/hotline/server_test.go index bebde8f..b3f4dce 100644 --- a/hotline/server_test.go +++ b/hotline/server_test.go @@ -764,66 +764,58 @@ func TestServer_CurrentStats(t *testing.T) { assert.Equal(t, 0, result.UploadsInProgress) } -func TestServer_sendTransaction(t *testing.T) { - t.Run("sends transaction to client connection", func(t *testing.T) { - wBuf := &bytes.Buffer{} +func TestServer_Send(t *testing.T) { + t.Run("enqueues transaction for the target client", func(t *testing.T) { + client := &ClientConn{ + Connection: &nopCloserRWC{Buffer: &bytes.Buffer{}}, + } mockMgr := &MockClientMgr{} - mockMgr.On("Get", ClientID{0, 1}).Return(&ClientConn{ - Connection: &nopCloserRWC{Buffer: wBuf}, - }) + mockMgr.On("Get", ClientID{0, 1}).Return(client) srv := &Server{ClientMgr: mockMgr} tran := NewTransaction(TranChatMsg, ClientID{0, 1}, NewField(FieldData, []byte("hello"))) - err := srv.sendTransaction(tran) + srv.Send(tran) - assert.NoError(t, err) - assert.Greater(t, wBuf.Len(), 0) + assert.Len(t, client.sendCh, 1) + queued := <-client.sendCh + assert.Equal(t, TranChatMsg, queued.Type) mockMgr.AssertExpectations(t) }) - t.Run("returns nil when client not found", func(t *testing.T) { + t.Run("drops transaction when client not found", func(t *testing.T) { mockMgr := &MockClientMgr{} mockMgr.On("Get", ClientID{0, 99}).Return((*ClientConn)(nil)) srv := &Server{ClientMgr: mockMgr} - tran := NewTransaction(TranChatMsg, ClientID{0, 99}) - err := srv.sendTransaction(tran) + srv.Send(NewTransaction(TranChatMsg, ClientID{0, 99})) - assert.NoError(t, err) mockMgr.AssertExpectations(t) }) } func TestServer_SendAll(t *testing.T) { - mockMgr := &MockClientMgr{} - mockMgr.On("List").Return([]*ClientConn{ + peers := []*ClientConn{ {ID: ClientID{0, 1}}, {ID: ClientID{0, 2}}, {ID: ClientID{0, 3}}, - }) - - outbox := make(chan Transaction, 10) - srv := &Server{ - ClientMgr: mockMgr, - outbox: outbox, } + mockMgr := &MockClientMgr{} + mockMgr.On("List").Return(peers) + + srv := &Server{ClientMgr: mockMgr} srv.SendAll(TranChatMsg, NewField(FieldData, []byte("broadcast"))) - assert.Len(t, outbox, 3) + // Verify each transaction was queued for its own client + for _, peer := range peers { + assert.Len(t, peer.sendCh, 1) - // Verify each transaction targets a different client - clientIDs := make(map[ClientID]bool) - for range 3 { - tran := <-outbox - clientIDs[tran.ClientID] = true + tran := <-peer.sendCh + assert.Equal(t, peer.ID, tran.ClientID) assert.Equal(t, TranChatMsg, tran.Type) } - assert.True(t, clientIDs[ClientID{0, 1}]) - assert.True(t, clientIDs[ClientID{0, 2}]) - assert.True(t, clientIDs[ClientID{0, 3}]) mockMgr.AssertExpectations(t) } |