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())) } }