aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJeff Halter <868228+jhalter@users.noreply.github.com>2026-06-12 08:26:25 -0700
committerJeff Halter <868228+jhalter@users.noreply.github.com>2026-06-12 08:26:25 -0700
commitb2c462a3a1353f0653a5964b3a6924538ce83523 (patch)
tree90968ec234bbfa44e58fcf887efd9ba15dc280dd
parent7ebc802d0a269218f05b3b51eb10ac66eacb4d1f (diff)
Shut down gracefully instead of exiting from library code
ListenAndServe previously started each listener in a goroutine that called log.Fatal on any error, which skipped deferred cleanup and made errors unobservable to callers, and Server.Shutdown terminated the process with os.Exit. Context cancellation was also ineffective: Serve only checked ctx between Accept calls, which block indefinitely. ListenAndServe now binds its listeners up front and returns bind errors, closes every listener when the context is canceled so accept loops unblock and return, and reports the first serve loop error to the caller. Shutdown closes a lazily-initialized channel that cancels ListenAndServe's context, so the shutdown API works race-free even though it starts before ListenAndServe. "Server shutting down" is logged once by ListenAndServe rather than per accept loop, which produced duplicate or missing lines depending on scheduling. main.go now treats context.Canceled as a clean exit, logs other server errors and exits nonzero, and runs deferred cleanup (e.g. Bonjour shutdown) on the way out.
-rw-r--r--cmd/mobius-hotline-server/main.go14
-rw-r--r--hotline/server.go178
-rw-r--r--hotline/server_test.go115
3 files changed, 227 insertions, 80 deletions
diff --git a/cmd/mobius-hotline-server/main.go b/cmd/mobius-hotline-server/main.go
index dced2c5..1c298a2 100644
--- a/cmd/mobius-hotline-server/main.go
+++ b/cmd/mobius-hotline-server/main.go
@@ -4,10 +4,10 @@ import (
"context"
"crypto/tls"
"embed"
+ "errors"
"flag"
"fmt"
"io"
- "log"
"os"
"os/signal"
"path"
@@ -207,9 +207,9 @@ func main() {
reloadFunc()
default:
+ // Canceling the context stops ListenAndServe, which unblocks main for a clean exit.
signal.Stop(sigChan)
cancel()
- os.Exit(0)
}
}
@@ -235,8 +235,14 @@ func main() {
defer s.Shutdown()
}
- // Serve Hotline requests until program exit
- log.Fatal(srv.ListenAndServe(ctx))
+ // Serve Hotline requests until shutdown is requested via signal, the shutdown API, or a
+ // server error.
+ if err := srv.ListenAndServe(ctx); err != nil && !errors.Is(err, context.Canceled) {
+ slogger.Error("Server error", "err", err)
+ os.Exit(1)
+ }
+
+ slogger.Info("Server shut down")
}
// findConfigPath searches for an existing config directory from the predefined search order.
diff --git a/hotline/server.go b/hotline/server.go
index a892b2e..3459346 100644
--- a/hotline/server.go
+++ b/hotline/server.go
@@ -7,12 +7,11 @@ import (
"crypto/rand"
"crypto/tls"
"encoding/binary"
+ "errors"
"fmt"
"io"
- "log"
"log/slog"
"net"
- "os"
"strings"
"sync"
"time"
@@ -71,6 +70,14 @@ type Server struct {
TLSConfig *tls.Config
TLSPort int
+
+ 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
+}
+
+func (s *Server) initShutdownCh() {
+ s.shutdownInit.Do(func() { s.shutdownCh = make(chan struct{}) })
}
type Option = func(s *Server)
@@ -158,63 +165,77 @@ func (s *Server) CurrentStats() StatValues {
return s.Stats.Values()
}
+// 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.
func (s *Server) ListenAndServe(ctx context.Context) error {
- go s.registerWithTrackers(ctx)
- go s.keepaliveHandler(ctx)
-
- var wg sync.WaitGroup
+ ctx, cancel := context.WithCancel(ctx)
+ defer cancel()
- wg.Add(1)
+ // Cancel the context when Shutdown is called.
+ s.initShutdownCh()
go func() {
- ln, err := net.Listen("tcp", fmt.Sprintf("%s:%v", s.NetInterface, s.Port))
- if err != nil {
- log.Fatal(err)
+ select {
+ case <-s.shutdownCh:
+ cancel()
+ case <-ctx.Done():
}
-
- log.Fatal(s.Serve(ctx, ln))
}()
- wg.Add(1)
- go func() {
- ln, err := net.Listen("tcp", fmt.Sprintf("%s:%v", s.NetInterface, s.Port+1))
+ go s.registerWithTrackers(ctx)
+ go s.keepaliveHandler(ctx)
+
+ errCh := make(chan error, 4)
+
+ listen := func(port int, serve func(context.Context, net.Listener) error) error {
+ ln, err := net.Listen("tcp", fmt.Sprintf("%s:%v", s.NetInterface, port))
if err != nil {
- log.Fatal(err)
+ return err
}
- log.Fatal(s.ServeFileTransfers(ctx, ln))
- }()
-
- if s.TLSConfig != nil {
- wg.Add(1)
+ // Close the listener when the context is canceled to unblock Accept.
go func() {
- ln, err := net.Listen("tcp", fmt.Sprintf("%s:%v", s.NetInterface, s.TLSPort))
- if err != nil {
- log.Fatal(err)
- }
-
- log.Fatal(s.ServeWithTLS(ctx, ln))
+ <-ctx.Done()
+ _ = ln.Close()
}()
- wg.Add(1)
- go func() {
- ln, err := net.Listen("tcp", fmt.Sprintf("%s:%v", s.NetInterface, s.TLSPort+1))
- if err != nil {
- log.Fatal(err)
- }
+ go func() { errCh <- serve(ctx, ln) }()
- log.Fatal(s.ServeFileTransfersWithTLS(ctx, ln))
- }()
+ return nil
}
- wg.Wait()
+ if err := listen(s.Port, s.Serve); err != nil {
+ return err
+ }
+ if err := listen(s.Port+1, s.ServeFileTransfers); err != nil {
+ return err
+ }
- return nil
+ if s.TLSConfig != nil {
+ if err := listen(s.TLSPort, s.ServeWithTLS); err != nil {
+ return err
+ }
+ if err := listen(s.TLSPort+1, s.ServeFileTransfersWithTLS); err != nil {
+ return err
+ }
+ }
+
+ // Block until the first serve loop returns. The deferred cancel closes the remaining
+ // listeners and stops their serve loops.
+ err := <-errCh
+ if ctx.Err() != nil {
+ s.Logger.Info("Server shutting down")
+ }
+ return err
}
func (s *Server) ServeFileTransfers(ctx context.Context, ln net.Listener) error {
for {
conn, err := ln.Accept()
if err != nil {
+ if ctx.Err() != nil {
+ return ctx.Err()
+ }
return err
}
@@ -255,52 +276,54 @@ const perIPRateLimit = rate.Limit(0.5)
func (s *Server) Serve(ctx context.Context, ln net.Listener) error {
for {
- select {
- case <-ctx.Done():
- s.Logger.Info("Server shutting down")
- return ctx.Err()
- default:
- conn, err := ln.Accept()
- if err != nil {
- s.Logger.Error("Error accepting connection", "err", err)
- continue
+ conn, err := ln.Accept()
+ if err != nil {
+ // Context cancellation closes the listener, unblocking Accept.
+ if ctx.Err() != nil {
+ return ctx.Err()
+ }
+ if errors.Is(err, net.ErrClosed) {
+ return err
}
- go func() {
- ipAddr, _, _ := net.SplitHostPort(conn.RemoteAddr().String())
+ s.Logger.Error("Error accepting connection", "err", err)
+ continue
+ }
- connCtx := context.WithValue(ctx, contextKeyReq, requestCtx{
- remoteAddr: conn.RemoteAddr().String(),
- })
+ go func() {
+ ipAddr, _, _ := net.SplitHostPort(conn.RemoteAddr().String())
- s.Logger.Info("Connection established", "ip", ipAddr)
- defer func() { _ = conn.Close() }()
+ connCtx := context.WithValue(ctx, contextKeyReq, requestCtx{
+ remoteAddr: conn.RemoteAddr().String(),
+ })
- // Check if we have an existing rate limit for the IP and create one if we do not.
- s.rateLimitersMu.Lock()
- rl, ok := s.rateLimiters[ipAddr]
- if !ok {
- rl = rate.NewLimiter(perIPRateLimit, 1)
- s.rateLimiters[ipAddr] = rl
- }
- s.rateLimitersMu.Unlock()
+ s.Logger.Info("Connection established", "ip", ipAddr)
+ defer func() { _ = conn.Close() }()
- // Check if the rate limit is exceeded and close the connection if so.
- if !rl.Allow() {
- s.Logger.Info("Rate limit exceeded", "remoteAddr", conn.RemoteAddr())
- _ = conn.Close()
- return
- }
+ // Check if we have an existing rate limit for the IP and create one if we do not.
+ s.rateLimitersMu.Lock()
+ rl, ok := s.rateLimiters[ipAddr]
+ if !ok {
+ rl = rate.NewLimiter(perIPRateLimit, 1)
+ s.rateLimiters[ipAddr] = rl
+ }
+ s.rateLimitersMu.Unlock()
+
+ // Check if the rate limit is exceeded and close the connection if so.
+ if !rl.Allow() {
+ s.Logger.Info("Rate limit exceeded", "remoteAddr", conn.RemoteAddr())
+ _ = conn.Close()
+ return
+ }
- if err := s.handleNewConnection(connCtx, conn, conn.RemoteAddr().String()); err != nil {
- if err == io.EOF {
- s.Logger.Info("Client disconnected", "remoteAddr", conn.RemoteAddr())
- } else {
- s.Logger.Error("Error serving request", "remoteAddr", conn.RemoteAddr(), "err", err)
- }
+ if err := s.handleNewConnection(connCtx, conn, conn.RemoteAddr().String()); err != nil {
+ if err == io.EOF {
+ s.Logger.Info("Client disconnected", "remoteAddr", conn.RemoteAddr())
+ } else {
+ s.Logger.Error("Error serving request", "remoteAddr", conn.RemoteAddr(), "err", err)
}
- }()
- }
+ }
+ }()
}
}
@@ -761,11 +784,14 @@ func (s *Server) SendAll(t TranType, fields ...Field) {
}
}
+// Shutdown sends msg to all connected clients and stops ListenAndServe.
func (s *Server) Shutdown(msg []byte) {
s.Logger.Info("Shutdown signal received")
s.SendAll(TranDisconnectMsg, NewField(FieldData, msg))
+ // Give the client writer goroutines a moment to flush the disconnect message.
time.Sleep(3 * time.Second)
- os.Exit(0)
+ s.initShutdownCh()
+ s.shutdownOnce.Do(func() { close(s.shutdownCh) })
}
diff --git a/hotline/server_test.go b/hotline/server_test.go
index b3f4dce..a6c5e19 100644
--- a/hotline/server_test.go
+++ b/hotline/server_test.go
@@ -7,6 +7,7 @@ import (
"fmt"
"io"
"log/slog"
+ "net"
"os"
"strings"
"testing"
@@ -14,6 +15,7 @@ import (
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
+ "github.com/stretchr/testify/require"
"golang.org/x/text/encoding"
"golang.org/x/text/encoding/charmap"
)
@@ -921,3 +923,116 @@ func TestSendBanMessage(t *testing.T) {
assert.Greater(t, buf.Len(), 0)
assert.Contains(t, buf.String(), "You are banned")
}
+
+// findFreePortPair returns a port p where both p and p+1 are free to listen on, as required by
+// ListenAndServe for the Hotline and file transfer listeners.
+func findFreePortPair(t *testing.T) int {
+ t.Helper()
+
+ for range 10 {
+ ln, err := net.Listen("tcp", "127.0.0.1:0")
+ require.NoError(t, err)
+ port := ln.Addr().(*net.TCPAddr).Port
+ _ = ln.Close()
+
+ ln2, err := net.Listen("tcp", fmt.Sprintf("127.0.0.1:%d", port+1))
+ if err != nil {
+ continue
+ }
+ _ = ln2.Close()
+
+ return port
+ }
+
+ t.Fatal("could not find a free port pair")
+ return 0
+}
+
+func TestServer_Serve_returnsOnContextCancel(t *testing.T) {
+ ln, err := net.Listen("tcp", "127.0.0.1:0")
+ require.NoError(t, err)
+
+ srv := &Server{Logger: NewTestLogger()}
+
+ ctx, cancel := context.WithCancel(context.Background())
+
+ errCh := make(chan error, 1)
+ go func() { errCh <- srv.Serve(ctx, ln) }()
+
+ // Cancel the context and close the listener, as ListenAndServe's watcher goroutine does.
+ cancel()
+ _ = ln.Close()
+
+ select {
+ case err := <-errCh:
+ assert.ErrorIs(t, err, context.Canceled)
+ case <-time.After(2 * time.Second):
+ t.Fatal("Serve did not return after context cancellation")
+ }
+}
+
+func TestServer_ListenAndServe_returnsOnContextCancel(t *testing.T) {
+ srv, err := NewServer(
+ WithLogger(NewTestLogger()),
+ WithInterface("127.0.0.1"),
+ WithPort(findFreePortPair(t)),
+ )
+ require.NoError(t, err)
+
+ ctx, cancel := context.WithCancel(context.Background())
+
+ errCh := make(chan error, 1)
+ go func() { errCh <- srv.ListenAndServe(ctx) }()
+
+ // Give the listeners a moment to start before canceling.
+ time.Sleep(100 * time.Millisecond)
+ cancel()
+
+ select {
+ case err := <-errCh:
+ assert.ErrorIs(t, err, context.Canceled)
+ case <-time.After(2 * time.Second):
+ t.Fatal("ListenAndServe did not return after context cancellation")
+ }
+}
+
+func TestServer_ListenAndServe_returnsErrorWhenPortUnavailable(t *testing.T) {
+ // Occupy a port so ListenAndServe fails to bind to it.
+ blocker, err := net.Listen("tcp", "127.0.0.1:0")
+ require.NoError(t, err)
+ defer func() { _ = blocker.Close() }()
+
+ srv, err := NewServer(
+ WithLogger(NewTestLogger()),
+ WithInterface("127.0.0.1"),
+ WithPort(blocker.Addr().(*net.TCPAddr).Port),
+ )
+ require.NoError(t, err)
+
+ err = srv.ListenAndServe(context.Background())
+ assert.Error(t, err)
+ assert.NotErrorIs(t, err, context.Canceled)
+}
+
+func TestServer_Shutdown_stopsListenAndServe(t *testing.T) {
+ srv, err := NewServer(
+ WithLogger(NewTestLogger()),
+ WithInterface("127.0.0.1"),
+ WithPort(findFreePortPair(t)),
+ )
+ require.NoError(t, err)
+
+ errCh := make(chan error, 1)
+ go func() { errCh <- srv.ListenAndServe(context.Background()) }()
+
+ // Give the listeners a moment to start before requesting shutdown.
+ time.Sleep(100 * time.Millisecond)
+ go srv.Shutdown([]byte("goodbye"))
+
+ select {
+ case err := <-errCh:
+ assert.ErrorIs(t, err, context.Canceled)
+ case <-time.After(5 * time.Second): // Shutdown sleeps 3 seconds to flush client queues
+ t.Fatal("ListenAndServe did not return after Shutdown")
+ }
+}