aboutsummaryrefslogtreecommitdiff
path: root/hotline/server_test.go
diff options
context:
space:
mode:
Diffstat (limited to 'hotline/server_test.go')
-rw-r--r--hotline/server_test.go115
1 files changed, 115 insertions, 0 deletions
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")
+ }
+}