aboutsummaryrefslogtreecommitdiff
path: root/hotline/server.go
diff options
context:
space:
mode:
Diffstat (limited to 'hotline/server.go')
-rw-r--r--hotline/server.go96
1 files changed, 88 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))