diff options
Diffstat (limited to 'hotline')
| -rw-r--r-- | hotline/server.go | 178 | ||||
| -rw-r--r-- | hotline/server_test.go | 115 |
2 files changed, 217 insertions, 76 deletions
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") + } +} |