From 4aa935c66ff28e21a13c6b3bae436c3e1e2aadeb Mon Sep 17 00:00:00 2001 From: Paul Wells Date: Wed, 16 Sep 2026 13:48:27 -0700 Subject: [PATCH] agent endpoints: give the http3 listener a server type ListenWebTransport returned a stop closure, so every caller had to park it next to the server it belonged to. In cloud that meant a second field on the node base, written at start and read at stop, for one listener. WebTransportServer embeds webtransport.Server and owns the sockets and accept loops, so Listen and Shutdown are methods and the shape matches http.Server. Shutdown takes a context: it stops accepting, sends GOAWAY, and closes whatever has not drained by the deadline, where the closure closed everything at once with no deadline of its own. It is safe on a server that never listened, which is what lets a holder key teardown off the field alone. The embedded Close still releases sessions without releasing the sockets; that is what the type comment warns about. --- pkg/service/agentendpoint_test.go | 9 ++- pkg/service/server.go | 12 ++-- pkg/service/webtransport.go | 95 +++++++++++++++++--------- pkg/service/webtransportlisten_test.go | 5 +- 4 files changed, 78 insertions(+), 43 deletions(-) diff --git a/pkg/service/agentendpoint_test.go b/pkg/service/agentendpoint_test.go index 95b48ce3c..6fd47eecc 100644 --- a/pkg/service/agentendpoint_test.go +++ b/pkg/service/agentendpoint_test.go @@ -125,9 +125,14 @@ func newEndpointStack(t *testing.T, endpointsCfg agent.EndpointsConfig) *endpoin wtMux.Handle("/agent", service.NewAgentWTService(h)) wt := service.NewWebTransportServer(selfSignedTLS(t)) wt.H3.Handler = service.NewWebTransportHandler(keyProvider, wt, wtMux) - bound, stopWT, err := service.ListenWebTransport(wt, []string{"127.0.0.1"}, 0) + bound, err := wt.Listen([]string{"127.0.0.1"}, 0) require.NoError(t, err) - t.Cleanup(stopWT) + t.Cleanup(func() { + // t.Context() is already cancelled by the time cleanups run + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + _ = wt.Shutdown(ctx) + }) wtURL := "https://" + bound[0].String() + "/agent" return &endpointStack{t: t, ts: ts, handler: h, scopes: scopes, wtURL: wtURL} diff --git a/pkg/service/server.go b/pkg/service/server.go index a648ca348..bc6a35326 100644 --- a/pkg/service/server.go +++ b/pkg/service/server.go @@ -30,7 +30,6 @@ import ( "github.com/pion/turn/v5" "github.com/prometheus/client_golang/prometheus/promhttp" - "github.com/quic-go/webtransport-go" "github.com/rs/cors" "github.com/twitchtv/twirp" "github.com/urfave/negroni/v3" @@ -55,7 +54,7 @@ type LivekitServer struct { httpServer *http.Server promServer *http.Server debugServer *http.Server - webtransportServer *webtransport.Server + webtransportServer *WebTransportServer router routing.Router roomManager *RoomManager signalServer *SignalServer @@ -267,13 +266,10 @@ func (s *LivekitServer) Start() error { } } - stopWebTransport := func() {} if s.webtransportServer != nil { - _, stop, err := ListenWebTransport(s.webtransportServer, s.config.BindAddresses, s.config.WebTransport.Port) - if err != nil { + if _, err := s.webtransportServer.Listen(s.config.BindAddresses, s.config.WebTransport.Port); err != nil { return err } - stopWebTransport = stop } values := []any{ @@ -354,7 +350,9 @@ func (s *LivekitServer) Start() error { if s.debugServer != nil { _ = s.debugServer.Shutdown(ctx) } - stopWebTransport() + if s.webtransportServer != nil { + _ = s.webtransportServer.Shutdown(ctx) + } if s.turnServer != nil { _ = s.turnServer.Close() diff --git a/pkg/service/webtransport.go b/pkg/service/webtransport.go index 2d0ad276d..42db2c454 100644 --- a/pkg/service/webtransport.go +++ b/pkg/service/webtransport.go @@ -75,19 +75,31 @@ func WebTransportTLS(certFile, keyFile string, dev bool) (*tls.Config, error) { }, nil } +// WebTransportServer's embedded Close releases sessions but not the sockets; +// Shutdown is the teardown path. +type WebTransportServer struct { + *webtransport.Server + + mu sync.Mutex + conns []*net.UDPConn + lns []*quic.EarlyListener + cancel context.CancelFunc + serving sync.WaitGroup +} + // NewWebTransportServer wraps an HTTP/3 WebTransport server around tlsConf (the // h3 ALPN is set here if absent). The caller must assign wt.H3.Handler. -func NewWebTransportServer(tlsConf *tls.Config) *webtransport.Server { +func NewWebTransportServer(tlsConf *tls.Config) *WebTransportServer { tlsConf = tlsConf.Clone() if len(tlsConf.NextProtos) == 0 { tlsConf.NextProtos = []string{http3.NextProtoH3} } - return &webtransport.Server{H3: &http3.Server{TLSConfig: tlsConf}} + return &WebTransportServer{Server: &webtransport.Server{H3: &http3.Server{TLSConfig: tlsConf}}} } // NewWebTransportHandler builds the listener's handler: mux behind api-key auth, // with wt in each request's context. -func NewWebTransportHandler(keyProvider auth.KeyProvider, wt *webtransport.Server, mux http.Handler) http.Handler { +func NewWebTransportHandler(keyProvider auth.KeyProvider, wt *WebTransportServer, mux http.Handler) http.Handler { middlewares := []negroni.Handler{negroni.NewRecovery()} if keyProvider != nil { middlewares = append(middlewares, NewAPIKeyAuthMiddleware(keyProvider)) @@ -99,7 +111,7 @@ type webTransportServerKey struct{} // WithWebTransportServer puts wt in each request's context, where a route that // upgrades reads it. Apply it outermost, ahead of the listener's middleware chain. -func WithWebTransportServer(wt *webtransport.Server, next http.Handler) http.Handler { +func WithWebTransportServer(wt *WebTransportServer, next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), webTransportServerKey{}, wt))) }) @@ -107,8 +119,8 @@ func WithWebTransportServer(wt *webtransport.Server, next http.Handler) http.Han // GetWebTransportServer returns the server serving this request, or nil when the // request did not arrive over a WebTransport listener. -func GetWebTransportServer(ctx context.Context) *webtransport.Server { - wt, _ := ctx.Value(webTransportServerKey{}).(*webtransport.Server) +func GetWebTransportServer(ctx context.Context) *WebTransportServer { + wt, _ := ctx.Value(webTransportServerKey{}).(*WebTransportServer) return wt } @@ -134,20 +146,17 @@ func unwrapResponseWriter(w http.ResponseWriter) http.ResponseWriter { } } -// ListenWebTransport binds a UDP socket per address and serves wt on each, -// returning the bound addresses and a stop func. Empty addrs binds all interfaces. -func ListenWebTransport(wt *webtransport.Server, addrs []string, port uint32) ([]net.Addr, func(), error) { +// Listen binds one UDP socket per address. Empty addrs binds all interfaces. +func (s *WebTransportServer) Listen(addrs []string, port uint32) ([]net.Addr, error) { if len(addrs) == 0 { addrs = []string{""} } - // webtransport.Server.Serve takes a reference on a WaitGroup that - // Server.Close waits on, so running the two concurrently is a race by the - // WaitGroup's own rules. Accepting here instead leaves that counter at zero: - // Server.ServeQUICConn never touches it. + // Server.Serve takes a reference on the WaitGroup Server.Close waits on; + // Server.ServeQUICConn does not. quicConf := &quic.Config{} - if wt.H3.QUICConfig != nil { - quicConf = wt.H3.QUICConfig.Clone() + if s.H3.QUICConfig != nil { + quicConf = s.H3.QUICConfig.Clone() } quicConf.EnableDatagrams = true quicConf.EnableStreamResetPartialDelivery = true @@ -169,50 +178,72 @@ func ListenWebTransport(wt *webtransport.Server, addrs []string, port uint32) ([ udpAddr, err := net.ResolveUDPAddr("udp", net.JoinHostPort(addr, strconv.Itoa(int(port)))) if err != nil { closeAll() - return nil, nil, err + return nil, err } udp, err := net.ListenUDP("udp", udpAddr) if err != nil { closeAll() - return nil, nil, err + return nil, err } conns = append(conns, udp) bound = append(bound, udp.LocalAddr()) - ln, err := quic.ListenEarly(udp, wt.H3.TLSConfig, quicConf) + ln, err := quic.ListenEarly(udp, s.H3.TLSConfig, quicConf) if err != nil { closeAll() - return nil, nil, err + return nil, err } lns = append(lns, ln) } ctx, cancel := context.WithCancel(context.Background()) - var serving sync.WaitGroup + s.mu.Lock() + s.conns, s.lns, s.cancel = conns, lns, cancel + s.mu.Unlock() + for _, ln := range lns { - serving.Go(func() { acceptWebTransport(ctx, wt, ln, &serving) }) + s.serving.Go(func() { s.accept(ctx, ln) }) } logger.Infow("webtransport listener started", "addresses", bound) - return bound, func() { - // the sockets stay open until everything has drained, so every - // CONNECTION_CLOSE frame still reaches its peer - _ = wt.Close() - cancel() - serving.Wait() - closeAll() - }, nil + return bound, nil } -func acceptWebTransport(ctx context.Context, wt *webtransport.Server, ln *quic.EarlyListener, serving *sync.WaitGroup) { +// Shutdown is safe on a server that never listened, and safe to call more than +// once. Peers get a GOAWAY, then whatever has not drained by the ctx deadline is +// closed under it. +func (s *WebTransportServer) Shutdown(ctx context.Context) error { + s.mu.Lock() + cancel, conns, lns := s.cancel, s.conns, s.lns + s.cancel, s.conns, s.lns = nil, nil, nil + s.mu.Unlock() + + if cancel != nil { + cancel() + } + err := s.H3.Shutdown(ctx) + _ = s.Server.Close() + s.serving.Wait() + + // the sockets close last, so every CONNECTION_CLOSE frame still reaches its peer + for _, ln := range lns { + _ = ln.Close() + } + for _, c := range conns { + _ = c.Close() + } + return err +} + +func (s *WebTransportServer) accept(ctx context.Context, ln *quic.EarlyListener) { for { conn, err := ln.Accept(ctx) if err != nil { logger.Infow("webtransport listener stopped", "error", err) return } - serving.Go(func() { - if err := wt.ServeQUICConn(conn); err != nil && !errors.Is(err, http.ErrServerClosed) { + s.serving.Go(func() { + if err := s.ServeQUICConn(conn); err != nil && !errors.Is(err, http.ErrServerClosed) { logger.Infow("webtransport connection stopped", "error", err) } }) diff --git a/pkg/service/webtransportlisten_test.go b/pkg/service/webtransportlisten_test.go index fb0654707..1ed41531e 100644 --- a/pkg/service/webtransportlisten_test.go +++ b/pkg/service/webtransportlisten_test.go @@ -15,6 +15,7 @@ package service_test import ( + "context" "net/http" "testing" @@ -30,8 +31,8 @@ func TestWebTransportStopBeforeFirstConnection(t *testing.T) { for range 20 { wt := service.NewWebTransportServer(selfSignedTLS(t)) wt.H3.Handler = service.NewWebTransportHandler(nil, wt, http.NewServeMux()) - _, stop, err := service.ListenWebTransport(wt, []string{"127.0.0.1"}, 0) + _, err := wt.Listen([]string{"127.0.0.1"}, 0) require.NoError(t, err) - stop() + require.NoError(t, wt.Shutdown(context.Background())) } }