diff --git a/relay/server/listener/quic/listener.go b/relay/server/listener/quic/listener.go index 4c3b07571..115a2d0c2 100644 --- a/relay/server/listener/quic/listener.go +++ b/relay/server/listener/quic/listener.go @@ -25,21 +25,28 @@ type Listener struct { listener *quic.Listener } -func (l *Listener) Listen(acceptFn func(conn relaylistener.Conn)) error { +func (l *Listener) Bind() error { quicCfg := &quic.Config{ EnableDatagrams: true, InitialPacketSize: nbRelay.QUICInitialPacketSize, } listener, err := quic.ListenAddr(l.Address, l.TLSConfig, quicCfg) if err != nil { - return fmt.Errorf("failed to create QUIC listener: %v", err) + return err } l.listener = listener log.Infof("QUIC server listening on address: %s", l.Address) + return nil +} + +func (l *Listener) Serve(acceptFn func(conn relaylistener.Conn)) error { + if l.listener == nil { + return errors.New("listener is not bound") + } for { - session, err := listener.Accept(context.Background()) + session, err := l.listener.Accept(context.Background()) if err != nil { if errors.Is(err, quic.ErrServerClosed) { return nil diff --git a/relay/server/listener/quic/listener_test.go b/relay/server/listener/quic/listener_test.go new file mode 100644 index 000000000..6a2e9ce75 --- /dev/null +++ b/relay/server/listener/quic/listener_test.go @@ -0,0 +1,130 @@ +package quic + +import ( + "context" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "math/big" + "net" + "testing" + "time" + + "github.com/quic-go/quic-go" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + relaylistener "github.com/netbirdio/netbird/relay/server/listener" + quictls "github.com/netbirdio/netbird/shared/relay/tls" +) + +func TestListener_ShutdownBeforeServe(t *testing.T) { + l := &Listener{Address: "127.0.0.1:0", TLSConfig: testTLSConfig(t)} + require.NoError(t, l.Bind()) + addr := l.listener.Addr().String() + require.NoError(t, l.Shutdown(context.Background())) + + errChan := make(chan error, 1) + go func() { + errChan <- l.Serve(func(relaylistener.Conn) {}) + }() + + assert.NoError(t, waitForServeToReturn(t, errChan)) + requireUDPAddressFree(t, addr) +} + +func TestListener_ShutdownStopsServe(t *testing.T) { + l := &Listener{Address: "127.0.0.1:0", TLSConfig: testTLSConfig(t)} + require.NoError(t, l.Bind()) + addr := l.listener.Addr().String() + + accepted := make(chan relaylistener.Conn, 1) + errChan := make(chan error, 1) + go func() { + errChan <- l.Serve(func(conn relaylistener.Conn) { accepted <- conn }) + }() + + // quic-go completes handshakes on a bound socket before Accept is called, so + // only a session handed to acceptFn proves Serve is blocked in Accept when + // Shutdown arrives. + dialTestSession(t, addr) + var conn relaylistener.Conn + select { + case conn = <-accepted: + case <-time.After(5 * time.Second): + t.Fatal("listener did not accept the test session") + } + + require.NoError(t, l.Shutdown(context.Background())) + assert.NoError(t, waitForServeToReturn(t, errChan)) + + // Shutdown leaves accepted sessions alive, as the relay closes its peers itself, + // and quic-go keeps the UDP socket bound until the last session is gone. + require.NoError(t, conn.Close()) + requireUDPAddressFree(t, addr) +} + +func TestListener_Unbound(t *testing.T) { + l := &Listener{Address: "127.0.0.1:0", TLSConfig: testTLSConfig(t)} + assert.Error(t, l.Serve(func(relaylistener.Conn) {}), "Serve must refuse an unbound listener") + assert.NoError(t, l.Shutdown(context.Background()), "Shutdown of an unbound listener is a no-op") +} + +func testTLSConfig(t *testing.T) *tls.Config { + t.Helper() + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + require.NoError(t, err) + + template := &x509.Certificate{ + SerialNumber: big.NewInt(1), + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(time.Hour), + IPAddresses: []net.IP{net.ParseIP("127.0.0.1")}, + } + certDER, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key) + require.NoError(t, err) + + return &tls.Config{ + Certificates: []tls.Certificate{{Certificate: [][]byte{certDER}, PrivateKey: key}}, + NextProtos: []string{quictls.NBalpn}, + } +} + +func dialTestSession(t *testing.T, addr string) { + t.Helper() + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + tlsCfg := &tls.Config{InsecureSkipVerify: true, NextProtos: []string{quictls.NBalpn}} + session, err := quic.DialAddr(ctx, addr, tlsCfg, &quic.Config{EnableDatagrams: true}) + require.NoError(t, err) + t.Cleanup(func() { _ = session.CloseWithError(0, "") }) +} + +func waitForServeToReturn(t *testing.T, errChan <-chan error) error { + t.Helper() + select { + case err := <-errChan: + return err + case <-time.After(5 * time.Second): + t.Fatal("Serve did not return") + return nil + } +} + +// quic-go retires a closed session's connection IDs on a timer and closes the UDP +// socket from its read loop once the last one is gone, so the address is polled +// rather than checked once. +func requireUDPAddressFree(t *testing.T, addr string) { + t.Helper() + require.Eventually(t, func() bool { + conn, err := net.ListenPacket("udp", addr) + if err != nil { + return false + } + _ = conn.Close() + return true + }, 5*time.Second, 10*time.Millisecond, "udp address %s must be released", addr) +} diff --git a/relay/server/listener/ws/listener.go b/relay/server/listener/ws/listener.go index 208b9186e..a1a26668a 100644 --- a/relay/server/listener/ws/listener.go +++ b/relay/server/listener/ws/listener.go @@ -32,32 +32,57 @@ type Listener struct { // headers are trusted. Headers from any other immediate peer are ignored. TrustedProxies *trustedproxy.List + listener net.Listener server *http.Server acceptFn func(conn relaylistener.Conn) } -func (l *Listener) Listen(acceptFn func(conn relaylistener.Conn)) error { - l.acceptFn = acceptFn +func (l *Listener) Bind() error { + addr := l.Address + if addr == "" { + addr = ":http" + if l.TLSConfig != nil { + addr = ":https" + } + } + + listener, err := net.Listen("tcp", addr) + if err != nil { + return err + } + mux := http.NewServeMux() mux.HandleFunc(URLPath, l.onAccept) + l.listener = listener l.server = &http.Server{ - Addr: l.Address, Handler: mux, TLSConfig: l.TLSConfig, ReadHeaderTimeout: 5 * time.Second, } - log.Infof("WS server listening address: %s", l.Address) + log.Infof("WS server listening address: %s", addr) + return nil +} + +func (l *Listener) Serve(acceptFn func(conn relaylistener.Conn)) error { + if l.listener == nil { + return errors.New("listener is not bound") + } + + l.acceptFn = acceptFn var err error if l.TLSConfig != nil { - err = l.server.ListenAndServeTLS("", "") + err = l.server.ServeTLS(l.listener, "", "") } else { - err = l.server.ListenAndServe() + err = l.server.Serve(l.listener) } if errors.Is(err, http.ErrServerClosed) { return nil } + if closeErr := l.listener.Close(); closeErr != nil && !errors.Is(closeErr, net.ErrClosed) { + log.Debugf("failed to close WS listener: %v", closeErr) + } return err } @@ -66,7 +91,7 @@ func (l *Listener) Protocol() protocol.Protocol { } func (l *Listener) Shutdown(ctx context.Context) error { - if l.server == nil { + if l.listener == nil { return nil } @@ -74,6 +99,9 @@ func (l *Listener) Shutdown(ctx context.Context) error { if err := l.server.Shutdown(ctx); err != nil { return fmt.Errorf("server shutdown failed: %v", err) } + if err := l.listener.Close(); err != nil && !errors.Is(err, net.ErrClosed) { + return fmt.Errorf("close listener: %w", err) + } log.Infof("WS listener stopped") return nil } diff --git a/relay/server/listener/ws/listener_test.go b/relay/server/listener/ws/listener_test.go new file mode 100644 index 000000000..2373fb5ed --- /dev/null +++ b/relay/server/listener/ws/listener_test.go @@ -0,0 +1,104 @@ +package ws + +import ( + "context" + "crypto/tls" + "net" + "net/http" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + relaylistener "github.com/netbirdio/netbird/relay/server/listener" +) + +func TestListener_ShutdownBeforeServe(t *testing.T) { + l := &Listener{Address: "127.0.0.1:0"} + require.NoError(t, l.Bind()) + addr := l.listener.Addr().String() + require.NoError(t, l.Shutdown(context.Background())) + + errChan := make(chan error, 1) + go func() { + errChan <- l.Serve(func(relaylistener.Conn) {}) + }() + + assert.NoError(t, waitForServeToReturn(t, errChan)) + requireTCPAddressFree(t, addr) +} + +func TestListener_ShutdownStopsServe(t *testing.T) { + l := &Listener{Address: "127.0.0.1:0"} + require.NoError(t, l.Bind()) + addr := l.listener.Addr().String() + + errChan := make(chan error, 1) + go func() { + errChan <- l.Serve(func(relaylistener.Conn) {}) + }() + + // A bound socket accepts TCP connections before Serve runs, so only a served + // HTTP response proves the accept loop is running. + require.Eventually(t, func() bool { + return httpGetSucceeds("http://" + addr + "/") + }, 5*time.Second, 10*time.Millisecond, "listener did not start serving") + + require.NoError(t, l.Shutdown(context.Background())) + assert.NoError(t, waitForServeToReturn(t, errChan)) + requireTCPAddressFree(t, addr) +} + +func TestListener_ServeErrorReleasesSocket(t *testing.T) { + // A TLS config without certificates makes ServeTLS fail before it takes + // ownership of the listener, so only Serve itself can close the socket. + l := &Listener{Address: "127.0.0.1:0", TLSConfig: &tls.Config{}} + require.NoError(t, l.Bind()) + addr := l.listener.Addr().String() + + assert.Error(t, l.Serve(func(relaylistener.Conn) {}), "ServeTLS must fail without a certificate") + requireTCPAddressFree(t, addr) + assert.NoError(t, l.Shutdown(context.Background()), "Shutdown after a failed Serve must succeed") +} + +func TestListener_Unbound(t *testing.T) { + l := &Listener{Address: "127.0.0.1:0"} + assert.Error(t, l.Serve(func(relaylistener.Conn) {}), "Serve must refuse an unbound listener") + assert.NoError(t, l.Shutdown(context.Background()), "Shutdown of an unbound listener is a no-op") +} + +func waitForServeToReturn(t *testing.T, errChan <-chan error) error { + t.Helper() + select { + case err := <-errChan: + return err + case <-time.After(5 * time.Second): + t.Fatal("Serve did not return") + return nil + } +} + +func httpGetSucceeds(url string) bool { + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, http.NoBody) + if err != nil { + return false + } + + resp, err := http.DefaultClient.Do(req) + if err != nil { + return false + } + _ = resp.Body.Close() + return true +} + +func requireTCPAddressFree(t *testing.T, addr string) { + t.Helper() + ln, err := net.Listen("tcp", addr) + require.NoError(t, err, "tcp address %s must be released", addr) + require.NoError(t, ln.Close()) +} diff --git a/relay/server/relay.go b/relay/server/relay.go index 84c424b8e..58f9db840 100644 --- a/relay/server/relay.go +++ b/relay/server/relay.go @@ -21,7 +21,8 @@ import ( ) type Listener interface { - Listen(func(conn listener.Conn)) error + Bind() error + Serve(func(conn listener.Conn)) error Shutdown(ctx context.Context) error Protocol() protocol.Protocol } @@ -122,6 +123,9 @@ func (r *Relay) Accept(conn listener.Conn) { r.closeMu.RLock() defer r.closeMu.RUnlock() if r.closed { + if err := conn.Close(); err != nil { + log.Debugf("failed to close connection after shutdown, %s: %s", conn.RemoteAddr(), err) + } return } diff --git a/relay/server/relay_test.go b/relay/server/relay_test.go new file mode 100644 index 000000000..aea0f2a62 --- /dev/null +++ b/relay/server/relay_test.go @@ -0,0 +1,55 @@ +package server + +import ( + "context" + "net" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/shared/relay/auth/allow" +) + +type closeRecorderConn struct { + reads int + closed bool +} + +func (c *closeRecorderConn) Read(context.Context, []byte) (int, error) { + c.reads++ + return 0, net.ErrClosed +} + +func (c *closeRecorderConn) Write(context.Context, []byte) (int, error) { + return 0, net.ErrClosed +} + +func (c *closeRecorderConn) RemoteAddr() net.Addr { + return &net.TCPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 1} +} + +func (c *closeRecorderConn) Close() error { + c.closed = true + return nil +} + +func (c *closeRecorderConn) Protocol() string { + return "test" +} + +func TestRelay_AcceptAfterShutdownClosesConn(t *testing.T) { + relay, err := NewRelay(Config{ + ExposedAddress: "rel://127.0.0.1:1234", + AuthValidator: &allow.Auth{}, + }) + require.NoError(t, err) + relay.Shutdown(context.Background()) + + conn := &closeRecorderConn{} + relay.Accept(conn) + assert.True(t, conn.closed, "a connection accepted after shutdown must be closed") + // The handshake error path closes the connection too, so only an untouched Read + // proves the closed guard rejected it before any handshake. + assert.Zero(t, conn.reads, "a connection accepted after shutdown must not be read") +} diff --git a/relay/server/server.go b/relay/server/server.go index 8d303e9e4..4636eed65 100644 --- a/relay/server/server.go +++ b/relay/server/server.go @@ -3,6 +3,7 @@ package server import ( "context" "crypto/tls" + "fmt" "net/url" "sync" @@ -35,6 +36,7 @@ type Server struct { relay *Relay listeners []Listener listenerMux sync.Mutex + closed bool } // NewServer creates and returns a new relay server instance. @@ -62,39 +64,30 @@ func NewServer(config Config) (*Server, error) { }, nil } -// Listen starts the relay server. +// Listen binds the relay listeners and serves them until Shutdown is called. func (r *Server) Listen(cfg ListenerConfig) error { - wSListener := &ws.Listener{ - Address: cfg.Address, - TLSConfig: cfg.TLSConfig, - TrustedProxies: cfg.TrustedProxies, - } - r.listenerMux.Lock() - r.listeners = append(r.listeners, wSListener) - - tlsConfigQUIC, err := quictls.ServerQUICTLSConfig(cfg.TLSConfig) - if err != nil { - log.Warnf("Not starting QUIC listener: %v", err) - } else { - quicListener := &quic.Listener{ - Address: cfg.Address, - TLSConfig: tlsConfigQUIC, - } - - r.listeners = append(r.listeners, quicListener) + if r.closed { + r.listenerMux.Unlock() + return nil } - errChan := make(chan error, len(r.listeners)) + listeners, err := bindListeners(newListeners(cfg)) + if err != nil { + r.listenerMux.Unlock() + return err + } + r.listeners = append(r.listeners, listeners...) + + errChan := make(chan error, len(listeners)) wg := sync.WaitGroup{} - for _, l := range r.listeners { + for _, l := range listeners { wg.Add(1) go func(listener Listener) { defer wg.Done() - errChan <- listener.Listen(r.relay.Accept) + errChan <- listener.Serve(r.relay.Accept) }(l) } - r.listenerMux.Unlock() wg.Wait() @@ -110,18 +103,14 @@ func (r *Server) Listen(cfg ListenerConfig) error { // Shutdown stops the relay server. If there are active connections, they will be closed gracefully. In case of a context, // the connections will be forcefully closed. func (r *Server) Shutdown(ctx context.Context) error { - r.relay.Shutdown(ctx) - r.listenerMux.Lock() - var multiErr *multierror.Error - for _, l := range r.listeners { - if err := l.Shutdown(ctx); err != nil { - multiErr = multierror.Append(multiErr, err) - } - } - r.listeners = r.listeners[:0] + r.closed = true + listeners := r.listeners + r.listeners = nil r.listenerMux.Unlock() - return nberrors.FormatErrorOrNil(multiErr) + + r.relay.Shutdown(ctx) + return shutdownListeners(ctx, listeners) } func (r *Server) ListenerProtocols() []protocol.Protocol { @@ -145,3 +134,48 @@ func (r *Server) InstanceURL() url.URL { func (r *Server) RelayAccept() func(conn listener.Conn) { return r.relay.Accept } + +func newListeners(cfg ListenerConfig) []Listener { + listeners := []Listener{ + &ws.Listener{ + Address: cfg.Address, + TLSConfig: cfg.TLSConfig, + TrustedProxies: cfg.TrustedProxies, + }, + } + + tlsConfigQUIC, err := quictls.ServerQUICTLSConfig(cfg.TLSConfig) + if err != nil { + log.Warnf("Not starting QUIC listener: %v", err) + return listeners + } + + return append(listeners, &quic.Listener{ + Address: cfg.Address, + TLSConfig: tlsConfigQUIC, + }) +} + +func bindListeners(listeners []Listener) ([]Listener, error) { + bound := make([]Listener, 0, len(listeners)) + for _, l := range listeners { + if err := l.Bind(); err != nil { + if shutdownErr := shutdownListeners(context.Background(), bound); shutdownErr != nil { + log.Warnf("failed to close listeners after bind error: %v", shutdownErr) + } + return nil, fmt.Errorf("%s listener: %w", l.Protocol(), err) + } + bound = append(bound, l) + } + return bound, nil +} + +func shutdownListeners(ctx context.Context, listeners []Listener) error { + var multiErr *multierror.Error + for _, l := range listeners { + if err := l.Shutdown(ctx); err != nil { + multiErr = multierror.Append(multiErr, err) + } + } + return nberrors.FormatErrorOrNil(multiErr) +} diff --git a/relay/server/server_test.go b/relay/server/server_test.go new file mode 100644 index 000000000..6ec802b94 --- /dev/null +++ b/relay/server/server_test.go @@ -0,0 +1,230 @@ +package server + +import ( + "context" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "fmt" + "math/big" + "net" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/relay/server/listener/quic" + "github.com/netbirdio/netbird/shared/relay/auth/allow" +) + +func TestServer_ShutdownBeforeListen(t *testing.T) { + addr := freeAddress(t) + srv := newTestServer(t, addr) + require.NoError(t, srv.Shutdown(context.Background())) + + errChan := make(chan error, 1) + go func() { + errChan <- srv.Listen(ListenerConfig{Address: addr}) + }() + + assert.NoError(t, waitForReturn(t, "Listen", errChan)) + assert.Empty(t, srv.ListenerProtocols(), "a shut down server must not register listeners") + requireAddressFree(t, addr) +} + +func TestServer_ShutdownStopsListen(t *testing.T) { + addr := freeAddress(t) + srv := newTestServer(t, addr) + + errChan := make(chan error, 1) + go func() { + errChan <- srv.Listen(ListenerConfig{Address: addr}) + }() + + waitForListeners(t, srv, errChan) + + require.NoError(t, srv.Shutdown(context.Background())) + assert.NoError(t, waitForReturn(t, "Listen", errChan)) + requireAddressFree(t, addr) +} + +func TestServer_ConcurrentListenAndShutdown(t *testing.T) { + tlsCfg := testTLSConfig(t) + for round := 0; round < 20; round++ { + t.Run(fmt.Sprintf("round-%d", round), func(t *testing.T) { + addr := freeAddress(t) + srv := newTestServer(t, addr) + + start := make(chan struct{}) + listenErr := make(chan error, 1) + shutdownErr := make(chan error, 1) + go func() { + <-start + listenErr <- srv.Listen(ListenerConfig{Address: addr, TLSConfig: tlsCfg}) + }() + go func() { + <-start + shutdownErr <- srv.Shutdown(context.Background()) + }() + close(start) + + // Either side may take the lock first: Listen then returns without binding, + // or Shutdown stops its accept loops. Listen must return in both cases. + assert.NoError(t, waitForReturn(t, "Listen", listenErr)) + assert.NoError(t, waitForReturn(t, "Shutdown", shutdownErr)) + assert.Empty(t, srv.ListenerProtocols(), "no listener may stay registered after shutdown") + requireAddressFree(t, addr) + }) + } +} + +func TestServer_ListenReturnsBindError(t *testing.T) { + blocker, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + t.Cleanup(func() { _ = blocker.Close() }) + addr := blocker.Addr().String() + srv := newTestServer(t, addr) + + errChan := make(chan error, 1) + go func() { + errChan <- srv.Listen(ListenerConfig{Address: addr}) + }() + + assert.Error(t, waitForReturn(t, "Listen", errChan), "Listen must return the bind error instead of serving") + assert.Empty(t, srv.ListenerProtocols(), "a failed Listen must not register listeners") + + require.NoError(t, blocker.Close()) + requireAddressFree(t, addr) +} + +func TestServer_ListenRollsBackOnBindError(t *testing.T) { + addr := freeAddress(t) + blocker, err := net.ListenPacket("udp", addr) + require.NoError(t, err) + t.Cleanup(func() { _ = blocker.Close() }) + srv := newTestServer(t, addr) + tlsCfg := testTLSConfig(t) + + errChan := make(chan error, 1) + go func() { + errChan <- srv.Listen(ListenerConfig{Address: addr, TLSConfig: tlsCfg}) + }() + + err = waitForReturn(t, "Listen", errChan) + require.Error(t, err, "Listen must fail when the QUIC port is taken") + assert.ErrorContains(t, err, string(quic.Proto)+" listener", "the QUIC listener must be the one that failed to bind") + assert.Empty(t, srv.ListenerProtocols(), "a failed Listen must not register listeners") + + // The WS listener binds first, so a free TCP port proves the rollback closed it. + ln, err := net.Listen("tcp", addr) + require.NoError(t, err, "the WS socket must be released after the QUIC bind failure") + require.NoError(t, ln.Close()) +} + +func newTestServer(t *testing.T, addr string) *Server { + t.Helper() + srv, err := NewServer(Config{ + ExposedAddress: "rel://" + addr, + AuthValidator: &allow.Auth{}, + }) + require.NoError(t, err) + t.Cleanup(func() { + assert.NoError(t, srv.Shutdown(context.Background())) + }) + return srv +} + +// The server binds TCP and UDP on the same port, so a port is only picked when +// both are free. The probes are closed before Listen binds, which leaves a small +// window for another process to take the port; waitForListeners then reports the +// bind error instead of a timeout. +func freeAddress(t *testing.T) string { + t.Helper() + for attempt := 0; attempt < 10; attempt++ { + ln, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + addr := ln.Addr().String() + require.NoError(t, ln.Close()) + + conn, err := net.ListenPacket("udp", addr) + if err != nil { + continue + } + require.NoError(t, conn.Close()) + return addr + } + t.Fatal("no address with both tcp and udp ports free") + return "" +} + +// Listen registers its listeners under the same lock that starts serving, so a +// non-empty ListenerProtocols means the sockets are bound. A Listen that returns +// before that has failed to bind, and its error is reported instead of a timeout. +func waitForListeners(t *testing.T, srv *Server, errChan <-chan error) { + t.Helper() + deadline := time.After(5 * time.Second) + for len(srv.ListenerProtocols()) == 0 { + select { + case err := <-errChan: + t.Fatalf("Listen returned before binding: %v", err) + case <-deadline: + t.Fatal("listeners were not bound") + case <-time.After(10 * time.Millisecond): + } + } +} + +func waitForReturn(t *testing.T, op string, errChan <-chan error) error { + t.Helper() + select { + case err := <-errChan: + return err + case <-time.After(5 * time.Second): + t.Fatalf("%s did not return", op) + return nil + } +} + +// The QUIC listener releases its UDP socket from the read loop after Close returns, +// so both sockets are polled instead of checked once. +func requireAddressFree(t *testing.T, addr string) { + t.Helper() + require.Eventually(t, func() bool { + ln, err := net.Listen("tcp", addr) + if err != nil { + return false + } + _ = ln.Close() + + conn, err := net.ListenPacket("udp", addr) + if err != nil { + return false + } + _ = conn.Close() + return true + }, 5*time.Second, 10*time.Millisecond, "address %s must be released", addr) +} + +// A nil TLS config only yields a QUIC listener in the devcert build, so the tests +// that need both listeners pass a real one and bind them regardless of build tags. +func testTLSConfig(t *testing.T) *tls.Config { + t.Helper() + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + require.NoError(t, err) + + template := &x509.Certificate{ + SerialNumber: big.NewInt(1), + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(time.Hour), + IPAddresses: []net.IP{net.ParseIP("127.0.0.1")}, + } + certDER, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key) + require.NoError(t, err) + + return &tls.Config{ + Certificates: []tls.Certificate{{Certificate: [][]byte{certDER}, PrivateKey: key}}, + } +}