fix(transport): harden QUIC admission boundaries

This commit is contained in:
sechmachine
2026-08-12 20:04:51 +07:00
parent 111092becb
commit 09f40eb9c7
7 changed files with 508 additions and 78 deletions
+212 -1
View File
@@ -536,6 +536,164 @@ func TestStableErrorFramePrecedesConnectionTeardown(t *testing.T) {
}
}
func TestStableErrorDrainOutlivesExpiredHandlerContext(t *testing.T) {
serverTLS, clientTLS := testTLS(t)
clientTLS.NextProtos = []string{"versevdi-gateway-v1"}
fake := NewFakeApollo(FakeApolloConfig{Now: time.Now()})
authority := protocol.SessionAuthority{Version: "1", SessionID: "session-1", GatewayID: "gateway-1", Audience: "versevdi-gateway", ExpiresAt: time.Now().Add(time.Minute).UTC().Format(time.RFC3339Nano), Capabilities: DefaultCapabilities(), ProviderProfile: ProviderProfileApollo, ProviderIdentity: fake.config.Identity.Key()}
admission := AdmissionFunc(func(ctx context.Context, _ protocol.TunnelAdmissionRequest) (protocol.SessionAuthority, error) {
<-ctx.Done()
return protocol.SessionAuthority{}, ctx.Err()
})
server, err := NewServer(ServerConfig{ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: authority.GatewayID, Admission: admission, Provider: fake})
if err != nil {
t.Fatal(err)
}
server.helloTimeout = 20 * time.Millisecond
ctx, cancel := context.WithCancel(context.Background())
go func() { _ = server.Serve(ctx) }()
connection, err := quic.DialAddr(context.Background(), server.Addr().String(), clientTLS, &quic.Config{EnableDatagrams: true})
if err != nil {
t.Fatal(err)
}
stream, err := connection.OpenStreamSync(context.Background())
if err != nil {
t.Fatal(err)
}
request := protocol.TunnelAdmissionRequest{Version: "1", SessionID: authority.SessionID, GatewayID: authority.GatewayID, Audience: authority.Audience, Grant: strings.Repeat("g", 64), ClientNonce: "nonce-0000000001", DeviceSignature: strings.Repeat("s", 86), Capabilities: DefaultCapabilities()}
payload, err := protocol.EncodeTunnelAdmissionRequest(request)
if err != nil {
t.Fatal(err)
}
if err := writeWire(stream, payload, defaultHelloLimit); err != nil {
t.Fatal(err)
}
response, err := readWire(stream, defaultHelloLimit)
if err != nil {
t.Fatalf("stable response lost with expired handler context: %v", err)
}
if stable, err := protocol.DecodeStableError(response); err != nil || stable.Code != "admission_rejected" {
t.Fatalf("stable response = %#v, %v", stable, err)
}
_ = connection.CloseWithError(applicationError, "done")
cancel()
_ = server.Close()
}
func TestServerCloseInterruptsPartialHelloClient(t *testing.T) {
serverTLS, clientTLS := testTLS(t)
clientTLS.NextProtos = []string{"versevdi-gateway-v1"}
fake := NewFakeApollo(FakeApolloConfig{Now: time.Now()})
authority := protocol.SessionAuthority{Version: "1", SessionID: "session-1", GatewayID: "gateway-1", Audience: "versevdi-gateway", ExpiresAt: time.Now().Add(time.Minute).UTC().Format(time.RFC3339Nano), Capabilities: DefaultCapabilities(), ProviderProfile: ProviderProfileApollo, ProviderIdentity: fake.config.Identity.Key()}
server, err := NewServer(ServerConfig{ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: authority.GatewayID, Admission: &oneTimeAdmission{authority: authority, released: make(chan struct{})}, Provider: fake})
if err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithCancel(context.Background())
serveDone := make(chan error, 1)
go func() { serveDone <- server.Serve(ctx) }()
connection, err := quic.DialAddr(context.Background(), server.Addr().String(), clientTLS, &quic.Config{EnableDatagrams: true})
if err != nil {
t.Fatal(err)
}
stream, err := connection.OpenStreamSync(context.Background())
if err != nil {
t.Fatal(err)
}
if _, err := stream.Write([]byte{0, 0}); err != nil {
t.Fatal(err)
}
closed := make(chan error, 1)
go func() { closed <- server.Close() }()
select {
case err := <-closed:
if err != nil {
t.Fatal(err)
}
case <-time.After(time.Second):
t.Fatal("Server.Close blocked on partial hello")
}
select {
case <-connection.Context().Done():
case <-time.After(time.Second):
t.Fatal("partial hello connection remained open")
}
cancel()
if err := <-serveDone; err != nil {
t.Fatal(err)
}
}
func TestHelloUsesOneAbsoluteDeadlineAgainstTrickle(t *testing.T) {
serverTLS, clientTLS := testTLS(t)
clientTLS.NextProtos = []string{"versevdi-gateway-v1"}
fake := NewFakeApollo(FakeApolloConfig{Now: time.Now()})
authority := protocol.SessionAuthority{Version: "1", SessionID: "session-1", GatewayID: "gateway-1", Audience: "versevdi-gateway", ExpiresAt: time.Now().Add(time.Minute).UTC().Format(time.RFC3339Nano), Capabilities: DefaultCapabilities(), ProviderProfile: ProviderProfileApollo, ProviderIdentity: fake.config.Identity.Key()}
server, err := NewServer(ServerConfig{ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: authority.GatewayID, Admission: &oneTimeAdmission{authority: authority, released: make(chan struct{})}, Provider: fake})
if err != nil {
t.Fatal(err)
}
server.helloTimeout = 40 * time.Millisecond
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
go func() { _ = server.Serve(ctx) }()
connection, err := quic.DialAddr(context.Background(), server.Addr().String(), clientTLS, &quic.Config{EnableDatagrams: true})
if err != nil {
t.Fatal(err)
}
stream, err := connection.OpenStreamSync(context.Background())
if err != nil {
t.Fatal(err)
}
if _, err := stream.Write([]byte{0}); err != nil {
t.Fatal(err)
}
started := time.Now()
time.Sleep(25 * time.Millisecond)
_, _ = stream.Write([]byte{0})
response, err := readWire(stream, defaultHelloLimit)
if err != nil {
t.Fatalf("absolute hello deadline did not produce a stable error: %v", err)
}
stable, err := protocol.DecodeStableError(response)
if err != nil || stable.Code != "invalid_hello" {
t.Fatalf("stable response = %#v, %v", stable, err)
}
if time.Since(started) > 250*time.Millisecond {
t.Fatal("trickling hello reset or escaped its absolute deadline")
}
_ = connection.CloseWithError(applicationError, "done")
_ = server.Close()
}
func TestHelloDeadlineIsClearedAfterAdmission(t *testing.T) {
serverTLS, clientTLS := testTLS(t)
fake := NewFakeApollo(FakeApolloConfig{Now: time.Now()})
authority := protocol.SessionAuthority{Version: "1", SessionID: "session-1", GatewayID: "gateway-1", Audience: "versevdi-gateway", ExpiresAt: time.Now().Add(time.Minute).UTC().Format(time.RFC3339Nano), Capabilities: DefaultCapabilities(), ProviderProfile: ProviderProfileApollo, ProviderIdentity: fake.config.Identity.Key()}
admission := &oneTimeAdmission{authority: authority, released: make(chan struct{}), disableClipboard: true}
server, err := NewServer(ServerConfig{ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: authority.GatewayID, Admission: admission, Provider: fake})
if err != nil {
t.Fatal(err)
}
server.helloTimeout = 40 * time.Millisecond
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
go func() { _ = server.Serve(ctx) }()
request := protocol.TunnelAdmissionRequest{Version: "1", SessionID: authority.SessionID, GatewayID: authority.GatewayID, Audience: authority.Audience, Grant: strings.Repeat("g", 64), ClientNonce: "nonce-0000000001", DeviceSignature: strings.Repeat("s", 86), Capabilities: DefaultCapabilities()}
client, err := Dial(context.Background(), server.Addr().String(), clientTLS, request)
if err != nil {
t.Fatal(err)
}
time.Sleep(80 * time.Millisecond)
select {
case <-client.connection.Context().Done():
t.Fatal("hello deadline leaked into the admitted session")
default:
}
_ = client.Close()
_ = server.Close()
}
func TestStableErrorNeverLeaksInternalProviderDetails(t *testing.T) {
var wire bytes.Buffer
if err := writeStableError(&wire, "provider_unavailable", errors.New("https://provider.invalid/launch?rikey=secret-sentinel"), true); err != nil {
@@ -554,6 +712,58 @@ func TestStableErrorNeverLeaksInternalProviderDetails(t *testing.T) {
}
}
func TestPostAdmissionProviderFailureIsTerminalEvenWhenReleaseFails(t *testing.T) {
serverTLS, clientTLS := testTLS(t)
clientTLS.NextProtos = []string{"versevdi-gateway-v1"}
fake := NewFakeApollo(FakeApolloConfig{Now: time.Now()})
authority := protocol.SessionAuthority{Version: "1", SessionID: "session-1", GatewayID: "gateway-1", Audience: "versevdi-gateway", ExpiresAt: time.Now().Add(time.Minute).UTC().Format(time.RFC3339Nano), Capabilities: DefaultCapabilities(), ProviderProfile: ProviderProfileApollo, ProviderIdentity: fake.config.Identity.Key()}
admission := &oneTimeAdmission{authority: authority, released: make(chan struct{}), disableClipboard: true, releaseErr: errors.New("release failed")}
provider := providerStartFunc(func(context.Context, LaunchRequest) (ProviderSession, error) {
return nil, context.DeadlineExceeded
})
server, err := NewServer(ServerConfig{ListenAddress: "127.0.0.1:0", TLSConfig: serverTLS, GatewayID: authority.GatewayID, Admission: admission, Provider: provider})
if err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
go func() { _ = server.Serve(ctx) }()
connection, err := quic.DialAddr(context.Background(), server.Addr().String(), clientTLS, &quic.Config{EnableDatagrams: true})
if err != nil {
t.Fatal(err)
}
stream, err := connection.OpenStreamSync(context.Background())
if err != nil {
t.Fatal(err)
}
request := protocol.TunnelAdmissionRequest{Version: "1", SessionID: authority.SessionID, GatewayID: authority.GatewayID, Audience: authority.Audience, Grant: strings.Repeat("g", 64), ClientNonce: "nonce-0000000001", DeviceSignature: strings.Repeat("s", 86), Capabilities: DefaultCapabilities()}
payload, err := protocol.EncodeTunnelAdmissionRequest(request)
if err != nil {
t.Fatal(err)
}
if err := writeWire(stream, payload, defaultHelloLimit); err != nil {
t.Fatal(err)
}
response, err := readWire(stream, defaultHelloLimit)
if err != nil {
t.Fatal(err)
}
stable, err := protocol.DecodeStableError(response)
if err != nil || stable.Code != "provider_timeout" || stable.Retryable {
t.Fatalf("post-admission stable response = %#v, %v", stable, err)
}
select {
case <-admission.released:
case <-time.After(time.Second):
t.Fatal("post-admission failure did not attempt release")
}
if got := admission.releases.Load(); got != 1 {
t.Fatalf("release attempts = %d, want 1", got)
}
_ = connection.CloseWithError(applicationError, "done")
_ = server.Close()
}
func TestGatewayTelemetrySeparatesQueueProcessingAndPacing(t *testing.T) {
serverTLS, clientTLS := testTLS(t)
session := &fakeSession{
@@ -1668,6 +1878,7 @@ type oneTimeAdmission struct {
streamPolicy protocol.ProviderStreamPolicy
providerWork *protocol.ProviderSessionWork
disableClipboard bool
releaseErr error
}
type recordingProviderStateReporter struct {
@@ -1742,7 +1953,7 @@ func (a *oneTimeAdmission) Release(_ context.Context, authority protocol.Session
a.releaseAuthority = authority
close(a.released)
}
return nil
return a.releaseErr
}
func mustRead(t *testing.T, path string) []byte {
+59 -24
View File
@@ -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)