fix(transport): harden QUIC admission boundaries
This commit is contained in:
+59
-24
@@ -21,6 +21,7 @@ import (
|
||||
|
||||
const (
|
||||
defaultHelloLimit = 16 * 1024
|
||||
defaultHelloTimeout = 10 * time.Second
|
||||
defaultControlLimit = 128 * 1024
|
||||
clientControlBacklog = 64
|
||||
terminalAckTimeout = 2 * time.Second
|
||||
@@ -88,16 +89,18 @@ type mediaTimingObservation struct {
|
||||
}
|
||||
|
||||
type Server struct {
|
||||
listener *quic.Listener
|
||||
config ServerConfig
|
||||
metrics *Metrics
|
||||
pacer *fairPacer
|
||||
mu sync.Mutex
|
||||
sessions map[*gatewaySession]struct{}
|
||||
draining atomic.Bool
|
||||
closed atomic.Bool
|
||||
closeOnce sync.Once
|
||||
workers sync.WaitGroup
|
||||
listener *quic.Listener
|
||||
config ServerConfig
|
||||
metrics *Metrics
|
||||
pacer *fairPacer
|
||||
mu sync.Mutex
|
||||
sessions map[*gatewaySession]struct{}
|
||||
connections map[*quic.Conn]struct{}
|
||||
helloTimeout time.Duration
|
||||
draining atomic.Bool
|
||||
closed atomic.Bool
|
||||
closeOnce sync.Once
|
||||
workers sync.WaitGroup
|
||||
}
|
||||
|
||||
func NewServer(config ServerConfig) (*Server, error) {
|
||||
@@ -138,7 +141,7 @@ func NewServer(config ServerConfig) (*Server, error) {
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &Server{listener: listener, config: config, metrics: &Metrics{}, pacer: newFairPacer(config.PacerKbps), sessions: make(map[*gatewaySession]struct{})}, nil
|
||||
return &Server{listener: listener, config: config, metrics: &Metrics{}, pacer: newFairPacer(config.PacerKbps), sessions: make(map[*gatewaySession]struct{}), connections: make(map[*quic.Conn]struct{}), helloTimeout: defaultHelloTimeout}, nil
|
||||
}
|
||||
|
||||
func validateServerTLS(config *tls.Config) error {
|
||||
@@ -174,9 +177,22 @@ func (s *Server) Serve(ctx context.Context) error {
|
||||
}
|
||||
return err
|
||||
}
|
||||
s.mu.Lock()
|
||||
if s.closed.Load() {
|
||||
s.mu.Unlock()
|
||||
_ = connection.CloseWithError(applicationError, "server closed")
|
||||
continue
|
||||
}
|
||||
s.connections[connection] = struct{}{}
|
||||
s.workers.Add(1)
|
||||
s.mu.Unlock()
|
||||
go func() {
|
||||
defer s.workers.Done()
|
||||
defer func() {
|
||||
s.mu.Lock()
|
||||
delete(s.connections, connection)
|
||||
s.mu.Unlock()
|
||||
}()
|
||||
s.handleConnection(ctx, connection)
|
||||
}()
|
||||
}
|
||||
@@ -186,13 +202,24 @@ func (s *Server) Close() error {
|
||||
var err error
|
||||
s.closeOnce.Do(func() {
|
||||
s.BeginDrain()
|
||||
s.closed.Store(true)
|
||||
err = s.listener.Close()
|
||||
s.mu.Lock()
|
||||
s.closed.Store(true)
|
||||
sessions := make([]*gatewaySession, 0, len(s.sessions))
|
||||
for session := range s.sessions {
|
||||
session.cancel()
|
||||
sessions = append(sessions, session)
|
||||
}
|
||||
connections := make([]*quic.Conn, 0, len(s.connections))
|
||||
for connection := range s.connections {
|
||||
connections = append(connections, connection)
|
||||
}
|
||||
s.mu.Unlock()
|
||||
err = s.listener.Close()
|
||||
for _, session := range sessions {
|
||||
session.cancel()
|
||||
}
|
||||
for _, connection := range connections {
|
||||
_ = connection.CloseWithError(applicationError, "server closed")
|
||||
}
|
||||
})
|
||||
s.workers.Wait()
|
||||
return err
|
||||
@@ -200,15 +227,23 @@ func (s *Server) Close() error {
|
||||
|
||||
func (s *Server) handleConnection(parent context.Context, connection *quic.Conn) {
|
||||
defer connection.CloseWithError(applicationError, "connection closed")
|
||||
ctx, cancel := context.WithTimeout(parent, 10*time.Second)
|
||||
helloDeadline := time.Now().Add(s.helloTimeout)
|
||||
ctx, cancel := context.WithDeadline(parent, helloDeadline)
|
||||
defer cancel()
|
||||
stream, err := connection.AcceptStream(ctx)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
if err := stream.SetDeadline(helloDeadline); err != nil {
|
||||
return
|
||||
}
|
||||
writeError := func(code string, err error, retryable bool) {
|
||||
responseDeadline := time.Now().Add(time.Second)
|
||||
if stream.SetDeadline(responseDeadline) != nil {
|
||||
return
|
||||
}
|
||||
if writeStableError(stream, code, err, retryable) == nil && stream.Close() == nil {
|
||||
responseCtx, responseCancel := context.WithTimeout(ctx, time.Second)
|
||||
responseCtx, responseCancel := context.WithDeadline(context.Background(), responseDeadline)
|
||||
defer responseCancel()
|
||||
select {
|
||||
case <-connection.Context().Done():
|
||||
@@ -237,12 +272,12 @@ func (s *Server) handleConnection(parent context.Context, connection *quic.Conn)
|
||||
authority, err := s.config.Admission.Admit(ctx, request)
|
||||
if err != nil {
|
||||
s.metrics.AdmissionRejects.Add(1)
|
||||
writeError(stableAdmissionCode(err), err, errors.Is(err, context.DeadlineExceeded))
|
||||
writeError(stableAdmissionCode(err), err, false)
|
||||
return
|
||||
}
|
||||
if s.Draining() {
|
||||
_ = s.config.Admission.Release(context.Background(), authority)
|
||||
writeError("gateway_draining", ErrGatewayDraining, true)
|
||||
writeError("gateway_draining", ErrGatewayDraining, false)
|
||||
return
|
||||
}
|
||||
if err := s.validateAuthority(authority, request); err != nil {
|
||||
@@ -253,7 +288,7 @@ func (s *Server) handleConnection(parent context.Context, connection *quic.Conn)
|
||||
work, err := s.config.Admission.ProviderWork(ctx, authority)
|
||||
if err != nil || s.validateProviderWork(work, authority) != nil {
|
||||
_ = s.config.Admission.Release(context.Background(), authority)
|
||||
writeError("provider_work_unavailable", ErrAdmissionRejected, err != nil)
|
||||
writeError("provider_work_unavailable", ErrAdmissionRejected, false)
|
||||
return
|
||||
}
|
||||
selected, err := IntersectCapabilities(s.config.Capabilities, s.config.ProviderCapabilities, request.Capabilities, authority.Capabilities)
|
||||
@@ -278,12 +313,12 @@ func (s *Server) handleConnection(parent context.Context, connection *quic.Conn)
|
||||
}
|
||||
if (work.ClipboardPolicy.ClientToProviderEnabled || work.ClipboardPolicy.ProviderToClientEnabled) && s.config.ClipboardAuditReporter == nil {
|
||||
_ = s.config.Admission.Release(context.Background(), authority)
|
||||
writeError("clipboard_audit_unavailable", ErrAdmissionRejected, true)
|
||||
writeError("clipboard_audit_unavailable", ErrAdmissionRejected, false)
|
||||
return
|
||||
}
|
||||
if err := s.reportProviderState(ctx, protocol.ProviderState{Version: "1", SessionID: request.SessionID, State: ProviderStateStarting, CleanupPending: false, Channels: []string{"video", "audio", "input", "feedback"}}); err != nil {
|
||||
_ = s.config.Admission.Release(context.Background(), authority)
|
||||
writeError("provider_state_unavailable", err, true)
|
||||
writeError("provider_state_unavailable", err, false)
|
||||
return
|
||||
}
|
||||
providerSession, err := s.config.Provider.Start(ctx, LaunchRequest{SessionID: request.SessionID, Capabilities: selected, ProviderProfile: authority.ProviderProfile, ProviderIdentity: work.ProviderIdentity, ProviderWork: work})
|
||||
@@ -291,14 +326,14 @@ func (s *Server) handleConnection(parent context.Context, connection *quic.Conn)
|
||||
s.metrics.ProviderErrors.Add(1)
|
||||
_ = s.reportProviderState(context.Background(), protocol.ProviderState{Version: "1", SessionID: request.SessionID, State: ProviderStateFailed, CleanupPending: false, Channels: []string{"video", "audio", "input", "feedback"}})
|
||||
_ = s.config.Admission.Release(context.Background(), authority)
|
||||
writeError(stableProviderCode(err), err, errors.Is(err, context.DeadlineExceeded))
|
||||
writeError(stableProviderCode(err), err, false)
|
||||
return
|
||||
}
|
||||
if err := s.reportProviderState(ctx, providerSession.State()); err != nil {
|
||||
_ = providerSession.ReleaseAll(context.Background())
|
||||
_ = providerSession.Terminate(context.Background())
|
||||
_ = s.config.Admission.Release(context.Background(), authority)
|
||||
writeError("provider_state_unavailable", err, true)
|
||||
writeError("provider_state_unavailable", err, false)
|
||||
return
|
||||
}
|
||||
clientAuthority := protocol.ClientSessionAuthority{
|
||||
@@ -306,7 +341,7 @@ func (s *Server) handleConnection(parent context.Context, connection *quic.Conn)
|
||||
ReconnectSequence: authority.ReconnectSequence, ExpiresAt: authority.ExpiresAt, Capabilities: selected,
|
||||
}
|
||||
authorityBytes, err := protocol.EncodeClientSessionAuthority(clientAuthority)
|
||||
if err != nil || writeWire(stream, authorityBytes, defaultHelloLimit) != nil {
|
||||
if err != nil || writeWire(stream, authorityBytes, defaultHelloLimit) != nil || stream.SetDeadline(time.Time{}) != nil {
|
||||
_ = providerSession.ReleaseAll(context.Background())
|
||||
_ = providerSession.Terminate(context.Background())
|
||||
_ = s.config.Admission.Release(context.Background(), authority)
|
||||
|
||||
Reference in New Issue
Block a user