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
+56 -5
View File
@@ -1,6 +1,6 @@
use std::fmt; use std::fmt;
use std::io::Cursor; use std::io::Cursor;
use std::sync::Arc; use std::sync::{Arc, Mutex};
use rustls::client::ResolvesClientCert; use rustls::client::ResolvesClientCert;
use rustls::pki_types::CertificateDer; 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) type SignCallback = dyn Fn(&[u8]) -> Result<[u8; 64]> + Send + Sync;
pub(crate) struct CallbackSigningKey { pub(crate) struct CallbackSigningKey {
callback: Arc<SignCallback>, callback: Arc<Mutex<Option<Arc<SignCallback>>>>,
} }
impl fmt::Debug for CallbackSigningKey { impl fmt::Debug for CallbackSigningKey {
@@ -24,7 +24,9 @@ impl fmt::Debug for CallbackSigningKey {
impl CallbackSigningKey { impl CallbackSigningKey {
pub(crate) fn new(callback: Arc<SignCallback>) -> Self { pub(crate) fn new(callback: Arc<SignCallback>) -> Self {
Self { callback } Self {
callback: Arc::new(Mutex::new(Some(callback))),
}
} }
} }
@@ -40,7 +42,7 @@ impl SigningKey for CallbackSigningKey {
} }
} }
struct CallbackSigner(Arc<SignCallback>); struct CallbackSigner(Arc<Mutex<Option<Arc<SignCallback>>>>);
impl fmt::Debug for CallbackSigner { impl fmt::Debug for CallbackSigner {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
@@ -50,7 +52,13 @@ impl fmt::Debug for CallbackSigner {
impl Signer for CallbackSigner { impl Signer for CallbackSigner {
fn sign(&self, message: &[u8]) -> std::result::Result<Vec<u8>, rustls::Error> { fn sign(&self, message: &[u8]) -> std::result::Result<Vec<u8>, 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(|signature| signature.to_vec())
.map_err(|_| rustls::Error::General("client signing failed".to_owned())) .map_err(|_| rustls::Error::General("client signing failed".to_owned()))
} }
@@ -112,3 +120,46 @@ pub(crate) fn client_config(
config.enable_early_data = false; config.enable_early_data = false;
Ok(config) Ok(config)
} }
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicBool, Ordering};
use super::*;
struct DropMarker(Arc<AtomicBool>);
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<SignCallback> = Arc::new(move |_| {
let _ = &marker;
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"
);
}
}
+87 -17
View File
@@ -1,4 +1,5 @@
use std::fmt; use std::fmt;
use std::io;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}; use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc; use std::sync::Arc;
@@ -157,15 +158,6 @@ pub async fn connect_with_cancellation(
} }
Ok(signature) 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 let remaining = deadline
.checked_sub(started.elapsed()) .checked_sub(started.elapsed())
.ok_or(CoreError::Cancelled)?; .ok_or(CoreError::Cancelled)?;
@@ -173,16 +165,25 @@ pub async fn connect_with_cancellation(
let address_count = manifest.addresses().len(); let address_count = manifest.addresses().len();
for (address_index, address) in manifest.addresses().iter().enumerate() { for (address_index, address) in manifest.addresses().iter().enumerate() {
cancellation.check()?; cancellation.check()?;
let Ok(resolved) = tokio::net::lookup_host(address).await else {
continue;
};
let resolved = resolved.take(8).collect::<Vec<_>>();
let address_budget = deadline.saturating_sub(started.elapsed()) let address_budget = deadline.saturating_sub(started.elapsed())
/ u32::try_from(address_count.saturating_sub(address_index).max(1)).unwrap_or(1); / u32::try_from(address_count.saturating_sub(address_index).max(1)).unwrap_or(1);
let address_deadline = Instant::now() + address_budget; 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(); let remote_count = resolved.len();
for (remote_index, remote) in resolved.into_iter().enumerate() { for (remote_index, remote) in resolved.into_iter().enumerate() {
cancellation.check()?; 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 remaining_remotes = remote_count.saturating_sub(remote_index).max(1);
let attempt_budget = address_deadline.saturating_duration_since(Instant::now()) let attempt_budget = address_deadline.saturating_duration_since(Instant::now())
/ u32::try_from(remaining_remotes).unwrap_or(1); / 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<SignCallback>,
) -> Result<quinn::ClientConfig> {
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> { struct AttemptContext<'a> {
manifest: &'a ConnectionManifest, manifest: &'a ConnectionManifest,
offered: &'a CapabilityProfile, offered: &'a CapabilityProfile,
@@ -253,6 +269,20 @@ async fn cancellable_timeout<T>(
} }
} }
async fn bounded_lookup<T>(
duration: Duration,
cancellation: &Cancellation,
future: impl std::future::Future<Output = io::Result<T>>,
) -> Result<Vec<SocketAddr>>
where
T: Iterator<Item = SocketAddr>,
{
let resolved = cancellable_timeout(duration, cancellation, future)
.await?
.map_err(|_| CoreError::Transport)?;
Ok(resolved.take(8).collect())
}
async fn dial( async fn dial(
remote: SocketAddr, remote: SocketAddr,
config: quinn::ClientConfig, config: quinn::ClientConfig,
@@ -270,7 +300,7 @@ async fn dial(
let negotiated = connecting let negotiated = connecting
.handshake_data() .handshake_data()
.await .await
.map_err(|_| AttemptError::Terminal(CoreError::Tls))? .map_err(|_| AttemptError::Terminal(tls_error(context)))?
.downcast::<quinn::crypto::rustls::HandshakeData>() .downcast::<quinn::crypto::rustls::HandshakeData>()
.ok() .ok()
.and_then(|data| data.protocol.clone()); .and_then(|data| data.protocol.clone());
@@ -279,7 +309,7 @@ async fn dial(
} }
let connection = connecting let connection = connecting
.await .await
.map_err(|_| AttemptError::Terminal(CoreError::Tls))?; .map_err(|_| AttemptError::Terminal(tls_error(context)))?;
let payload = admission_payload( let payload = admission_payload(
context.manifest, context.manifest,
context.offered, context.offered,
@@ -302,7 +332,7 @@ async fn dial(
.map_err(|_| AttemptError::Terminal(connection_error(&connection)))?; .map_err(|_| AttemptError::Terminal(connection_error(&connection)))?;
let Ok(authority) = ClientSessionAuthority::decode(&response) else { let Ok(authority) = ClientSessionAuthority::decode(&response) else {
let stable = decode_stable_error(&response)?; let stable = decode_stable_error(&response)?;
return if stable.retryable { return if stable.retryable && stable.code == "gateway_draining" {
Err(AttemptError::Retry) Err(AttemptError::Retry)
} else { } else {
Err(AttemptError::Terminal(stable.error)) 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( fn admission_payload(
manifest: &ConnectionManifest, manifest: &ConnectionManifest,
offered: &CapabilityProfile, offered: &CapabilityProfile,
@@ -410,7 +448,11 @@ fn hello_length(length: usize) -> Result<u32> {
#[cfg(test)] #[cfg(test)]
mod tests { 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; use crate::error::CoreError;
#[test] #[test]
@@ -419,4 +461,32 @@ mod tests {
assert_eq!(hello_length(0), Err(CoreError::Protocol)); assert_eq!(hello_length(0), Err(CoreError::Protocol));
assert_eq!(hello_length(HELLO_LIMIT + 1), 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::<io::Result<std::vec::IntoIter<SocketAddr>>>(),
));
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::<io::Result<std::vec::IntoIter<SocketAddr>>>(),
));
assert_eq!(result.err(), Some(CoreError::Cancelled));
}
} }
+3
View File
@@ -388,6 +388,7 @@ impl ConnectionManifest {
.iter() .iter()
.any(|address| !bounded(address, 1, 256)) .any(|address| !bounded(address, 1, 256))
|| !valid_dns_name(&self.gateway.public_identity) || !valid_dns_name(&self.gateway.public_identity)
|| self.gateway.public_identity == self.gateway.id
|| !(1..=4).contains(&self.tunnel.versions.len()) || !(1..=4).contains(&self.tunnel.versions.len())
|| self || self
.tunnel .tunnel
@@ -621,6 +622,7 @@ struct StableError {
pub(crate) struct DecodedStableError { pub(crate) struct DecodedStableError {
pub(crate) error: CoreError, pub(crate) error: CoreError,
pub(crate) code: String,
pub(crate) retryable: bool, pub(crate) retryable: bool,
} }
@@ -703,6 +705,7 @@ pub(crate) fn decode_stable_error(bytes: &[u8]) -> Result<DecodedStableError> {
}; };
Ok(DecodedStableError { Ok(DecodedStableError {
error, error,
code: stable.code,
retryable: stable.retryable, retryable: stable.retryable,
}) })
} }
+83 -31
View File
@@ -46,6 +46,7 @@ type admission struct {
lastTranscript []byte lastTranscript []byte
reusable bool reusable bool
failWorkOnce bool failWorkOnce bool
releaseFails bool
authority protocol.SessionAuthority authority protocol.SessionAuthority
work protocol.ProviderSessionWork work protocol.ProviderSessionWork
public ed25519.PublicKey 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 } if a.failWorkOnce { a.failWorkOnce = false; return protocol.ProviderSessionWork{}, context.DeadlineExceeded }
return a.work, nil 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 } 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 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 } func (s *session) Ready(context.Context) error { return nil }
@@ -125,9 +130,9 @@ func main() {
_, firstReplayErr := replayGuard.Admit(context.Background(),replayRequest) _, firstReplayErr := replayGuard.Admit(context.Background(),replayRequest)
_, secondReplayErr := replayGuard.Admit(context.Background(),replayRequest) _, secondReplayErr := replayGuard.Admit(context.Background(),replayRequest)
replayRejected := firstReplayErr == nil && errors.Is(secondReplayErr,gateway.ErrAdmissionRejected) 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 { 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 return service
} }
service := newService() service := newService()
@@ -139,7 +144,7 @@ func main() {
services = []*gateway.Server{draining, service} services = []*gateway.Server{draining, service}
addresses = []string{draining.Addr().String(), service.Addr().String()} 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() second := newService()
services = []*gateway.Server{service, second} services = []*gateway.Server{service, second}
addresses = []string{service.Addr().String(), second.Addr().String()} 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!(session.authority().session_id(), "session-1");
assert_eq!(admission_calls.load(Ordering::SeqCst), 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 assert!(admission_inputs
.lock() .lock()
.expect("admission transcript") .expect("admission transcript")
@@ -598,6 +603,41 @@ fn overall_deadline_includes_synchronous_admission_signing() {
assert!(started.elapsed() < Duration::from_secs(2)); 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] #[test]
fn explicit_cancellation_interrupts_network_wait() { fn explicit_cancellation_interrupts_network_wait() {
let oracle = Oracle::start(""); 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() { fn retryable_stable_error_advances_to_live_manifest_address() {
let oracle = Oracle::start("retryable"); let oracle = Oracle::start("retryable");
let manifest = ConnectionManifest::decode(oracle.ready.manifest.as_bytes()).expect("manifest"); 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 = let credential =
NativeTunnelCredential::decode(oracle.ready.credential.as_bytes()).expect("credential"); NativeTunnelCredential::decode(oracle.ready.credential.as_bytes()).expect("credential");
let transcripts = Arc::new(Mutex::new(Vec::new())); 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", "2026-08-12T00:00:00Z",
Duration::from_secs(5), 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()); runtime.block_on(session.close());
let transcripts = transcripts.lock().expect("transcripts"); let transcripts = transcripts.lock().expect("transcripts");
assert_eq!(transcripts.len(), 2); assert_eq!(transcripts.len(), 2);
assert_ne!(transcripts[0], transcripts[1]); 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] #[test]
fn repeated_production_gateway_admissions_remain_bounded() { fn repeated_production_gateway_admissions_remain_bounded() {
let runtime = tokio::runtime::Runtime::new().expect("runtime"); let runtime = tokio::runtime::Runtime::new().expect("runtime");
+8
View File
@@ -145,6 +145,14 @@ fn manifest_public_identity_requires_dns_sni_not_ip_or_uuid() {
"non-DNS SNI accepted: {invalid_identity}" "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] #[test]
+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) { func TestStableErrorNeverLeaksInternalProviderDetails(t *testing.T) {
var wire bytes.Buffer var wire bytes.Buffer
if err := writeStableError(&wire, "provider_unavailable", errors.New("https://provider.invalid/launch?rikey=secret-sentinel"), true); err != nil { 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) { func TestGatewayTelemetrySeparatesQueueProcessingAndPacing(t *testing.T) {
serverTLS, clientTLS := testTLS(t) serverTLS, clientTLS := testTLS(t)
session := &fakeSession{ session := &fakeSession{
@@ -1668,6 +1878,7 @@ type oneTimeAdmission struct {
streamPolicy protocol.ProviderStreamPolicy streamPolicy protocol.ProviderStreamPolicy
providerWork *protocol.ProviderSessionWork providerWork *protocol.ProviderSessionWork
disableClipboard bool disableClipboard bool
releaseErr error
} }
type recordingProviderStateReporter struct { type recordingProviderStateReporter struct {
@@ -1742,7 +1953,7 @@ func (a *oneTimeAdmission) Release(_ context.Context, authority protocol.Session
a.releaseAuthority = authority a.releaseAuthority = authority
close(a.released) close(a.released)
} }
return nil return a.releaseErr
} }
func mustRead(t *testing.T, path string) []byte { func mustRead(t *testing.T, path string) []byte {
+49 -14
View File
@@ -21,6 +21,7 @@ import (
const ( const (
defaultHelloLimit = 16 * 1024 defaultHelloLimit = 16 * 1024
defaultHelloTimeout = 10 * time.Second
defaultControlLimit = 128 * 1024 defaultControlLimit = 128 * 1024
clientControlBacklog = 64 clientControlBacklog = 64
terminalAckTimeout = 2 * time.Second terminalAckTimeout = 2 * time.Second
@@ -94,6 +95,8 @@ type Server struct {
pacer *fairPacer pacer *fairPacer
mu sync.Mutex mu sync.Mutex
sessions map[*gatewaySession]struct{} sessions map[*gatewaySession]struct{}
connections map[*quic.Conn]struct{}
helloTimeout time.Duration
draining atomic.Bool draining atomic.Bool
closed atomic.Bool closed atomic.Bool
closeOnce sync.Once closeOnce sync.Once
@@ -138,7 +141,7 @@ func NewServer(config ServerConfig) (*Server, error) {
if err != nil { if err != nil {
return nil, err 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 { func validateServerTLS(config *tls.Config) error {
@@ -174,9 +177,22 @@ func (s *Server) Serve(ctx context.Context) error {
} }
return err 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.workers.Add(1)
s.mu.Unlock()
go func() { go func() {
defer s.workers.Done() defer s.workers.Done()
defer func() {
s.mu.Lock()
delete(s.connections, connection)
s.mu.Unlock()
}()
s.handleConnection(ctx, connection) s.handleConnection(ctx, connection)
}() }()
} }
@@ -186,13 +202,24 @@ func (s *Server) Close() error {
var err error var err error
s.closeOnce.Do(func() { s.closeOnce.Do(func() {
s.BeginDrain() s.BeginDrain()
s.closed.Store(true)
err = s.listener.Close()
s.mu.Lock() s.mu.Lock()
s.closed.Store(true)
sessions := make([]*gatewaySession, 0, len(s.sessions))
for session := range 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() 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() s.workers.Wait()
return err return err
@@ -200,15 +227,23 @@ func (s *Server) Close() error {
func (s *Server) handleConnection(parent context.Context, connection *quic.Conn) { func (s *Server) handleConnection(parent context.Context, connection *quic.Conn) {
defer connection.CloseWithError(applicationError, "connection closed") 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() defer cancel()
stream, err := connection.AcceptStream(ctx) stream, err := connection.AcceptStream(ctx)
if err != nil { if err != nil {
return return
} }
if err := stream.SetDeadline(helloDeadline); err != nil {
return
}
writeError := func(code string, err error, retryable bool) { 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 { 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() defer responseCancel()
select { select {
case <-connection.Context().Done(): 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) authority, err := s.config.Admission.Admit(ctx, request)
if err != nil { if err != nil {
s.metrics.AdmissionRejects.Add(1) s.metrics.AdmissionRejects.Add(1)
writeError(stableAdmissionCode(err), err, errors.Is(err, context.DeadlineExceeded)) writeError(stableAdmissionCode(err), err, false)
return return
} }
if s.Draining() { if s.Draining() {
_ = s.config.Admission.Release(context.Background(), authority) _ = s.config.Admission.Release(context.Background(), authority)
writeError("gateway_draining", ErrGatewayDraining, true) writeError("gateway_draining", ErrGatewayDraining, false)
return return
} }
if err := s.validateAuthority(authority, request); err != nil { 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) work, err := s.config.Admission.ProviderWork(ctx, authority)
if err != nil || s.validateProviderWork(work, authority) != nil { if err != nil || s.validateProviderWork(work, authority) != nil {
_ = s.config.Admission.Release(context.Background(), authority) _ = s.config.Admission.Release(context.Background(), authority)
writeError("provider_work_unavailable", ErrAdmissionRejected, err != nil) writeError("provider_work_unavailable", ErrAdmissionRejected, false)
return return
} }
selected, err := IntersectCapabilities(s.config.Capabilities, s.config.ProviderCapabilities, request.Capabilities, authority.Capabilities) 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 { if (work.ClipboardPolicy.ClientToProviderEnabled || work.ClipboardPolicy.ProviderToClientEnabled) && s.config.ClipboardAuditReporter == nil {
_ = s.config.Admission.Release(context.Background(), authority) _ = s.config.Admission.Release(context.Background(), authority)
writeError("clipboard_audit_unavailable", ErrAdmissionRejected, true) writeError("clipboard_audit_unavailable", ErrAdmissionRejected, false)
return 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 { 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) _ = s.config.Admission.Release(context.Background(), authority)
writeError("provider_state_unavailable", err, true) writeError("provider_state_unavailable", err, false)
return return
} }
providerSession, err := s.config.Provider.Start(ctx, LaunchRequest{SessionID: request.SessionID, Capabilities: selected, ProviderProfile: authority.ProviderProfile, ProviderIdentity: work.ProviderIdentity, ProviderWork: work}) 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.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.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) _ = s.config.Admission.Release(context.Background(), authority)
writeError(stableProviderCode(err), err, errors.Is(err, context.DeadlineExceeded)) writeError(stableProviderCode(err), err, false)
return return
} }
if err := s.reportProviderState(ctx, providerSession.State()); err != nil { if err := s.reportProviderState(ctx, providerSession.State()); err != nil {
_ = providerSession.ReleaseAll(context.Background()) _ = providerSession.ReleaseAll(context.Background())
_ = providerSession.Terminate(context.Background()) _ = providerSession.Terminate(context.Background())
_ = s.config.Admission.Release(context.Background(), authority) _ = s.config.Admission.Release(context.Background(), authority)
writeError("provider_state_unavailable", err, true) writeError("provider_state_unavailable", err, false)
return return
} }
clientAuthority := protocol.ClientSessionAuthority{ 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, ReconnectSequence: authority.ReconnectSequence, ExpiresAt: authority.ExpiresAt, Capabilities: selected,
} }
authorityBytes, err := protocol.EncodeClientSessionAuthority(clientAuthority) 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.ReleaseAll(context.Background())
_ = providerSession.Terminate(context.Background()) _ = providerSession.Terminate(context.Background())
_ = s.config.Admission.Release(context.Background(), authority) _ = s.config.Admission.Release(context.Background(), authority)