aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-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")
+ }
+}