diff --git a/Cargo.lock b/Cargo.lock index 8a5a2a3dd6..cde21630b3 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -286,9 +286,9 @@ checksum = "c08606f8c3cbf4ce6ec8e28fb0014a2c086708fe954eaa885384a6165172e7e8" [[package]] name = "aws-lc-fips-sys" -version = "0.13.16" +version = "0.14.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "37b00953a69b2cfb471d13d72538d2e66930832340b0f31deadd404b48c573c5" +checksum = "118303cd75f63d1933a90c2ceb7e697281ac6acbdbcc490b46419f25a527ab90" dependencies = [ "bindgen", "cc", @@ -301,9 +301,9 @@ dependencies = [ [[package]] name = "aws-lc-rs" -version = "1.17.3" +version = "1.18.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "00bdb5da18dac48ca2cc7cd4a98e533e8635a58e2361d13a1a4ee3888e0d72f1" +checksum = "ce2b2dcc879c3bae0d371e77c99f2238400ef24ec001394befa67b6e543add9e" dependencies = [ "aws-lc-fips-sys", "aws-lc-sys", @@ -312,9 +312,9 @@ dependencies = [ [[package]] name = "aws-lc-sys" -version = "0.43.0" +version = "0.44.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "43103168cc76fe62678a375e722fc9cb3a0146159ac5828bc4f0dfd755c2224c" +checksum = "f09fae7be8bb3174e05c6afdb34199e6dc0c7c04ba9fa237b1967adfbde27483" dependencies = [ "cc", "cmake", @@ -383,6 +383,7 @@ dependencies = [ "quinn", "rcgen", "rustls", + "rustls-aws-lc-rs", "tokio", "tracing", "tracing-subscriber", @@ -451,6 +452,7 @@ dependencies = [ "quinn", "rcgen", "rustls", + "rustls-ring", ] [[package]] @@ -1816,6 +1818,8 @@ dependencies = [ "quinn-proto", "rcgen", "rustls", + "rustls-aws-lc-rs", + "rustls-util", "serde", "serde_json", "socket2", @@ -1985,6 +1989,9 @@ dependencies = [ "rcgen", "rustc-hash", "rustls", + "rustls-aws-lc-rs", + "rustls-ring", + "rustls-util", "smol", "socket2", "thiserror 2.0.19", @@ -2014,8 +2021,11 @@ dependencies = [ "ring", "rustc-hash", "rustls", + "rustls-aws-lc-rs", "rustls-pki-types", "rustls-platform-verifier", + "rustls-ring", + "rustls-util", "slab", "thiserror 2.0.19", "tinyvec", @@ -2202,17 +2212,26 @@ dependencies = [ [[package]] name = "rustls" -version = "0.23.42" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3c54fcab019b409d04215d3a17cb438fd7fbf192ee61461f20f4fe18704bc138" +version = "0.24.0-dev.1" +source = "git+https://github.com/rustls/rustls.git?branch=main#3925f65934364edafe8d6b20707d9e5e6183648e" dependencies = [ - "aws-lc-rs", - "log", "once_cell", - "ring", "rustls-pki-types", "rustls-webpki", "subtle", + "tracing", + "zeroize", +] + +[[package]] +name = "rustls-aws-lc-rs" +version = "0.1.0-dev.1" +source = "git+https://github.com/rustls/rustls.git?branch=main#3925f65934364edafe8d6b20707d9e5e6183648e" +dependencies = [ + "aws-lc-rs", + "rustls", + "rustls-pki-types", + "subtle", "zeroize", ] @@ -2240,9 +2259,8 @@ dependencies = [ [[package]] name = "rustls-platform-verifier" -version = "0.7.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "26d1e2536ce4f35f4846aa13bff16bd0ff40157cdb14cc056c7b14ba41233ba0" +version = "0.8.0" +source = "git+https://github.com/iadev09/rustls-platform-verifier.git?rev=df094724adf95136d5d09cf6d54296a3e809fff8#df094724adf95136d5d09cf6d54296a3e809fff8" dependencies = [ "core-foundation", "core-foundation-sys", @@ -2262,17 +2280,34 @@ dependencies = [ [[package]] name = "rustls-platform-verifier-android" version = "0.1.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f87165f0995f63a9fbeea62b64d10b4d9d8e78ec6d7d51fb2125fda7bb36788f" +source = "git+https://github.com/iadev09/rustls-platform-verifier.git?rev=df094724adf95136d5d09cf6d54296a3e809fff8#df094724adf95136d5d09cf6d54296a3e809fff8" + +[[package]] +name = "rustls-ring" +version = "0.1.0-dev.1" +source = "git+https://github.com/rustls/rustls.git?branch=main#3925f65934364edafe8d6b20707d9e5e6183648e" +dependencies = [ + "ring", + "rustls", + "rustls-pki-types", + "subtle", +] + +[[package]] +name = "rustls-util" +version = "0.1.0-dev.1" +source = "git+https://github.com/rustls/rustls.git?branch=main#3925f65934364edafe8d6b20707d9e5e6183648e" +dependencies = [ + "rustls", + "tracing", +] [[package]] name = "rustls-webpki" -version = "0.103.13" +version = "0.104.0-alpha.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "61c429a8649f110dddef65e2a5ad240f747e85f7758a6bccc7e5777bd33f756e" +checksum = "bea702cca24d344fc70973022bf7eb920c224e318466eb49784272337dd24b1a" dependencies = [ - "aws-lc-rs", - "ring", "rustls-pki-types", "untrusted", ] diff --git a/Cargo.toml b/Cargo.toml index c5b8381436..f9a1438581 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -35,8 +35,11 @@ rand = "0.10.1" rcgen = "0.14" ring = "0.17" rustc-hash = "2" -rustls = { version = "0.23.5", default-features = false, features = ["std"] } -rustls-platform-verifier = "0.7" +rustls = { version = "0.24.0-dev.1", git = "https://github.com/rustls/rustls.git", branch = "main", default-features = false, features = ["webpki"] } +rustls-aws-lc-rs = { version = "0.1.0-dev.1", git = "https://github.com/rustls/rustls.git", branch = "main", default-features = false, features = ["aws-lc-sys", "std"] } +rustls-platform-verifier = { version = "0.8.0", git = "https://github.com/iadev09/rustls-platform-verifier.git", rev = "df094724adf95136d5d09cf6d54296a3e809fff8", default-features = false } +rustls-ring = { version = "0.1.0-dev.1", git = "https://github.com/rustls/rustls.git", branch = "main", default-features = false, features = ["std"] } +rustls-util = { version = "0.1.0-dev.1", git = "https://github.com/rustls/rustls.git", branch = "main" } rustls-pki-types = "1.7" serde = { version = "1.0", features = ["derive"] } serde_json = "1" diff --git a/bench/Cargo.toml b/bench/Cargo.toml index 3712da30d3..3b87b3cdb7 100644 --- a/bench/Cargo.toml +++ b/bench/Cargo.toml @@ -10,9 +10,10 @@ anyhow = { workspace = true } bytes = { workspace = true } clap = { workspace = true } hdrhistogram = { workspace = true } -quinn = { path = "../quinn", features = ["ring"] } +quinn = { path = "../quinn" } rcgen = { workspace = true } rustls = { workspace = true } +rustls-aws-lc-rs = { workspace = true } tokio = { workspace = true, features = ["rt"] } tracing = { workspace = true } tracing-subscriber = { workspace = true } diff --git a/bench/src/lib.rs b/bench/src/lib.rs index 86c30c8c13..5e0420ee80 100644 --- a/bench/src/lib.rs +++ b/bench/src/lib.rs @@ -63,17 +63,12 @@ pub async fn connect_client( let mut roots = RootCertStore::empty(); roots.add(server_cert)?; - let default_provider = rustls::crypto::ring::default_provider(); - let provider = rustls::crypto::CryptoProvider { - cipher_suites: vec![opt.cipher.as_rustls()], - ..default_provider - }; + let mut provider = rustls_aws_lc_rs::DEFAULT_PROVIDER; + provider.tls13_cipher_suites = vec![opt.cipher.as_rustls()].into(); - let crypto = rustls::ClientConfig::builder_with_provider(provider.into()) - .with_protocol_versions(&[&rustls::version::TLS13]) - .unwrap() + let crypto = rustls::ClientConfig::builder(Arc::new(provider)) .with_root_certificates(roots) - .with_no_client_auth(); + .with_no_client_auth()?; let mut client_config = quinn::ClientConfig::new(Arc::new(QuicClientConfig::try_from(crypto)?)); client_config.transport_config(Arc::new(transport_config(&opt))); @@ -230,8 +225,8 @@ pub enum CipherSuite { } impl CipherSuite { - pub fn as_rustls(self) -> rustls::SupportedCipherSuite { - use rustls::crypto::ring::cipher_suite; + pub fn as_rustls(self) -> &'static rustls::Tls13CipherSuite { + use rustls_aws_lc_rs::cipher_suite; match self { Self::Aes128 => cipher_suite::TLS13_AES_128_GCM_SHA256, Self::Aes256 => cipher_suite::TLS13_AES_256_GCM_SHA384, diff --git a/docs/book/Cargo.toml b/docs/book/Cargo.toml index 484f57050d..4dbb14b70b 100644 --- a/docs/book/Cargo.toml +++ b/docs/book/Cargo.toml @@ -14,3 +14,4 @@ bytes = { workspace = true } quinn = { version = "0.12.0", path = "../../quinn" } rcgen.workspace = true rustls.workspace = true +rustls-ring.workspace = true diff --git a/docs/book/src/bin/certificate.rs b/docs/book/src/bin/certificate.rs index eb3d83fd41..d78da7f115 100644 --- a/docs/book/src/bin/certificate.rs +++ b/docs/book/src/bin/certificate.rs @@ -1,16 +1,10 @@ -use std::{error::Error, sync::Arc}; +use std::{error::Error, hash::Hasher, sync::Arc}; -use quinn::{ - ClientConfig, - crypto::rustls::{NoInitialCipherSuite, QuicClientConfig}, -}; +use quinn::{ClientConfig, crypto::rustls::QuicClientConfig}; use rustls::{ - DigitallySignedStruct, SignatureScheme, client::danger, crypto::{CryptoProvider, verify_tls12_signature, verify_tls13_signature}, - pki_types::{ - CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer, ServerName, UnixTime, pem::PemObject, - }, + pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer, pem::PemObject}, }; #[allow(unused_variables)] @@ -22,69 +16,64 @@ fn main() { } #[allow(dead_code)] // Included in `certificate.md` -fn configure_client() -> Result { - let crypto = rustls::ClientConfig::builder() +fn configure_client() -> Result> { + let crypto = rustls::ClientConfig::builder(Arc::new(rustls_ring::DEFAULT_PROVIDER)) .dangerous() .with_custom_certificate_verifier(SkipServerVerification::new()) - .with_no_client_auth(); + .with_no_client_auth()?; Ok(ClientConfig::new(Arc::new(QuicClientConfig::try_from( crypto, )?))) } -// Implementation of `ServerCertVerifier` that verifies everything as trustworthy. +// Implementation of `ServerVerifier` that verifies everything as trustworthy. #[derive(Debug)] struct SkipServerVerification(Arc); impl SkipServerVerification { fn new() -> Arc { - Arc::new(Self(Arc::new(rustls::crypto::ring::default_provider()))) + Arc::new(Self(Arc::new(rustls_ring::DEFAULT_PROVIDER))) } } -impl danger::ServerCertVerifier for SkipServerVerification { - fn verify_server_cert( +impl danger::ServerVerifier for SkipServerVerification { + fn verify_identity<'a>( &self, - _end_entity: &CertificateDer<'_>, - _intermediates: &[CertificateDer<'_>], - _server_name: &ServerName<'_>, - _ocsp: &[u8], - _now: UnixTime, - ) -> Result { - Ok(danger::ServerCertVerified::assertion()) + identity: &danger::ServerIdentity<'a, '_>, + ) -> Result, rustls::Error> { + Ok(rustls::crypto::VerifiedIdentity::assertion( + identity.identity.clone(), + )) } + fn verify_tls12_signature( &self, - message: &[u8], - cert: &CertificateDer<'_>, - dss: &DigitallySignedStruct, + input: &danger::SignatureVerificationInput<'_>, ) -> Result { - verify_tls12_signature( - message, - cert, - dss, - &self.0.signature_verification_algorithms, - ) + verify_tls12_signature(input, &self.0.signature_verification_algorithms) } fn verify_tls13_signature( &self, - message: &[u8], - cert: &CertificateDer<'_>, - dss: &DigitallySignedStruct, + input: &danger::SignatureVerificationInput<'_>, ) -> Result { - verify_tls13_signature( - message, - cert, - dss, - &self.0.signature_verification_algorithms, - ) + verify_tls13_signature(input, &self.0.signature_verification_algorithms) } - fn supported_verify_schemes(&self) -> Vec { + fn supported_verify_schemes(&self) -> Vec { self.0.signature_verification_algorithms.supported_schemes() } + + fn request_ocsp_response(&self) -> bool { + false + } + + fn hash_config(&self, h: &mut dyn Hasher) { + for scheme in self.supported_verify_schemes() { + h.write_u16(scheme.0); + } + } } fn generate_self_signed_cert() diff --git a/docs/book/src/quinn/certificate.md b/docs/book/src/quinn/certificate.md index 1c52a8c694..77053da4ca 100644 --- a/docs/book/src/quinn/certificate.md +++ b/docs/book/src/quinn/certificate.md @@ -7,25 +7,26 @@ As QUIC uses TLS 1.3 for authentication of connections, the server needs to prov ## Insecure Connection For our example use case, the easiest way to allow the client to trust our server is to disable certificate verification (don't do this in production!). -When the [rustls][3] `dangerous_configuration` feature flag is enabled, a client can be configured to trust any server. +With [rustls][3]'s dangerous client configuration API, a client can be configured to trust any server. -Start by adding a [rustls][3] dependency with the `dangerous_configuration` feature flag to your `Cargo.toml` file. +Start by adding [rustls][3] and provider dependencies to your `Cargo.toml` file. ```toml -quinn = "0.11" -rustls = "0.23" +quinn = "0.12" +rustls = "0.24" +rustls-ring = "0.1" ``` -Then, allow the client to skip the certificate validation by implementing [ServerCertVerifier][ServerCertVerifier] and letting it assert verification for any server. +Then, allow the client to skip the certificate validation by implementing [ServerVerifier][ServerVerifier] and letting it assert verification for any server. ```rust -{{#include ../bin/certificate.rs:36:88}} +{{#include ../bin/certificate.rs:30:77}} ``` -After that, modify the [ClientConfig][ClientConfig] to use this [ServerCertVerifier][ServerCertVerifier] implementation. +After that, modify the [ClientConfig][ClientConfig] to use this [ServerVerifier][ServerVerifier] implementation. ```rust -{{#include ../bin/certificate.rs:25:34}} +{{#include ../bin/certificate.rs:19:28}} ``` Finally, if you plug this [ClientConfig][ClientConfig] into the [Endpoint::set_default_client_config()][set_default_client_config] your client endpoint should verify all connections as trustworthy. @@ -45,7 +46,7 @@ This example uses [rcgen][4] to generate a certificate. Let's look at an example: ```rust -{{#include ../bin/certificate.rs:90:96}} +{{#include ../bin/certificate.rs:79:85}} ``` _Note that [generate_simple_self_signed][generate_simple_self_signed] returns a [Certificate][2] that can be serialized to both `.der` and `.pem` formats._ @@ -68,7 +69,7 @@ certbot asks for the required data and writes the certificates to `fullchain.pem These files can then be referenced in code. ```rust -{{#include ../bin/certificate.rs:98:106}} +{{#include ../bin/certificate.rs:87:95}} ``` ### Configuring Certificates @@ -79,7 +80,7 @@ After configuring plug the configuration into the `Endpoint`. **Configure Server** ```rust -{{#include ../bin/certificate.rs:20}} +{{#include ../bin/certificate.rs:14}} ``` This is the only thing you need to do for your server to be secured. @@ -87,7 +88,7 @@ This is the only thing you need to do for your server to be secured. **Configure Client** ```rust -{{#include ../bin/certificate.rs:21}} +{{#include ../bin/certificate.rs:15}} ``` This is the only thing you need to do for your client to trust a server certificate signed by a conventional certificate authority. @@ -104,7 +105,7 @@ This is the only thing you need to do for your client to trust a server certific [6]: https://letsencrypt.org/getting-started/ [7]: https://certbot.eff.org/instructions [ClientConfig]: https://docs.rs/quinn/latest/quinn/struct.ClientConfig.html -[ServerCertVerifier]: https://docs.rs/rustls/latest/rustls/client/trait.ServerCertVerifier.html +[ServerVerifier]: https://docs.rs/rustls/latest/rustls/client/danger/trait.ServerVerifier.html [set_default_client_config]: https://docs.rs/quinn/latest/quinn/struct.Endpoint.html#method.set_default_client_config [generate_simple_self_signed]: https://docs.rs/rcgen/latest/rcgen/fn.generate_simple_self_signed.html [Certificate]: https://docs.rs/rcgen/latest/rcgen/struct.Certificate.html diff --git a/perf/Cargo.toml b/perf/Cargo.toml index 7c6fe7e119..9ce3bed502 100644 --- a/perf/Cargo.toml +++ b/perf/Cargo.toml @@ -34,6 +34,8 @@ quinn = { path = "../quinn" } quinn-proto = { path = "../quinn-proto" } rcgen = { workspace = true } rustls = { workspace = true } +rustls-aws-lc-rs = { workspace = true } +rustls-util = { workspace = true } serde = { workspace = true, optional = true } serde_json = { workspace = true, optional = true } socket2 = { workspace = true } diff --git a/perf/src/client.rs b/perf/src/client.rs index ff25d7d1f3..c848a310bd 100644 --- a/perf/src/client.rs +++ b/perf/src/client.rs @@ -1,6 +1,7 @@ #[cfg(feature = "json-output")] use std::path::PathBuf; use std::{ + hash::Hasher, net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr}, path::Path, sync::Arc, @@ -11,7 +12,7 @@ use anyhow::{Context, Result}; use bytes::Bytes; use clap::Parser; use quinn::{TokioRuntime, crypto::rustls::QuicClientConfig}; -use rustls::pki_types::{CertificateDer, ServerName, UnixTime}; +use rustls::enums::ApplicationProtocol; use tokio::sync::Semaphore; use tracing::{debug, error, info}; @@ -104,22 +105,18 @@ pub async fn run(opt: Opt) -> Result<()> { let endpoint = quinn::Endpoint::new(endpoint_cfg, None, socket, Arc::new(TokioRuntime))?; - let default_provider = rustls::crypto::ring::default_provider(); - let provider = Arc::new(rustls::crypto::CryptoProvider { - cipher_suites: PERF_CIPHER_SUITES.into(), - ..default_provider - }); + let mut provider = rustls_aws_lc_rs::DEFAULT_PROVIDER; + provider.tls13_cipher_suites = PERF_CIPHER_SUITES.into(); + let provider = Arc::new(provider); - let mut crypto = rustls::ClientConfig::builder_with_provider(provider.clone()) - .with_protocol_versions(&[&rustls::version::TLS13]) - .unwrap() + let mut crypto = rustls::ClientConfig::builder(provider.clone()) .dangerous() .with_custom_certificate_verifier(SkipServerVerification::new(provider)) - .with_no_client_auth(); - crypto.alpn_protocols = vec![b"perf".to_vec()]; + .with_no_client_auth()?; + crypto.alpn_protocols = vec![ApplicationProtocol::from(b"perf")]; if opt.common.keylog { - crypto.key_log = Arc::new(rustls::KeyLogFile::new()); + crypto.key_log = Arc::new(rustls_util::KeyLogFile::new()); } let transport = opt.common.build_transport_config( @@ -381,47 +378,41 @@ impl SkipServerVerification { } } -impl rustls::client::danger::ServerCertVerifier for SkipServerVerification { - fn verify_server_cert( +impl rustls::client::danger::ServerVerifier for SkipServerVerification { + fn verify_identity<'a>( &self, - _end_entity: &CertificateDer<'_>, - _intermediates: &[CertificateDer<'_>], - _server_name: &ServerName<'_>, - _ocsp: &[u8], - _now: UnixTime, - ) -> Result { - Ok(rustls::client::danger::ServerCertVerified::assertion()) + identity: &rustls::client::danger::ServerIdentity<'a, '_>, + ) -> Result, rustls::Error> { + Ok(rustls::crypto::VerifiedIdentity::assertion( + identity.identity.clone(), + )) } fn verify_tls12_signature( &self, - message: &[u8], - cert: &CertificateDer<'_>, - dss: &rustls::DigitallySignedStruct, + input: &rustls::client::danger::SignatureVerificationInput<'_>, ) -> Result { - rustls::crypto::verify_tls12_signature( - message, - cert, - dss, - &self.0.signature_verification_algorithms, - ) + rustls::crypto::verify_tls12_signature(input, &self.0.signature_verification_algorithms) } fn verify_tls13_signature( &self, - message: &[u8], - cert: &CertificateDer<'_>, - dss: &rustls::DigitallySignedStruct, + input: &rustls::client::danger::SignatureVerificationInput<'_>, ) -> Result { - rustls::crypto::verify_tls13_signature( - message, - cert, - dss, - &self.0.signature_verification_algorithms, - ) + rustls::crypto::verify_tls13_signature(input, &self.0.signature_verification_algorithms) } - fn supported_verify_schemes(&self) -> Vec { + fn supported_verify_schemes(&self) -> Vec { self.0.signature_verification_algorithms.supported_schemes() } + + fn request_ocsp_response(&self) -> bool { + false + } + + fn hash_config(&self, h: &mut dyn Hasher) { + for scheme in self.supported_verify_schemes() { + h.write_u16(scheme.0); + } + } } diff --git a/perf/src/lib.rs b/perf/src/lib.rs index f547643cb0..1aa95273f4 100644 --- a/perf/src/lib.rs +++ b/perf/src/lib.rs @@ -11,7 +11,7 @@ use quinn::{ congestion::{self, ControllerFactory}, udp::UdpSocketState, }; -use rustls::crypto::ring::cipher_suite; +use rustls_aws_lc_rs::cipher_suite; use socket2::{Domain, Protocol, Socket, Type}; use tracing::warn; @@ -214,7 +214,7 @@ impl CongestionAlgorithm { } } -pub static PERF_CIPHER_SUITES: &[rustls::SupportedCipherSuite] = &[ +pub static PERF_CIPHER_SUITES: &[&rustls::Tls13CipherSuite] = &[ cipher_suite::TLS13_AES_128_GCM_SHA256, cipher_suite::TLS13_AES_256_GCM_SHA384, cipher_suite::TLS13_CHACHA20_POLY1305_SHA256, diff --git a/perf/src/noprotection.rs b/perf/src/noprotection.rs index 37eb118f42..2c4521a517 100644 --- a/perf/src/noprotection.rs +++ b/perf/src/noprotection.rs @@ -120,7 +120,7 @@ impl crypto::Session for NoProtectionSession { } fn export_keying_material( - &self, + &mut self, output: &mut [u8], label: &[u8], context: &[u8], diff --git a/perf/src/server.rs b/perf/src/server.rs index 18d1a4d8ef..0fb64508a8 100644 --- a/perf/src/server.rs +++ b/perf/src/server.rs @@ -4,7 +4,11 @@ use anyhow::{Context, Result}; use bytes::Bytes; use clap::Parser; use quinn::{TokioRuntime, crypto::rustls::QuicServerConfig}; -use rustls::pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer, pem::PemObject}; +use rustls::{ + crypto::Identity, + enums::ApplicationProtocol, + pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer, pem::PemObject}, +}; use tracing::{debug, error, info}; use crate::{CommonOpt, PERF_CIPHER_SUITES, noprotection::NoProtectionServerConfig}; @@ -44,22 +48,16 @@ pub async fn run(opt: Opt) -> Result<()> { } }; - let default_provider = rustls::crypto::ring::default_provider(); - let provider = rustls::crypto::CryptoProvider { - cipher_suites: PERF_CIPHER_SUITES.into(), - ..default_provider - }; + let mut provider = rustls_aws_lc_rs::DEFAULT_PROVIDER; + provider.tls13_cipher_suites = PERF_CIPHER_SUITES.into(); - let mut crypto = rustls::ServerConfig::builder_with_provider(provider.into()) - .with_protocol_versions(&[&rustls::version::TLS13]) - .unwrap() + let mut crypto = rustls::ServerConfig::builder(Arc::new(provider)) .with_no_client_auth() - .with_single_cert(cert, key) - .unwrap(); - crypto.alpn_protocols = vec![b"perf".to_vec()]; + .with_single_cert(Arc::new(Identity::from_cert_chain(cert)?), key)?; + crypto.alpn_protocols = vec![ApplicationProtocol::from(b"perf")]; if opt.common.keylog { - crypto.key_log = Arc::new(rustls::KeyLogFile::new()); + crypto.key_log = Arc::new(rustls_util::KeyLogFile::new()); } let transport = opt.common.build_transport_config( diff --git a/quinn-proto/Cargo.toml b/quinn-proto/Cargo.toml index c018ef1541..af4eee0638 100644 --- a/quinn-proto/Cargo.toml +++ b/quinn-proto/Cargo.toml @@ -21,26 +21,31 @@ bloom = ["dep:fastbloom"] # For backwards compatibility, `rustls` forwards to `rustls-ring` rustls = ["rustls-ring"] # Enable rustls with the `aws-lc-rs` crypto provider -rustls-aws-lc-rs = ["dep:rustls", "rustls?/aws-lc-rs", "aws-lc-rs"] -rustls-aws-lc-rs-fips = ["rustls-aws-lc-rs", "aws-lc-rs-fips"] +rustls-aws-lc-rs = ["dep:rustls", "dep:rustls-aws-lc-rs", "aws-lc-rs"] +rustls-aws-lc-rs-fips = ["rustls-aws-lc-rs", "aws-lc-rs-fips", "rustls-aws-lc-rs?/fips"] # Enable rustls with the `ring` crypto provider -rustls-ring = ["dep:rustls", "rustls?/ring", "ring"] +rustls-ring = ["dep:rustls", "dep:rustls-ring", "ring"] ring = ["dep:ring"] # Enable rustls ring provider and direct ring usage # Provides `ClientConfig::with_platform_verifier()` convenience method platform-verifier = ["dep:rustls-platform-verifier"] # Configure `tracing` to log events via `log` if no `tracing` subscriber exists. tracing-log = ["tracing/log"] -# Enable rustls logging -rustls-log = ["rustls?/logging"] +# Enable rustls tracing (feature name retained for backwards compatibility) +rustls-log = ["rustls?/tracing"] # Enable qlog support qlog = ["dep:qlog"] -# Internal (PRIVATE!) features used to aid testing. +# Internal (PRIVATE!) features. # Don't rely on these whatsoever. They may disappear at any time. +# Used only to aid testing. __rustls-post-quantum-test = [] +# Re-exports split-accept state types required by the `quinn` crate. +# This doc-hidden API is an implementation detail. +__internal_split_accept = [] + [dependencies] arbitrary = { workspace = true, optional = true } aws-lc-rs = { workspace = true, optional = true } @@ -53,7 +58,9 @@ rand = { workspace = true } rand_pcg = "0.10" ring = { workspace = true, optional = true } rustls = { workspace = true, optional = true } +rustls-aws-lc-rs = { workspace = true, optional = true } rustls-platform-verifier = { workspace = true, optional = true } +rustls-ring = { workspace = true, optional = true } slab = { workspace = true } thiserror = { workspace = true } tinyvec = { workspace = true, features = ["alloc"] } @@ -71,6 +78,7 @@ web-time = { workspace = true } assert_matches = { workspace = true } hex-literal = { workspace = true } rcgen = { workspace = true } +rustls-util = { workspace = true } tracing-subscriber = { workspace = true } wasm-bindgen-test = { workspace = true } diff --git a/quinn-proto/src/config/mod.rs b/quinn-proto/src/config/mod.rs index 129ee1dbaa..f3663b78da 100644 --- a/quinn-proto/src/config/mod.rs +++ b/quinn-proto/src/config/mod.rs @@ -310,13 +310,13 @@ impl ServerConfig { self } - /// Maximum number of [`Incoming`][crate::Incoming] to allow to exist at a time + /// Maximum number of incoming connection attempts to hold before they become active /// - /// An [`Incoming`][crate::Incoming] comes into existence when an incoming connection attempt - /// is received and stops existing when the application either accepts it or otherwise disposes - /// of it. While this limit is reached, new incoming connection attempts are immediately - /// refused. Larger values have greater worst-case memory consumption, but accommodate greater - /// application latency in handling incoming connection attempts. + /// An attempt counts toward this limit from the time its initial packet is received until it is + /// either registered as an active connection or otherwise disposed of. While this limit is + /// reached, new incoming connection attempts are not admitted. Larger values have greater + /// worst-case memory consumption, but accommodate greater application latency in handling + /// incoming connection attempts. /// /// The default value is set to 65536. With a typical Ethernet MTU of 1500 bytes, this limits /// memory consumption from this to under 100 MiB--a generous amount that still prevents memory @@ -326,13 +326,11 @@ impl ServerConfig { self } - /// Maximum number of received bytes to buffer for each [`Incoming`][crate::Incoming] + /// Maximum number of received bytes to buffer for each incoming connection attempt /// - /// An [`Incoming`][crate::Incoming] comes into existence when an incoming connection attempt - /// is received and stops existing when the application either accepts it or otherwise disposes - /// of it. This limit governs only packets received within that period, and does not include - /// the first packet. Packets received in excess of this limit are dropped, which may cause - /// 0-RTT or handshake data to have to be retransmitted. + /// This limit governs packets received after the first packet and before the attempt is either + /// registered as an active connection or otherwise disposed of. Packets received in excess of + /// this limit are dropped, which may cause 0-RTT or handshake data to have to be retransmitted. /// /// The default value is set to 10 MiB--an amount such that in most situations a client would /// not transmit that much 0-RTT data faster than the server handles the corresponding @@ -342,14 +340,12 @@ impl ServerConfig { self } - /// Maximum number of received bytes to buffer for all [`Incoming`][crate::Incoming] - /// collectively + /// Maximum number of received bytes to buffer collectively for all incoming connection attempts /// - /// An [`Incoming`][crate::Incoming] comes into existence when an incoming connection attempt - /// is received and stops existing when the application either accepts it or otherwise disposes - /// of it. This limit governs only packets received within that period, and does not include - /// the first packet. Packets received in excess of this limit are dropped, which may cause - /// 0-RTT or handshake data to have to be retransmitted. + /// This limit governs packets received after each attempt's first packet and before the attempt + /// is either registered as an active connection or otherwise disposed of. Packets received in + /// excess of this limit are dropped, which may cause 0-RTT or handshake data to have to be + /// retransmitted. /// /// The default value is set to 100 MiB--a generous amount that still prevents memory /// exhaustion in most contexts. @@ -632,8 +628,9 @@ impl ClientConfig { pub fn with_root_certificates( roots: Arc, ) -> Result { + let provider = configured_provider(); Ok(Self::new(Arc::new(crypto::rustls::QuicClientConfig::new( - WebPkiServerVerifier::builder_with_provider(roots, configured_provider()).build()?, + Arc::new(WebPkiServerVerifier::builder(roots, &provider).build()?), )))) } } diff --git a/quinn-proto/src/connection/assembler.rs b/quinn-proto/src/connection/assembler.rs index 994e0068ec..ed065bf9d5 100644 --- a/quinn-proto/src/connection/assembler.rs +++ b/quinn-proto/src/connection/assembler.rs @@ -10,7 +10,7 @@ use crate::range_set::RangeSet; /// Helper to assemble unordered stream frames into an ordered stream #[derive(Debug, Default)] -pub(super) struct Assembler { +pub(crate) struct Assembler { state: State, data: BinaryHeap, /// Total number of buffered bytes, including duplicates in ordered mode. @@ -25,7 +25,7 @@ pub(super) struct Assembler { } impl Assembler { - pub(super) fn new() -> Self { + pub(crate) fn new() -> Self { Self::default() } @@ -57,7 +57,7 @@ impl Assembler { } /// Get the the next chunk - pub(super) fn read(&mut self, max_length: usize, ordered: bool) -> Option { + pub(crate) fn read(&mut self, max_length: usize, ordered: bool) -> Option { loop { let mut chunk = self.data.peek_mut()?; @@ -147,7 +147,7 @@ impl Assembler { // Note: If a packet contains many frames from the same stream, the estimated over-allocation // will be much higher because we are counting the same allocation multiple times. - pub(super) fn insert( + pub(crate) fn insert( &mut self, mut offset: u64, mut bytes: Bytes, @@ -221,10 +221,16 @@ impl Assembler { } /// Number of bytes consumed by the application - pub(super) fn bytes_read(&self) -> u64 { + pub(crate) fn bytes_read(&self) -> u64 { self.bytes_read } + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + pub(crate) fn skip_to(&mut self, offset: u64) { + self.bytes_read = self.bytes_read.max(offset); + self.end = self.end.max(offset); + } + /// Discard all buffered data pub(super) fn clear(&mut self) { self.data.clear(); diff --git a/quinn-proto/src/connection/mod.rs b/quinn-proto/src/connection/mod.rs index 745640baec..6fd614fd56 100644 --- a/quinn-proto/src/connection/mod.rs +++ b/quinn-proto/src/connection/mod.rs @@ -40,7 +40,7 @@ use crate::{ mod ack_frequency; use ack_frequency::AckFrequencyState; -mod assembler; +pub(crate) mod assembler; pub use assembler::Chunk; mod cid_state; @@ -361,6 +361,9 @@ impl Connection { if path_validated { this.on_path_validated(); } + if this.crypto.handshake_data().is_some() { + this.events.push_back(Event::HandshakeDataReady); + } if side.is_client() { // Kick off the connection this.write_crypto(); @@ -1106,6 +1109,10 @@ impl Connection { first_decode, remaining, }) => { + if self.is_handshaking() && remote != self.path.remote { + debug!("discarding packet with unexpected remote during handshake"); + return; + } // If this packet could initiate a migration and we're a client or a server that // forbids migration, drop the datagram. This could be relaxed to heuristically // permit NAT-rebinding-like migration. @@ -1308,6 +1315,28 @@ impl Connection { &*self.crypto } + /// Get a mutable session reference + pub fn crypto_session_mut(&mut self) -> &mut dyn crypto::Session { + &mut *self.crypto + } + + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + pub(crate) fn skip_initial_crypto(&mut self, offset: u64) { + self.spaces[SpaceId::Initial].crypto_stream.skip_to(offset); + } + + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + pub(crate) fn seed_staged_initial_sends(&mut self, next: u64, bytes: u64, datagrams: u64) { + let space = &mut self.spaces[SpaceId::Initial]; + space.next_packet_number = space.next_packet_number.max(next); + self.path.total_sent = self.path.total_sent.saturating_add(bytes); + self.stats.udp_tx.datagrams = self.stats.udp_tx.datagrams.saturating_add(datagrams); + self.stats.udp_tx.bytes = self.stats.udp_tx.bytes.saturating_add(bytes); + self.stats.udp_tx.ios = self.stats.udp_tx.ios.saturating_add(datagrams); + self.stats.frame_tx.acks = self.stats.frame_tx.acks.saturating_add(datagrams); + self.stats.path.sent_packets = self.stats.path.sent_packets.saturating_add(datagrams); + } + /// Whether the connection is in the process of being established /// /// If this returns `false`, the connection may be either established or closed, signaled by the @@ -2331,6 +2360,39 @@ impl Connection { return; } + if self.side.is_server() { + if let Some(Packet { + header: Header::Initial(header), + .. + }) = packet.as_ref() + { + if header.version != self.version { + debug!( + version = header.version, + expected = self.version, + "discarding Initial with unexpected version" + ); + return; + } + if let State::Handshake(handshake) = &self.state { + if header.token != handshake.expected_token { + // Initial packets can be spoofed, so discard rather than killing the + // connection. + warn!("discarding Initial with invalid retry token"); + return; + } + if handshake.rem_cid_set && header.src_cid != self.rem_handshake_cid { + debug!( + expected = %self.rem_handshake_cid, + actual = %header.src_cid, + "discarding Initial with mismatched remote CID" + ); + return; + } + } + } + } + let was_closed = self.state.is_closed(); let was_drained = self.state.is_drained(); @@ -2381,18 +2443,6 @@ impl Connection { trace!("dropping short packet during handshake"); return; } else { - if let Header::Initial(InitialHeader { ref token, .. }) = packet.header { - if let State::Handshake(ref hs) = self.state { - if self.side.is_server() && token != &hs.expected_token { - // Clients must send the same retry token in every Initial. Initial - // packets can be spoofed, so we discard rather than killing the - // connection. - warn!("discarding Initial with invalid retry token"); - return; - } - } - } - if !self.state.is_closed() { let spin = match packet.header { Header::Short { spin, .. } => spin, diff --git a/quinn-proto/src/crypto.rs b/quinn-proto/src/crypto.rs index 2ac40fc1ee..3087f2b930 100644 --- a/quinn-proto/src/crypto.rs +++ b/quinn-proto/src/crypto.rs @@ -86,7 +86,7 @@ pub trait Session: Send + Sync + 'static { /// This function will fail, returning [ExportKeyingMaterialError], /// if the requested output length is too large. fn export_keying_material( - &self, + &mut self, output: &mut [u8], label: &[u8], context: &[u8], @@ -139,6 +139,19 @@ pub trait ServerConfig: Send + Sync { version: u32, params: &TransportParameters, ) -> Box; + + /// Continue a server session after rustls has read the ClientHello. + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + fn start_session_from_accepted( + self: Arc, + _version: u32, + _params: &TransportParameters, + _accepted: rustls::Accepted, + ) -> Result, TransportError> { + Err(TransportError::PROTOCOL_VIOLATION( + "server crypto config does not support rustls acceptor", + )) + } } /// Keys used to protect packet payloads diff --git a/quinn-proto/src/crypto/rustls.rs b/quinn-proto/src/crypto/rustls.rs index b1fd9da1e4..e7392d3da8 100644 --- a/quinn-proto/src/crypto/rustls.rs +++ b/quinn-proto/src/crypto/rustls.rs @@ -1,31 +1,35 @@ -use std::{any::Any, io, str, sync::Arc}; +use std::{any::Any, collections::VecDeque, io, str, sync::Arc}; -#[cfg(all(feature = "aws-lc-rs", not(feature = "ring")))] -use aws_lc_rs::aead; +use crate::{ + ConnectError, ConnectionId, Side, TransportError, TransportErrorCode, + crypto::{ + self, CryptoError, ExportKeyingMaterialError, HeaderKey, KeyPair, Keys, UnsupportedVersion, + }, + transport_parameters::TransportParameters, +}; use bytes::BytesMut; -#[cfg(feature = "ring")] -use ring::aead; pub use rustls::Error; #[cfg(feature = "__rustls-post-quantum-test")] -use rustls::NamedGroup; +use rustls::crypto::kx::NamedGroup; use rustls::{ - self, CipherSuite, - client::danger::ServerCertVerifier, + self, TlsInputBuffer, + client::danger::ServerVerifier, + crypto::{ + CipherSuite, CryptoProvider, Identity, + cipher::{AeadKey, Iv}, + }, + error::AlertDescription, pki_types::{CertificateDer, PrivateKeyDer, ServerName}, - quic::{Connection, HeaderProtectionKey, KeyChange, PacketKey, Secrets, Suite, Version}, + quic::{ + ClientConnection, Connection as _, DirectionalKeys, HeaderProtectionKey, KeyChange, + NeedsInput, PacketKey, QuicEvent, Secrets, ServerConnection, ServerHandshake, + Side as QuicSide, Suite, Version, + }, }; #[cfg(feature = "platform-verifier")] use rustls_platform_verifier::BuilderVerifierExt; -use crate::{ - ConnectError, ConnectionId, Side, TransportError, TransportErrorCode, - crypto::{ - self, CryptoError, ExportKeyingMaterialError, HeaderKey, KeyPair, Keys, UnsupportedVersion, - }, - transport_parameters::TransportParameters, -}; - -impl From for rustls::Side { +impl From for QuicSide { fn from(s: Side) -> Self { match s { Side::Client => Self::Client, @@ -39,15 +43,170 @@ pub struct TlsSession { version: Version, got_handshake_data: bool, next_secrets: Option, - inner: Connection, + exporter: Option, + inner: QuicConnection, + input: HandshakeInput, + pending_events: VecDeque, suite: Suite, } +#[derive(Default)] +struct HandshakeInput { + bytes: Vec, + offset: usize, + consumed: usize, +} + +impl HandshakeInput { + fn extend_from_slice(&mut self, bytes: &[u8]) { + if self.offset != 0 { + self.bytes.drain(..self.offset); + self.offset = 0; + } + self.bytes.extend_from_slice(bytes); + } + + fn len(&self) -> usize { + self.bytes.len() - self.offset + } + + fn is_empty(&self) -> bool { + self.len() == 0 + } + + fn consumed(&self) -> usize { + self.consumed + } +} + +impl TlsInputBuffer for HandshakeInput { + fn slice_mut(&mut self) -> &mut [u8] { + &mut self.bytes[self.offset..] + } + + fn discard(&mut self, num_bytes: usize) { + assert!(num_bytes <= self.len()); + self.offset += num_bytes; + self.consumed += num_bytes; + if self.offset == self.bytes.len() { + self.bytes.clear(); + self.offset = 0; + } + } + + fn received_close_notify(&mut self) {} + + fn has_seen_eof(&self) -> bool { + false + } +} + +fn transport_error_from_rustls(e: Error) -> TransportError { + if let Ok(alert) = AlertDescription::try_from(&e) { + TransportError { + code: TransportErrorCode::crypto(alert.into()), + frame: None, + reason: e.to_string(), + crypto: Some(Arc::new(e)), + } + } else { + TransportError::PROTOCOL_VIOLATION(format!("TLS error: {e}")) + } +} + +/// A rustls QUIC handshake paused after reading ClientHello. +pub struct Accepted { + inner: rustls::quic::Accepted, + input: HandshakeInput, + initial_crypto_offset: u64, +} + +impl Accepted { + /// Get the ClientHello for this connection. + pub fn client_hello(&self) -> rustls::server::ClientHello<'_> { + self.inner.client_hello() + } + + pub(crate) fn initial_crypto_offset(&self) -> u64 { + self.initial_crypto_offset + } +} + +/// Reads a rustls QUIC ClientHello before a server configuration is selected. +pub(crate) struct Acceptor { + state: Option, + input: HandshakeInput, +} + +impl Acceptor { + pub(crate) fn new(version: u32) -> Result { + Ok(Self { + state: Some(ServerHandshake::start(interpret_version(version)?)), + input: HandshakeInput::default(), + }) + } + + pub(crate) fn read_hs(&mut self, plaintext: &[u8]) -> Result, TransportError> { + self.input.extend_from_slice(plaintext); + loop { + let before = self.input.len(); + let Some(state) = self.state.take() else { + return Err(TransportError::INTERNAL_ERROR( + "rustls acceptor used after completion", + )); + }; + let mut events = Vec::new(); + let state = state + .process(&mut self.input, &mut events) + .map_err(transport_error_from_rustls)?; + if !events.is_empty() { + return Err(TransportError::INTERNAL_ERROR( + "rustls emitted data before config selection", + )); + } + match state { + ServerHandshake::NeedsInput(state) => { + self.state = Some(state); + if self.input.is_empty() || self.input.len() == before { + return Ok(None); + } + } + ServerHandshake::Accepted(inner) => { + let initial_crypto_offset = (self.input.consumed() + self.input.len()) as u64; + return Ok(Some(Accepted { + inner, + input: std::mem::take(&mut self.input), + initial_crypto_offset, + })); + } + ServerHandshake::VerifyClientIdentity(_) | ServerHandshake::Complete(_) => { + return Err(TransportError::INTERNAL_ERROR( + "rustls advanced past ClientHello without a server config", + )); + } + _ => { + return Err(TransportError::INTERNAL_ERROR( + "rustls returned an unsupported server handshake state", + )); + } + } + } + } +} + impl TlsSession { fn side(&self) -> Side { - match self.inner { - Connection::Client(_) => Side::Client, - Connection::Server(_) => Side::Server, + self.inner.side() + } + + fn required_handshake_data_is_ready(&self) -> bool { + #[cfg(feature = "__rustls-post-quantum-test")] + { + self.inner.negotiated_key_exchange_group().is_some() + } + #[cfg(not(feature = "__rustls-post-quantum-test"))] + { + true } } } @@ -61,40 +220,38 @@ impl crypto::Session for TlsSession { if !self.got_handshake_data { return None; } + #[cfg(feature = "__rustls-post-quantum-test")] + let negotiated_key_exchange_group = self + .inner + .negotiated_key_exchange_group() + .expect("key exchange group is negotiated"); + Some(Box::new(HandshakeData { protocol: self.inner.alpn_protocol().map(|x| x.into()), - server_name: match &self.inner { - Connection::Client(_) => None, - Connection::Server(session) => session.server_name().map(|x| x.into()), - }, + server_name: self.inner.server_name().map(str::to_owned), protocol_version: match &self.inner { - Connection::Client(session) => session.protocol_version(), - Connection::Server(session) => session.protocol_version(), + QuicConnection::Client(session) => session.protocol_version(), + QuicConnection::Server(session) => session.protocol_version(), + QuicConnection::ServerHandshake(session) => session.protocol_version(), } .map(|x| -> Box { Box::new(x) }), cipher_suite: match &self.inner { - Connection::Client(session) => session.negotiated_cipher_suite(), - Connection::Server(session) => session.negotiated_cipher_suite(), + QuicConnection::Client(session) => session.negotiated_cipher_suite(), + QuicConnection::Server(session) => session.negotiated_cipher_suite(), + QuicConnection::ServerHandshake(session) => session.negotiated_cipher_suite(), } .map(|suite| -> Box { Box::new(suite.suite()) }), #[cfg(feature = "__rustls-post-quantum-test")] - negotiated_key_exchange_group: self - .inner - .negotiated_key_exchange_group() - .expect("key exchange group is negotiated") - .name(), + negotiated_key_exchange_group, })) } - /// For the rustls `TlsSession`, the `Any` type is `Vec` + /// For the rustls `TlsSession`, the `Any` type is `rustls::crypto::Identity<'static>` fn peer_identity(&self) -> Option> { - self.inner.peer_certificates().map(|v| -> Box { - Box::new( - v.iter() - .map(|v| v.clone().into_owned()) - .collect::>>(), - ) - }) + self.inner + .peer_identity() + .cloned() + .map(|identity| -> Box { Box::new(identity) }) } fn early_crypto(&self) -> Option<(Box, Box)> { @@ -103,10 +260,7 @@ impl crypto::Session for TlsSession { } fn early_data_accepted(&self) -> Option { - match self.inner { - Connection::Client(ref session) => Some(session.is_early_data_accepted()), - _ => None, - } + self.inner.is_early_data_accepted() } fn is_handshaking(&self) -> bool { @@ -114,27 +268,25 @@ impl crypto::Session for TlsSession { } fn read_handshake(&mut self, buf: &[u8]) -> Result { - self.inner.read_hs(buf).map_err(|e| { - if let Some(alert) = self.inner.alert() { - TransportError { - code: TransportErrorCode::crypto(alert.into()), - frame: None, - reason: e.to_string(), - crypto: Some(Arc::new(e)), - } - } else { - TransportError::PROTOCOL_VIOLATION(format!("TLS error: {e}")) + self.input.extend_from_slice(buf); + loop { + let before = self.input.len(); + self.inner + .read_hs(&mut self.input) + .map_err(transport_error_from_rustls)?; + self.inner.drain_events(&mut self.pending_events); + if self.input.is_empty() || self.input.len() == before { + break; } - })?; + } if !self.got_handshake_data { // Hack around the lack of an explicit signal from rustls to reflect ClientHello being // ready on incoming connections, or ALPN negotiation completing on outgoing // connections. - let have_server_name = match self.inner { - Connection::Client(_) => false, - Connection::Server(ref session) => session.server_name().is_some(), - }; - if self.inner.alpn_protocol().is_some() || have_server_name || !self.is_handshaking() { + let have_server_name = self.inner.server_name().is_some(); + if (self.inner.alpn_protocol().is_some() || have_server_name || !self.is_handshaking()) + && self.required_handshake_data_is_ready() + { self.got_handshake_data = true; return Ok(true); } @@ -153,11 +305,20 @@ impl crypto::Session for TlsSession { } fn write_handshake(&mut self, buf: &mut Vec) -> Option { - let keys = match self.inner.write_hs(buf)? { - KeyChange::Handshake { keys } => keys, - KeyChange::OneRtt { keys, next } => { - self.next_secrets = Some(next); - keys + self.inner.drain_events(&mut self.pending_events); + let keys = loop { + match self.pending_events.pop_front()? { + QuicEvent::Message(message) => buf.extend_from_slice(&message), + QuicEvent::KeyChange(key_change) => { + break match key_change { + KeyChange::Handshake { keys } => keys, + KeyChange::OneRtt { keys, next } => { + self.next_secrets = Some(next); + keys + } + }; + } + event => unreachable!("unsupported rustls QUIC event: {event:?}"), } }; @@ -195,45 +356,371 @@ impl crypto::Session for TlsSession { let tag_start = tag_start + pseudo_packet.len(); pseudo_packet.extend_from_slice(payload); - let (nonce, key) = match self.version { - Version::V1 => (RETRY_INTEGRITY_NONCE_V1, RETRY_INTEGRITY_KEY_V1), - Version::V1Draft => (RETRY_INTEGRITY_NONCE_DRAFT, RETRY_INTEGRITY_KEY_DRAFT), - _ => unreachable!(), - }; - - let nonce = aead::Nonce::assume_unique_for_key(nonce); - let key = aead::LessSafeKey::new(aead::UnboundKey::new(&aead::AES_128_GCM, &key).unwrap()); - let (aad, tag) = pseudo_packet.split_at_mut(tag_start); - key.open_in_place(nonce, aead::Aad::from(aad), tag).is_ok() + retry_key_for_version(self.version, &self.suite) + .decrypt_in_place(0, aad, tag, None) + .is_ok() } fn export_keying_material( - &self, + &mut self, output: &mut [u8], label: &[u8], context: &[u8], ) -> Result<(), ExportKeyingMaterialError> { - self.inner - .export_keying_material(output, label, Some(context)) + if self.exporter.is_none() { + self.exporter = Some( + self.inner + .exporter() + .map_err(|_| ExportKeyingMaterialError)?, + ); + } + + self.exporter + .as_ref() + .expect("exporter is set") + .derive(label, Some(context), output) .map_err(|_| ExportKeyingMaterialError)?; Ok(()) } } -const RETRY_INTEGRITY_KEY_DRAFT: [u8; 16] = [ - 0xcc, 0xce, 0x18, 0x7e, 0xd0, 0x9a, 0x09, 0xd0, 0x57, 0x28, 0x15, 0x5a, 0x6c, 0xb9, 0x6b, 0xe1, -]; -const RETRY_INTEGRITY_NONCE_DRAFT: [u8; 12] = [ - 0xe5, 0x49, 0x30, 0xf9, 0x7f, 0x21, 0x36, 0xf0, 0x53, 0x0a, 0x8c, 0x1c, -]; +enum QuicConnection { + Client(ClientConnection), + Server(ServerConnection), + ServerHandshake(ServerHandshakeConnection), +} -const RETRY_INTEGRITY_KEY_V1: [u8; 16] = [ - 0xbe, 0x0c, 0x69, 0x0b, 0x9f, 0x66, 0x57, 0x5a, 0x1d, 0x76, 0x6b, 0x54, 0xe3, 0x68, 0xc8, 0x4e, -]; -const RETRY_INTEGRITY_NONCE_V1: [u8; 12] = [ - 0x46, 0x15, 0x99, 0xd3, 0x5d, 0x63, 0x2b, 0xf2, 0x23, 0x98, 0x25, 0xbb, -]; +enum ServerHandshakeState { + NeedsInput(NeedsInput), + Complete(ServerConnection), + Failed(Error), +} + +struct ServerHandshakeConnection { + state: Option, + pending_events: VecDeque, + snapshot: ServerHandshakeSnapshot, +} + +#[derive(Default)] +struct ServerHandshakeSnapshot { + alpn_protocol: Option>, + peer_identity: Option>, + quic_transport_parameters: Option>, + server_name: Option, + protocol_version: Option, + negotiated_cipher_suite: Option, + #[cfg(feature = "__rustls-post-quantum-test")] + negotiated_key_exchange_group: Option, +} + +impl ServerHandshakeSnapshot { + fn update_from_needs_input(&mut self, state: &NeedsInput) { + if let Some(protocol) = state.alpn_protocol() { + self.alpn_protocol = Some(protocol.as_ref().to_vec()); + } + if let Some(identity) = state.peer_identity() { + self.peer_identity = Some(identity.identity().clone()); + } + if let Some(params) = state.quic_transport_parameters() { + self.quic_transport_parameters = Some(params.to_vec()); + } + if let Some(server_name) = state.server_name() { + self.server_name = Some(server_name.as_ref().to_owned()); + } + if let Some(version) = state.protocol_version() { + self.protocol_version = Some(version); + } + if let Some(suite) = state.negotiated_cipher_suite() { + self.negotiated_cipher_suite = Some(suite); + } + #[cfg(feature = "__rustls-post-quantum-test")] + if let Some(group) = state.negotiated_key_exchange_group() { + self.negotiated_key_exchange_group = Some(group.name()); + } + } + + fn update_from_complete(&mut self, state: &ServerConnection) { + if let Some(protocol) = state.alpn_protocol() { + self.alpn_protocol = Some(protocol.as_ref().to_vec()); + } + if let Some(identity) = state.peer_identity() { + self.peer_identity = Some(identity.identity().clone()); + } + if let Some(params) = state.quic_transport_parameters() { + self.quic_transport_parameters = Some(params.to_vec()); + } + if let Some(server_name) = state.server_name() { + self.server_name = Some(server_name.as_ref().to_owned()); + } + if let Some(version) = state.protocol_version() { + self.protocol_version = Some(version); + } + if let Some(suite) = state.negotiated_cipher_suite() { + self.negotiated_cipher_suite = Some(suite); + } + #[cfg(feature = "__rustls-post-quantum-test")] + if let Some(group) = state.negotiated_key_exchange_group() { + self.negotiated_key_exchange_group = Some(group.name()); + } + } +} + +impl ServerHandshakeConnection { + fn new(state: ServerHandshake, events: Vec) -> Result { + let mut this = Self { + state: None, + pending_events: events.into(), + snapshot: ServerHandshakeSnapshot::default(), + }; + let _ = this.set_state(state)?; + Ok(this) + } + + /// Store `state`, returning whether synchronous client identity verification advanced the + /// handshake without consuming input. + fn set_state(&mut self, mut state: ServerHandshake) -> Result { + let mut verified_client_identity = false; + loop { + match state { + ServerHandshake::NeedsInput(state) => { + self.snapshot.update_from_needs_input(&state); + self.state = Some(ServerHandshakeState::NeedsInput(state)); + return Ok(verified_client_identity); + } + ServerHandshake::VerifyClientIdentity(verify) => { + verified_client_identity = true; + state = match verify.use_verifier_trait() { + Ok(state) => state, + Err(error) => return self.fail(error), + }; + } + ServerHandshake::Complete(state) => { + self.snapshot.update_from_complete(&state); + self.state = Some(ServerHandshakeState::Complete(state)); + return Ok(false); + } + ServerHandshake::Accepted(_) => { + return self.fail(Error::General( + "server config was requested more than once".into(), + )); + } + _ => { + return self.fail(Error::General( + "rustls returned an unsupported server handshake state".into(), + )); + } + } + } + } + + fn fail(&mut self, error: Error) -> Result { + self.pending_events.clear(); + self.state = Some(ServerHandshakeState::Failed(error.clone())); + Err(error) + } + + fn read_hs(&mut self, input: &mut dyn TlsInputBuffer) -> Result<(), Error> { + let Some(state) = self.state.take() else { + return self.fail(Error::General("rustls handshake state missing".into())); + }; + match state { + ServerHandshakeState::NeedsInput(state) => { + let mut events = Vec::new(); + match state.process(input, &mut events) { + Ok(state) => { + self.pending_events.extend(events); + if self.set_state(state)? { + // `NeedsInput::process()` stops when client identity verification is + // required. The verifier transition itself consumes no input, and + // rustls can already have CertificateVerify/Finished buffered from the + // same CRYPTO chunk. Give the resulting state one immediate chance to + // process that buffered data even when Quinn's input buffer is empty. + self.read_hs(input) + } else { + Ok(()) + } + } + Err(error) => self.fail(error), + } + } + ServerHandshakeState::Complete(mut state) => { + let result = state.read_hs(input); + self.snapshot.update_from_complete(&state); + match result { + Ok(()) => { + self.state = Some(ServerHandshakeState::Complete(state)); + Ok(()) + } + Err(error) => self.fail(error), + } + } + ServerHandshakeState::Failed(error) => self.fail(error), + } + } + + fn drain_events(&mut self, events: &mut VecDeque) { + events.append(&mut self.pending_events); + if let Some(ServerHandshakeState::Complete(state)) = &mut self.state { + events.extend(state.events()); + } + } + + fn alpn_protocol(&self) -> Option<&[u8]> { + self.snapshot.alpn_protocol.as_deref() + } + + fn peer_identity(&self) -> Option<&Identity<'static>> { + self.snapshot.peer_identity.as_ref() + } + + fn zero_rtt_keys(&self) -> Option { + match self.state.as_ref()? { + ServerHandshakeState::NeedsInput(state) => state.zero_rtt_keys(), + ServerHandshakeState::Complete(state) => state.zero_rtt_keys(), + ServerHandshakeState::Failed(_) => None, + } + } + + fn is_handshaking(&self) -> bool { + match self.state.as_ref() { + None => false, + Some(ServerHandshakeState::Failed(_)) => false, + Some(ServerHandshakeState::NeedsInput(_)) => true, + Some(ServerHandshakeState::Complete(state)) => state.is_handshaking(), + } + } + + fn quic_transport_parameters(&self) -> Option<&[u8]> { + self.snapshot.quic_transport_parameters.as_deref() + } + + fn server_name(&self) -> Option<&str> { + self.snapshot.server_name.as_deref() + } + + fn protocol_version(&self) -> Option { + self.snapshot.protocol_version + } + + fn negotiated_cipher_suite(&self) -> Option { + self.snapshot.negotiated_cipher_suite + } + + #[cfg(feature = "__rustls-post-quantum-test")] + fn negotiated_key_exchange_group(&self) -> Option { + self.snapshot.negotiated_key_exchange_group + } + + fn exporter(&mut self) -> Result { + match self.state.as_mut() { + Some(ServerHandshakeState::NeedsInput(_)) | None => Err(Error::HandshakeNotComplete), + Some(ServerHandshakeState::Complete(state)) => state.exporter(), + Some(ServerHandshakeState::Failed(error)) => Err(error.clone()), + } + } +} + +impl QuicConnection { + fn side(&self) -> Side { + match self { + Self::Client(_) => Side::Client, + Self::Server(_) | Self::ServerHandshake(_) => Side::Server, + } + } + + fn alpn_protocol(&self) -> Option<&[u8]> { + match self { + Self::Client(session) => session.alpn_protocol(), + Self::Server(session) => session.alpn_protocol(), + Self::ServerHandshake(session) => return session.alpn_protocol(), + } + .map(AsRef::as_ref) + } + + fn peer_identity(&self) -> Option<&Identity<'static>> { + match self { + Self::Client(session) => session.peer_identity(), + Self::Server(session) => session.peer_identity(), + Self::ServerHandshake(session) => return session.peer_identity(), + } + .map(|identity| identity.identity()) + } + + fn zero_rtt_keys(&self) -> Option { + match self { + Self::Client(session) => session.zero_rtt_keys(), + Self::Server(session) => session.zero_rtt_keys(), + Self::ServerHandshake(session) => session.zero_rtt_keys(), + } + } + + fn is_early_data_accepted(&self) -> Option { + match self { + Self::Client(session) => Some(session.is_early_data_accepted()), + Self::Server(_) | Self::ServerHandshake(_) => None, + } + } + + fn is_handshaking(&self) -> bool { + match self { + Self::Client(session) => session.is_handshaking(), + Self::Server(session) => session.is_handshaking(), + Self::ServerHandshake(session) => session.is_handshaking(), + } + } + + fn read_hs(&mut self, input: &mut dyn TlsInputBuffer) -> Result<(), Error> { + match self { + Self::Client(session) => session.read_hs(input), + Self::Server(session) => session.read_hs(input), + Self::ServerHandshake(session) => session.read_hs(input), + } + } + + fn drain_events(&mut self, events: &mut VecDeque) { + match self { + Self::Client(session) => events.extend(session.events()), + Self::Server(session) => events.extend(session.events()), + Self::ServerHandshake(session) => session.drain_events(events), + } + } + + fn quic_transport_parameters(&self) -> Option<&[u8]> { + match self { + Self::Client(session) => session.quic_transport_parameters(), + Self::Server(session) => session.quic_transport_parameters(), + Self::ServerHandshake(session) => session.quic_transport_parameters(), + } + } + + fn server_name(&self) -> Option<&str> { + match self { + Self::Client(_) => None, + Self::Server(session) => session.server_name().map(AsRef::as_ref), + Self::ServerHandshake(session) => session.server_name(), + } + } + + #[cfg(feature = "__rustls-post-quantum-test")] + fn negotiated_key_exchange_group(&self) -> Option { + match self { + Self::Client(session) => session.negotiated_key_exchange_group(), + Self::Server(session) => session.negotiated_key_exchange_group(), + Self::ServerHandshake(session) => return session.negotiated_key_exchange_group(), + } + .map(|group| group.name()) + } + + fn exporter(&mut self) -> Result { + match self { + Self::Client(session) => session.exporter(), + Self::Server(session) => session.exporter(), + Self::ServerHandshake(session) => session.exporter(), + } + } +} impl HeaderKey for Box { fn decrypt(&self, pn_offset: usize, packet: &mut [u8]) { @@ -288,7 +775,8 @@ pub struct HandshakeData { /// A QUIC-compatible TLS client configuration /// /// Quinn implicitly constructs a `QuicClientConfig` with reasonable defaults within -/// [`ClientConfig::with_root_certificates()`][root_certs] and [`ClientConfig::try_with_platform_verifier()`][platform]. +/// [`ClientConfig::with_root_certificates()`][root_certs] and +/// [`ClientConfig::try_with_platform_verifier()`][platform]. /// Alternatively, `QuicClientConfig`'s [`TryFrom`] implementation can be used to wrap around a /// custom [`rustls::ClientConfig`], in which case care should be taken around certain points: /// @@ -311,17 +799,15 @@ pub struct QuicClientConfig { impl QuicClientConfig { #[cfg(feature = "platform-verifier")] pub(crate) fn with_platform_verifier() -> Result { - // Keep in sync with `inner()` below - let mut inner = rustls::ClientConfig::builder_with_provider(configured_provider()) - .with_protocol_versions(&[&rustls::version::TLS13]) - .unwrap() // The default providers support TLS 1.3 + let mut inner = rustls::ClientConfig::builder(configured_provider()) .with_platform_verifier()? - .with_no_client_auth(); + .with_no_client_auth() + .expect("default providers are valid for QUIC"); inner.enable_early_data = true; Ok(Self { - // We're confident that the *ring* default provider contains TLS13_AES_128_GCM_SHA256 - initial: initial_suite_from_provider(inner.crypto_provider()) + // We're confident that the default providers contain TLS13_AES_128_GCM_SHA256 + initial: initial_suite_from_provider(inner.provider()) .expect("no initial cipher suite found"), inner: Arc::new(inner), }) @@ -331,11 +817,11 @@ impl QuicClientConfig { /// /// QUIC requires that TLS 1.3 be enabled. Advanced users can use any [`rustls::ClientConfig`] that /// satisfies this requirement. - pub(crate) fn new(verifier: Arc) -> Self { + pub(crate) fn new(verifier: Arc) -> Self { let inner = Self::inner(verifier); Self { - // We're confident that the *ring* default provider contains TLS13_AES_128_GCM_SHA256 - initial: initial_suite_from_provider(inner.crypto_provider()) + // We're confident that the default providers contain TLS13_AES_128_GCM_SHA256 + initial: initial_suite_from_provider(inner.provider()) .expect("no initial cipher suite found"), inner: Arc::new(inner), } @@ -348,20 +834,25 @@ impl QuicClientConfig { inner: Arc, initial: Suite, ) -> Result { - match initial.suite.common.suite { + match initial.inner.common.suite { CipherSuite::TLS13_AES_128_GCM_SHA256 => Ok(Self { inner, initial }), _ => Err(NoInitialCipherSuite { specific: true }), } } - pub(crate) fn inner(verifier: Arc) -> rustls::ClientConfig { - // Keep in sync with `with_platform_verifier()` above - let mut config = rustls::ClientConfig::builder_with_provider(configured_provider()) - .with_protocol_versions(&[&rustls::version::TLS13]) - .unwrap() // The default providers support TLS 1.3 + pub(crate) fn inner(verifier: Arc) -> rustls::ClientConfig { + Self::inner_with_provider(verifier, configured_provider()) + } + + pub(crate) fn inner_with_provider( + verifier: Arc, + provider: Arc, + ) -> rustls::ClientConfig { + let mut config = rustls::ClientConfig::builder(provider) .dangerous() .with_custom_certificate_verifier(verifier) - .with_no_client_auth(); + .with_no_client_auth() + .expect("default providers are valid for QUIC"); config.enable_early_data = true; config @@ -380,8 +871,9 @@ impl crypto::ClientConfig for QuicClientConfig { version, got_handshake_data: false, next_secrets: None, - inner: Connection::Client( - rustls::quic::ClientConnection::new( + exporter: None, + inner: QuicConnection::Client( + ClientConnection::new( self.inner.clone(), version, ServerName::try_from(server_name) @@ -391,6 +883,8 @@ impl crypto::ClientConfig for QuicClientConfig { ) .unwrap(), ), + input: HandshakeInput::default(), + pending_events: VecDeque::new(), suite: self.initial, })) } @@ -409,7 +903,7 @@ impl TryFrom> for QuicClientConfig { fn try_from(inner: Arc) -> Result { Ok(Self { - initial: initial_suite_from_provider(inner.crypto_provider()) + initial: initial_suite_from_provider(inner.provider()) .ok_or(NoInitialCipherSuite { specific: false })?, inner, }) @@ -420,9 +914,7 @@ impl TryFrom> for QuicClientConfig { /// /// When the cipher suite is supplied `with_initial()`, it must be /// [`CipherSuite::TLS13_AES_128_GCM_SHA256`]. When the cipher suite is derived from a config's -/// [`CryptoProvider`][provider], that provider must reference a cipher suite with the same ID. -/// -/// [provider]: rustls::crypto::CryptoProvider +/// [`CryptoProvider`], that provider must reference a cipher suite with the same ID. #[derive(Clone, Debug)] pub struct NoInitialCipherSuite { /// Whether the initial cipher suite was supplied by the caller @@ -464,8 +956,8 @@ impl QuicServerConfig { ) -> Result { let inner = Self::inner(cert_chain, key)?; Ok(Self { - // We're confident that the *ring* default provider contains TLS13_AES_128_GCM_SHA256 - initial: initial_suite_from_provider(inner.crypto_provider()) + // We're confident that the default providers contain TLS13_AES_128_GCM_SHA256 + initial: initial_suite_from_provider(inner.provider()) .expect("no initial cipher suite found"), inner: Arc::new(inner), }) @@ -478,7 +970,7 @@ impl QuicServerConfig { inner: Arc, initial: Suite, ) -> Result { - match initial.suite.common.suite { + match initial.inner.common.suite { CipherSuite::TLS13_AES_128_GCM_SHA256 => Ok(Self { inner, initial }), _ => Err(NoInitialCipherSuite { specific: true }), } @@ -493,11 +985,17 @@ impl QuicServerConfig { cert_chain: Vec>, key: PrivateKeyDer<'static>, ) -> Result { - let mut inner = rustls::ServerConfig::builder_with_provider(configured_provider()) - .with_protocol_versions(&[&rustls::version::TLS13]) - .unwrap() // The *ring* default provider supports TLS 1.3 + Self::inner_with_provider(cert_chain, key, configured_provider()) + } + + pub(crate) fn inner_with_provider( + cert_chain: Vec>, + key: PrivateKeyDer<'static>, + provider: Arc, + ) -> Result { + let mut inner = rustls::ServerConfig::builder(provider) .with_no_client_auth() - .with_single_cert(cert_chain, key)?; + .with_single_cert(Arc::new(Identity::from_cert_chain(cert_chain)?), key)?; inner.max_early_data_size = u32::MAX; Ok(inner) @@ -517,7 +1015,7 @@ impl TryFrom> for QuicServerConfig { fn try_from(inner: Arc) -> Result { Ok(Self { - initial: initial_suite_from_provider(inner.crypto_provider()) + initial: initial_suite_from_provider(inner.provider()) .ok_or(NoInitialCipherSuite { specific: false })?, inner, }) @@ -536,14 +1034,49 @@ impl crypto::ServerConfig for QuicServerConfig { version, got_handshake_data: false, next_secrets: None, - inner: Connection::Server( - rustls::quic::ServerConnection::new(self.inner.clone(), version, to_vec(params)) - .unwrap(), + exporter: None, + inner: QuicConnection::Server( + ServerConnection::new(self.inner.clone(), version, to_vec(params)).unwrap(), ), + input: HandshakeInput::default(), + pending_events: VecDeque::new(), suite: self.initial, }) } + fn start_session_from_accepted( + self: Arc, + version: u32, + params: &TransportParameters, + accepted: Accepted, + ) -> Result, TransportError> { + // Safe: `start_session_from_accepted()` is never called if `initial_keys()` rejected + // `version`. + let version = interpret_version(version).unwrap(); + let Accepted { inner, input, .. } = accepted; + let mut events = Vec::new(); + let state = inner + .choose_config(self.inner.clone(), to_vec(params), &mut events) + .map_err(transport_error_from_rustls)?; + let inner = + ServerHandshakeConnection::new(state, events).map_err(transport_error_from_rustls)?; + + let mut session = TlsSession { + version, + got_handshake_data: false, + next_secrets: None, + exporter: None, + inner: QuicConnection::ServerHandshake(inner), + input, + pending_events: VecDeque::new(), + suite: self.initial, + }; + // The staged acceptor already consumed ClientHello, but HelloRetryRequest can require + // another ClientHello before all required handshake data is available. + session.got_handshake_data = session.required_handshake_data_is_ready(); + Ok(Box::new(session)) + } + fn initial_keys( &self, version: u32, @@ -556,22 +1089,13 @@ impl crypto::ServerConfig for QuicServerConfig { fn retry_tag(&self, version: u32, orig_dst_cid: ConnectionId, packet: &[u8]) -> [u8; 16] { // Safe: `start_session()` is never called if `initial_keys()` rejected `version` let version = interpret_version(version).unwrap(); - let (nonce, key) = match version { - Version::V1 => (RETRY_INTEGRITY_NONCE_V1, RETRY_INTEGRITY_KEY_V1), - Version::V1Draft => (RETRY_INTEGRITY_NONCE_DRAFT, RETRY_INTEGRITY_KEY_DRAFT), - _ => unreachable!(), - }; - let mut pseudo_packet = Vec::with_capacity(packet.len() + orig_dst_cid.len() + 1); pseudo_packet.push(orig_dst_cid.len() as u8); pseudo_packet.extend_from_slice(&orig_dst_cid); pseudo_packet.extend_from_slice(packet); - let nonce = aead::Nonce::assume_unique_for_key(nonce); - let key = aead::LessSafeKey::new(aead::UnboundKey::new(&aead::AES_128_GCM, &key).unwrap()); - - let tag = key - .seal_in_place_separate_tag(nonce, aead::Aad::from(pseudo_packet), &mut []) + let tag = retry_key_for_version(version, &self.initial) + .encrypt_in_place(0, &pseudo_packet, &mut [], None) .unwrap(); let mut result = [0; 16]; result.copy_from_slice(tag.as_ref()); @@ -579,24 +1103,39 @@ impl crypto::ServerConfig for QuicServerConfig { } } -pub(crate) fn initial_suite_from_provider( - provider: &Arc, -) -> Option { +fn retry_key_for_version(version: Version, initial_suite: &Suite) -> Box { + let (nonce, key) = match version { + Version::V1 => (RETRY_INTEGRITY_NONCE_V1, RETRY_INTEGRITY_KEY_V1), + _ => unreachable!(), + }; + + initial_suite + .quic + .packet_key(AeadKey::from(key), Iv::from(nonce)) +} + +const RETRY_INTEGRITY_KEY_V1: [u8; 16] = [ + 0xbe, 0x0c, 0x69, 0x0b, 0x9f, 0x66, 0x57, 0x5a, 0x1d, 0x76, 0x6b, 0x54, 0xe3, 0x68, 0xc8, 0x4e, +]; +const RETRY_INTEGRITY_NONCE_V1: [u8; 12] = [ + 0x46, 0x15, 0x99, 0xd3, 0x5d, 0x63, 0x2b, 0xf2, 0x23, 0x98, 0x25, 0xbb, +]; + +pub(crate) fn initial_suite_from_provider(provider: &Arc) -> Option { provider - .cipher_suites + .tls13_cipher_suites .iter() - .find_map(|cs| match (cs.suite(), cs.tls13()) { - (CipherSuite::TLS13_AES_128_GCM_SHA256, Some(suite)) => Some(suite.quic_suite()), + .find_map(|&suite| match suite.common.suite { + CipherSuite::TLS13_AES_128_GCM_SHA256 => Suite::try_from(suite).ok(), _ => None, }) - .flatten() } -pub(crate) fn configured_provider() -> Arc { +pub(crate) fn configured_provider() -> Arc { #[cfg(all(feature = "rustls-aws-lc-rs", not(feature = "rustls-ring")))] - let provider = rustls::crypto::aws_lc_rs::default_provider(); + let provider = rustls_aws_lc_rs::DEFAULT_PROVIDER; #[cfg(feature = "rustls-ring")] - let provider = rustls::crypto::ring::default_provider(); + let provider = rustls_ring::DEFAULT_PROVIDER; Arc::new(provider) } @@ -629,7 +1168,9 @@ impl crypto::PacketKey for Box { fn encrypt(&self, packet: u64, buf: &mut [u8], header_len: usize) { let (header, payload_tag) = buf.split_at_mut(header_len); let (payload, tag_storage) = payload_tag.split_at_mut(payload_tag.len() - self.tag_len()); - let tag = self.encrypt_in_place(packet, &*header, payload).unwrap(); + let tag = self + .encrypt_in_place(packet, &*header, payload, None) + .unwrap(); tag_storage.copy_from_slice(tag.as_ref()); } @@ -640,7 +1181,7 @@ impl crypto::PacketKey for Box { payload: &mut BytesMut, ) -> Result<(), CryptoError> { let plain = self - .decrypt_in_place(packet, header, payload.as_mut()) + .decrypt_in_place(packet, header, payload.as_mut(), None) .map_err(|_| CryptoError)?; let plain_len = plain.len(); payload.truncate(plain_len); @@ -662,8 +1203,53 @@ impl crypto::PacketKey for Box { fn interpret_version(version: u32) -> Result { match version { - 0xff00_001d..=0xff00_0020 => Ok(Version::V1Draft), 0x0000_0001 | 0xff00_0021..=0xff00_0022 => Ok(Version::V1), _ => Err(UnsupportedVersion), } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn failed_server_handshake_is_a_stable_terminal_state() { + let mut connection = ServerHandshakeConnection { + state: None, + pending_events: VecDeque::from([QuicEvent::Message(vec![1, 2, 3])]), + snapshot: ServerHandshakeSnapshot { + alpn_protocol: Some(b"test".to_vec()), + quic_transport_parameters: Some(vec![4, 5, 6]), + server_name: Some("localhost".into()), + protocol_version: Some(rustls::enums::ProtocolVersion::TLSv1_3), + ..ServerHandshakeSnapshot::default() + }, + }; + let error = Error::General("fatal test error".into()); + + assert_eq!(connection.fail::<()>(error.clone()), Err(error.clone())); + assert!(!connection.is_handshaking()); + assert_eq!(connection.alpn_protocol(), Some(b"test".as_slice())); + assert_eq!(connection.server_name(), Some("localhost")); + assert_eq!( + connection.protocol_version(), + Some(rustls::enums::ProtocolVersion::TLSv1_3) + ); + assert_eq!( + connection.quic_transport_parameters(), + Some([4, 5, 6].as_slice()) + ); + assert!(connection.peer_identity().is_none()); + assert!(connection.zero_rtt_keys().is_none()); + match connection.exporter() { + Err(actual) => assert_eq!(actual, error), + Ok(_) => panic!("failed handshake unexpectedly produced an exporter"), + } + + let mut input = HandshakeInput::default(); + assert_eq!(connection.read_hs(&mut input), Err(error)); + let mut events = VecDeque::new(); + connection.drain_events(&mut events); + assert!(events.is_empty()); + } +} diff --git a/quinn-proto/src/endpoint.rs b/quinn-proto/src/endpoint.rs index 529a5bbaec..1a75f3d576 100644 --- a/quinn-proto/src/endpoint.rs +++ b/quinn-proto/src/endpoint.rs @@ -7,6 +7,9 @@ use std::{ sync::Arc, }; +#[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] +use std::collections::VecDeque; + use bytes::{BufMut, Bytes, BytesMut}; use rand::{ Rng, RngExt, SeedableRng, @@ -38,6 +41,9 @@ use crate::{ transport_parameters::{PreferredAddress, TransportParameters}, }; +#[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] +use crate::{connection::assembler::Assembler, range_set::ArrayRangeSet}; + /// The main entry point to the library /// /// This object performs no I/O whatsoever. Instead, it consumes incoming packets and @@ -46,6 +52,7 @@ pub struct Endpoint { rng: StdRng, index: ConnectionIndex, connections: Slab, + pending_accepts: usize, local_cid_generator: Box, config: Arc, server_config: Option>, @@ -79,6 +86,7 @@ impl Endpoint { }, index: ConnectionIndex::default(), connections: Slab::new(), + pending_accepts: 0, local_cid_generator: (config.connection_id_generator_factory.as_ref())(), config, server_config, @@ -210,20 +218,29 @@ impl Endpoint { match route_to { RouteDatagramTo::Incoming(incoming_idx) => { let incoming_buffer = &mut self.incoming_buffers[incoming_idx]; - let config = &self.server_config.as_ref().unwrap(); + + if incoming_buffer.addresses != addresses { + debug!( + remote = %addresses.remote, + expected = %incoming_buffer.addresses.remote, + "discarding pending incoming datagram from unexpected path" + ); + return None; + } if incoming_buffer .total_bytes .checked_add(datagram_len as u64) - .is_some_and(|n| n <= config.incoming_buffer_size) + .is_some_and(|n| n <= incoming_buffer.size_limit) && self .all_incoming_buffers_total_bytes .checked_add(datagram_len as u64) - .is_some_and(|n| n <= config.incoming_buffer_size_total) + .is_some_and(|n| n <= incoming_buffer.total_size_limit) { incoming_buffer.datagrams.push(event); incoming_buffer.total_bytes += datagram_len as u64; self.all_incoming_buffers_total_bytes += datagram_len as u64; + return Some(DatagramEvent::IncomingData(incoming_idx)); } None @@ -505,7 +522,13 @@ impl Endpoint { } }; - let incoming_idx = self.incoming_buffers.insert(IncomingBuffer::default()); + let incoming_idx = self.incoming_buffers.insert(IncomingBuffer { + addresses, + datagrams: Vec::new(), + total_bytes: 0, + size_limit: server_config.incoming_buffer_size, + total_size_limit: server_config.incoming_buffer_size_total, + }); self.index .insert_initial_incoming(header.dst_cid, incoming_idx); @@ -530,15 +553,31 @@ impl Endpoint { // box err to avoid clippy::result_large_err pub fn accept( &mut self, - mut incoming: Incoming, + incoming: Incoming, now: Instant, buf: &mut Vec, server_config: Option>, ) -> Result<(ConnectionHandle, Connection), Box> { + let accepting = self.start_accept(incoming, now, buf, server_config)?; + match accepting.finish_without_endpoint() { + Ok(accepted) => Ok(self.finish_accept(accepted)), + Err(error) => Err(self.finish_accept_error(error, buf)), + } + } + + /// First phase of connection acceptance: everything that requires `&mut Endpoint`. + /// Reserves CIDs and routing state, but does NOT create the connection, process the first + /// packet, or replay buffered datagrams. This is the minimum work that must happen under the + /// endpoint lock. + #[doc(hidden)] + pub fn start_accept( + &mut self, + mut incoming: Incoming, + now: Instant, + buf: &mut Vec, + server_config: Option>, + ) -> Result> { let remote_address_validated = incoming.remote_address_validated(); - incoming.improper_drop_warner.dismiss(); - let incoming_buffer = self.incoming_buffers.remove(incoming.incoming_idx); - self.all_incoming_buffers_total_bytes -= incoming_buffer.total_bytes; let packet_number = incoming.packet.header.number.expand(0); let InitialHeader { @@ -554,30 +593,57 @@ impl Endpoint { .transport .max_idle_timeout .is_some_and(|timeout| { - incoming.received_at + Duration::from_millis(timeout.into()) <= now + timeout.into_inner() != 0 + && incoming.received_at + Duration::from_millis(timeout.into()) <= now }) { debug!("abandoning accept of stale initial"); - self.index.remove_initial(dst_cid); + self.ignore(incoming); return Err(Box::new(AcceptError { cause: ConnectionError::TimedOut, response: None, })); } + if self.local_cid_generator.cid_len() == 0 + && self + .index + .incoming_connection_remotes + .contains_key(&incoming.addresses) + { + debug!( + remote = %incoming.addresses.remote, + "refusing connection because its zero-length CID route is already in use" + ); + let response = self.initial_close( + version, + incoming.addresses, + &incoming.crypto, + src_cid, + TransportError::CONNECTION_REFUSED("zero-length connection ID route in use"), + buf, + ); + self.ignore(incoming); + return Err(Box::new(AcceptError { + cause: ConnectionError::CidsExhausted, + response: Some(response), + })); + } + if self.cids_exhausted() { debug!("refusing connection"); - self.index.remove_initial(dst_cid); + let response = self.initial_close( + version, + incoming.addresses, + &incoming.crypto, + src_cid, + TransportError::CONNECTION_REFUSED(""), + buf, + ); + self.ignore(incoming); return Err(Box::new(AcceptError { cause: ConnectionError::CidsExhausted, - response: Some(self.initial_close( - version, - incoming.addresses, - &incoming.crypto, - src_cid, - TransportError::CONNECTION_REFUSED(""), - buf, - )), + response: Some(response), })); } @@ -593,15 +659,15 @@ impl Endpoint { .is_err() { debug!(packet_number, "failed to authenticate initial packet"); - self.index.remove_initial(dst_cid); + self.ignore(incoming); return Err(Box::new(AcceptError { cause: TransportError::PROTOCOL_VIOLATION("authentication failed").into(), response: None, })); }; - let ch = ConnectionHandle(self.connections.vacant_key()); - let loc_cid = self.new_cid(RouteDatagramTo::Connection(ch)); + let accepting_idx = incoming.incoming_idx; + let loc_cid = self.new_cid(RouteDatagramTo::Incoming(accepting_idx)); let mut params = TransportParameters::new( &server_config.transport, &self.config, @@ -615,7 +681,7 @@ impl Endpoint { params.retry_src_cid = incoming.token.retry_src_cid; let mut pref_addr_cid = None; if server_config.has_preferred_address() { - let cid = self.new_cid(RouteDatagramTo::Connection(ch)); + let cid = self.new_cid(RouteDatagramTo::Incoming(accepting_idx)); pref_addr_cid = Some(cid); params.preferred_address = Some(PreferredAddress { address_v4: server_config.preferred_address_v4, @@ -625,60 +691,212 @@ impl Endpoint { }); } - let tls = server_config.crypto.clone().start_session(version, ¶ms); - let transport_config = server_config.transport.clone(); - let mut conn = self.add_connection( - ch, - version, - dst_cid, + // Gather everything needed to create the Connection outside the lock. + let mut rng_seed = [0; 32]; + self.rng.fill_bytes(&mut rng_seed); + let endpoint_config = self.config.clone(); + let cid_len = self.local_cid_generator.cid_len(); + let cid_lifetime = self.local_cid_generator.cid_lifetime(); + let allow_mtud = self.allow_mtud; + + let reservation = AcceptReservation { + incoming_idx: accepting_idx, + init_cid: dst_cid, + addresses: incoming.addresses, loc_cid, - src_cid, - incoming.addresses, - incoming.received_at, - tls, - transport_config, - SideArgs::Server { - server_config, - pref_addr_cid, - path_validated: remote_address_validated, - }, - ); - self.index.insert_initial(dst_cid, ch); + pref_addr_cid, + }; + if loc_cid.is_empty() { + let previous = self + .index + .incoming_connection_remotes + .insert(incoming.addresses, RouteDatagramTo::Incoming(accepting_idx)); + debug_assert!(previous.is_none()); + } + self.pending_accepts += 1; - match conn.handle_first_packet( - incoming.received_at, - incoming.addresses.remote, - incoming.ecn, + Ok(Accepting { + reservation, + version, + src_cid, packet_number, - incoming.packet, - incoming.rest, - ) { - Ok(()) => { - trace!(id = ch.0, icid = %dst_cid, "new connection"); + last_activity: incoming.received_at, + incoming, + // Deferred connection creation state + server_config, + params, + remote_address_validated, + rng_seed, + endpoint_config, + cid_len, + cid_lifetime, + allow_mtud, + initial_sends: InitialSendState::default(), + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + rustls_buffered_datagrams: 0, + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + rustls_input: VecDeque::new(), + }) + } - for event in incoming_buffer.datagrams { - conn.handle_event(ConnectionEvent(ConnectionEventInner::Datagram(event))) - } + /// Copy newly buffered datagrams into a split accept for processing outside the endpoint lock. + /// + /// Cloning a datagram only clones its reference-counted byte buffers. Authentication, frame + /// parsing, rustls processing, and response encoding are deferred to + /// [`Accepting::poll_rustls_acceptor`]. + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + #[doc(hidden)] + pub fn buffer_rustls_acceptor_input(&self, accepting: &mut Accepting) { + let incoming_buffer = &self.incoming_buffers[accepting.reservation.incoming_idx]; + accepting.rustls_input.extend( + incoming_buffer + .datagrams + .iter() + .skip(accepting.rustls_buffered_datagrams) + .cloned(), + ); + accepting.rustls_buffered_datagrams = incoming_buffer.datagrams.len(); + } + + /// Replace the config captured by `start_accept` after ClientHello-based selection. + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + #[doc(hidden)] + pub fn select_accepting_config( + &mut self, + accepting: &mut Accepting, + server_config: Option>, + now: Instant, + ) -> Result<(), ConnectionError> { + let server_config = server_config.unwrap_or_else(|| accepting.server_config.clone()); + if server_config + .transport + .max_idle_timeout + .is_some_and(|timeout| { + timeout.into_inner() != 0 + && accepting.last_activity + Duration::from_millis(timeout.into()) <= now + }) + { + return Err(ConnectionError::TimedOut); + } - Ok((ch, conn)) + let wants_preferred_address = server_config.has_preferred_address(); + match (accepting.reservation.pref_addr_cid, wants_preferred_address) { + (Some(cid), false) => { + self.index.retire(cid); + accepting.reservation.pref_addr_cid = None; } - Err(e) => { - debug!("handshake failed: {}", e); - self.handle_event(ch, EndpointEvent(EndpointEventInner::Drained)); - let response = match e { - ConnectionError::TransportError(ref e) => Some(self.initial_close( - version, - incoming.addresses, - &incoming.crypto, - src_cid, - e.clone(), - buf, - )), - _ => None, - }; - Err(Box::new(AcceptError { cause: e, response })) + (None, true) => { + if self.cids_exhausted() { + return Err(ConnectionError::CidsExhausted); + } + accepting.reservation.pref_addr_cid = Some(self.new_cid( + RouteDatagramTo::Incoming(accepting.reservation.incoming_idx), + )); } + _ => {} + } + + let mut params = TransportParameters::new( + &server_config.transport, + &self.config, + self.local_cid_generator.as_ref(), + accepting.reservation.loc_cid, + Some(&server_config), + &mut self.rng, + ); + params.stateless_reset_token = Some(ResetToken::new( + &*self.config.reset_key, + accepting.reservation.loc_cid, + )); + params.original_dst_cid = Some(accepting.incoming.token.orig_dst_cid); + params.retry_src_cid = accepting.incoming.token.retry_src_cid; + if let Some(cid) = accepting.reservation.pref_addr_cid { + params.preferred_address = Some(PreferredAddress { + address_v4: server_config.preferred_address_v4, + address_v6: server_config.preferred_address_v6, + connection_id: cid, + stateless_reset_token: ResetToken::new(&*self.config.reset_key, cid), + }); } + + accepting.server_config = server_config; + accepting.params = params; + Ok(()) + } + + #[doc(hidden)] + pub fn finish_accept(&mut self, accepted: Accepted) -> (ConnectionHandle, Connection) { + let Accepted { + reservation, + mut conn, + guard, + } = accepted; + guard.dismiss(); + let accepting_buffer = self.remove_accept_reservation(&reservation); + let ch = ConnectionHandle(self.connections.vacant_key()); + self.register_connection( + ch, + reservation.init_cid, + reservation.loc_cid, + reservation.pref_addr_cid, + reservation.addresses, + Side::Server, + ); + trace!(id = ch.0, icid = %reservation.init_cid, "new connection"); + + for event in accepting_buffer.datagrams { + conn.handle_event(ConnectionEvent(ConnectionEventInner::Datagram(event))) + } + + (ch, conn) + } + + /// Clean up after a failed [`Accepting::finish_without_endpoint`] and optionally generate a + /// close response. + #[doc(hidden)] + pub fn finish_accept_error( + &mut self, + error: Box, + buf: &mut Vec, + ) -> Box { + let AcceptingError { + cause, + reservation, + version, + src_cid, + crypto, + initial_sends, + guard, + } = *error; + guard.dismiss(); + debug!("handshake failed: {}", cause); + let response = match cause { + ConnectionError::TransportError(ref e) => Some(Self::initial_close_with( + version, + reservation.addresses, + &crypto, + src_cid, + reservation.loc_cid, + initial_sends.next_packet_number, + e.clone(), + buf, + )), + _ => None, + }; + self.remove_accept_reservation(&reservation); + Box::new(AcceptError { cause, response }) + } + + /// Clean up a staged rustls accept before it has created a connection. + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + #[doc(hidden)] + pub fn fail_accepting( + &mut self, + accepting: Accepting, + cause: ConnectionError, + buf: &mut Vec, + ) -> Box { + self.finish_accept_error(Box::new(accepting.into_error(cause)), buf) } /// Check if we should refuse a connection attempt regardless of the packet's contents @@ -708,7 +926,7 @@ impl Endpoint { /// Reject this incoming connection attempt pub fn refuse(&mut self, incoming: Incoming, buf: &mut Vec) -> Transmit { - self.clean_up_incoming(&incoming); + self.remove_incoming_state(&incoming); incoming.improper_drop_warner.dismiss(); self.initial_close( @@ -729,7 +947,7 @@ impl Endpoint { return Err(RetryError(Box::new(incoming))); } - self.clean_up_incoming(&incoming); + self.remove_incoming_state(&incoming); incoming.improper_drop_warner.dismiss(); let server_config = self.server_config.as_ref().unwrap(); @@ -778,15 +996,44 @@ impl Endpoint { /// Doing this actively, rather than merely dropping the [`Incoming`], is necessary to prevent /// memory leaks due to state within [`Endpoint`] tracking the incoming connection. pub fn ignore(&mut self, incoming: Incoming) { - self.clean_up_incoming(&incoming); + self.remove_incoming_state(&incoming); incoming.improper_drop_warner.dismiss(); } - /// Clean up endpoint data structures associated with an `Incoming`. - fn clean_up_incoming(&mut self, incoming: &Incoming) { + /// Remove endpoint state associated with an `Incoming`. + fn remove_incoming_state(&mut self, incoming: &Incoming) { self.index.remove_initial(incoming.packet.header.dst_cid); - let incoming_buffer = self.incoming_buffers.remove(incoming.incoming_idx); + self.remove_incoming_buffer(incoming.incoming_idx); + } + + fn remove_accept_reservation(&mut self, reservation: &AcceptReservation) -> IncomingBuffer { + self.index.remove_initial(reservation.init_cid); + if reservation.loc_cid.is_empty() + && matches!( + self.index + .incoming_connection_remotes + .get(&reservation.addresses), + Some(RouteDatagramTo::Incoming(incoming_idx)) + if *incoming_idx == reservation.incoming_idx + ) + { + self.index + .incoming_connection_remotes + .remove(&reservation.addresses); + } + self.index.retire(reservation.loc_cid); + if let Some(cid) = reservation.pref_addr_cid { + self.index.retire(cid); + } + debug_assert!(self.pending_accepts > 0); + self.pending_accepts -= 1; + self.remove_incoming_buffer(reservation.incoming_idx) + } + + fn remove_incoming_buffer(&mut self, incoming_idx: usize) -> IncomingBuffer { + let incoming_buffer = self.incoming_buffers.remove(incoming_idx); self.all_incoming_buffers_total_bytes -= incoming_buffer.total_bytes; + incoming_buffer } fn add_connection( @@ -861,7 +1108,30 @@ impl Endpoint { }); debug_assert_eq!(id, ch.0, "connection handle allocation out of sync"); - self.index.insert_conn(addresses, loc_cid, ch, side); + let conn_meta = &self.connections[ch]; + if conn_meta.side.is_server() { + self.index.insert_initial(conn_meta.init_cid, ch); + } + for cid in conn_meta.loc_cids.values() { + if cid.is_empty() { + match conn_meta.side { + Side::Server => { + self.index + .incoming_connection_remotes + .insert(conn_meta.addresses, RouteDatagramTo::Connection(ch)); + } + Side::Client => { + self.index + .outgoing_connection_remotes + .insert(conn_meta.addresses.remote, ch); + } + } + } else { + self.index + .connection_ids + .insert(*cid, RouteDatagramTo::Connection(ch)); + } + } } fn initial_close( @@ -877,7 +1147,23 @@ impl Endpoint { // shouldn't respond, and if it does, and the CID collides, we'll just drop the // unexpected response. let local_id = self.local_cid_generator.generate_cid(); - let number = PacketNumber::U8(0); + Self::initial_close_with( + version, addresses, crypto, remote_id, local_id, 0, reason, buf, + ) + } + + #[allow(clippy::too_many_arguments)] + fn initial_close_with( + version: u32, + addresses: FourTuple, + crypto: &Keys, + remote_id: ConnectionId, + local_id: ConnectionId, + packet_number: u64, + reason: TransportError, + buf: &mut Vec, + ) -> Transmit { + let number = PacketNumber::new(packet_number, 0); let header = Header::Initial(InitialHeader { dst_cid: remote_id, src_cid: local_id, @@ -891,7 +1177,11 @@ impl Endpoint { INITIAL_MTU as usize - partial_encode.header_len - crypto.packet.local.tag_len(); frame::Close::from(reason).encode(buf, max_len); buf.resize(buf.len() + crypto.packet.local.tag_len(), 0); - partial_encode.finish(buf, &*crypto.header.local, Some((0, &*crypto.packet.local))); + partial_encode.finish( + buf, + &*crypto.header.local, + Some((packet_number, &*crypto.packet.local)), + ); Transmit { destination: addresses.remote, ecn: None, @@ -911,8 +1201,16 @@ impl Endpoint { self.connections.len() } + /// Number of incoming accepts that have reserved endpoint state but have not yet been + /// finalized into active connections. + #[doc(hidden)] + pub fn pending_accepts(&self) -> usize { + self.pending_accepts + } + /// Counter for the number of bytes currently used /// in the buffers for Initial and 0-RTT messages for pending incoming connections + /// and accepts that are still being finalized pub fn incoming_buffer_bytes(&self) -> u64 { self.all_incoming_buffers_total_bytes } @@ -924,7 +1222,14 @@ impl Endpoint { // Not all connections have known reset tokens debug_assert!(x >= self.index.connection_reset_tokens.0.len()); // Not all connections have unique remotes, and 0-length CIDs might not be in use. - debug_assert!(x >= self.index.incoming_connection_remotes.len()); + debug_assert!( + x >= self + .index + .incoming_connection_remotes + .values() + .filter(|route| matches!(route, RouteDatagramTo::Connection(_))) + .count() + ); debug_assert!(x >= self.index.outgoing_connection_remotes.len()); x } @@ -960,6 +1265,7 @@ impl fmt::Debug for Endpoint { .field("rng", &self.rng) .field("index", &self.index) .field("connections", &self.connections) + .field("pending_accepts", &self.pending_accepts) .field("config", &self.config) .field("server_config", &self.server_config) // incoming_buffers too large @@ -973,10 +1279,14 @@ impl fmt::Debug for Endpoint { } /// Buffered Initial and 0-RTT messages for a pending incoming connection -#[derive(Default)] struct IncomingBuffer { + addresses: FourTuple, datagrams: Vec, total_bytes: u64, + /// Limits captured when the attempt first arrived, so replacing the endpoint's server + /// configuration only affects new incoming attempts. + size_limit: u64, + total_size_limit: u64, } /// Part of protocol state incoming datagrams can be routed to @@ -1002,7 +1312,7 @@ struct ConnectionIndex { /// Identifies incoming connections with zero-length CIDs /// /// Uses a standard `HashMap` to protect against hash collision attacks. - incoming_connection_remotes: HashMap, + incoming_connection_remotes: HashMap, /// Identifies outgoing connections with zero-length CIDs /// /// We don't yet support explicit source addresses for client connections, and zero-length CIDs @@ -1047,33 +1357,6 @@ impl ConnectionIndex { .insert(dst_cid, RouteDatagramTo::Connection(connection)); } - /// Associate a connection with its first locally-chosen destination CID if used, or otherwise - /// its current 4-tuple - fn insert_conn( - &mut self, - addresses: FourTuple, - dst_cid: ConnectionId, - connection: ConnectionHandle, - side: Side, - ) { - match dst_cid.len() { - 0 => match side { - Side::Server => { - self.incoming_connection_remotes - .insert(addresses, connection); - } - Side::Client => { - self.outgoing_connection_remotes - .insert(addresses.remote, connection); - } - }, - _ => { - self.connection_ids - .insert(dst_cid, RouteDatagramTo::Connection(connection)); - } - } - } - /// Discard a connection ID fn retire(&mut self, dst_cid: ConnectionId) { self.connection_ids.remove(&dst_cid); @@ -1108,8 +1391,8 @@ impl ConnectionIndex { } } if datagram.dst_cid().is_empty() { - if let Some(&ch) = self.incoming_connection_remotes.get(addresses) { - return Some(RouteDatagramTo::Connection(ch)); + if let Some(&route) = self.incoming_connection_remotes.get(addresses) { + return Some(route); } if let Some(&ch) = self.outgoing_connection_remotes.get(&addresses.remote) { return Some(RouteDatagramTo::Connection(ch)); @@ -1170,6 +1453,11 @@ impl IndexMut for Slab { pub enum DatagramEvent { /// The datagram is redirected to its `Connection` ConnectionEvent(ConnectionHandle, ConnectionEvent), + /// The datagram was buffered for an incoming connection whose acceptance is in progress + /// + /// The contained value identifies the pending accept and can be used to wake its exact waiter. + #[doc(hidden)] + IncomingData(usize), /// The datagram may result in starting a new `Connection` NewConnection(Incoming), /// Response generated directly by the endpoint @@ -1258,6 +1546,29 @@ impl Drop for IncomingImproperDropWarner { } } +/// Warns if a finalized split-accept state (`Accepted`/`AcceptingError`) is dropped without being +/// passed back to `Endpoint::finish_accept`/`finish_accept_error`. Doing so leaks the reserved +/// CIDs, the buffered Initial/0-RTT slot, and the pending-accept count, since that cleanup needs +/// the endpoint and cannot run from `Drop`. The earlier `Accepting` state is instead covered by +/// the `Incoming` it still holds, whose own warner stays armed until `finish_without_endpoint`. +struct AcceptDropGuard; + +impl AcceptDropGuard { + fn dismiss(self) { + mem::forget(self); + } +} + +impl Drop for AcceptDropGuard { + fn drop(&mut self) { + warn!( + "quinn_proto split-accept state dropped without passing to \ + Endpoint::finish_accept/finish_accept_error (leaks reserved CIDs, buffered packets, \ + and the pending-accept slot)" + ); + } +} + /// Errors in the parameters being used to create a new connection /// /// These arise before any I/O has been performed. @@ -1300,6 +1611,659 @@ pub struct AcceptError { pub response: Option, } +#[derive(Copy, Clone, Debug)] +struct AcceptReservation { + incoming_idx: usize, + init_cid: ConnectionId, + addresses: FourTuple, + loc_cid: ConnectionId, + pref_addr_cid: Option, +} + +#[derive(Copy, Clone, Debug, Default)] +struct InitialSendState { + next_packet_number: u64, + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + bytes: u64, + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + datagrams: u64, +} + +/// Internal split-accept success state used by `quinn`. +#[doc(hidden)] +#[allow(unnameable_types)] // internal split-accept API; re-exported only with __internal_split_accept +pub struct Accepted { + reservation: AcceptReservation, + conn: Connection, + guard: AcceptDropGuard, +} + +/// State for reading a rustls QUIC ClientHello before choosing a server config. +#[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] +#[allow(unnameable_types)] // internal split-accept API +pub struct RustlsAcceptor { + tls: crypto::rustls::Acceptor, + crypto_stream: Assembler, + processed_first_packet: bool, + processed_first_datagram_tail: bool, + rx_packet: u64, + authenticated_initials: ArrayRangeSet, + authenticated_initial_count: u64, + pending_initial_acks: ArrayRangeSet, + ack_pending: bool, +} + +#[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] +impl RustlsAcceptor { + /// Create a rustls acceptor for an incoming connection. + pub fn new(incoming: &Incoming) -> Result { + Ok(Self { + tls: crypto::rustls::Acceptor::new(incoming.packet.header.version)?, + crypto_stream: Assembler::new(), + processed_first_packet: false, + processed_first_datagram_tail: false, + rx_packet: 0, + authenticated_initials: ArrayRangeSet::new(), + authenticated_initial_count: 0, + pending_initial_acks: ArrayRangeSet::new(), + ack_pending: false, + }) + } + + #[cfg(test)] + pub(crate) fn test_read_initial_payload( + &mut self, + payload: Bytes, + packet_number: u64, + ) -> Result<(), ConnectionError> { + let payload_len = payload.len(); + self.read_initial_payload(payload, payload_len, usize::MAX, packet_number)?; + Ok(()) + } + + #[cfg(test)] + pub(crate) fn test_ack_pending(&self) -> bool { + self.ack_pending + } + + #[cfg(test)] + pub(crate) fn test_read_initial_decode( + &mut self, + partial_decode: PartialDecode, + crypto: &Keys, + expected_src_cid: ConnectionId, + expected_token: Bytes, + expected_version: u32, + ) -> Result { + Ok(self + .read_initial_decode( + partial_decode, + crypto, + usize::MAX, + expected_src_cid, + &expected_token, + expected_version, + )? + .is_some()) + } + + fn read_initial_datagram( + &mut self, + first_decode: Option, + mut remaining: Option, + crypto: &Keys, + crypto_buffer_size: usize, + cid_len: usize, + version: u32, + grease_quic_bit: bool, + expected_src_cid: ConnectionId, + expected_token: &Bytes, + ) -> Result, ConnectionError> { + if let Some(partial_decode) = first_decode { + if partial_decode.initial_header().is_some() { + if let Some(accepted) = self.read_initial_decode( + partial_decode, + crypto, + crypto_buffer_size, + expected_src_cid, + expected_token, + version, + )? { + return Ok(Some(accepted)); + } + } + } + + while let Some(data) = remaining { + let Ok((partial_decode, rest)) = PartialDecode::new( + data, + &FixedLengthConnectionIdParser::new(cid_len), + &[version], + grease_quic_bit, + ) else { + break; + }; + remaining = rest; + if partial_decode.initial_header().is_some() { + if let Some(accepted) = self.read_initial_decode( + partial_decode, + crypto, + crypto_buffer_size, + expected_src_cid, + expected_token, + version, + )? { + return Ok(Some(accepted)); + } + } + } + + Ok(None) + } + + fn read_initial_decode( + &mut self, + partial_decode: PartialDecode, + crypto: &Keys, + crypto_buffer_size: usize, + expected_src_cid: ConnectionId, + expected_token: &Bytes, + expected_version: u32, + ) -> Result, ConnectionError> { + let Ok(packet) = partial_decode.finish(Some(&*crypto.header.remote)) else { + return Ok(None); + }; + if !packet.reserved_bits_valid() { + return Ok(None); + } + let Header::Initial(header) = packet.header else { + return Ok(None); + }; + if header.version != expected_version { + debug!( + version = header.version, + expected_version, "discarding staged Initial with mismatched version" + ); + return Ok(None); + } + let packet_number = header.number.expand(self.rx_packet + 1); + let payload_len = packet.payload.len(); + let mut payload = packet.payload; + if crypto + .packet + .remote + .decrypt(packet_number, &packet.header_data, &mut payload) + .is_err() + { + return Ok(None); + } + if header.src_cid != expected_src_cid { + debug!( + packet_number, + "discarding staged Initial with mismatched client connection ID" + ); + return Ok(None); + } + if header.token != expected_token { + debug!( + packet_number, + "discarding staged Initial with mismatched Retry token" + ); + return Ok(None); + } + if !self.authenticated_initials.insert_one(packet_number) { + debug!(packet_number, "discarding duplicate staged Initial"); + return Ok(None); + } + self.authenticated_initial_count += 1; + self.rx_packet = self.rx_packet.max(packet_number); + self.read_authenticated_initial_payload( + payload.freeze(), + payload_len, + crypto_buffer_size, + packet_number, + ) + } + + fn read_initial_payload( + &mut self, + payload: Bytes, + payload_len: usize, + crypto_buffer_size: usize, + packet_number: u64, + ) -> Result, ConnectionError> { + if !self.authenticated_initials.insert_one(packet_number) { + debug!(packet_number, "discarding duplicate staged Initial"); + return Ok(None); + } + self.authenticated_initial_count += 1; + self.rx_packet = self.rx_packet.max(packet_number); + self.read_authenticated_initial_payload( + payload, + payload_len, + crypto_buffer_size, + packet_number, + ) + } + + fn read_authenticated_initial_payload( + &mut self, + payload: Bytes, + payload_len: usize, + crypto_buffer_size: usize, + packet_number: u64, + ) -> Result, ConnectionError> { + let mut accepted = None; + let mut ack_eliciting = false; + for result in frame::Iter::new(payload)? { + let frame = result.map_err(TransportError::from)?; + ack_eliciting |= frame.is_ack_eliciting(); + match frame { + frame::Frame::Crypto(frame) => { + let end = frame.offset + frame.data.len() as u64; + let max = end.saturating_sub(self.crypto_stream.bytes_read()); + if max > crypto_buffer_size as u64 { + return Err(TransportError::CRYPTO_BUFFER_EXCEEDED("").into()); + } + + self.crypto_stream + .insert(frame.offset, frame.data.clone(), payload_len) + .map_err(|_| { + TransportError::INTERNAL_ERROR("too many gaps in crypto stream buffer") + })?; + while accepted.is_none() { + let Some(chunk) = self.crypto_stream.read(usize::MAX, true) else { + break; + }; + accepted = self.tls.read_hs(&chunk.bytes)?.map(|tls| RustlsAccepted { + initial_crypto_offset: tls.initial_crypto_offset(), + tls, + }); + } + } + frame::Frame::Padding | frame::Frame::Ping | frame::Frame::Ack(_) => {} + frame::Frame::Close(frame::Close::Connection(reason)) => { + return Err(ConnectionError::ConnectionClosed(reason)); + } + frame => { + let mut err = + TransportError::PROTOCOL_VIOLATION("illegal frame type in handshake"); + err.frame = Some(frame.ty()); + return Err(err.into()); + } + } + } + + self.pending_initial_acks.insert_one(packet_number); + self.ack_pending |= ack_eliciting; + Ok(accepted) + } +} + +/// A rustls ClientHello and the state required to continue the QUIC handshake. +#[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] +#[allow(unnameable_types)] // internal split-accept API +pub struct RustlsAccepted { + tls: crypto::rustls::Accepted, + initial_crypto_offset: u64, +} + +#[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] +impl RustlsAccepted { + /// Get the rustls ClientHello for this connection. + pub fn client_hello(&self) -> rustls::server::ClientHello<'_> { + self.tls.client_hello() + } +} + +/// Internal split-accept handle used by `quinn`. +#[doc(hidden)] +#[allow(unnameable_types)] // internal split-accept API; re-exported only with __internal_split_accept +pub struct Accepting { + reservation: AcceptReservation, + version: u32, + src_cid: ConnectionId, + packet_number: u64, + incoming: Incoming, + // State for deferred Connection creation + server_config: Arc, + params: TransportParameters, + remote_address_validated: bool, + rng_seed: [u8; 32], + endpoint_config: Arc, + cid_len: usize, + cid_lifetime: Option, + allow_mtud: bool, + initial_sends: InitialSendState, + last_activity: Instant, + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + rustls_buffered_datagrams: usize, + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + rustls_input: VecDeque, +} + +impl Accepting { + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + fn observe_authenticated_activity( + &mut self, + received_at: Instant, + ) -> Result<(), ConnectionError> { + if self + .idle_timeout_deadline() + .is_some_and(|deadline| received_at >= deadline) + { + return Err(ConnectionError::TimedOut); + } + self.last_activity = self.last_activity.max(received_at); + Ok(()) + } + + /// Identifier used to route datagrams and wake this exact pending accept. + #[doc(hidden)] + pub fn incoming_idx(&self) -> usize { + self.reservation.incoming_idx + } + + /// Deadline imposed by the currently selected server config's idle timeout. + /// + /// A missing timeout or a timeout of zero disables this deadline. + #[doc(hidden)] + pub fn idle_timeout_deadline(&self) -> Option { + let timeout = self.server_config.transport.max_idle_timeout?; + let millis = timeout.into_inner(); + (millis != 0).then(|| self.last_activity + Duration::from_millis(millis)) + } + + /// Authenticate newly buffered Initials, feed contiguous CRYPTO data to rustls, and encode an + /// optional ACK without holding the endpoint lock. + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + #[doc(hidden)] + pub fn poll_rustls_acceptor( + &mut self, + acceptor: &mut RustlsAcceptor, + buf: &mut Vec, + ) -> Result<(Option, Option), ConnectionError> { + let crypto_buffer_size = self.server_config.transport.crypto_buffer_size; + let mut accepted = None; + + if !acceptor.processed_first_packet { + acceptor.processed_first_packet = true; + let before = acceptor.authenticated_initial_count; + accepted = acceptor.read_initial_payload( + self.incoming.packet.payload.clone().freeze(), + self.incoming.packet.payload.len(), + crypto_buffer_size, + self.packet_number, + )?; + if acceptor.authenticated_initial_count != before { + self.observe_authenticated_activity(self.incoming.received_at)?; + } + } + + if accepted.is_none() && !acceptor.processed_first_datagram_tail { + acceptor.processed_first_datagram_tail = true; + let before = acceptor.authenticated_initial_count; + accepted = acceptor.read_initial_datagram( + None, + self.incoming.rest.clone(), + &self.incoming.crypto, + crypto_buffer_size, + self.cid_len, + self.version, + self.endpoint_config.grease_quic_bit, + self.src_cid, + &self.incoming.packet.header.token, + )?; + if acceptor.authenticated_initial_count != before { + self.observe_authenticated_activity(self.incoming.received_at)?; + } + } + + while accepted.is_none() { + let Some(event) = self.rustls_input.pop_front() else { + break; + }; + if event.remote != self.incoming.addresses.remote { + debug!( + remote = %event.remote, + expected = %self.incoming.addresses.remote, + "discarding staged Initial datagram from unexpected remote" + ); + continue; + } + let before = acceptor.authenticated_initial_count; + let received_at = event.now; + accepted = acceptor.read_initial_datagram( + Some(event.first_decode), + event.remaining, + &self.incoming.crypto, + crypto_buffer_size, + self.cid_len, + self.version, + self.endpoint_config.grease_quic_bit, + self.src_cid, + &self.incoming.packet.header.token, + )?; + if acceptor.authenticated_initial_count != before { + self.observe_authenticated_activity(received_at)?; + } + } + + let response = self.poll_rustls_initial_ack(acceptor, buf); + Ok((accepted, response)) + } + + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + fn poll_rustls_initial_ack( + &mut self, + acceptor: &mut RustlsAcceptor, + buf: &mut Vec, + ) -> Option { + if !acceptor.ack_pending || acceptor.pending_initial_acks.is_empty() { + return None; + } + + let packet_number = self.initial_sends.next_packet_number; + self.initial_sends.next_packet_number += 1; + let header = Header::Initial(InitialHeader { + dst_cid: self.src_cid, + src_cid: self.reservation.loc_cid, + number: PacketNumber::new(packet_number, 0), + token: Bytes::new(), + version: self.version, + }); + + let packet_start = buf.len(); + let partial_encode = header.encode(buf); + frame::Ack::encode(0, &acceptor.pending_initial_acks, None, buf); + buf.resize(buf.len() + self.incoming.crypto.packet.local.tag_len(), 0); + partial_encode.finish( + buf, + &*self.incoming.crypto.header.local, + Some((packet_number, &*self.incoming.crypto.packet.local)), + ); + acceptor.pending_initial_acks = ArrayRangeSet::new(); + acceptor.ack_pending = false; + + self.initial_sends.bytes = self + .initial_sends + .bytes + .saturating_add((buf.len() - packet_start) as u64); + self.initial_sends.datagrams = self.initial_sends.datagrams.saturating_add(1); + + Some(Transmit { + destination: self.incoming.addresses.remote, + ecn: None, + size: buf.len(), + segment_size: None, + src_ip: self.incoming.addresses.local_ip, + }) + } + + /// Complete computationally expensive connection setup steps without holding the endpoint lock. + /// + /// Creates the `Connection` and processes the first packet. + /// None of this requires `&mut Endpoint`. + /// + /// On success, returns the connection plus the reservation that still needs to be activated + /// under the endpoint lock. + #[doc(hidden)] + pub fn finish_without_endpoint(self) -> Result> { + self.incoming.improper_drop_warner.dismiss(); + + let transport_config = self.server_config.transport.clone(); + let tls = self + .server_config + .crypto + .clone() + .start_session(self.version, &self.params); + let mut conn = Connection::new( + self.endpoint_config, + transport_config, + self.reservation.init_cid, + self.reservation.loc_cid, + self.src_cid, + self.incoming.addresses.remote, + self.incoming.addresses.local_ip, + tls, + self.cid_len, + self.cid_lifetime, + self.incoming.received_at, + self.version, + self.allow_mtud, + self.rng_seed, + SideArgs::Server { + server_config: self.server_config, + pref_addr_cid: self.reservation.pref_addr_cid, + path_validated: self.remote_address_validated, + }, + ); + + match conn.handle_first_packet( + self.incoming.received_at, + self.incoming.addresses.remote, + self.incoming.ecn, + self.packet_number, + self.incoming.packet, + self.incoming.rest, + ) { + Ok(()) => Ok(Accepted { + reservation: self.reservation, + conn, + guard: AcceptDropGuard, + }), + Err(e) => Err(Box::new(AcceptingError { + cause: e, + reservation: self.reservation, + version: self.version, + src_cid: self.src_cid, + crypto: self.incoming.crypto, + initial_sends: self.initial_sends, + guard: AcceptDropGuard, + })), + } + } + + /// Continue a split accept from a rustls handshake paused after ClientHello. + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + #[doc(hidden)] + pub fn finish_from_rustls( + self, + accepted: RustlsAccepted, + ) -> Result> { + let tls = match self + .server_config + .crypto + .clone() + .start_session_from_accepted(self.version, &self.params, accepted.tls) + { + Ok(tls) => tls, + Err(error) => return Err(Box::new(self.into_error(error.into()))), + }; + + self.incoming.improper_drop_warner.dismiss(); + let transport_config = self.server_config.transport.clone(); + let mut conn = Connection::new( + self.endpoint_config, + transport_config, + self.reservation.init_cid, + self.reservation.loc_cid, + self.src_cid, + self.incoming.addresses.remote, + self.incoming.addresses.local_ip, + tls, + self.cid_len, + self.cid_lifetime, + self.incoming.received_at, + self.version, + self.allow_mtud, + self.rng_seed, + SideArgs::Server { + server_config: self.server_config, + pref_addr_cid: self.reservation.pref_addr_cid, + path_validated: self.remote_address_validated, + }, + ); + conn.skip_initial_crypto(accepted.initial_crypto_offset); + conn.seed_staged_initial_sends( + self.initial_sends.next_packet_number, + self.initial_sends.bytes, + self.initial_sends.datagrams, + ); + + match conn.handle_first_packet( + self.incoming.received_at, + self.incoming.addresses.remote, + self.incoming.ecn, + self.packet_number, + self.incoming.packet, + self.incoming.rest, + ) { + Ok(()) => Ok(Accepted { + reservation: self.reservation, + conn, + guard: AcceptDropGuard, + }), + Err(e) => Err(Box::new(AcceptingError { + cause: e, + reservation: self.reservation, + version: self.version, + src_cid: self.src_cid, + crypto: self.incoming.crypto, + initial_sends: self.initial_sends, + guard: AcceptDropGuard, + })), + } + } + + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + fn into_error(self, cause: ConnectionError) -> AcceptingError { + self.incoming.improper_drop_warner.dismiss(); + AcceptingError { + cause, + reservation: self.reservation, + version: self.version, + src_cid: self.src_cid, + crypto: self.incoming.crypto, + initial_sends: self.initial_sends, + guard: AcceptDropGuard, + } + } +} + +/// Internal split-accept failure state used by `quinn`. +#[doc(hidden)] +#[allow(unnameable_types)] // internal split-accept API; re-exported only with __internal_split_accept +pub struct AcceptingError { + cause: ConnectionError, + reservation: AcceptReservation, + version: u32, + src_cid: ConnectionId, + crypto: Keys, + initial_sends: InitialSendState, + guard: AcceptDropGuard, +} + /// Error for attempting to retry an [`Incoming`] which already bears a token from a previous retry #[derive(Debug, Error)] #[error("retry() with validated Incoming")] diff --git a/quinn-proto/src/lib.rs b/quinn-proto/src/lib.rs index 36ccb5677e..0743d3ab5a 100644 --- a/quinn-proto/src/lib.rs +++ b/quinn-proto/src/lib.rs @@ -73,6 +73,15 @@ mod endpoint; pub use crate::endpoint::{ AcceptError, ConnectError, ConnectionHandle, DatagramEvent, Endpoint, Incoming, RetryError, }; +#[cfg(feature = "__internal_split_accept")] +#[doc(hidden)] +pub use crate::endpoint::{Accepted, Accepting, AcceptingError}; +#[cfg(all( + feature = "__internal_split_accept", + any(feature = "rustls-aws-lc-rs", feature = "rustls-ring") +))] +#[doc(hidden)] +pub use crate::endpoint::{RustlsAccepted, RustlsAcceptor}; mod packet; pub use packet::{ @@ -157,15 +166,7 @@ pub mod fuzzing { } /// The QUIC protocol version implemented. -pub const DEFAULT_SUPPORTED_VERSIONS: &[u32] = &[ - 0x00000001, - 0xff00_001d, - 0xff00_001e, - 0xff00_001f, - 0xff00_0020, - 0xff00_0021, - 0xff00_0022, -]; +pub const DEFAULT_SUPPORTED_VERSIONS: &[u32] = &[0x00000001, 0xff00_0021, 0xff00_0022]; /// Whether an endpoint was the initiator of a connection #[cfg_attr(feature = "arbitrary", derive(Arbitrary))] diff --git a/quinn-proto/src/packet.rs b/quinn-proto/src/packet.rs index e8ab1aabb5..6abf42dc85 100644 --- a/quinn-proto/src/packet.rs +++ b/quinn-proto/src/packet.rs @@ -21,8 +21,7 @@ use crate::{ /// across QUIC versions), which gives us the destination CID and allows us /// to inspect the version and packet type (which depends on the version). /// This information allows us to fully decode and decrypt the packet. -#[cfg_attr(test, derive(Clone))] -#[derive(Debug)] +#[derive(Clone, Debug)] pub struct PartialDecode { plain_header: ProtectedHeader, buf: io::Cursor, @@ -938,17 +937,15 @@ mod tests { #[test] fn header_encoding() { use crate::Side; - use crate::crypto::rustls::{initial_keys, initial_suite_from_provider}; - #[cfg(all(feature = "rustls-aws-lc-rs", not(feature = "rustls-ring")))] - use rustls::crypto::aws_lc_rs::default_provider; - #[cfg(feature = "rustls-ring")] - use rustls::crypto::ring::default_provider; + use crate::crypto::rustls::{ + configured_provider, initial_keys, initial_suite_from_provider, + }; use rustls::quic::Version; let dcid = ConnectionId::new(&hex!("06b858ec6f80452b")); - let provider = default_provider(); + let provider = configured_provider(); - let suite = initial_suite_from_provider(&std::sync::Arc::new(provider)).unwrap(); + let suite = initial_suite_from_provider(&provider).unwrap(); let client = initial_keys(Version::V1, dcid, Side::Client, &suite); let mut buf = Vec::new(); let header = Header::Initial(InitialHeader { diff --git a/quinn-proto/src/shared.rs b/quinn-proto/src/shared.rs index f2d0ad5d2a..4a1aa1a608 100644 --- a/quinn-proto/src/shared.rs +++ b/quinn-proto/src/shared.rs @@ -17,7 +17,7 @@ pub(crate) enum ConnectionEventInner { } /// Variant of [`ConnectionEventInner`]. -#[derive(Debug)] +#[derive(Clone, Debug)] pub(crate) struct DatagramConnectionEvent { pub(crate) now: Instant, pub(crate) remote: SocketAddr, diff --git a/quinn-proto/src/tests/mod.rs b/quinn-proto/src/tests/mod.rs index 226fcf78e2..c250c565fe 100644 --- a/quinn-proto/src/tests/mod.rs +++ b/quinn-proto/src/tests/mod.rs @@ -13,12 +13,11 @@ use hex_literal::hex; use rand::Rng; #[cfg(feature = "ring")] use ring::hmac; -#[cfg(all(feature = "rustls-aws-lc-rs", not(feature = "rustls-ring")))] -use rustls::crypto::aws_lc_rs::default_provider; -#[cfg(feature = "rustls-ring")] -use rustls::crypto::ring::default_provider; use rustls::{ - AlertDescription, RootCertStore, + RootCertStore, + crypto::Identity, + enums::{ApplicationProtocol, ProtocolVersion}, + error::AlertDescription, pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer}, server::WebPkiClientVerifier, }; @@ -29,6 +28,7 @@ use crate::{ Duration, Instant, cid_generator::{ConnectionIdGenerator, RandomConnectionIdGenerator}, crypto::rustls::QuicServerConfig, + endpoint::RustlsAcceptor, frame::FrameStruct, transport_parameters::TransportParameters, }; @@ -217,11 +217,11 @@ fn stats_include_congestion_controller_bandwidth_estimate() { } #[test] -fn draft_version_compat() { +fn compatible_version() { let _guard = subscribe(); let mut client_config = client_config(); - client_config.version(0xff00_0020); + client_config.version(0xff00_0022); let mut pair = Pair::default(); let (client_ch, server_ch) = pair.connect_with(client_config); @@ -367,14 +367,14 @@ fn export_keying_material() { // client keying material let mut client_buf = [0u8; 64]; pair.client_conn_mut(client_ch) - .crypto_session() + .crypto_session_mut() .export_keying_material(&mut client_buf, LABEL, CONTEXT) .unwrap(); // server keying material let mut server_buf = [0u8; 64]; pair.server_conn_mut(server_ch) - .crypto_session() + .crypto_session_mut() .export_keying_material(&mut server_buf, LABEL, CONTEXT) .unwrap(); @@ -507,7 +507,7 @@ fn reject_self_signed_server_cert() { assert_matches!(pair.client_conn_mut(client_ch).poll(), Some(Event::ConnectionLost { reason: ConnectionError::TransportError(ref error)}) - if error.code == TransportErrorCode::crypto(AlertDescription::UnknownCA.into())); + if error.code == TransportErrorCode::crypto(AlertDescription::UnknownCa.into())); } #[test] @@ -522,16 +522,17 @@ fn reject_missing_client_cert() { let key = PrivatePkcs8KeyDer::from(CERTIFIED_KEY.signing_key.serialize_der()); let cert = CERTIFIED_KEY.cert.der().clone(); - let provider = Arc::new(default_provider()); - let config = rustls::ServerConfig::builder_with_provider(provider.clone()) - .with_protocol_versions(&[&rustls::version::TLS13]) - .unwrap() - .with_client_cert_verifier( - WebPkiClientVerifier::builder_with_provider(Arc::new(store), provider) + let provider = crypto::rustls::configured_provider(); + let config = rustls::ServerConfig::builder(provider.clone()) + .with_client_cert_verifier(Arc::new( + WebPkiClientVerifier::builder(Arc::new(store), &provider) .build() .unwrap(), + )) + .with_single_cert( + Arc::new(Identity::from_cert_chain(vec![cert]).unwrap()), + PrivateKeyDer::from(key), ) - .with_single_cert(vec![cert], PrivateKeyDer::from(key)) .unwrap(); let config = QuicServerConfig::try_from(config).unwrap(); @@ -751,7 +752,7 @@ fn zero_rtt_rejection() { // the existing `ClientConfig` and change the ALPN protocols to make that happen. let this = Arc::get_mut(&mut client_crypto).expect("QuicClientConfig is shared"); let inner = Arc::get_mut(&mut this.inner).expect("QuicClientConfig.inner is shared"); - inner.alpn_protocols = vec!["bar".into()]; + inner.alpn_protocols = vec![ApplicationProtocol::from(b"bar")]; // Changing protocols invalidates 0-RTT let client_config = ClientConfig::new(client_crypto); @@ -905,6 +906,131 @@ fn zero_rtt_incoming_buffer_size_total() { }); } +/// Verify that datagrams arriving while a connection is in the `Accepting` state (between +/// `start_accept` and `finish_accept`) are buffered in `incoming_buffers` and replayed into the +/// connection after `finish_accept`. Drives through the full handshake and clean shutdown to +/// confirm no endpoint state is leaked. +#[test] +fn accepting_state_buffers_retransmitted_initials() { + let _guard = subscribe(); + let mut pair = Pair::default(); + pair.server.handle_incoming = Box::new(|_| IncomingConnectionBehavior::Wait); + + let client_ch = pair.begin_connect(client_config()); + pair.drive_client(); + pair.drive_server(); + + let incoming = pair.server.pop_waiting_incoming(); + + let accepting = pair.server.start_split_accept(incoming, pair.time); + assert_eq!(pair.server.incoming_buffer_bytes(), 0); + assert_eq!(pair.server.open_connections(), 0); + assert_eq!(pair.server.pending_accepts(), 1); + + // Removing the server configuration only prevents new attempts; it does not close the + // endpoint. An accept that already captured its configuration must keep buffering. + pair.server.disable_new_connections(); + + // With no server response, the client's next wakeup is its loss timer. Advancing to it and + // driving the client emits a retransmitted Initial for the same connection attempt. + pair.time = pair.client.next_wakeup().unwrap(); + pair.drive_client(); + assert!(!pair.server.inbound.is_empty()); + pair.drive_server(); + + assert!(pair.server.waiting_incoming.is_empty()); + assert!(pair.server.incoming_buffer_bytes() > 0); + + let server_ch = pair.server.finish_split_accept(accepting); + assert_eq!(pair.server.incoming_buffer_bytes(), 0); + assert_eq!(pair.server.open_connections(), 1); + assert_eq!(pair.server.pending_accepts(), 0); + + pair.drive(); + pair.finish_connect(client_ch, server_ch); + + pair.client + .connections + .get_mut(&client_ch) + .unwrap() + .close(pair.time, VarInt(42), Bytes::new()); + pair.drive(); + assert_eq!(pair.client.known_connections(), 0); + assert_eq!(pair.client.known_cids(), 0); + assert_eq!(pair.server.known_connections(), 0); + assert_eq!(pair.server.known_cids(), 0); +} + +/// Verify that attempts in the `Accepting` state count toward `max_incoming`, so a second +/// connection attempt is refused while the first attempt is still between `start_accept` +/// and `finish_accept`. +#[test] +fn max_incoming_counts_accepts_in_progress() { + let _guard = subscribe(); + let mut server_config = server_config(); + server_config.max_incoming(1); + let mut pair = Pair::new(Arc::new(EndpointConfig::default()), server_config); + pair.server.handle_incoming = Box::new(|_| IncomingConnectionBehavior::Wait); + + let _client_ch = pair.begin_connect(client_config()); + pair.drive_client(); + pair.drive_server(); + + let incoming = pair.server.pop_waiting_incoming(); + + let accepting = pair.server.start_split_accept(incoming, pair.time); + assert_eq!(pair.server.open_connections(), 0); + assert_eq!(pair.server.pending_accepts(), 1); + + let _refused_ch = pair.begin_connect(client_config()); + pair.drive_client(); + pair.drive_server(); + assert!(pair.server.waiting_incoming.is_empty()); + assert_eq!(pair.server.open_connections(), 0); + assert_eq!(pair.server.pending_accepts(), 1); + + pair.server.finish_split_accept(accepting); + assert_eq!(pair.server.open_connections(), 1); + assert_eq!(pair.server.pending_accepts(), 0); +} + +/// Verify that when the off-lock handshake fails (here via ALPN mismatch) after `start_accept` +/// has reserved endpoint state, `finish_accept_error` releases the pending-accept slot and the +/// reserved CIDs/buffer, leaving no endpoint state behind. +#[test] +fn accepting_state_cleaned_up_on_handshake_failure() { + let _guard = subscribe(); + let server_config = + ServerConfig::with_crypto(Arc::new(server_crypto_with_alpn(vec!["foo".into()]))); + let mut pair = Pair::new(Arc::new(EndpointConfig::default()), server_config); + pair.server.handle_incoming = Box::new(|_| IncomingConnectionBehavior::Wait); + + let _client_ch = + pair.begin_connect(ClientConfig::new(Arc::new(client_crypto_with_alpn(vec![ + "bar".into(), + ])))); + pair.drive_client(); + pair.drive_server(); + + let incoming = pair.server.pop_waiting_incoming(); + let accepting = pair.server.start_split_accept(incoming, pair.time); + assert_eq!(pair.server.pending_accepts(), 1); + + // The TLS handshake runs in finish_without_endpoint and fails on the ALPN mismatch. + let cause = pair.server.finish_split_accept_error(accepting); + assert_matches!( + cause, + ConnectionError::TransportError(ref e) if e.code == TransportErrorCode::crypto(0x78) + ); + + // The failed accept must leave no reserved endpoint state behind. + assert_eq!(pair.server.pending_accepts(), 0); + assert_eq!(pair.server.open_connections(), 0); + assert_eq!(pair.server.incoming_buffer_bytes(), 0); + assert_eq!(pair.server.known_connections(), 0); + assert_eq!(pair.server.known_cids(), 0); +} + #[test] fn alpn_success() { let _guard = subscribe(); @@ -949,13 +1075,13 @@ fn alpn_success() { assert_eq!( hd.protocol_version .unwrap() - .downcast_ref::(), - Some(&rustls::ProtocolVersion::TLSv1_3) + .downcast_ref::(), + Some(&ProtocolVersion::TLSv1_3) ); assert!( hd.cipher_suite .unwrap() - .downcast_ref::() + .downcast_ref::() .is_some() ); } @@ -2278,6 +2404,596 @@ fn large_initial() { ); } +#[test] +fn staged_acceptor_reassembles_reordered_initials() { + let _guard = subscribe(); + let server_config = + ServerConfig::with_crypto(Arc::new(server_crypto_with_alpn(vec![vec![0, 0, 0, 42]]))); + let mut pair = Pair::new(Arc::new(EndpointConfig::default()), server_config); + pair.server.handle_incoming = Box::new(|_| IncomingConnectionBehavior::Wait); + + let client_crypto = + client_crypto_with_alpn((0..1000u32).map(|x| x.to_be_bytes().to_vec()).collect()); + let client_ch = pair.begin_connect(ClientConfig::new(Arc::new(client_crypto))); + pair.drive_client(); + while pair.server.inbound.len() < 5 { + pair.time = pair.client.next_wakeup().unwrap(); + pair.drive_client(); + } + assert!(pair.server.inbound.len() > 1); + pair.server.inbound.make_contiguous().reverse(); + pair.drive_server(); + + let incoming = pair.server.pop_waiting_incoming(); + let mut acceptor = RustlsAcceptor::new(&incoming).unwrap(); + let mut accepting = pair.server.start_split_accept(incoming, pair.time); + assert!(pair.server.incoming_buffer_bytes() > 0); + + let mut buf = Vec::new(); + pair.server + .endpoint + .buffer_rustls_acceptor_input(&mut accepting); + let (accepted, ack) = accepting + .poll_rustls_acceptor(&mut acceptor, &mut buf) + .unwrap(); + let accepted = accepted.expect("reordered Initials should complete ClientHello"); + if let Some(transmit) = ack { + let size = transmit.size; + pair.server + .outbound + .push_back((transmit, Bytes::copy_from_slice(&buf[..size]))); + } + + pair.server + .endpoint + .select_accepting_config(&mut accepting, None, pair.time) + .unwrap(); + let Ok(accepted) = accepting.finish_from_rustls(accepted) else { + panic!("staged rustls accept unexpectedly failed") + }; + let (server_ch, conn) = pair.server.endpoint.finish_accept(accepted); + pair.server.connections.insert(server_ch, conn); + + pair.drive(); + pair.finish_connect(client_ch, server_ch); +} + +#[test] +fn staged_acceptor_ignores_initials_from_spoofed_remote() { + let _guard = subscribe(); + let server_config = + ServerConfig::with_crypto(Arc::new(server_crypto_with_alpn(vec![vec![0, 0, 0, 42]]))); + let mut pair = Pair::new(Arc::new(EndpointConfig::default()), server_config); + pair.server.handle_incoming = Box::new(|_| IncomingConnectionBehavior::Wait); + + let client_crypto = + client_crypto_with_alpn((0..1000u32).map(|x| x.to_be_bytes().to_vec()).collect()); + let client_ch = pair.begin_connect(ClientConfig::new(Arc::new(client_crypto))); + pair.drive_client(); + while pair.server.inbound.len() < 5 { + pair.time = pair.client.next_wakeup().unwrap(); + pair.drive_client(); + } + + let first = pair.server.inbound.pop_front().unwrap(); + let remaining: Vec<_> = pair.server.inbound.drain(..).collect(); + pair.server.inbound.push_back(first); + pair.drive_server(); + + let incoming = pair.server.pop_waiting_incoming(); + let mut acceptor = RustlsAcceptor::new(&incoming).unwrap(); + let mut accepting = pair.server.start_split_accept(incoming, pair.time); + let incoming_idx = accepting.incoming_idx(); + + let mut spoofed_remote = pair.client.addr; + spoofed_remote.set_port(spoofed_remote.port() + 1); + let mut endpoint_buf = Vec::new(); + let buffered_before_spoof = pair.server.incoming_buffer_bytes(); + for (now, ecn, packet) in &remaining { + let event = pair.server.endpoint.handle( + *now, + spoofed_remote, + None, + *ecn, + packet.clone(), + &mut endpoint_buf, + ); + assert!( + event.is_none(), + "spoofed datagrams must be dropped before pending-incoming routing" + ); + assert_eq!( + pair.server.incoming_buffer_bytes(), + buffered_before_spoof, + "spoofed datagrams must not consume the pending-incoming byte quota" + ); + endpoint_buf.clear(); + } + + pair.server + .endpoint + .buffer_rustls_acceptor_input(&mut accepting); + let mut ack_buf = Vec::new(); + let (accepted, ack) = accepting + .poll_rustls_acceptor(&mut acceptor, &mut ack_buf) + .unwrap(); + assert!( + accepted.is_none(), + "spoofed datagrams must not complete ClientHello" + ); + if let Some(transmit) = ack { + let size = transmit.size; + pair.server + .outbound + .push_back((transmit, Bytes::copy_from_slice(&ack_buf[..size]))); + } + + for (now, ecn, packet) in remaining { + let event = pair.server.endpoint.handle( + now, + pair.client.addr, + None, + ecn, + packet, + &mut endpoint_buf, + ); + assert!(matches!(event, Some(DatagramEvent::IncomingData(idx)) if idx == incoming_idx)); + endpoint_buf.clear(); + } + pair.server + .endpoint + .buffer_rustls_acceptor_input(&mut accepting); + ack_buf.clear(); + let (accepted, ack) = accepting + .poll_rustls_acceptor(&mut acceptor, &mut ack_buf) + .unwrap(); + let accepted = accepted.expect("legitimate datagrams should complete ClientHello"); + if let Some(transmit) = ack { + let size = transmit.size; + pair.server + .outbound + .push_back((transmit, Bytes::copy_from_slice(&ack_buf[..size]))); + } + + pair.server + .endpoint + .select_accepting_config(&mut accepting, None, pair.time) + .unwrap(); + let Ok(accepted) = accepting.finish_from_rustls(accepted) else { + panic!("staged rustls accept unexpectedly failed") + }; + let (server_ch, conn) = pair.server.endpoint.finish_accept(accepted); + pair.server.connections.insert(server_ch, conn); + + pair.drive(); + pair.finish_connect(client_ch, server_ch); +} + +#[test] +fn staged_acceptor_ignores_initials_with_mismatched_client_cid() { + let _guard = subscribe(); + let server_cfg = server_config(); + let server_crypto = server_cfg.crypto.clone(); + let mut pair = Pair::new(Arc::new(EndpointConfig::default()), server_cfg); + pair.server.handle_incoming = Box::new(|_| IncomingConnectionBehavior::Wait); + + let initial_dcid = ConnectionId::new(&[7; 16]); + let mut first_config = client_config(); + first_config.initial_dst_cid_provider(Arc::new(move || initial_dcid)); + let client_ch = pair.begin_connect(first_config); + pair.drive_client(); + let raw_initial = pair.server.inbound.front().unwrap().2.clone(); + pair.drive_server(); + + let incoming = pair.server.pop_waiting_incoming(); + let mut acceptor = RustlsAcceptor::new(&incoming).unwrap(); + let mut accepting = pair.server.start_split_accept(incoming, pair.time); + + // A distinct connection using the same client-chosen Initial DCID derives the same public + // Initial keys, but uses an empty client SCID. Feed its packet number 0 before the legitimate + // packet number 0 to ensure the rejected packet cannot poison duplicate detection. + let cid_generator_factory: fn() -> Box = + || Box::new(RandomConnectionIdGenerator::new(0)); + let mut attacker = Pair::new( + Arc::new(EndpointConfig { + connection_id_generator_factory: Arc::new(cid_generator_factory), + ..EndpointConfig::default() + }), + server_config(), + ); + let mut attacker_config = client_config(); + attacker_config.initial_dst_cid_provider(Arc::new(move || initial_dcid)); + attacker.begin_connect(attacker_config); + attacker.drive_client(); + let spoofed_initial = attacker.server.inbound.front().unwrap().2.clone(); + + let keys = server_crypto.initial_keys(1, initial_dcid).unwrap(); + let decode = |packet| { + PartialDecode::new( + packet, + &FixedLengthConnectionIdParser::new(8), + &[1], + pair.server.config().grease_quic_bit, + ) + .unwrap() + .0 + }; + let expected_src_cid = decode(raw_initial).initial_header().unwrap().src_cid; + let spoofed_decode = decode(spoofed_initial); + assert_ne!( + spoofed_decode.initial_header().unwrap().src_cid, + expected_src_cid + ); + assert!( + !acceptor + .test_read_initial_decode(spoofed_decode, &keys, expected_src_cid, Bytes::new(), 1,) + .unwrap(), + "mismatched client SCID must be rejected" + ); + + let mut ack_buf = Vec::new(); + let (accepted, ack) = accepting + .poll_rustls_acceptor(&mut acceptor, &mut ack_buf) + .unwrap(); + let accepted = accepted + .expect("matching packet with the same packet number should still complete ClientHello"); + if let Some(transmit) = ack { + let size = transmit.size; + pair.server + .outbound + .push_back((transmit, Bytes::copy_from_slice(&ack_buf[..size]))); + } + + pair.server + .endpoint + .select_accepting_config(&mut accepting, None, pair.time) + .unwrap(); + let Ok(accepted) = accepting.finish_from_rustls(accepted) else { + panic!("staged rustls accept unexpectedly failed") + }; + let (server_ch, conn) = pair.server.endpoint.finish_accept(accepted); + pair.server.connections.insert(server_ch, conn); + + pair.drive(); + pair.finish_connect(client_ch, server_ch); +} + +#[test] +fn staged_acceptor_ignores_mismatched_token_without_consuming_packet_number() { + let _guard = subscribe(); + let server_config = server_config(); + let server_crypto = server_config.crypto.clone(); + let mut pair = Pair::new(Arc::new(EndpointConfig::default()), server_config); + pair.server.handle_incoming = Box::new(|_| IncomingConnectionBehavior::Wait); + + let initial_dcid = ConnectionId::new(&[9; 16]); + let mut client_config = client_config(); + client_config.initial_dst_cid_provider(Arc::new(move || initial_dcid)); + let client_ch = pair.begin_connect(client_config); + pair.drive_client(); + let raw_initial = pair.server.inbound.front().unwrap().2.clone(); + pair.drive_server(); + + let incoming = pair.server.pop_waiting_incoming(); + let mut acceptor = RustlsAcceptor::new(&incoming).unwrap(); + let keys = server_crypto.initial_keys(1, initial_dcid).unwrap(); + let decode = || { + PartialDecode::new( + raw_initial.clone(), + &FixedLengthConnectionIdParser::new(8), + &[1], + pair.server.config().grease_quic_bit, + ) + .unwrap() + .0 + }; + let expected_src_cid = decode().initial_header().unwrap().src_cid; + + assert!( + !acceptor + .test_read_initial_decode( + decode(), + &keys, + expected_src_cid, + Bytes::from_static(b"wrong token"), + 1, + ) + .unwrap() + ); + let mut accepting = pair.server.start_split_accept(incoming, pair.time); + let mut ack_buf = Vec::new(); + let (accepted, ack) = accepting + .poll_rustls_acceptor(&mut acceptor, &mut ack_buf) + .unwrap(); + let accepted = accepted.expect( + "invalid-token packet must not consume the legitimate packet number or ClientHello", + ); + if let Some(transmit) = ack { + let size = transmit.size; + pair.server + .outbound + .push_back((transmit, Bytes::copy_from_slice(&ack_buf[..size]))); + } + + pair.server + .endpoint + .select_accepting_config(&mut accepting, None, pair.time) + .unwrap(); + let Ok(accepted) = accepting.finish_from_rustls(accepted) else { + panic!("staged rustls accept unexpectedly failed") + }; + let (server_ch, conn) = pair.server.endpoint.finish_accept(accepted); + pair.server.connections.insert(server_ch, conn); + + pair.drive(); + pair.finish_connect(client_ch, server_ch); +} + +#[test] +fn staged_acceptor_peer_connection_close_is_silent() { + let _guard = subscribe(); + let mut pair = Pair::default(); + pair.server.handle_incoming = Box::new(|_| IncomingConnectionBehavior::Wait); + pair.begin_connect(client_config()); + pair.drive_client(); + pair.drive_server(); + + let incoming = pair.server.pop_waiting_incoming(); + let mut acceptor = RustlsAcceptor::new(&incoming).unwrap(); + let accepting = pair.server.start_split_accept(incoming, pair.time); + assert_eq!(pair.server.pending_accepts(), 1); + + let close = ConnectionClose { + error_code: TransportErrorCode::NO_ERROR, + frame_type: None, + reason: Bytes::from_static(b"peer closed while staged"), + }; + let mut payload = Vec::new(); + frame::Close::Connection(close.clone()).encode(&mut payload, usize::MAX); + let error = acceptor + .test_read_initial_payload(payload.into(), 1) + .unwrap_err(); + let ConnectionError::ConnectionClosed(actual) = &error else { + panic!("expected peer connection close") + }; + assert_eq!(*actual, close); + + let mut response_buf = Vec::new(); + let error = pair + .server + .endpoint + .fail_accepting(accepting, error, &mut response_buf); + assert!(error.response.is_none(), "peer close must not be echoed"); + assert!(response_buf.is_empty()); + assert_eq!(pair.server.pending_accepts(), 0); + assert_eq!(pair.server.incoming_buffer_bytes(), 0); + assert_eq!(pair.server.open_connections(), 0); +} + +#[test] +fn staged_acceptor_idle_deadline_tracks_only_timely_authenticated_activity() { + let _guard = subscribe(); + const IDLE_TIMEOUT: u64 = 60_000; + + for activity_after_deadline in [false, true] { + let mut server_config = + ServerConfig::with_crypto(Arc::new(server_crypto_with_alpn(vec![vec![0, 0, 0, 42]]))); + Arc::get_mut(&mut server_config.transport) + .unwrap() + .max_idle_timeout(Some( + Duration::from_millis(IDLE_TIMEOUT).try_into().unwrap(), + )); + let mut pair = Pair::new(Arc::new(EndpointConfig::default()), server_config); + pair.server.handle_incoming = Box::new(|_| IncomingConnectionBehavior::Wait); + + let client_crypto = + client_crypto_with_alpn((0..1000u32).map(|x| x.to_be_bytes().to_vec()).collect()); + pair.begin_connect(ClientConfig::new(Arc::new(client_crypto))); + pair.drive_client(); + while pair.server.inbound.len() < 2 { + pair.time = pair.client.next_wakeup().unwrap(); + pair.drive_client(); + } + + let first = pair.server.inbound.pop_front().unwrap(); + let (_, second_ecn, second_packet) = pair.server.inbound.pop_front().unwrap(); + pair.server.inbound.clear(); + pair.server.inbound.push_back(first); + pair.drive_server(); + + let incoming = pair.server.pop_waiting_incoming(); + let mut acceptor = RustlsAcceptor::new(&incoming).unwrap(); + let mut accepting = pair.server.start_split_accept(incoming, pair.time); + let initial_deadline = accepting.idle_timeout_deadline().unwrap(); + + let mut ack_buf = Vec::new(); + let (accepted, _) = accepting + .poll_rustls_acceptor(&mut acceptor, &mut ack_buf) + .unwrap(); + assert!( + accepted.is_none(), + "the first fragment must not complete the large ClientHello" + ); + + let activity_time = if activity_after_deadline { + initial_deadline + } else { + initial_deadline - Duration::from_millis(1) + }; + let mut endpoint_buf = Vec::new(); + let event = pair.server.endpoint.handle( + activity_time, + pair.client.addr, + None, + second_ecn, + second_packet, + &mut endpoint_buf, + ); + assert!(matches!(event, Some(DatagramEvent::IncomingData(_)))); + pair.server + .endpoint + .buffer_rustls_acceptor_input(&mut accepting); + ack_buf.clear(); + let result = accepting.poll_rustls_acceptor(&mut acceptor, &mut ack_buf); + + if activity_after_deadline { + assert!(matches!(result, Err(ConnectionError::TimedOut))); + assert_eq!(accepting.idle_timeout_deadline(), Some(initial_deadline)); + } else { + result.unwrap(); + assert_eq!( + accepting.idle_timeout_deadline(), + Some(activity_time + Duration::from_millis(IDLE_TIMEOUT)), + "timely authenticated activity must extend the idle deadline" + ); + } + + let mut close_buf = Vec::new(); + pair.server.endpoint.fail_accepting( + accepting, + ConnectionError::LocallyClosed, + &mut close_buf, + ); + } +} + +#[test] +fn staged_acceptor_routes_zero_length_cid_initials() { + let _guard = subscribe(); + let cid_generator_factory: fn() -> Box = + || Box::new(RandomConnectionIdGenerator::new(0)); + let mut pair = Pair::new( + Arc::new(EndpointConfig { + connection_id_generator_factory: Arc::new(cid_generator_factory), + ..EndpointConfig::default() + }), + server_config(), + ); + pair.server.handle_incoming = Box::new(|_| IncomingConnectionBehavior::Wait); + + let client_ch = pair.begin_connect(client_config()); + pair.drive_client(); + pair.drive_server(); + + let incoming = pair.server.pop_waiting_incoming(); + let mut acceptor = RustlsAcceptor::new(&incoming).unwrap(); + let mut accepting = pair.server.start_split_accept(incoming, pair.time); + let incoming_idx = accepting.incoming_idx(); + + let mut buf = Vec::new(); + pair.server + .endpoint + .buffer_rustls_acceptor_input(&mut accepting); + let (accepted, ack) = accepting + .poll_rustls_acceptor(&mut acceptor, &mut buf) + .unwrap(); + let accepted = accepted.expect("first Initial should contain ClientHello"); + let ack = ack.expect("ClientHello Initial should be acknowledged"); + let size = ack.size; + pair.server + .outbound + .push_back((ack, Bytes::copy_from_slice(&buf[..size]))); + + // Receiving the staged ACK switches the client's Initial DCID to the server's empty SCID. + // Keep the accept pending until the client's next handshake probe reaches the endpoint. + pair.drive_server(); + pair.drive_client(); + if pair.server.inbound.is_empty() { + pair.time = pair.client.next_wakeup().unwrap(); + pair.drive_client(); + } + assert!(!pair.server.inbound.is_empty()); + pair.drive_server(); + assert!(pair.server.incoming_buffer_bytes() > 0); + assert_eq!(accepting.incoming_idx(), incoming_idx); + + pair.server + .endpoint + .select_accepting_config(&mut accepting, None, pair.time) + .unwrap(); + let Ok(accepted) = accepting.finish_from_rustls(accepted) else { + panic!("staged rustls accept unexpectedly failed") + }; + let (server_ch, conn) = pair.server.endpoint.finish_accept(accepted); + pair.server.connections.insert(server_ch, conn); + + pair.drive(); + pair.finish_connect(client_ch, server_ch); +} + +#[test] +fn staged_acceptor_error_close_uses_reserved_cid_and_next_packet_number() { + let _guard = subscribe(); + let mut pair = Pair::default(); + pair.server.handle_incoming = Box::new(|_| IncomingConnectionBehavior::Wait); + + let client_ch = pair.begin_connect(client_config()); + pair.drive_client(); + pair.drive_server(); + + let incoming = pair.server.pop_waiting_incoming(); + let mut acceptor = RustlsAcceptor::new(&incoming).unwrap(); + let mut accepting = pair.server.start_split_accept(incoming, pair.time); + let mut ack_buf = Vec::new(); + let (_, ack) = accepting + .poll_rustls_acceptor(&mut acceptor, &mut ack_buf) + .unwrap(); + let ack = ack.expect("ClientHello Initial should be acknowledged"); + let ack_size = ack.size; + pair.server + .outbound + .push_back((ack, Bytes::copy_from_slice(&ack_buf[..ack_size]))); + + let mut close_buf = Vec::new(); + let error = pair.server.endpoint.fail_accepting( + accepting, + TransportError::CONNECTION_REFUSED("test refusal").into(), + &mut close_buf, + ); + let close = error + .response + .expect("transport error should generate close"); + let close_size = close.size; + pair.server + .outbound + .push_back((close, Bytes::copy_from_slice(&close_buf[..close_size]))); + + // If the close reused staged ACK packet number 0, the client would discard it as a duplicate. + // If it changed the reserved server SCID, the client would discard it as a CID mismatch. + pair.drive_server(); + pair.drive_client(); + assert_matches!( + pair.client_conn_mut(client_ch).poll(), + Some(Event::ConnectionLost { + reason: ConnectionError::ConnectionClosed(close) + }) if close.error_code == TransportErrorCode::CONNECTION_REFUSED + ); + assert_eq!(pair.server.pending_accepts(), 0); + assert_eq!(pair.server.incoming_buffer_bytes(), 0); +} + +#[test] +fn staged_acceptor_does_not_ack_ack_only_initial() { + let _guard = subscribe(); + let mut pair = Pair::default(); + pair.server.handle_incoming = Box::new(|_| IncomingConnectionBehavior::Wait); + pair.begin_connect(client_config()); + pair.drive_client(); + pair.drive_server(); + + let incoming = pair.server.pop_waiting_incoming(); + let mut acceptor = RustlsAcceptor::new(&incoming).unwrap(); + let mut ranges = range_set::ArrayRangeSet::new(); + ranges.insert_one(0); + let mut payload = Vec::new(); + frame::Ack::encode(0, &ranges, None, &mut payload); + acceptor + .test_read_initial_payload(payload.into(), 1) + .unwrap(); + assert!(!acceptor.test_ack_pending()); + pair.server.endpoint.ignore(incoming); +} + #[test] /// Ensure that we don't yield a finish event before the actual FIN is acked so the peer isn't left /// hanging diff --git a/quinn-proto/src/tests/util.rs b/quinn-proto/src/tests/util.rs index 24c13d2a47..c5fcc6950b 100644 --- a/quinn-proto/src/tests/util.rs +++ b/quinn-proto/src/tests/util.rs @@ -13,15 +13,17 @@ use std::{ use assert_matches::assert_matches; use bytes::BytesMut; use rustls::{ - KeyLogFile, client::WebPkiServerVerifier, + crypto::CryptoProvider, + enums::ApplicationProtocol, pki_types::{CertificateDer, PrivateKeyDer}, }; +use rustls_util::KeyLogFile; use tracing::{info_span, trace}; -use super::crypto::rustls::{QuicClientConfig, QuicServerConfig, configured_provider}; +use super::crypto::rustls::{QuicClientConfig, QuicServerConfig}; use super::*; -use crate::{Duration, Instant}; +use crate::{Duration, Instant, endpoint::Accepting}; pub(super) const DEFAULT_MTU: usize = 1452; @@ -212,7 +214,11 @@ impl Pair { client_ch } - fn finish_connect(&mut self, client_ch: ConnectionHandle, server_ch: ConnectionHandle) { + pub(super) fn finish_connect( + &mut self, + client_ch: ConnectionHandle, + server_ch: ConnectionHandle, + ) { assert_matches!( self.client_conn_mut(client_ch).poll(), Some(Event::HandshakeDataReady) @@ -396,6 +402,7 @@ impl TestEndpoint { self.conn_events.entry(ch).or_default().push_back(event); } + DatagramEvent::IncomingData(_) => {} DatagramEvent::Response(transmit) => { let size = transmit.size; self.outbound.extend(split_transmit(transmit, &buf[..size])); @@ -490,6 +497,42 @@ impl TestEndpoint { } } + pub(super) fn pop_waiting_incoming(&mut self) -> Incoming { + let incoming = self.waiting_incoming.pop().unwrap(); + assert!(self.waiting_incoming.is_empty()); + incoming + } + + pub(super) fn start_split_accept(&mut self, incoming: Incoming, now: Instant) -> Accepting { + let mut buf = Vec::new(); + self.endpoint + .start_accept(incoming, now, &mut buf, None) + .unwrap() + } + + pub(super) fn disable_new_connections(&mut self) { + self.endpoint.set_server_config(None); + } + + pub(super) fn finish_split_accept(&mut self, accepting: Accepting) -> ConnectionHandle { + let Ok(accepted) = accepting.finish_without_endpoint() else { + panic!("split accept unexpectedly failed") + }; + let (ch, conn) = self.endpoint.finish_accept(accepted); + self.connections.insert(ch, conn); + ch + } + + /// Like `finish_split_accept`, but expects the off-lock handshake to fail. Runs the + /// error-cleanup path (`finish_accept_error`) and returns the resulting cause. + pub(super) fn finish_split_accept_error(&mut self, accepting: Accepting) -> ConnectionError { + let mut buf = Vec::new(); + let Err(error) = accepting.finish_without_endpoint() else { + panic!("split accept unexpectedly succeeded") + }; + self.endpoint.finish_accept_error(error, &mut buf).cause + } + pub(super) fn retry(&mut self, incoming: Incoming) { let mut buf = Vec::new(); let transmit = self.endpoint.retry(incoming, &mut buf).unwrap(); @@ -610,9 +653,10 @@ fn server_crypto_inner( ) }); - let mut config = QuicServerConfig::inner(vec![cert], key).unwrap(); + let mut config = + QuicServerConfig::inner_with_provider(vec![cert], key, test_provider()).unwrap(); if let Some(alpn) = alpn { - config.alpn_protocols = alpn; + config.alpn_protocols = alpn.into_iter().map(ApplicationProtocol::from).collect(); } config.try_into().unwrap() @@ -651,19 +695,44 @@ fn client_crypto_inner( roots.add(cert).unwrap(); } - let mut inner = QuicClientConfig::inner( - WebPkiServerVerifier::builder_with_provider(Arc::new(roots), configured_provider()) - .build() - .unwrap(), - ); + let provider = test_provider(); + let verifier = WebPkiServerVerifier::builder(Arc::new(roots), &provider) + .build() + .unwrap(); + let mut inner = QuicClientConfig::inner_with_provider(Arc::new(verifier), provider); inner.key_log = Arc::new(KeyLogFile::new()); if let Some(alpn) = alpn { - inner.alpn_protocols = alpn; + inner.alpn_protocols = alpn.into_iter().map(ApplicationProtocol::from).collect(); } inner.try_into().unwrap() } +#[cfg(all(feature = "rustls-aws-lc-rs-fips", not(feature = "rustls-ring")))] +fn test_provider() -> Arc { + Arc::new(CryptoProvider { + kx_groups: std::borrow::Cow::Owned(vec![rustls_aws_lc_rs::kx_group::SECP256R1]), + ..rustls_aws_lc_rs::DEFAULT_FIPS_PROVIDER + }) +} + +#[cfg(all( + feature = "rustls-aws-lc-rs", + not(feature = "rustls-aws-lc-rs-fips"), + not(feature = "rustls-ring") +))] +fn test_provider() -> Arc { + Arc::new(CryptoProvider { + kx_groups: std::borrow::Cow::Owned(vec![rustls_aws_lc_rs::kx_group::X25519]), + ..rustls_aws_lc_rs::DEFAULT_PROVIDER + }) +} + +#[cfg(feature = "rustls-ring")] +fn test_provider() -> Arc { + Arc::new(rustls_ring::DEFAULT_PROVIDER) +} + pub(super) fn min_opt(x: Option, y: Option) -> Option { match (x, y) { (Some(x), Some(y)) => Some(cmp::min(x, y)), diff --git a/quinn/Cargo.toml b/quinn/Cargo.toml index 4f27ab99bf..a5089b0eae 100644 --- a/quinn/Cargo.toml +++ b/quinn/Cargo.toml @@ -28,10 +28,10 @@ platform-verifier = ["proto/platform-verifier"] # For backwards compatibility, `rustls` forwards to `rustls-ring` rustls = ["rustls-ring"] # Enable rustls with the `aws-lc-rs` crypto provider -rustls-aws-lc-rs = ["dep:rustls", "aws-lc-rs", "proto/rustls-aws-lc-rs", "proto/aws-lc-rs"] -rustls-aws-lc-rs-fips = ["dep:rustls", "aws-lc-rs-fips", "proto/rustls-aws-lc-rs-fips", "proto/aws-lc-rs-fips"] +rustls-aws-lc-rs = ["__rustls", "dep:rustls-aws-lc-rs", "aws-lc-rs", "proto/rustls-aws-lc-rs", "proto/aws-lc-rs"] +rustls-aws-lc-rs-fips = ["rustls-aws-lc-rs", "aws-lc-rs-fips", "proto/rustls-aws-lc-rs-fips", "proto/aws-lc-rs-fips"] # Enable rustls with the `ring` crypto provider -rustls-ring = ["dep:rustls", "ring", "proto/rustls-ring", "proto/ring"] +rustls-ring = ["__rustls", "dep:rustls-ring", "ring", "proto/rustls-ring", "proto/ring"] # Enable the `ring` crypto provider. # Outside wasm*-unknown-unknown targets, this enables `Endpoint::client` and `Endpoint::server` conveniences. ring = ["proto/ring"] @@ -40,15 +40,16 @@ runtime-smol = ["dep:async-io", "dep:smol"] # Configure `tracing` to log events via `log` if no `tracing` subscriber exists. tracing-log = ["tracing/log", "proto/tracing-log", "udp/tracing-log"] -# Enable rustls logging -rustls-log = ["rustls?/logging"] +# Enable rustls tracing (feature name retained for backwards compatibility) +rustls-log = ["rustls?/tracing"] # Enable qlog support qlog = ["proto/qlog"] # Internal (PRIVATE!) features used to aid testing. # Don't rely on these whatsoever. They may disappear at any time. -__rustls-post-quantum-test = ["rustls/prefer-post-quantum", "rustls-aws-lc-rs", "proto/__rustls-post-quantum-test"] +__rustls-post-quantum-test = ["rustls-aws-lc-rs", "proto/__rustls-post-quantum-test"] +__rustls = ["dep:rustls"] [dependencies] async-io = { workspace = true, optional = true } @@ -57,8 +58,10 @@ bytes = { workspace = true } futures-io = { workspace = true, optional = true } rustc-hash = { workspace = true } pin-project-lite = { workspace = true } -proto = { package = "quinn-proto", path = "../quinn-proto", version = "0.12.0", default-features = false } +proto = { package = "quinn-proto", path = "../quinn-proto", version = "0.12.0", default-features = false, features = ["__internal_split_accept"] } rustls = { workspace = true, optional = true } +rustls-aws-lc-rs = { workspace = true, optional = true } +rustls-ring = { workspace = true, optional = true } smol = { workspace = true, optional = true } thiserror = { workspace = true } tracing = { workspace = true } @@ -78,6 +81,7 @@ bencher = { workspace = true } directories-next = { workspace = true } rand = { workspace = true } rcgen = { workspace = true } +rustls-util = { workspace = true } clap = { workspace = true } tokio = { workspace = true, features = ["rt", "rt-multi-thread", "time", "macros", "test-util"] } tracing-subscriber = { workspace = true } @@ -91,23 +95,23 @@ workspace = true [[example]] name = "server" -required-features = ["rustls-ring"] +required-features = ["__rustls"] [[example]] name = "client" -required-features = ["rustls-ring"] +required-features = ["__rustls"] [[example]] name = "insecure_connection" -required-features = ["rustls-ring"] +required-features = ["__rustls"] [[example]] name = "single_socket" -required-features = ["rustls-ring"] +required-features = ["__rustls"] [[example]] name = "connection" -required-features = ["rustls-ring"] +required-features = ["__rustls"] [[test]] name = "post_quantum" @@ -116,8 +120,8 @@ required-features = ["__rustls-post-quantum-test"] [[bench]] name = "bench" harness = false -required-features = ["rustls-ring"] +required-features = ["__rustls"] [package.metadata.docs.rs] # all non-default features except fips (cannot build on docs.rs environment) -features = ["lock_tracking", "rustls-aws-lc-rs", "rustls-ring", "runtime-tokio", "runtime-smol", "tracing-log", "rustls-log"] +features = ["lock_tracking", "rustls-aws-lc-rs", "rustls-ring", "platform-verifier", "runtime-tokio", "runtime-smol", "tracing-log", "rustls-log"] diff --git a/quinn/examples/client.rs b/quinn/examples/client.rs index 40d29c9dbe..40763534ac 100644 --- a/quinn/examples/client.rs +++ b/quinn/examples/client.rs @@ -20,6 +20,16 @@ use url::Url; mod common; +#[cfg(feature = "rustls-aws-lc-rs")] +fn default_provider() -> rustls::crypto::CryptoProvider { + rustls_aws_lc_rs::DEFAULT_PROVIDER +} + +#[cfg(all(not(feature = "rustls-aws-lc-rs"), feature = "rustls-ring"))] +fn default_provider() -> rustls::crypto::CryptoProvider { + rustls_ring::DEFAULT_PROVIDER +} + /// HTTP/0.9 over QUIC client #[derive(Parser, Debug)] #[clap(name = "client")] @@ -92,13 +102,13 @@ async fn run(options: Opt) -> Result<()> { } } } - let mut client_crypto = rustls::ClientConfig::builder() + let mut client_crypto = rustls::ClientConfig::builder(Arc::new(default_provider())) .with_root_certificates(roots) - .with_no_client_auth(); + .with_no_client_auth()?; client_crypto.alpn_protocols = common::ALPN_QUIC_HTTP.iter().map(|&x| x.into()).collect(); if options.keylog { - client_crypto.key_log = Arc::new(rustls::KeyLogFile::new()); + client_crypto.key_log = Arc::new(rustls_util::KeyLogFile::new()); } let client_config = diff --git a/quinn/examples/insecure_connection.rs b/quinn/examples/insecure_connection.rs index 39fa5ba231..20d9e9b250 100644 --- a/quinn/examples/insecure_connection.rs +++ b/quinn/examples/insecure_connection.rs @@ -4,17 +4,27 @@ use std::{ error::Error, + hash::Hasher, net::{IpAddr, Ipv4Addr, SocketAddr}, sync::Arc, }; use proto::crypto::rustls::QuicClientConfig; use quinn::{ClientConfig, Endpoint}; -use rustls::pki_types::{CertificateDer, ServerName, UnixTime}; mod common; use common::make_server_endpoint; +#[cfg(feature = "rustls-aws-lc-rs")] +fn default_provider() -> rustls::crypto::CryptoProvider { + rustls_aws_lc_rs::DEFAULT_PROVIDER +} + +#[cfg(all(not(feature = "rustls-aws-lc-rs"), feature = "rustls-ring"))] +fn default_provider() -> rustls::crypto::CryptoProvider { + rustls_ring::DEFAULT_PROVIDER +} + #[tokio::main] async fn main() -> Result<(), Box> { // server and client are running on the same thread asynchronously @@ -40,10 +50,10 @@ async fn run_client(server_addr: SocketAddr) -> Result<(), Box); impl SkipServerVerification { fn new() -> Arc { - Arc::new(Self(Arc::new(rustls::crypto::ring::default_provider()))) + Arc::new(Self(Arc::new(default_provider()))) } } -impl rustls::client::danger::ServerCertVerifier for SkipServerVerification { - fn verify_server_cert( +impl rustls::client::danger::ServerVerifier for SkipServerVerification { + fn verify_identity<'a>( &self, - _end_entity: &CertificateDer<'_>, - _intermediates: &[CertificateDer<'_>], - _server_name: &ServerName<'_>, - _ocsp: &[u8], - _now: UnixTime, - ) -> Result { - Ok(rustls::client::danger::ServerCertVerified::assertion()) + identity: &rustls::client::danger::ServerIdentity<'a, '_>, + ) -> Result, rustls::Error> { + Ok(rustls::crypto::VerifiedIdentity::assertion( + identity.identity.clone(), + )) } fn verify_tls12_signature( &self, - message: &[u8], - cert: &CertificateDer<'_>, - dss: &rustls::DigitallySignedStruct, + input: &rustls::client::danger::SignatureVerificationInput<'_>, ) -> Result { - rustls::crypto::verify_tls12_signature( - message, - cert, - dss, - &self.0.signature_verification_algorithms, - ) + rustls::crypto::verify_tls12_signature(input, &self.0.signature_verification_algorithms) } fn verify_tls13_signature( &self, - message: &[u8], - cert: &CertificateDer<'_>, - dss: &rustls::DigitallySignedStruct, + input: &rustls::client::danger::SignatureVerificationInput<'_>, ) -> Result { - rustls::crypto::verify_tls13_signature( - message, - cert, - dss, - &self.0.signature_verification_algorithms, - ) + rustls::crypto::verify_tls13_signature(input, &self.0.signature_verification_algorithms) } - fn supported_verify_schemes(&self) -> Vec { + fn supported_verify_schemes(&self) -> Vec { self.0.signature_verification_algorithms.supported_schemes() } + + fn request_ocsp_response(&self) -> bool { + false + } + + fn hash_config(&self, h: &mut dyn Hasher) { + for scheme in self.supported_verify_schemes() { + h.write_u16(scheme.0); + } + } } diff --git a/quinn/examples/server.rs b/quinn/examples/server.rs index b0787c50aa..352b94cd53 100644 --- a/quinn/examples/server.rs +++ b/quinn/examples/server.rs @@ -13,12 +13,25 @@ use std::{ use anyhow::{Context, Result, anyhow, bail}; use clap::Parser; use proto::crypto::rustls::QuicServerConfig; -use rustls::pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer, pem::PemObject}; +use rustls::{ + crypto::Identity, + pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer, pem::PemObject}, +}; use tracing::instrument::Instrument as _; use tracing::{error, info, info_span}; mod common; +#[cfg(feature = "rustls-aws-lc-rs")] +fn default_provider() -> rustls::crypto::CryptoProvider { + rustls_aws_lc_rs::DEFAULT_PROVIDER +} + +#[cfg(all(not(feature = "rustls-aws-lc-rs"), feature = "rustls-ring"))] +fn default_provider() -> rustls::crypto::CryptoProvider { + rustls_ring::DEFAULT_PROVIDER +} + #[derive(Parser, Debug)] #[clap(name = "server")] struct Opt { @@ -119,12 +132,12 @@ async fn run(options: Opt) -> Result<()> { (vec![cert], key) }; - let mut server_crypto = rustls::ServerConfig::builder() + let mut server_crypto = rustls::ServerConfig::builder(Arc::new(default_provider())) .with_no_client_auth() - .with_single_cert(certs, key)?; + .with_single_cert(Arc::new(Identity::from_cert_chain(certs)?), key)?; server_crypto.alpn_protocols = common::ALPN_QUIC_HTTP.iter().map(|&x| x.into()).collect(); if options.keylog { - server_crypto.key_log = Arc::new(rustls::KeyLogFile::new()); + server_crypto.key_log = Arc::new(rustls_util::KeyLogFile::new()); } let mut server_config = diff --git a/quinn/src/connection.rs b/quinn/src/connection.rs index 812eada8ea..89976dc56d 100644 --- a/quinn/src/connection.rs +++ b/quinn/src/connection.rs @@ -654,7 +654,7 @@ impl Connection { /// /// The dynamic type returned is determined by the configured /// [`Session`](proto::crypto::Session). For the default `rustls` session, the return value can - /// be [`downcast`](Box::downcast) to a Vec<[rustls::pki_types::CertificateDer]> + /// be [`downcast`](Box::downcast) to rustls::crypto::Identity<'static>. pub fn peer_identity(&self) -> Option> { self.0 .state @@ -701,7 +701,7 @@ impl Connection { .state .lock("export_keying_material") .inner - .crypto_session() + .crypto_session_mut() .export_keying_material(output, label, context) } diff --git a/quinn/src/endpoint.rs b/quinn/src/endpoint.rs index 1df429183c..81a5d02d8a 100644 --- a/quinn/src/endpoint.rs +++ b/quinn/src/endpoint.rs @@ -334,6 +334,8 @@ impl Endpoint { }); } self.inner.shared.incoming.notify_waiters(); + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + self.inner.shared.notify_all_acceptors(); } /// Wait for all connections on the endpoint to be cleanly shut down @@ -350,7 +352,7 @@ impl Endpoint { loop { { let endpoint = &mut *self.inner.state.lock().unwrap(); - if endpoint.recv_state.connections.is_empty() { + if endpoint.is_idle() { break; } // Construct future while lock is held to avoid race @@ -400,16 +402,14 @@ impl Future for EndpointDriver { let now = endpoint.runtime.now(); let mut keep_going = false; - keep_going |= endpoint.drive_recv(cx, now)?; + keep_going |= endpoint.drive_recv(cx, now, &self.0.shared)?; keep_going |= endpoint.handle_events(cx, &self.0.shared); if !endpoint.recv_state.incoming.is_empty() { self.0.shared.incoming.notify_waiters(); } - if self.0.shared.ref_count.load(Ordering::Relaxed) == 0 - && endpoint.recv_state.connections.is_empty() - { + if self.0.shared.ref_count.load(Ordering::Relaxed) == 0 && endpoint.is_idle() { Poll::Ready(Ok(())) } else { drop(endpoint); @@ -429,6 +429,8 @@ impl Drop for EndpointDriver { let mut endpoint = self.0.state.lock().unwrap(); endpoint.driver_lost = true; self.0.shared.incoming.notify_waiters(); + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + self.0.shared.notify_all_acceptors(); // Drop all outgoing channels, signaling the termination of the endpoint to the associated // connections. endpoint.recv_state.connections.senders.clear(); @@ -441,37 +443,235 @@ pub(crate) struct EndpointInner { pub(crate) shared: Shared, } +#[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] +pub(crate) struct RustlsAcceptorStart { + pub(crate) accepting: proto::Accepting, + pub(crate) acceptor: proto::RustlsAcceptor, + pub(crate) incoming_idx: usize, + pub(crate) notify: Arc, + pub(crate) deadline: Option, + pub(crate) timer: Option>>, +} + impl EndpointInner { + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + pub(crate) fn runtime_now(&self) -> Instant { + self.state.lock().unwrap().runtime.now() + } + pub(crate) fn accept( &self, incoming: proto::Incoming, server_config: Option>, ) -> Result { + let mut response_buffer = Vec::new(); + + // Phase 1: acquire lock, do the minimum work that needs endpoint state. + let accepting = { + let mut state = self.state.lock().unwrap(); + let now = state.runtime.now(); + match state + .inner + .start_accept(incoming, now, &mut response_buffer, server_config) + { + Ok(accepting) => accepting, + Err(error) => { + if let Some(transmit) = error.response { + respond(transmit, &response_buffer, &mut state.sender); + } + return Err(error.cause); + } + } + }; + + // Phase 2: do the expensive connection construction and first-packet handling + // without holding the lock. + let result = accepting.finish_without_endpoint(); + + // Phase 3: re-acquire the lock to finalize the reserved endpoint state. let mut state = self.state.lock().unwrap(); + match result { + Ok(accepted) => { + state.stats.accepted_handshakes += 1; + let sender = state.socket.create_sender(); + let runtime = state.runtime.clone(); + let (handle, conn) = state.inner.finish_accept(accepted); + let connecting = state + .recv_state + .connections + .insert(handle, conn, sender, runtime); + Ok(connecting) + } + Err(error) => { + let error = state.inner.finish_accept_error(error, &mut response_buffer); + if let Some(transmit) = error.response { + respond(transmit, &response_buffer, &mut state.sender); + } + // If this failed accept was the endpoint's last in-flight work, wake idle + // waiters explicitly: the accept never became a connection, so no Drained + // event will fire to do it. + if state.is_idle() { + self.shared.idle.notify_waiters(); + } + Err(error.cause) + } + } + } + + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + pub(crate) fn start_rustls_acceptor( + &self, + incoming: proto::Incoming, + ) -> Result { + let Ok(acceptor) = proto::RustlsAcceptor::new(&incoming) else { + self.ignore(incoming); + return Err(ConnectionError::VersionMismatch); + }; let mut response_buffer = Vec::new(); - let now = state.runtime.now(); - match state - .inner - .accept(incoming, now, &mut response_buffer, server_config) + let (accepting, deadline, timer) = { + let mut state = self.state.lock().unwrap(); + let now = state.runtime.now(); + let accepting = + match state + .inner + .start_accept(incoming, now, &mut response_buffer, None) + { + Ok(accepting) => accepting, + Err(error) => { + if let Some(transmit) = error.response { + respond(transmit, &response_buffer, &mut state.sender); + } + return Err(error.cause); + } + }; + let deadline = accepting.idle_timeout_deadline(); + let timer = deadline.map(|deadline| state.runtime.new_timer(deadline)); + (accepting, deadline, timer) + }; + let incoming_idx = accepting.incoming_idx(); + let notify = self.shared.register_acceptor(incoming_idx); + Ok(RustlsAcceptorStart { + accepting, + acceptor, + incoming_idx, + notify, + deadline, + timer, + }) + } + + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + pub(crate) fn poll_rustls_acceptor( + &self, + accepting: &mut proto::Accepting, + acceptor: &mut proto::RustlsAcceptor, + ) -> Result, ConnectionError> { { - Ok((handle, conn)) => { + let state = self.state.lock().unwrap(); + if state.driver_lost || state.recv_state.connections.close.is_some() { + return Err(ConnectionError::LocallyClosed); + } + state.inner.buffer_rustls_acceptor_input(accepting); + } + + let mut response_buffer = Vec::new(); + let (accepted, response) = + accepting.poll_rustls_acceptor(acceptor, &mut response_buffer)?; + + let mut state = self.state.lock().unwrap(); + if state.driver_lost || state.recv_state.connections.close.is_some() { + return Err(ConnectionError::LocallyClosed); + } + if let Some(transmit) = response { + respond(transmit, &response_buffer, &mut state.sender); + } + Ok(accepted) + } + + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + pub(crate) fn accept_rustls( + &self, + mut accepting: proto::Accepting, + accepted: proto::RustlsAccepted, + server_config: Option>, + ) -> Result { + let mut response_buffer = Vec::new(); + + let selection = { + let mut state = self.state.lock().unwrap(); + if state.driver_lost || state.recv_state.connections.close.is_some() { + Err(ConnectionError::LocallyClosed) + } else { + let now = state.runtime.now(); + state + .inner + .select_accepting_config(&mut accepting, server_config, now) + } + }; + if let Err(cause) = selection { + self.fail_rustls_accepting(accepting, cause.clone()); + return Err(cause); + } + + let result = accepting.finish_from_rustls(accepted); + let mut state = self.state.lock().unwrap(); + match result { + Ok(accepted) => { state.stats.accepted_handshakes += 1; let sender = state.socket.create_sender(); let runtime = state.runtime.clone(); + let (handle, conn) = state.inner.finish_accept(accepted); Ok(state .recv_state .connections .insert(handle, conn, sender, runtime)) } Err(error) => { + let error = state.inner.finish_accept_error(error, &mut response_buffer); if let Some(transmit) = error.response { respond(transmit, &response_buffer, &mut state.sender); } + if state.is_idle() { + self.shared.idle.notify_waiters(); + } Err(error.cause) } } } + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + pub(crate) fn fail_rustls_accepting( + &self, + accepting: proto::Accepting, + cause: ConnectionError, + ) { + self.shared.unregister_acceptor(accepting.incoming_idx()); + let mut state = self.state.lock().unwrap(); + state.stats.refused_handshakes += 1; + let mut response_buffer = Vec::new(); + let error = state + .inner + .fail_accepting(accepting, cause, &mut response_buffer); + if let Some(transmit) = error.response { + respond(transmit, &response_buffer, &mut state.sender); + } + if state.is_idle() { + self.shared.idle.notify_waiters(); + } + } + + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + pub(crate) fn refuse_rustls_accepting(&self, accepting: proto::Accepting) { + self.fail_rustls_accepting( + accepting, + proto::TransportError::new( + proto::TransportErrorCode::CONNECTION_REFUSED, + String::new(), + ) + .into(), + ); + } + pub(crate) fn refuse(&self, incoming: proto::Incoming) { let mut state = self.state.lock().unwrap(); state.stats.refused_handshakes += 1; @@ -516,17 +716,60 @@ pub(crate) struct State { #[derive(Debug)] pub(crate) struct Shared { incoming: Notify, + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + acceptors: Mutex>>, idle: Notify, /// Number of live handles that can be used to initiate or handle I/O; excludes the driver ref_count: AtomicUsize, } +#[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] +impl Shared { + fn register_acceptor(&self, incoming_idx: usize) -> Arc { + let notify = Arc::new(Notify::new()); + let old = self + .acceptors + .lock() + .unwrap() + .insert(incoming_idx, notify.clone()); + debug_assert!(old.is_none()); + notify + } + + pub(crate) fn unregister_acceptor(&self, incoming_idx: usize) { + self.acceptors.lock().unwrap().remove(&incoming_idx); + } + + fn notify_acceptor(&self, incoming_idx: usize) { + let notify = self.acceptors.lock().unwrap().get(&incoming_idx).cloned(); + if let Some(notify) = notify { + notify.notify_waiters(); + } + } + + fn notify_all_acceptors(&self) { + let notifiers: Vec<_> = self.acceptors.lock().unwrap().values().cloned().collect(); + for notify in notifiers { + notify.notify_waiters(); + } + } +} + impl State { - fn drive_recv(&mut self, cx: &mut Context<'_>, now: Instant) -> Result { + fn is_idle(&self) -> bool { + self.recv_state.connections.is_empty() && self.inner.pending_accepts() == 0 + } + + fn drive_recv( + &mut self, + cx: &mut Context<'_>, + now: Instant, + shared: &Shared, + ) -> Result { let get_time = || self.runtime.now(); self.recv_state.recv_limiter.start_cycle(get_time); + let mut previous_progress = PollProgress::default(); if let Some(socket) = &mut self.prev_socket { - // We don't care about the `PollProgress` from old sockets. let poll_res = self.recv_state.poll_socket( cx, &mut self.inner, @@ -534,9 +777,11 @@ impl State { &mut self.sender, &*self.runtime, now, + shared, ); - if poll_res.is_err() { - self.prev_socket = None; + match poll_res { + Ok(progress) => previous_progress = progress, + Err(_) => self.prev_socket = None, } }; let poll_res = self.recv_state.poll_socket( @@ -546,6 +791,7 @@ impl State { &mut self.sender, &*self.runtime, now, + shared, ); self.recv_state.recv_limiter.finish_cycle(get_time); let poll_res = poll_res?; @@ -554,7 +800,7 @@ impl State { // one anymore. TODO: Account for multiple outgoing connections. self.prev_socket = None; } - Ok(poll_res.keep_going) + Ok(previous_progress.keep_going || poll_res.keep_going) } fn handle_events(&mut self, cx: &mut Context<'_>, shared: &Shared) -> bool { @@ -569,7 +815,7 @@ impl State { if event.is_drained() { self.recv_state.connections.senders.remove(&ch); - if self.recv_state.connections.is_empty() { + if self.is_idle() { shared.idle.notify_waiters(); } } @@ -754,6 +1000,8 @@ impl EndpointRef { Self(Arc::new(EndpointInner { shared: Shared { incoming: Notify::new(), + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + acceptors: Mutex::new(FxHashMap::default()), idle: Notify::new(), ref_count: AtomicUsize::new(0), }, @@ -837,6 +1085,7 @@ impl RecvState { } } + #[allow(clippy::too_many_arguments)] fn poll_socket( &mut self, cx: &mut Context<'_>, @@ -845,6 +1094,7 @@ impl RecvState { sender: &mut Pin>, runtime: &dyn Runtime, now: Instant, + _shared: &Shared, ) -> Result { let mut received_connection_packet = false; let mut metas = [RecvMeta::default(); BATCH_SIZE]; @@ -908,6 +1158,13 @@ impl RecvState { .unwrap() .send(ConnectionEvent::Proto(event)); } + Some(DatagramEvent::IncomingData(_incoming_idx)) => { + #[cfg(any( + feature = "rustls-aws-lc-rs", + feature = "rustls-ring" + ))] + _shared.notify_acceptor(_incoming_idx); + } Some(DatagramEvent::Response(transmit)) => { respond(transmit, &response_buffer, sender); } diff --git a/quinn/src/incoming.rs b/quinn/src/incoming.rs index 47471bdfbf..353d9e6eff 100644 --- a/quinn/src/incoming.rs +++ b/quinn/src/incoming.rs @@ -6,8 +6,12 @@ use std::{ task::{Context, Poll}, }; +#[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] +use crate::runtime::AsyncTimer; use proto::{ConnectionError, ConnectionId, ServerConfig}; use thiserror::Error; +#[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] +use tokio::sync::futures::OwnedNotified; use crate::{ connection::{Connecting, Connection}, @@ -40,6 +44,35 @@ impl Incoming { state.endpoint.accept(state.inner, Some(server_config)) } + /// Start reading this connection's rustls ClientHello before choosing a server config. + /// + /// The returned future buffers Initial and 0-RTT datagrams for this connection while it waits + /// for enough ClientHello data. Once it resolves, inspect the ClientHello and continue with + /// [`Accepted::accept_with`]. Dropping either the future or the resulting [`Accepted`] refuses + /// the connection and releases its buffered protocol state. Resource limits applied before + /// selection, including ClientHello and incoming-datagram buffering limits, come from the + /// endpoint configuration captured when this method is called; selecting another configuration + /// does not retroactively change them. + #[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] + pub fn acceptor(mut self) -> Result { + let state = self.0.take().unwrap(); + let endpoint = state.endpoint; + let start = endpoint.start_rustls_acceptor(state.inner)?; + let notify = Box::pin(start.notify.clone().notified_owned()); + Ok(Acceptor { + state: Some(AcceptorState { + accepting: start.accepting, + endpoint, + acceptor: start.acceptor, + incoming_idx: start.incoming_idx, + notify_source: start.notify, + notify, + deadline: start.deadline, + timer: start.timer, + }), + }) + } + /// Reject this incoming connection attempt pub fn refuse(mut self) { let state = self.0.take().unwrap(); @@ -115,6 +148,185 @@ struct State { endpoint: EndpointRef, } +/// Future that resolves once rustls has read the incoming ClientHello. +/// +/// Creating an `Acceptor` reserves one pending incoming-connection slot. Initial and 0-RTT +/// datagrams received while it is pending are buffered. The future observes the server's idle +/// timeout, and dropping it refuses the attempt and releases the reservation. +#[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] +pub struct Acceptor { + state: Option, +} + +#[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] +struct AcceptorState { + accepting: proto::Accepting, + endpoint: EndpointRef, + acceptor: proto::RustlsAcceptor, + incoming_idx: usize, + notify_source: Arc, + notify: Pin>, + deadline: Option, + timer: Option>>, +} + +#[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] +impl Future for Acceptor { + type Output = Result; + + fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { + loop { + let state = self.state.as_mut().expect("polled after completion"); + // Register before inspecting endpoint buffers so a datagram received between the + // inspection and the pending return cannot be missed. + state.notify.as_mut().enable(); + match state + .endpoint + .poll_rustls_acceptor(&mut state.accepting, &mut state.acceptor) + { + Ok(Some(accepted)) => { + if state + .accepting + .idle_timeout_deadline() + .is_some_and(|deadline| state.endpoint.runtime_now() >= deadline) + { + let state = self.state.take().unwrap(); + state + .endpoint + .fail_rustls_accepting(state.accepting, ConnectionError::TimedOut); + return Poll::Ready(Err(ConnectionError::TimedOut)); + } + let state = self.state.take().unwrap(); + state + .endpoint + .shared + .unregister_acceptor(state.incoming_idx); + return Poll::Ready(Ok(Accepted { + state: Some(AcceptedState { + accepting: state.accepting, + endpoint: state.endpoint, + accepted, + }), + })); + } + Ok(None) => { + let deadline = state.accepting.idle_timeout_deadline(); + if deadline != state.deadline { + if let (Some(timer), Some(deadline)) = (&mut state.timer, deadline) { + timer.as_mut().reset(deadline); + } + state.deadline = deadline; + } + let timed_out = state + .deadline + .is_some_and(|deadline| state.endpoint.runtime_now() >= deadline) + || state + .timer + .as_mut() + .is_some_and(|timer| timer.as_mut().poll(cx).is_ready()); + if timed_out { + let state = self.state.take().unwrap(); + state + .endpoint + .fail_rustls_accepting(state.accepting, ConnectionError::TimedOut); + return Poll::Ready(Err(ConnectionError::TimedOut)); + } + if state.notify.as_mut().poll(cx).is_ready() { + state.notify = Box::pin(state.notify_source.clone().notified_owned()); + continue; + } + return Poll::Pending; + } + Err(error) => { + let state = self.state.take().unwrap(); + state + .endpoint + .fail_rustls_accepting(state.accepting, error.clone()); + return Poll::Ready(Err(error)); + } + } + } + } +} + +#[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] +impl Drop for Acceptor { + fn drop(&mut self) { + if let Some(state) = self.state.take() { + state.endpoint.refuse_rustls_accepting(state.accepting); + } + } +} + +/// An incoming connection whose ClientHello has been read. +/// +/// Inspect [`client_hello()`](Self::client_hello), asynchronously choose a [`ServerConfig`], then +/// call [`accept_with()`](Self::accept_with) to continue the same TLS and QUIC handshake. Incoming +/// datagrams remain buffered during selection. Holding this value also keeps one pending incoming +/// slot reserved, so applications should bound slow configuration lookups and drop or refuse the +/// attempt if selection takes too long. Dropping this value refuses the connection and releases +/// all reserved state. +#[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] +pub struct Accepted { + state: Option, +} + +#[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] +struct AcceptedState { + accepting: proto::Accepting, + endpoint: EndpointRef, + accepted: proto::RustlsAccepted, +} + +#[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] +impl Accepted { + /// Get the rustls ClientHello for this connection. + pub fn client_hello(&self) -> rustls::server::ClientHello<'_> { + self.state.as_ref().unwrap().accepted.client_hello() + } + + /// Continue the QUIC handshake using the endpoint's configured server configuration. + /// + /// The configuration's cryptographic implementation must support continuing a staged rustls + /// handshake, as Quinn's rustls-backed configurations do. + pub fn accept(mut self) -> Result { + let state = self.state.take().unwrap(); + state + .endpoint + .accept_rustls(state.accepting, state.accepted, None) + } + + /// Continue the QUIC handshake using a custom server configuration. + /// + /// This selects the complete Quinn configuration, including both TLS and transport policy. + /// Its cryptographic implementation must support continuing a staged rustls handshake, as + /// Quinn's rustls-backed configurations do. + pub fn accept_with( + mut self, + server_config: Arc, + ) -> Result { + let state = self.state.take().unwrap(); + state + .endpoint + .accept_rustls(state.accepting, state.accepted, Some(server_config)) + } + + /// Reject this incoming connection attempt. + pub fn refuse(mut self) { + let state = self.state.take().unwrap(); + state.endpoint.refuse_rustls_accepting(state.accepting); + } +} + +#[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] +impl Drop for Accepted { + fn drop(&mut self) { + if let Some(state) = self.state.take() { + state.endpoint.refuse_rustls_accepting(state.accepting); + } + } +} + /// Error for attempting to retry an [`Incoming`] which already bears a token from a previous retry #[derive(Debug, Error)] #[error("retry() with validated Incoming")] diff --git a/quinn/src/lib.rs b/quinn/src/lib.rs index 649801fff8..d65c409690 100644 --- a/quinn/src/lib.rs +++ b/quinn/src/lib.rs @@ -77,6 +77,8 @@ pub use crate::connection::{ RemoteAddressWatcher, SendDatagram, SendDatagramError, }; pub use crate::endpoint::{Accept, Endpoint, EndpointStats}; +#[cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] +pub use crate::incoming::{Accepted, Acceptor}; pub use crate::incoming::{Incoming, IncomingFuture, RetryError}; pub use crate::recv_stream::{ReadError, ReadExactError, ReadToEndError, RecvStream, ResetError}; #[cfg(feature = "runtime-smol")] diff --git a/quinn/src/tests.rs b/quinn/src/tests.rs index 359e6f98bd..30a55507ed 100755 --- a/quinn/src/tests.rs +++ b/quinn/src/tests.rs @@ -1,10 +1,5 @@ #![cfg(any(feature = "rustls-aws-lc-rs", feature = "rustls-ring"))] -#[cfg(all(feature = "rustls-aws-lc-rs", not(feature = "rustls-ring")))] -use rustls::crypto::aws_lc_rs::default_provider; -#[cfg(feature = "rustls-ring")] -use rustls::crypto::ring::default_provider; - use std::{ convert::TryInto, future::Future, @@ -13,7 +8,7 @@ use std::{ pin::pin, str, sync::{ - Arc, + Arc, Condvar, Mutex, atomic::{AtomicUsize, Ordering}, }, task::{Context, Poll, RawWaker, RawWakerVTable, Waker}, @@ -22,11 +17,18 @@ use std::{ use crate::runtime::TokioRuntime; use crate::{Duration, Instant}; use bytes::Bytes; -use proto::{RandomConnectionIdGenerator, crypto::rustls::QuicClientConfig}; +use proto::{ + ConnectionId, RandomConnectionIdGenerator, + crypto::rustls::QuicClientConfig, + crypto::{Keys, ServerConfig as ProtoServerConfig, Session, UnsupportedVersion}, + transport_parameters::TransportParameters, +}; use rand::{Rng, SeedableRng, rngs::StdRng}; use rustls::{ RootCertStore, + crypto::Identity, pki_types::{CertificateDer, PrivateKeyDer, PrivatePkcs8KeyDer}, + server::WebPkiClientVerifier, }; use tokio::time::{sleep, timeout}; use tokio::{ @@ -39,6 +41,31 @@ use tracing_subscriber::EnvFilter; use super::{ClientConfig, Endpoint, EndpointConfig, RecvStream, SendStream, TransportConfig}; +#[cfg(all(feature = "rustls-aws-lc-rs-fips", not(feature = "rustls-ring")))] +fn default_provider() -> rustls::crypto::CryptoProvider { + rustls::crypto::CryptoProvider { + kx_groups: std::borrow::Cow::Owned(vec![rustls_aws_lc_rs::kx_group::SECP256R1]), + ..rustls_aws_lc_rs::DEFAULT_FIPS_PROVIDER + } +} + +#[cfg(all( + feature = "rustls-aws-lc-rs", + not(feature = "rustls-aws-lc-rs-fips"), + not(feature = "rustls-ring") +))] +fn default_provider() -> rustls::crypto::CryptoProvider { + rustls::crypto::CryptoProvider { + kx_groups: std::borrow::Cow::Owned(vec![rustls_aws_lc_rs::kx_group::X25519]), + ..rustls_aws_lc_rs::DEFAULT_PROVIDER + } +} + +#[cfg(feature = "rustls-ring")] +fn default_provider() -> rustls::crypto::CryptoProvider { + rustls_ring::DEFAULT_PROVIDER +} + #[test] fn handshake_timeout() { let _guard = subscribe(); @@ -272,6 +299,95 @@ fn endpoint_with_config(transport_config: TransportConfig) -> Endpoint { EndpointFactory::new().endpoint_with_config(transport_config) } +#[tokio::test] +async fn incoming_acceptor_selects_crypto_and_transport_config() { + let _guard = subscribe(); + let cert = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap(); + let cert_der = CertificateDer::from(cert.cert); + let selected_alpn = b"selected"; + + let server_config = |alpn: &[u8], max_uni: u32| { + let mut crypto = rustls::ServerConfig::builder(Arc::new(default_provider())) + .with_no_client_auth() + .with_single_cert( + Arc::new(Identity::from_cert_chain(vec![cert_der.clone()]).unwrap()), + PrivatePkcs8KeyDer::from(cert.signing_key.serialize_der()).into(), + ) + .unwrap(); + crypto.alpn_protocols = vec![alpn.to_vec().into()]; + let mut config = crate::ServerConfig::with_crypto(Arc::new( + crate::crypto::rustls::QuicServerConfig::try_from(crypto).unwrap(), + )); + let mut transport = TransportConfig::default(); + transport.max_concurrent_uni_streams(max_uni.into()); + config.transport_config(Arc::new(transport)); + config + }; + + let default_server_config = server_config(b"default", 0); + let selected_server_config = Arc::new(server_config(selected_alpn, 1)); + let server = Endpoint::new( + EndpointConfig::default(), + Some(default_server_config), + UdpSocket::bind(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0)).unwrap(), + Arc::new(TokioRuntime), + ) + .unwrap(); + let server_addr = server.local_addr().unwrap(); + + let mut roots = RootCertStore::empty(); + roots.add(cert_der).unwrap(); + let mut client_crypto = rustls::ClientConfig::builder(Arc::new(default_provider())) + .with_root_certificates(roots) + .with_no_client_auth() + .unwrap(); + client_crypto.alpn_protocols = vec![selected_alpn.as_slice().into()]; + let client = Endpoint::client(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0)).unwrap(); + client.set_default_client_config(ClientConfig::new(Arc::new( + QuicClientConfig::try_from(client_crypto).unwrap(), + ))); + + let server_task = async { + let accepted = server + .accept() + .await + .unwrap() + .acceptor() + .unwrap() + .await + .unwrap(); + assert_eq!( + accepted.client_hello().server_name().map(AsRef::as_ref), + Some("localhost") + ); + accepted + .accept_with(selected_server_config) + .unwrap() + .await + .unwrap() + }; + let client_task = async { + client + .connect(server_addr, "localhost") + .unwrap() + .await + .unwrap() + }; + let (server_conn, client_conn) = join!(server_task, client_task); + + let mut send = timeout(Duration::from_secs(2), client_conn.open_uni()) + .await + .expect("selected transport parameters should allow a unidirectional stream") + .unwrap(); + send.write_all(b"selected").await.unwrap(); + send.finish().unwrap(); + let mut recv = server_conn.accept_uni().await.unwrap(); + assert_eq!(recv.read_to_end(8).await.unwrap(), b"selected"); + + client_conn.close(0u32.into(), b"done"); + server_conn.closed().await; +} + /// Constructs endpoints suitable for connecting to themselves and each other struct EndpointFactory { cert: rcgen::CertifiedKey, @@ -290,15 +406,56 @@ impl EndpointFactory { self.endpoint_with_config(TransportConfig::default()) } - fn endpoint_with_config(&self, transport_config: TransportConfig) -> Endpoint { + fn endpoint_with_max_incoming(&self, max_incoming: usize) -> Endpoint { + let transport_config = Arc::new(TransportConfig::default()); + let mut server_config = self.server_config(transport_config.clone()); + server_config.max_incoming(max_incoming); + let endpoint = Endpoint::new( + self.endpoint_config.clone(), + Some(server_config), + UdpSocket::bind(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0)).unwrap(), + Arc::new(TokioRuntime), + ) + .unwrap(); + endpoint.set_default_client_config(self.client_config(transport_config)); + endpoint + } + + fn server_config(&self, transport_config: Arc) -> crate::ServerConfig { + let cert = self.cert.cert.der().clone(); let key = PrivateKeyDer::Pkcs8(self.cert.signing_key.serialize_der().into()); - let transport_config = Arc::new(transport_config); - let mut server_config = - crate::ServerConfig::with_single_cert(vec![self.cert.cert.der().clone()], key).unwrap(); - server_config.transport_config(transport_config.clone()); + let mut server_crypto = rustls::ServerConfig::builder(Arc::new(default_provider())) + .with_no_client_auth() + .with_single_cert( + Arc::new(Identity::from_cert_chain(vec![cert.clone()]).unwrap()), + key, + ) + .unwrap(); + server_crypto.max_early_data_size = u32::MAX; + let mut server_config = crate::ServerConfig::with_crypto(Arc::new( + crate::crypto::rustls::QuicServerConfig::try_from(server_crypto).unwrap(), + )); + server_config.transport_config(transport_config); + server_config + } + fn client_config(&self, transport_config: Arc) -> ClientConfig { let mut roots = RootCertStore::empty(); roots.add(self.cert.cert.der().clone()).unwrap(); + let mut client_crypto = rustls::ClientConfig::builder(Arc::new(default_provider())) + .with_root_certificates(roots) + .with_no_client_auth() + .unwrap(); + client_crypto.enable_early_data = true; + let mut client_config = + ClientConfig::new(Arc::new(QuicClientConfig::try_from(client_crypto).unwrap())); + client_config.transport_config(transport_config); + client_config + } + + fn endpoint_with_config(&self, transport_config: TransportConfig) -> Endpoint { + let transport_config = Arc::new(transport_config); + let server_config = self.server_config(transport_config.clone()); let endpoint = Endpoint::new( self.endpoint_config.clone(), Some(server_config), @@ -306,14 +463,549 @@ impl EndpointFactory { Arc::new(TokioRuntime), ) .unwrap(); - let mut client_config = ClientConfig::with_root_certificates(Arc::new(roots)).unwrap(); - client_config.transport_config(transport_config); + let client_config = self.client_config(transport_config); endpoint.set_default_client_config(client_config); endpoint } } +fn client_endpoint(client_config: ClientConfig) -> Endpoint { + let client = Endpoint::client(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0)).unwrap(); + client.set_default_client_config(client_config); + client +} + +fn fragmented_client_endpoint(factory: &EndpointFactory) -> Endpoint { + let mut roots = RootCertStore::empty(); + roots.add(factory.cert.cert.der().clone()).unwrap(); + let mut crypto = rustls::ClientConfig::builder(Arc::new(default_provider())) + .with_root_certificates(roots) + .with_no_client_auth() + .unwrap(); + // This stays below the default 16 KiB CRYPTO buffer, but spans more than the client's + // initial congestion window so the server cannot receive the complete ClientHello before + // the staged acceptor emits its first ACK. + crypto.alpn_protocols = (0..3000u32) + .map(|i| i.to_be_bytes().to_vec().into()) + .collect(); + client_endpoint(ClientConfig::new(Arc::new( + QuicClientConfig::try_from(crypto).unwrap(), + ))) +} + +async fn begin_staged_accept( + client: &Endpoint, + server: &Endpoint, +) -> (crate::Accepted, crate::Connecting) { + let client_addr = client.local_addr().unwrap(); + let connecting = client + .connect(server.local_addr().unwrap(), "localhost") + .unwrap(); + let accepted = timeout(Duration::from_secs(5), async { + loop { + let incoming = server.accept().await.unwrap(); + if incoming.remote_address() == client_addr { + break incoming.acceptor().unwrap().await; + } + incoming.ignore(); + } + }) + .await + .expect("timed out reading ClientHello") + .unwrap(); + (accepted, connecting) +} + +async fn establish_staged_connection( + client: &Endpoint, + server: &Endpoint, +) -> (crate::Connection, crate::Connection) { + let (accepted, client_connecting) = begin_staged_accept(client, server).await; + let server_connecting = accepted.accept().unwrap(); + let (client_conn, server_conn) = join!(client_connecting, server_connecting); + (client_conn.unwrap(), server_conn.unwrap()) +} + +#[tokio::test] +async fn staged_acceptor_buffers_zero_rtt_until_config_selection() { + let _guard = subscribe(); + const TICKET_CONFIRMED: &[u8] = b"ticket"; + const EARLY_DATA: &[u8] = b"buffered zero rtt"; + + let factory = EndpointFactory::new(); + let transport = Arc::new(TransportConfig::default()); + let selected_config = Arc::new(factory.server_config(transport.clone())); + let server = Endpoint::new( + factory.endpoint_config.clone(), + Some((*selected_config).clone()), + UdpSocket::bind(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0)).unwrap(), + Arc::new(TokioRuntime), + ) + .unwrap(); + let client = client_endpoint(factory.client_config(transport)); + + // Complete one connection and exchange 1-RTT data so the client processes the session ticket. + let client_connecting = client + .connect(server.local_addr().unwrap(), "localhost") + .unwrap(); + let server_connecting = server.accept().await.unwrap().accept().unwrap(); + let (client_conn, server_conn) = join!(client_connecting, server_connecting); + let client_conn = client_conn.unwrap(); + let server_conn = server_conn.unwrap(); + let mut ticket_stream = server_conn.open_uni().await.unwrap(); + ticket_stream.write_all(TICKET_CONFIRMED).await.unwrap(); + ticket_stream.finish().unwrap(); + let mut ticket_stream = client_conn.accept_uni().await.unwrap(); + assert_eq!( + ticket_stream.read_to_end(usize::MAX).await.unwrap(), + TICKET_CONFIRMED + ); + client_conn.close(0u32.into(), b"resume"); + let _ = wait_closed(&server_conn).await; + wait_idle(&server).await; + wait_idle(&client).await; + + let client_conn = client + .connect(server.local_addr().unwrap(), "localhost") + .unwrap() + .into_0rtt() + .unwrap_or_else(|_| panic!("missing 0-RTT keys after receiving a session ticket")); + let accepted = timeout(Duration::from_secs(5), async { + server.accept().await.unwrap().acceptor().unwrap().await + }) + .await + .expect("timed out reading resumed ClientHello") + .unwrap(); + + // Keep `Accepted` alive while the client sends early data. The endpoint must buffer the + // resulting datagrams until configuration selection creates the connection. + let mut early_stream = client_conn.open_uni().await.unwrap(); + early_stream.write_all(EARLY_DATA).await.unwrap(); + early_stream.finish().unwrap(); + sleep(Duration::from_millis(25)).await; + assert_eq!(server.open_connections(), 0); + + let server_conn = accepted + .accept_with(selected_config) + .unwrap() + .into_0rtt() + .unwrap_or_else(|_| unreachable!("servers always have 0.5-RTT keys")); + let mut received = timeout(Duration::from_secs(5), server_conn.accept_uni()) + .await + .expect("buffered 0-RTT stream was not replayed") + .unwrap(); + assert!(received.is_0rtt()); + assert_eq!(received.read_to_end(usize::MAX).await.unwrap(), EARLY_DATA); + client_conn.authenticated().await.unwrap(); + early_stream.stopped().await.unwrap(); + + client_conn.close(0u32.into(), b"done"); + let _ = wait_closed(&server_conn).await; + wait_idle(&server).await; + wait_idle(&client).await; +} + +#[tokio::test] +async fn staged_acceptor_verifies_client_identity() { + let _guard = subscribe(); + let server_identity = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap(); + let client_identity = rcgen::generate_simple_self_signed(vec!["client".into()]).unwrap(); + + let provider = Arc::new(default_provider()); + let mut client_roots = RootCertStore::empty(); + client_roots + .add(client_identity.cert.der().clone()) + .unwrap(); + let verifier = WebPkiClientVerifier::builder(Arc::new(client_roots), provider.as_ref()) + .build() + .unwrap(); + let server_crypto = rustls::ServerConfig::builder(provider) + .with_client_cert_verifier(Arc::new(verifier)) + .with_single_cert( + Arc::new(Identity::from_cert_chain(vec![server_identity.cert.der().clone()]).unwrap()), + PrivatePkcs8KeyDer::from(server_identity.signing_key.serialize_der()).into(), + ) + .unwrap(); + let selected_config = Arc::new(crate::ServerConfig::with_crypto(Arc::new( + crate::crypto::rustls::QuicServerConfig::try_from(server_crypto).unwrap(), + ))); + let server = Endpoint::new( + EndpointConfig::default(), + Some((*selected_config).clone()), + UdpSocket::bind(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0)).unwrap(), + Arc::new(TokioRuntime), + ) + .unwrap(); + + let client_config = |identity: Option<&rcgen::CertifiedKey>| { + let mut roots = RootCertStore::empty(); + roots.add(server_identity.cert.der().clone()).unwrap(); + let builder = rustls::ClientConfig::builder(Arc::new(default_provider())) + .with_root_certificates(roots); + let crypto = if let Some(identity) = identity { + builder + .with_client_auth_cert( + Arc::new(Identity::from_cert_chain(vec![identity.cert.der().clone()]).unwrap()), + PrivatePkcs8KeyDer::from(identity.signing_key.serialize_der()).into(), + ) + .unwrap() + } else { + builder.with_no_client_auth().unwrap() + }; + ClientConfig::new(Arc::new(QuicClientConfig::try_from(crypto).unwrap())) + }; + + let authenticated_client = client_endpoint(client_config(Some(&client_identity))); + let (accepted, client_connecting) = begin_staged_accept(&authenticated_client, &server).await; + let server_connecting = accepted.accept_with(selected_config.clone()).unwrap(); + let (client_conn, server_conn) = join!(client_connecting, server_connecting); + let client_conn = client_conn.unwrap(); + let server_conn = server_conn.unwrap(); + assert!(server_conn.peer_identity().is_some()); + client_conn.close(0u32.into(), b"authenticated"); + let _ = wait_closed(&server_conn).await; + wait_idle(&server).await; + wait_idle(&authenticated_client).await; + + let anonymous_client = client_endpoint(client_config(None)); + let (accepted, client_connecting) = begin_staged_accept(&anonymous_client, &server).await; + let server_connecting = accepted.accept_with(selected_config).unwrap(); + let (client_result, server_result) = join!(client_connecting, server_connecting); + assert!( + server_result.is_err(), + "missing client certificate reached the application" + ); + if let Ok(client_conn) = client_result { + assert!(matches!( + wait_closed(&client_conn).await, + crate::ConnectionError::ConnectionClosed(_) + | crate::ConnectionError::ApplicationClosed(_) + )); + } + wait_idle(&server).await; + wait_idle(&anonymous_client).await; + assert_eq!(server.open_connections(), 0); +} + +#[tokio::test] +async fn dropping_pending_staged_acceptor_releases_incoming_slot() { + let _guard = subscribe(); + let factory = EndpointFactory::new(); + let server = factory.endpoint_with_max_incoming(1); + let fragmented_client = fragmented_client_endpoint(&factory); + let connecting = fragmented_client + .connect(server.local_addr().unwrap(), "localhost") + .unwrap(); + let acceptor = server.accept().await.unwrap().acceptor().unwrap(); + + assert_eq!(server.open_connections(), 0); + assert_wait_idle_pending(&server).await; + drop(acceptor); + assert!( + timeout(Duration::from_secs(5), connecting) + .await + .unwrap() + .is_err() + ); + wait_idle(&server).await; + wait_idle(&fragmented_client).await; + + // With max_incoming=1, a subsequent successful staged accept proves the dropped future + // released its reservation. + let client = factory.endpoint(); + let (client_conn, server_conn) = establish_staged_connection(&client, &server).await; + client_conn.close(0u32.into(), b"done"); + let _ = wait_closed(&server_conn).await; + wait_idle(&server).await; + wait_idle(&client).await; +} + +#[tokio::test] +async fn dropping_or_refusing_accepted_releases_incoming_slot() { + let _guard = subscribe(); + let factory = EndpointFactory::new(); + let server = factory.endpoint_with_max_incoming(1); + + for refuse in [false, true] { + let client = factory.endpoint(); + let (accepted, connecting) = begin_staged_accept(&client, &server).await; + assert_eq!(server.open_connections(), 0); + assert_wait_idle_pending(&server).await; + if refuse { + accepted.refuse(); + } else { + drop(accepted); + } + client.close(0u32.into(), b"test cleanup"); + assert!( + timeout(Duration::from_secs(5), connecting) + .await + .unwrap() + .is_err() + ); + wait_idle(&server).await; + wait_idle(&client).await; + } + + let client = factory.endpoint(); + let (client_conn, server_conn) = establish_staged_connection(&client, &server).await; + assert_eq!(server.open_connections(), 1); + client_conn.close(0u32.into(), b"done"); + let _ = wait_closed(&server_conn).await; + wait_idle(&server).await; + wait_idle(&client).await; +} + +#[tokio::test] +async fn endpoint_close_cancels_pending_staged_acceptor() { + let _guard = subscribe(); + let factory = EndpointFactory::new(); + let server = factory.endpoint(); + let client = fragmented_client_endpoint(&factory); + let connecting = client + .connect(server.local_addr().unwrap(), "localhost") + .unwrap(); + let acceptor = server.accept().await.unwrap().acceptor().unwrap(); + tokio::pin!(acceptor); + + // Poll once to register the staged acceptor's dedicated notification and confirm the + // fragmented ClientHello has not completed yet. + tokio::select! { + biased; + result = &mut acceptor => panic!("fragmented ClientHello unexpectedly completed: {}", result.is_ok()), + _ = std::future::ready(()) => {} + } + assert_wait_idle_pending(&server).await; + server.close(0u32.into(), b"closing"); + let result = timeout(Duration::from_secs(5), &mut acceptor) + .await + .expect("endpoint close did not wake staged acceptor"); + let Err(error) = result else { + panic!("staged acceptor completed after endpoint close"); + }; + assert!(matches!(error, crate::ConnectionError::LocallyClosed)); + client.close(0u32.into(), b"test cleanup"); + assert!( + timeout(Duration::from_secs(5), connecting) + .await + .unwrap() + .is_err() + ); + assert_eq!(server.open_connections(), 0); + wait_idle(&server).await; + wait_idle(&client).await; +} + +#[tokio::test] +async fn endpoint_close_rejects_staged_accept_after_client_hello() { + let _guard = subscribe(); + let factory = EndpointFactory::new(); + let server = factory.endpoint(); + let client = factory.endpoint(); + let (accepted, connecting) = begin_staged_accept(&client, &server).await; + + assert_eq!(server.open_connections(), 0); + assert_wait_idle_pending(&server).await; + server.close(0u32.into(), b"closing"); + let Err(error) = accepted.accept() else { + panic!("staged accept succeeded after endpoint close"); + }; + assert!(matches!(error, crate::ConnectionError::LocallyClosed)); + + client.close(0u32.into(), b"test cleanup"); + assert!( + timeout(Duration::from_secs(5), connecting) + .await + .unwrap() + .is_err() + ); + wait_idle(&server).await; + wait_idle(&client).await; + assert_eq!(server.open_connections(), 0); +} + +#[derive(Default)] +struct HandshakeBlocker { + state: Mutex, + changed: Condvar, +} + +#[derive(Default)] +struct HandshakeBlockerState { + started: bool, + released: bool, +} + +impl HandshakeBlocker { + fn block(&self) { + let mut state = self.state.lock().unwrap(); + state.started = true; + self.changed.notify_all(); + while !state.released { + state = self.changed.wait(state).unwrap(); + } + } + + fn wait_until_started(&self) { + let state = self.state.lock().unwrap(); + let (state, _) = self + .changed + .wait_timeout_while(state, Duration::from_secs(5), |state| !state.started) + .unwrap(); + assert!(state.started, "timed out waiting for handshake to start"); + } + + fn release(&self) { + let mut state = self.state.lock().unwrap(); + state.released = true; + self.changed.notify_all(); + } +} + +struct HandshakeReleaseGuard(Arc); + +impl Drop for HandshakeReleaseGuard { + fn drop(&mut self) { + self.0.release(); + } +} + +struct BlockingServerConfig { + inner: Arc, + blocker: Arc, +} + +impl ProtoServerConfig for BlockingServerConfig { + fn initial_keys( + &self, + version: u32, + dst_cid: ConnectionId, + ) -> Result { + self.inner.initial_keys(version, dst_cid) + } + + fn retry_tag(&self, version: u32, orig_dst_cid: ConnectionId, packet: &[u8]) -> [u8; 16] { + self.inner.retry_tag(version, orig_dst_cid, packet) + } + + fn start_session( + self: Arc, + version: u32, + params: &TransportParameters, + ) -> Box { + self.blocker.block(); + self.inner.clone().start_session(version, params) + } +} + +fn blocking_server_pair() -> (Endpoint, Endpoint, Arc) { + let factory = EndpointFactory::new(); + let key = PrivateKeyDer::Pkcs8(factory.cert.signing_key.serialize_der().into()); + let mut server_config = + crate::ServerConfig::with_single_cert(vec![factory.cert.cert.der().clone()], key).unwrap(); + let blocker = Arc::new(HandshakeBlocker::default()); + server_config.crypto = Arc::new(BlockingServerConfig { + inner: server_config.crypto.clone(), + blocker: blocker.clone(), + }); + + let server = Endpoint::new( + factory.endpoint_config.clone(), + Some(server_config), + UdpSocket::bind(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0)).unwrap(), + Arc::new(TokioRuntime), + ) + .unwrap(); + + let mut roots = RootCertStore::empty(); + roots.add(factory.cert.cert.der().clone()).unwrap(); + let client = Endpoint::client(SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 0)).unwrap(); + client + .set_default_client_config(ClientConfig::with_root_certificates(Arc::new(roots)).unwrap()); + + (client, server, blocker) +} + +async fn wait_for_blocked_handshake(blocker: Arc) { + tokio::task::spawn_blocking(move || blocker.wait_until_started()) + .await + .unwrap(); +} + +async fn accept_with_blocked_start_session( + server: Endpoint, +) -> Result { + let incoming = server.accept().await.unwrap(); + // The test crypto provider intentionally blocks in `start_session`, so run + // accept on the blocking pool rather than parking a Tokio worker thread. + let connecting = tokio::task::spawn_blocking(move || incoming.accept()) + .await + .unwrap()?; + connecting.await +} + +struct BlockedHandshake { + client: Endpoint, + server: Endpoint, + blocker: Arc, + release_guard: HandshakeReleaseGuard, + accept: tokio::task::JoinHandle>, + connect: tokio::task::JoinHandle>, +} + +fn blocked_handshake() -> BlockedHandshake { + let (client, server, blocker) = blocking_server_pair(); + let server_addr = server.local_addr().unwrap(); + + let accept = tokio::spawn({ + let server = server.clone(); + async move { accept_with_blocked_start_session(server).await } + }); + let connect = tokio::spawn({ + let client = client.clone(); + async move { client.connect(server_addr, "localhost").unwrap().await } + }); + + BlockedHandshake { + client, + server, + blocker: blocker.clone(), + release_guard: HandshakeReleaseGuard(blocker), + accept, + connect, + } +} + +async fn await_connection( + task: tokio::task::JoinHandle>, +) -> Result { + timeout(Duration::from_secs(5), task) + .await + .unwrap() + .unwrap() +} + +async fn assert_wait_idle_pending(endpoint: &Endpoint) { + assert!( + timeout(Duration::from_millis(50), endpoint.wait_idle()) + .await + .is_err() + ); +} + +async fn wait_idle(endpoint: &Endpoint) { + timeout(Duration::from_secs(5), endpoint.wait_idle()) + .await + .unwrap(); +} + +async fn wait_closed(conn: &crate::Connection) -> crate::ConnectionError { + timeout(Duration::from_secs(5), conn.closed()) + .await + .unwrap() +} + #[tokio::test] async fn zero_rtt() { let _guard = subscribe(); @@ -528,13 +1220,11 @@ fn run_echo(args: EchoArgs) { let mut roots = RootCertStore::empty(); roots.add(cert).unwrap(); - let mut client_crypto = - rustls::ClientConfig::builder_with_provider(default_provider().into()) - .with_safe_default_protocol_versions() - .unwrap() - .with_root_certificates(roots) - .with_no_client_auth(); - client_crypto.key_log = Arc::new(rustls::KeyLogFile::new()); + let mut client_crypto = rustls::ClientConfig::builder(Arc::new(default_provider())) + .with_root_certificates(roots) + .with_no_client_auth() + .unwrap(); + client_crypto.key_log = Arc::new(rustls_util::KeyLogFile::new()); let client = { let _guard = runtime.enter(); @@ -758,6 +1448,82 @@ async fn rebind_recv() { server.await.unwrap(); } +/// Verify endpoint behavior while a connection is between `start_accept` and `finish_accept`: +/// `open_connections` must be 0, and `wait_idle` must not complete while `pending_accepts > 0`. +/// After the accept completes and the connection is closed, `wait_idle` must resolve. +#[tokio::test] +async fn split_accept_wait_idle_and_open_connections() { + let _guard = subscribe(); + let BlockedHandshake { + client, + server, + blocker, + release_guard: _release_guard, + accept, + connect, + } = blocked_handshake(); + + wait_for_blocked_handshake(blocker.clone()).await; + assert_eq!(server.open_connections(), 0); + assert_wait_idle_pending(&server).await; + + blocker.release(); + + let client_conn = await_connection(connect).await.unwrap(); + let server_conn = await_connection(accept).await.unwrap(); + assert_eq!(server.open_connections(), 1); + + client_conn.close(0u32.into(), b"done"); + let _ = wait_closed(&client_conn).await; + let _ = wait_closed(&server_conn).await; + wait_idle(&server).await; + wait_idle(&client).await; +} + +/// Verify that calling `Endpoint::close` while a connection is in the `Accepting` state +/// (between `start_accept` and `finish_accept`) produces `LocallyClosed` on the server side +/// and `ConnectionClosed` on the client side, and that `wait_idle` resolves afterward. +#[tokio::test] +async fn split_accept_close_during_pending_accept() { + let _guard = subscribe(); + let BlockedHandshake { + client, + server, + blocker, + release_guard: _release_guard, + accept, + connect, + } = blocked_handshake(); + + wait_for_blocked_handshake(blocker.clone()).await; + server.close(0u32.into(), b"closing"); + assert_wait_idle_pending(&server).await; + + blocker.release(); + + match await_connection(accept).await { + Ok(conn) => { + let err = wait_closed(&conn).await; + assert!(matches!(err, crate::ConnectionError::LocallyClosed)); + } + Err(err) => assert!(matches!(err, crate::ConnectionError::LocallyClosed)), + } + + match await_connection(connect).await { + Ok(conn) => { + let err = wait_closed(&conn).await; + assert!(matches!(err, crate::ConnectionError::ConnectionClosed(_))); + } + Err(err) => { + assert!(matches!(err, crate::ConnectionError::ConnectionClosed(_))); + } + } + + wait_idle(&server).await; + wait_idle(&client).await; + assert_eq!(server.open_connections(), 0); +} + #[tokio::test] async fn remote_address_monitoring() { let _guard = subscribe(); diff --git a/quinn/tests/post_quantum.rs b/quinn/tests/post_quantum.rs index 1d2970971b..986a530aa4 100644 --- a/quinn/tests/post_quantum.rs +++ b/quinn/tests/post_quantum.rs @@ -7,7 +7,7 @@ use std::{ }; use rustls::{ - NamedGroup, + crypto::{CryptoProvider, Identity, kx::NamedGroup}, pki_types::{CertificateDer, PrivatePkcs8KeyDer}, }; use tracing::info; @@ -27,6 +27,16 @@ async fn post_quantum_key_exchange_large_mtu() { check_post_quantum_key_exchange(1433).await; } +#[tokio::test] +async fn post_quantum_handshake_data_after_hello_retry_request() { + check_post_quantum_handshake_data_after_hello_retry_request(false).await; +} + +#[tokio::test] +async fn staged_post_quantum_handshake_data_after_hello_retry_request() { + check_post_quantum_handshake_data_after_hello_retry_request(true).await; +} + async fn check_post_quantum_key_exchange(min_mtu: u16) { let _ = tracing_subscriber::FmtSubscriber::builder() .with_env_filter(tracing_subscriber::EnvFilter::from_default_env()) @@ -35,7 +45,12 @@ async fn check_post_quantum_key_exchange(min_mtu: u16) { let server_addr = SocketAddr::from((Ipv4Addr::LOCALHOST, 0)); - let (endpoint, server_cert) = make_server_endpoint(server_addr, min_mtu).unwrap(); + let (endpoint, server_cert) = make_server_endpoint( + server_addr, + min_mtu, + Arc::new(rustls_aws_lc_rs::DEFAULT_PROVIDER), + ) + .unwrap(); let server_addr = endpoint.local_addr().unwrap(); // accept a single connection let jh = tokio::spawn(async move { @@ -55,8 +70,12 @@ async fn check_post_quantum_key_exchange(min_mtu: u16) { ) }); - let endpoint = - make_client_endpoint(SocketAddr::from((Ipv4Addr::UNSPECIFIED, 0)), server_cert).unwrap(); + let endpoint = make_client_endpoint( + SocketAddr::from((Ipv4Addr::UNSPECIFIED, 0)), + server_cert, + Arc::new(rustls_aws_lc_rs::DEFAULT_PROVIDER), + ) + .unwrap(); // connect to server let connection = endpoint .connect(server_addr, "localhost") @@ -73,19 +92,84 @@ async fn check_post_quantum_key_exchange(min_mtu: u16) { jh.await.unwrap(); } +async fn check_post_quantum_handshake_data_after_hello_retry_request(staged: bool) { + // The client initially shares X25519, while the server only supports X25519MLKEM768, forcing + // a HelloRetryRequest before the negotiated group becomes available as handshake data. + let server_provider = Arc::new(CryptoProvider { + kx_groups: std::borrow::Cow::Owned(vec![rustls_aws_lc_rs::kx_group::X25519MLKEM768]), + ..rustls_aws_lc_rs::DEFAULT_PROVIDER + }); + let (server, server_cert) = make_server_endpoint( + SocketAddr::from((Ipv4Addr::LOCALHOST, 0)), + 1433, + server_provider, + ) + .unwrap(); + let server_addr = server.local_addr().unwrap(); + + let client_provider = Arc::new(CryptoProvider { + kx_groups: std::borrow::Cow::Owned(vec![ + rustls_aws_lc_rs::kx_group::X25519, + rustls_aws_lc_rs::kx_group::X25519MLKEM768, + ]), + ..rustls_aws_lc_rs::DEFAULT_PROVIDER + }); + let client = make_client_endpoint( + SocketAddr::from((Ipv4Addr::UNSPECIFIED, 0)), + server_cert, + client_provider, + ) + .unwrap(); + + let server_task = async { + let incoming = server.accept().await.unwrap(); + let mut connecting = if staged { + incoming + .acceptor() + .unwrap() + .await + .unwrap() + .accept() + .unwrap() + } else { + incoming.accept().unwrap() + }; + let handshake_data = tokio::time::timeout( + std::time::Duration::from_secs(5), + connecting.handshake_data(), + ) + .await + .expect("timed out waiting for handshake data") + .unwrap() + .downcast::() + .unwrap(); + assert_eq!( + handshake_data.negotiated_key_exchange_group, + NamedGroup::X25519MLKEM768 + ); + connecting.await.unwrap() + }; + let client_task = client.connect(server_addr, "localhost").unwrap(); + let (server_connection, client_connection) = tokio::join!(server_task, client_task); + let client_connection = client_connection.unwrap(); + + client_connection.close(0u32.into(), b"done"); + server_connection.closed().await; + server.wait_idle().await; + client.wait_idle().await; +} + fn make_client_endpoint( bind_addr: SocketAddr, server_cert: CertificateDer<'static>, + provider: Arc, ) -> Result> { let mut certs = rustls::RootCertStore::empty(); certs.add(server_cert)?; - let rustls_config = rustls::ClientConfig::builder_with_provider(Arc::new( - rustls::crypto::aws_lc_rs::default_provider(), - )) - .with_safe_default_protocol_versions() - .unwrap() - .with_root_certificates(certs) - .with_no_client_auth(); + let rustls_config = rustls::ClientConfig::builder(provider) + .with_root_certificates(certs) + .with_no_client_auth() + .unwrap(); let client_cfg = quinn::ClientConfig::new(Arc::new(QuicClientConfig::try_from(rustls_config).unwrap())); @@ -97,20 +181,20 @@ fn make_client_endpoint( fn make_server_endpoint( bind_addr: SocketAddr, min_mtu: u16, + provider: Arc, ) -> Result<(Endpoint, CertificateDer<'static>), Box> { let cert = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap(); let key = PrivatePkcs8KeyDer::from(cert.signing_key.serialize_der()); let cert = CertificateDer::from(cert.cert); let mut server_config = quinn::ServerConfig::with_crypto(Arc::new( QuicServerConfig::try_from( - rustls::ServerConfig::builder_with_provider(Arc::new( - rustls::crypto::aws_lc_rs::default_provider(), - )) - .with_safe_default_protocol_versions() - .unwrap() - .with_no_client_auth() - .with_single_cert(vec![cert.clone()], key.into()) - .unwrap(), + rustls::ServerConfig::builder(provider) + .with_no_client_auth() + .with_single_cert( + Arc::new(Identity::from_cert_chain(vec![cert.clone()]).unwrap()), + key.into(), + ) + .unwrap(), ) .unwrap(), ));