diff --git a/core/src/tls.rs b/core/src/tls.rs index b147b50..41e0b7b 100644 --- a/core/src/tls.rs +++ b/core/src/tls.rs @@ -1,6 +1,6 @@ use std::fmt; use std::io::Cursor; -use std::sync::Arc; +use std::sync::{Arc, Mutex}; use rustls::client::ResolvesClientCert; use rustls::pki_types::CertificateDer; @@ -13,7 +13,7 @@ use crate::wire::NativeTunnelCredential; pub(crate) type SignCallback = dyn Fn(&[u8]) -> Result<[u8; 64]> + Send + Sync; pub(crate) struct CallbackSigningKey { - callback: Arc, + callback: Arc>>>, } impl fmt::Debug for CallbackSigningKey { @@ -24,7 +24,9 @@ impl fmt::Debug for CallbackSigningKey { impl CallbackSigningKey { pub(crate) fn new(callback: Arc) -> Self { - Self { callback } + Self { + callback: Arc::new(Mutex::new(Some(callback))), + } } } @@ -40,7 +42,7 @@ impl SigningKey for CallbackSigningKey { } } -struct CallbackSigner(Arc); +struct CallbackSigner(Arc>>>); impl fmt::Debug for CallbackSigner { fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { @@ -50,7 +52,13 @@ impl fmt::Debug for CallbackSigner { impl Signer for CallbackSigner { fn sign(&self, message: &[u8]) -> std::result::Result, rustls::Error> { - (self.0)(message) + let callback = self + .0 + .lock() + .map_err(|_| rustls::Error::General("client signing state failed".to_owned()))? + .take() + .ok_or_else(|| rustls::Error::General("client signer already used".to_owned()))?; + callback(message) .map(|signature| signature.to_vec()) .map_err(|_| rustls::Error::General("client signing failed".to_owned())) } @@ -112,3 +120,46 @@ pub(crate) fn client_config( config.enable_early_data = false; Ok(config) } + +#[cfg(test)] +mod tests { + use std::sync::atomic::{AtomicBool, Ordering}; + + use super::*; + + struct DropMarker(Arc); + + impl Drop for DropMarker { + fn drop(&mut self) { + self.0.store(true, Ordering::SeqCst); + } + } + + #[test] + fn tls_callback_is_consumed_and_released_after_one_signature() { + let dropped = Arc::new(AtomicBool::new(false)); + let marker = DropMarker(Arc::clone(&dropped)); + let callback: Arc = Arc::new(move |_| { + let _ = ▮ + Ok([7; 64]) + }); + let key = CallbackSigningKey::new(callback); + let signer = key + .choose_scheme(&[SignatureScheme::ED25519]) + .expect("ED25519 signer"); + let second_signer = key + .choose_scheme(&[SignatureScheme::ED25519]) + .expect("second ED25519 signer"); + drop(key); + + assert_eq!( + signer.sign(b"handshake").expect("first signature"), + vec![7; 64] + ); + assert!(dropped.load(Ordering::SeqCst), "callback remained retained"); + assert!( + second_signer.sign(b"second request").is_err(), + "signer was reusable" + ); + } +} diff --git a/core/src/transport.rs b/core/src/transport.rs index ab5d33f..ad8e6a7 100644 --- a/core/src/transport.rs +++ b/core/src/transport.rs @@ -1,4 +1,5 @@ use std::fmt; +use std::io; use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}; use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::Arc; @@ -157,15 +158,6 @@ pub async fn connect_with_cancellation( } Ok(signature) }); - let rustls = client_config(credential, checked_tls_signer)?; - let crypto = QuicClientConfig::try_from(rustls).map_err(|_| CoreError::Tls)?; - let mut quinn_config = quinn::ClientConfig::new(Arc::new(crypto)); - let mut transport = TransportConfig::default(); - transport.max_concurrent_bidi_streams(VarInt::from_u32(1)); - transport.max_concurrent_uni_streams(VarInt::from_u32(0)); - transport.datagram_receive_buffer_size(Some(64 * 1024)); - quinn_config.transport_config(Arc::new(transport)); - let remaining = deadline .checked_sub(started.elapsed()) .ok_or(CoreError::Cancelled)?; @@ -173,16 +165,25 @@ pub async fn connect_with_cancellation( let address_count = manifest.addresses().len(); for (address_index, address) in manifest.addresses().iter().enumerate() { cancellation.check()?; - let Ok(resolved) = tokio::net::lookup_host(address).await else { - continue; - }; - let resolved = resolved.take(8).collect::>(); let address_budget = deadline.saturating_sub(started.elapsed()) / u32::try_from(address_count.saturating_sub(address_index).max(1)).unwrap_or(1); let address_deadline = Instant::now() + address_budget; + let resolved = match bounded_lookup( + address_budget, + cancellation, + tokio::net::lookup_host(address), + ) + .await + { + Ok(resolved) => resolved, + Err(CoreError::Transport) => continue, + Err(error) => return Err(error), + }; let remote_count = resolved.len(); for (remote_index, remote) in resolved.into_iter().enumerate() { cancellation.check()?; + let quinn_config = + client_transport_config(credential, Arc::clone(&checked_tls_signer))?; let remaining_remotes = remote_count.saturating_sub(remote_index).max(1); let attempt_budget = address_deadline.saturating_duration_since(Instant::now()) / u32::try_from(remaining_remotes).unwrap_or(1); @@ -222,6 +223,21 @@ pub async fn connect_with_cancellation( } } +fn client_transport_config( + credential: &NativeTunnelCredential, + callback: Arc, +) -> Result { + let rustls = client_config(credential, callback)?; + let crypto = QuicClientConfig::try_from(rustls).map_err(|_| CoreError::Tls)?; + let mut config = quinn::ClientConfig::new(Arc::new(crypto)); + let mut transport = TransportConfig::default(); + transport.max_concurrent_bidi_streams(VarInt::from_u32(1)); + transport.max_concurrent_uni_streams(VarInt::from_u32(0)); + transport.datagram_receive_buffer_size(Some(64 * 1024)); + config.transport_config(Arc::new(transport)); + Ok(config) +} + struct AttemptContext<'a> { manifest: &'a ConnectionManifest, offered: &'a CapabilityProfile, @@ -253,6 +269,20 @@ async fn cancellable_timeout( } } +async fn bounded_lookup( + duration: Duration, + cancellation: &Cancellation, + future: impl std::future::Future>, +) -> Result> +where + T: Iterator, +{ + let resolved = cancellable_timeout(duration, cancellation, future) + .await? + .map_err(|_| CoreError::Transport)?; + Ok(resolved.take(8).collect()) +} + async fn dial( remote: SocketAddr, config: quinn::ClientConfig, @@ -270,7 +300,7 @@ async fn dial( let negotiated = connecting .handshake_data() .await - .map_err(|_| AttemptError::Terminal(CoreError::Tls))? + .map_err(|_| AttemptError::Terminal(tls_error(context)))? .downcast::() .ok() .and_then(|data| data.protocol.clone()); @@ -279,7 +309,7 @@ async fn dial( } let connection = connecting .await - .map_err(|_| AttemptError::Terminal(CoreError::Tls))?; + .map_err(|_| AttemptError::Terminal(tls_error(context)))?; let payload = admission_payload( context.manifest, context.offered, @@ -302,7 +332,7 @@ async fn dial( .map_err(|_| AttemptError::Terminal(connection_error(&connection)))?; let Ok(authority) = ClientSessionAuthority::decode(&response) else { let stable = decode_stable_error(&response)?; - return if stable.retryable { + return if stable.retryable && stable.code == "gateway_draining" { Err(AttemptError::Retry) } else { Err(AttemptError::Terminal(stable.error)) @@ -318,6 +348,14 @@ async fn dial( }) } +fn tls_error(context: &AttemptContext<'_>) -> CoreError { + if context.cancellation.check().is_err() || context.started.elapsed() >= context.deadline { + CoreError::Cancelled + } else { + CoreError::Tls + } +} + fn admission_payload( manifest: &ConnectionManifest, offered: &CapabilityProfile, @@ -410,7 +448,11 @@ fn hello_length(length: usize) -> Result { #[cfg(test)] mod tests { - use super::{hello_length, HELLO_LIMIT}; + use std::io; + use std::net::SocketAddr; + use std::time::{Duration, Instant}; + + use super::{bounded_lookup, hello_length, Cancellation, HELLO_LIMIT}; use crate::error::CoreError; #[test] @@ -419,4 +461,32 @@ mod tests { assert_eq!(hello_length(0), Err(CoreError::Protocol)); assert_eq!(hello_length(HELLO_LIMIT + 1), Err(CoreError::Protocol)); } + + #[test] + fn dns_lookup_is_bounded_by_address_share() { + let started = Instant::now(); + let result = tokio::runtime::Runtime::new() + .expect("runtime") + .block_on(bounded_lookup( + Duration::from_millis(20), + &Cancellation::new(), + std::future::pending::>>(), + )); + assert_eq!(result.err(), Some(CoreError::Transport)); + assert!(started.elapsed() < Duration::from_secs(1)); + } + + #[test] + fn cancellation_interrupts_dns_lookup() { + let cancellation = Cancellation::new(); + cancellation.cancel(); + let result = tokio::runtime::Runtime::new() + .expect("runtime") + .block_on(bounded_lookup( + Duration::from_secs(1), + &cancellation, + std::future::pending::>>(), + )); + assert_eq!(result.err(), Some(CoreError::Cancelled)); + } } diff --git a/core/src/wire.rs b/core/src/wire.rs index bbb8059..ec1d5d7 100644 --- a/core/src/wire.rs +++ b/core/src/wire.rs @@ -388,6 +388,7 @@ impl ConnectionManifest { .iter() .any(|address| !bounded(address, 1, 256)) || !valid_dns_name(&self.gateway.public_identity) + || self.gateway.public_identity == self.gateway.id || !(1..=4).contains(&self.tunnel.versions.len()) || self .tunnel @@ -621,6 +622,7 @@ struct StableError { pub(crate) struct DecodedStableError { pub(crate) error: CoreError, + pub(crate) code: String, pub(crate) retryable: bool, } @@ -703,6 +705,7 @@ pub(crate) fn decode_stable_error(bytes: &[u8]) -> Result { }; Ok(DecodedStableError { error, + code: stable.code, retryable: stable.retryable, }) } diff --git a/core/tests/gateway_oracle.rs b/core/tests/gateway_oracle.rs index 4f80666..a1de8fd 100644 --- a/core/tests/gateway_oracle.rs +++ b/core/tests/gateway_oracle.rs @@ -46,6 +46,7 @@ type admission struct { lastTranscript []byte reusable bool failWorkOnce bool + releaseFails bool authority protocol.SessionAuthority work protocol.ProviderSessionWork public ed25519.PublicKey @@ -67,11 +68,15 @@ func (a *admission) ProviderWork(context.Context, protocol.SessionAuthority) (pr if a.failWorkOnce { a.failWorkOnce = false; return protocol.ProviderSessionWork{}, context.DeadlineExceeded } return a.work, nil } -func (a *admission) Release(context.Context, protocol.SessionAuthority) error { return nil } +func (a *admission) Release(context.Context, protocol.SessionAuthority) error { + if a.releaseFails { return errors.New("release failed") } + return nil +} -type provider struct{} +type provider struct { failStart bool } type session struct { state protocol.ProviderState; video chan gateway.ProviderMedia; audio chan gateway.ProviderMedia; events chan gateway.ProviderEvent } -func (provider) Start(_ context.Context, request gateway.LaunchRequest) (gateway.ProviderSession, error) { +func (p provider) Start(_ context.Context, request gateway.LaunchRequest) (gateway.ProviderSession, error) { + if p.failStart { return nil, context.DeadlineExceeded } return &session{state: protocol.ProviderState{Version:"1", SessionID:request.SessionID, State:gateway.ProviderStateReady, Channels:[]string{"video","audio","input","feedback"}}, video:make(chan gateway.ProviderMedia), audio:make(chan gateway.ProviderMedia), events:make(chan gateway.ProviderEvent)}, nil } func (s *session) Ready(context.Context) error { return nil } @@ -125,9 +130,9 @@ func main() { _, firstReplayErr := replayGuard.Admit(context.Background(),replayRequest) _, secondReplayErr := replayGuard.Admit(context.Background(),replayRequest) replayRejected := firstReplayErr == nil && errors.Is(secondReplayErr,gateway.ErrAdmissionRejected) - admissionService := &admission{authority:authority,work:work,public:admissionKey.Public().(ed25519.PublicKey),reusable:mode == "reusable",failWorkOnce:mode == "post-retryable"} + admissionService := &admission{authority:authority,work:work,public:admissionKey.Public().(ed25519.PublicKey),reusable:mode == "reusable",failWorkOnce:mode == "post-retryable",releaseFails:mode == "cleanup-release-failure"} newService := func() *gateway.Server { - service, err := gateway.NewServer(gateway.ServerConfig{ListenAddress:"127.0.0.1:0",TLSConfig:server,GatewayID:authority.GatewayID,Admission:admissionService,Provider:provider{}}); if err != nil { panic(err) } + service, err := gateway.NewServer(gateway.ServerConfig{ListenAddress:"127.0.0.1:0",TLSConfig:server,GatewayID:authority.GatewayID,Admission:admissionService,Provider:provider{failStart:mode == "provider-start-lost-response" || mode == "cleanup-release-failure"}}); if err != nil { panic(err) } return service } service := newService() @@ -139,7 +144,7 @@ func main() { services = []*gateway.Server{draining, service} addresses = []string{draining.Addr().String(), service.Addr().String()} } - if mode == "post-retryable" { + if mode == "post-retryable" || mode == "provider-start-lost-response" || mode == "cleanup-release-failure" { second := newService() services = []*gateway.Server{service, second} addresses = []string{service.Addr().String(), second.Addr().String()} @@ -303,7 +308,7 @@ fn callback_ed25519_signer_completes_tls13_quic_admission_without_private_key_in assert_eq!(session.authority().session_id(), "session-1"); assert_eq!(admission_calls.load(Ordering::SeqCst), 1); - assert!(tls_calls.load(Ordering::SeqCst) >= 1); + assert_eq!(tls_calls.load(Ordering::SeqCst), 1); assert!(admission_inputs .lock() .expect("admission transcript") @@ -598,6 +603,41 @@ fn overall_deadline_includes_synchronous_admission_signing() { assert!(started.elapsed() < Duration::from_secs(2)); } +#[test] +fn tls_signer_finite_deadline_overrun_is_cancelled_not_tls() { + let oracle = Oracle::start(""); + let manifest = ConnectionManifest::decode(oracle.ready.manifest.as_bytes()).expect("manifest"); + let credential = + NativeTunnelCredential::decode(oracle.ready.credential.as_bytes()).expect("credential"); + let admission_key = test_key(&oracle.ready.admission_key); + let tls_key = test_key(&oracle.ready.client_key); + let cancellation = Cancellation::new(); + let trigger = cancellation.clone(); + let canceller = thread::spawn(move || { + thread::sleep(Duration::from_millis(25)); + trigger.cancel(); + }); + let result = + tokio::runtime::Runtime::new() + .expect("runtime") + .block_on(connect_with_cancellation( + &manifest, + &credential, + Signers::new( + AdmissionSigner::new(signer(admission_key)), + TlsEd25519Signer::new(move |message| { + thread::sleep(Duration::from_millis(100)); + signer(Arc::clone(&tls_key))(message) + }), + ), + "2026-08-12T00:00:00Z", + Duration::from_secs(5), + &cancellation, + )); + canceller.join().expect("canceller"); + assert_eq!(result.err(), Some(CoreError::Cancelled)); +} + #[test] fn explicit_cancellation_interrupts_network_wait() { let oracle = Oracle::start(""); @@ -733,29 +773,6 @@ fn invalid_first_manifest_address_does_not_block_live_second_address() { fn retryable_stable_error_advances_to_live_manifest_address() { let oracle = Oracle::start("retryable"); let manifest = ConnectionManifest::decode(oracle.ready.manifest.as_bytes()).expect("manifest"); - let credential = - NativeTunnelCredential::decode(oracle.ready.credential.as_bytes()).expect("credential"); - let runtime = tokio::runtime::Runtime::new().expect("runtime"); - let session = runtime - .block_on(connect( - &manifest, - &credential, - recording_signers( - test_key(&oracle.ready.admission_key), - test_key(&oracle.ready.client_key), - Arc::new(Mutex::new(Vec::new())), - ), - "2026-08-12T00:00:00Z", - Duration::from_secs(5), - )) - .expect("retryable draining response advanced to live address"); - runtime.block_on(session.close()); -} - -#[test] -fn post_admission_retry_uses_a_fresh_signed_nonce() { - let oracle = Oracle::start("post-retryable"); - let manifest = ConnectionManifest::decode(oracle.ready.manifest.as_bytes()).expect("manifest"); let credential = NativeTunnelCredential::decode(oracle.ready.credential.as_bytes()).expect("credential"); let transcripts = Arc::new(Mutex::new(Vec::new())); @@ -772,13 +789,48 @@ fn post_admission_retry_uses_a_fresh_signed_nonce() { "2026-08-12T00:00:00Z", Duration::from_secs(5), )) - .expect("post-admission retry used a fresh request"); + .expect("retryable draining response advanced to live address"); runtime.block_on(session.close()); let transcripts = transcripts.lock().expect("transcripts"); assert_eq!(transcripts.len(), 2); assert_ne!(transcripts[0], transcripts[1]); } +#[test] +fn post_admission_failure_is_terminal_and_does_not_mutate_twice() { + for mode in [ + "post-retryable", + "provider-start-lost-response", + "cleanup-release-failure", + ] { + let oracle = Oracle::start(mode); + let manifest = + ConnectionManifest::decode(oracle.ready.manifest.as_bytes()).expect("manifest"); + let credential = + NativeTunnelCredential::decode(oracle.ready.credential.as_bytes()).expect("credential"); + let transcripts = Arc::new(Mutex::new(Vec::new())); + let runtime = tokio::runtime::Runtime::new().expect("runtime"); + let result = runtime.block_on(connect( + &manifest, + &credential, + recording_signers( + test_key(&oracle.ready.admission_key), + test_key(&oracle.ready.client_key), + Arc::clone(&transcripts), + ), + "2026-08-12T00:00:00Z", + Duration::from_secs(5), + )); + assert!(result.is_err(), "{mode} was retried to success"); + let transcripts = transcripts.lock().expect("transcripts"); + assert_eq!( + transcripts.len(), + 1, + "{mode} attempted admission mutation twice" + ); + } +} + #[test] fn repeated_production_gateway_admissions_remain_bounded() { let runtime = tokio::runtime::Runtime::new().expect("runtime"); diff --git a/core/tests/protocol_fixtures.rs b/core/tests/protocol_fixtures.rs index 223cdbc..00fc09c 100644 --- a/core/tests/protocol_fixtures.rs +++ b/core/tests/protocol_fixtures.rs @@ -145,6 +145,14 @@ fn manifest_public_identity_requires_dns_sni_not_ip_or_uuid() { "non-DNS SNI accepted: {invalid_identity}" ); } + + let identity_equals_dns_shaped_gateway_id = String::from_utf8(valid_manifest().to_vec()) + .expect("fixture is UTF-8") + .replace("\"id\":\"gateway\"", "\"id\":\"gateway.test\""); + assert!( + ConnectionManifest::decode(identity_equals_dns_shaped_gateway_id.as_bytes()).is_err(), + "public SNI identity matched the logical gateway id" + ); } #[test] diff --git a/gateway/gateway_test.go b/gateway/gateway_test.go index 771ab00..0de8d8c 100644 --- a/gateway/gateway_test.go +++ b/gateway/gateway_test.go @@ -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 { diff --git a/gateway/transport.go b/gateway/transport.go index 7ca8c19..590f3c8 100644 --- a/gateway/transport.go +++ b/gateway/transport.go @@ -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)