diff options
| -rw-r--r-- | hotline/server.go | 96 | ||||
| -rw-r--r-- | hotline/server_test.go | 52 |
2 files changed, 140 insertions, 8 deletions
diff --git a/hotline/server.go b/hotline/server.go index 700cf19..4bf48c6 100644 --- a/hotline/server.go +++ b/hotline/server.go @@ -77,12 +77,59 @@ type Server struct { shutdownInit sync.Once // lazily creates shutdownCh so Shutdown works on test-constructed Servers shutdownOnce sync.Once // guards close(shutdownCh) shutdownCh chan struct{} // closed by Shutdown to stop ListenAndServe + + // activeConns tracks every accepted connection (sessions and file transfers) so shutdown can + // force-close them, which unblocks their read loops. connWG counts the goroutines serving + // them (including the client writer goroutines) so ListenAndServe can wait for all of them to + // finish before returning. activeConns is lazily initialized under activeConnsMu so + // test-constructed Servers keep working. + activeConns map[io.Closer]struct{} + activeConnsMu sync.Mutex + connsClosed bool // set once shutdown force-closes connections; later conns are closed on arrival + connWG sync.WaitGroup } func (s *Server) initShutdownCh() { s.shutdownInit.Do(func() { s.shutdownCh = make(chan struct{}) }) } +// trackConn registers an accepted connection for forced close on shutdown. It reports whether +// the connection was registered: once shutdown has begun the connection is closed instead, and +// the caller must not serve it. +func (s *Server) trackConn(conn io.Closer) bool { + s.activeConnsMu.Lock() + defer s.activeConnsMu.Unlock() + + if s.connsClosed { + _ = conn.Close() + return false + } + if s.activeConns == nil { + s.activeConns = make(map[io.Closer]struct{}) + } + s.activeConns[conn] = struct{}{} + return true +} + +func (s *Server) untrackConn(conn io.Closer) { + s.activeConnsMu.Lock() + defer s.activeConnsMu.Unlock() + + delete(s.activeConns, conn) +} + +// closeActiveConns force-closes all tracked connections, unblocking the read loops of the +// goroutines serving them. Connections accepted after this are closed on arrival by trackConn. +func (s *Server) closeActiveConns() { + s.activeConnsMu.Lock() + defer s.activeConnsMu.Unlock() + + s.connsClosed = true + for conn := range s.activeConns { + _ = conn.Close() + } +} + // Banner returns the server banner image. Callers must not modify the returned slice. func (s *Server) Banner() []byte { s.bannerMu.RLock() @@ -200,7 +247,8 @@ func (s *Server) CurrentStats() StatValues { // ListenAndServe starts the Hotline and file transfer listeners and blocks until the context is // canceled, Shutdown is called, or a serve loop fails. Canceling the context closes all -// listeners, which unblocks their accept loops. +// listeners, which unblocks their accept loops. Accepted connections are then force-closed and +// ListenAndServe waits for every session and file transfer goroutine to finish before returning. func (s *Server) ListenAndServe(ctx context.Context) error { ctx, cancel := context.WithCancel(ctx) defer cancel() @@ -253,12 +301,19 @@ func (s *Server) ListenAndServe(ctx context.Context) error { } } - // Block until the first serve loop returns. The deferred cancel closes the remaining - // listeners and stops their serve loops. + // Block until the first serve loop returns, then cancel to close the remaining listeners and + // stop their serve loops. err := <-errCh if ctx.Err() != nil { s.Logger.Info("Server shutting down") } + cancel() + + // Force-close accepted connections to unblock their read loops, then wait for every session, + // file transfer, and client writer goroutine to finish before returning. + s.closeActiveConns() + s.connWG.Wait() + return err } @@ -272,10 +327,17 @@ func (s *Server) ServeFileTransfers(ctx context.Context, ln net.Listener) error return err } + s.connWG.Add(1) go func() { + defer s.connWG.Done() + + if !s.trackConn(conn) { + return + } + defer s.untrackConn(conn) defer func() { _ = conn.Close() }() - err = s.handleFileTransfer( + err := s.handleFileTransfer( context.WithValue(ctx, contextKeyReq, requestCtx{remoteAddr: conn.RemoteAddr().String()}), conn, ) @@ -346,7 +408,15 @@ func (s *Server) Serve(ctx context.Context, ln net.Listener) error { continue } + s.connWG.Add(1) go func() { + defer s.connWG.Done() + + if !s.trackConn(conn) { + return + } + defer s.untrackConn(conn) + ipAddr, _, _ := net.SplitHostPort(conn.RemoteAddr().String()) connCtx := context.WithValue(ctx, contextKeyReq, requestCtx{ @@ -587,8 +657,13 @@ func (s *Server) handleNewConnection(ctx context.Context, rwc io.ReadWriteCloser c := s.NewClientConn(rwc, remoteAddr) // Start the client's writer goroutine: the single writer to the connection, which preserves - // transaction ordering and prevents interleaved writes. - go c.writeLoop() + // transaction ordering and prevents interleaved writes. It exits when the send queue is + // closed by Disconnect, so it is joined by the same WaitGroup as the session goroutines. + s.connWG.Add(1) + go func() { + defer s.connWG.Done() + c.writeLoop() + }() if s.Presence != nil { s.Presence.UserConnected(login, ipAddr) @@ -741,7 +816,11 @@ func (s *Server) handleFileTransfer(ctx context.Context, rwc io.ReadWriter) erro // Wait a few seconds before closing the connection: this is a workaround for problems // observed with Windows clients where the client must initiate close of the TCP connection before // the server does. This is gross and seems unnecessary. TODO: Revisit? - time.Sleep(3 * time.Second) + // Skipped during shutdown, when it would only delay process exit. + select { + case <-ctx.Done(): + case <-time.After(3 * time.Second): + } }() var remoteAddr string @@ -834,7 +913,8 @@ func (s *Server) SendAll(t TranType, fields ...Field) { } } -// Shutdown sends msg to all connected clients and stops ListenAndServe. +// Shutdown sends msg to all connected clients and stops ListenAndServe, which force-closes any +// remaining connections and waits for their goroutines to finish. func (s *Server) Shutdown(msg []byte) { s.Logger.Info("Shutdown signal received") s.SendAll(TranDisconnectMsg, NewField(FieldData, msg)) diff --git a/hotline/server_test.go b/hotline/server_test.go index a9842c8..fb9a9c7 100644 --- a/hotline/server_test.go +++ b/hotline/server_test.go @@ -1018,6 +1018,58 @@ func TestServer_ListenAndServe_returnsErrorWhenPortUnavailable(t *testing.T) { assert.NotErrorIs(t, err, context.Canceled) } +// TestServer_ListenAndServe_closesActiveConnsOnCancel verifies that shutdown force-closes +// accepted connections rather than only closing the listeners: a client blocked mid-handshake and +// a file transfer connection that never sends its header would otherwise keep their session +// goroutines alive past ListenAndServe's return. +func TestServer_ListenAndServe_closesActiveConnsOnCancel(t *testing.T) { + port := findFreePortPair(t) + srv, err := NewServer( + WithLogger(NewTestLogger()), + WithInterface("127.0.0.1"), + WithPort(port), + ) + require.NoError(t, err) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + errCh := make(chan error, 1) + go func() { errCh <- srv.ListenAndServe(ctx) }() + + // Give the listeners a moment to start. + time.Sleep(100 * time.Millisecond) + + // A session connection that stalls mid-handshake and a file transfer connection that never + // sends its 16-byte header: both park their serving goroutines in blocking reads. + sessionConn, err := net.Dial("tcp", fmt.Sprintf("127.0.0.1:%d", port)) + require.NoError(t, err) + defer func() { _ = sessionConn.Close() }() + + transferConn, err := net.Dial("tcp", fmt.Sprintf("127.0.0.1:%d", port+1)) + require.NoError(t, err) + defer func() { _ = transferConn.Close() }() + + // Give the accept loops a moment to hand the connections to their goroutines. + time.Sleep(100 * time.Millisecond) + + cancel() + + select { + case <-errCh: + case <-time.After(2 * time.Second): + t.Fatal("ListenAndServe did not return after context cancellation with active connections") + } + + // Both connections must have been closed by the server: reads unblock with EOF. + for _, conn := range []net.Conn{sessionConn, transferConn} { + require.NoError(t, conn.SetReadDeadline(time.Now().Add(2*time.Second))) + _, err = conn.Read(make([]byte, 1)) + assert.Error(t, err) + assert.NotErrorIs(t, err, os.ErrDeadlineExceeded) + } +} + func TestServer_Shutdown_stopsListenAndServe(t *testing.T) { srv, err := NewServer( WithLogger(NewTestLogger()), |