From 603bcef469fad3c025dc87e941b2f41d249309a9 Mon Sep 17 00:00:00 2001 From: peg Date: Fri, 25 Sep 2026 10:22:43 +0200 Subject: [PATCH] Experimental attestation resumption on reconnect --- attested-tls/src/lib.rs | 133 +++++++-- attested-tls/src/resumption/mod.rs | 235 +++++++++++++++ attested-tls/src/resumption/tests.rs | 426 +++++++++++++++++++++++++++ 3 files changed, 775 insertions(+), 19 deletions(-) create mode 100644 attested-tls/src/resumption/mod.rs create mode 100644 attested-tls/src/resumption/tests.rs diff --git a/attested-tls/src/lib.rs b/attested-tls/src/lib.rs index 7cc68c7..e89a9fd 100644 --- a/attested-tls/src/lib.rs +++ b/attested-tls/src/lib.rs @@ -8,7 +8,12 @@ pub mod attested_rpc; #[cfg(any(test, feature = "test-helpers"))] pub mod test_helpers; +mod resumption; + pub use attestation; +use resumption::{ + ClientCache, ClientStore, ConnectionRecord, ServerCache, ServerStore, VerifiedPeer, +}; use attestation::{ AttestationError, AttestationExchangeMessage, AttestationGenerator, AttestationType, @@ -34,10 +39,10 @@ use tokio_rustls::{ rustls::{ClientConfig, ServerConfig}, }; -/// This makes it possible to add breaking protocol changes and provide backwards compatibility. -/// When adding more supported versions, note that ordering is important. ALPN will pick the first -/// protocol which both parties support - so newer supported versions should come first. -pub const SUPPORTED_ALPN_PROTOCOL_VERSIONS: [&[u8]; 1] = [b"flashbots-ratls/1"]; +/// Experimental protocol identifier. This branch deliberately does not offer v1: +/// skipping the attestation exchange on resumption changes the wire protocol. +pub const SUPPORTED_ALPN_PROTOCOL_VERSIONS: [&[u8]; 1] = + [b"flashbots-ratls/experimental-resumption-1"]; /// The label used when exporting key material from a TLS session pub(crate) const EXPORTER_LABEL: &[u8; 24] = b"EXPORTER-Channel-Binding"; @@ -63,6 +68,9 @@ pub struct AttestedTlsServer { cert_chain: Vec>, /// For accepting TLS connections acceptor: TlsAcceptor, + sessions: Arc, + #[cfg(test)] + counts: Arc, } impl std::fmt::Debug for AttestedTlsServer { @@ -106,7 +114,8 @@ impl AttestedTlsServer { /// Start with preconfigured TLS /// - /// This allows dangerous configuration + /// This allows dangerous configuration. The experiment replaces session storage, + /// disables early data, and rejects enabled stateless ticket generators. pub fn new_with_tls_config( cert_chain: Vec>, mut server_config: ServerConfig, @@ -115,6 +124,10 @@ impl AttestedTlsServer { ) -> Result { #[cfg(feature = "mock")] tracing::warn!("AttestedTlsServer instantiated in MOCK mode - do NOT use in production"); + if server_config.ticketer.enabled() { + return Err(AttestedTlsError::StatelessResumptionUnsupported); + } + server_config.max_early_data_size = 0; // Ensure protocol version compatibility server_config.alpn_protocols = map_alpn_protocols(server_config.alpn_protocols); @@ -124,12 +137,16 @@ impl AttestedTlsServer { attestation_generator, attestation_verifier, acceptor, + sessions: Arc::default(), + #[cfg(test)] + counts: Arc::default(), cert_chain, }) } /// Handle an incoming connection from an [AttestedTlsClient] /// + /// Authenticated TLS resumptions reuse the previously verified peer metadata. /// This is transport agnostic and will work with any asynchronous stream pub async fn handle_connection( &self, @@ -147,8 +164,14 @@ impl AttestedTlsServer { { tracing::debug!("attested-tls-server accepted connection"); - // Do TLS handshake - let mut tls_stream = self.acceptor.accept(inbound).await?; + // Each connection gets a record for its issued and selected tickets. + let record = ConnectionRecord::default(); + let mut config = (**self.acceptor.config()).clone(); + config.session_storage = Arc::new(ServerStore { + cache: self.sessions.clone(), + record: record.clone(), + }); + let mut tls_stream = TlsAcceptor::from(Arc::new(config)).accept(inbound).await?; let (_io, connection) = tls_stream.get_ref(); // Ensure TLS 1.3 @@ -157,9 +180,15 @@ impl AttestedTlsServer { } // Ensure that we agreed a protocol - let _negotiated_protocol = connection - .alpn_protocol() - .ok_or(AttestedTlsError::AlpnFailed)?; + validate_alpn(connection.alpn_protocol())?; + + if connection.handshake_kind() == Some(rustls::HandshakeKind::Resumed) { + let peer = record + .verified_peer()? + .ok_or(AttestedTlsError::MissingResumptionAttestation)?; + record.authenticate(peer.clone())?; + return Ok((tls_stream, peer.measurements, peer.attestation_type)); + } // Compute an exporter unique to the session let mut exporter = [0u8; 32]; @@ -175,6 +204,10 @@ impl AttestedTlsServer { let remote_cert_chain = connection.peer_certificates().map(|c| c.to_owned()); // If we are in a CVM, generate an attestation off the async runtime thread. + #[cfg(test)] + self.counts + .generated + .fetch_add(1, std::sync::atomic::Ordering::Relaxed); let attestation = { let attestation_generator = self.attestation_generator.clone(); tokio::task::spawn_blocking(move || { @@ -200,12 +233,20 @@ impl AttestedTlsServer { // Validate every exchange, including policies that only accept no attestation, // before exposing the peer's attestation type to callers. let remote_input_data = compute_report_input(remote_cert_chain.as_deref(), exporter)?; + #[cfg(test)] + self.counts + .verified + .fetch_add(1, std::sync::atomic::Ordering::Relaxed); let measurements = self .attestation_verifier .verify_attestation(remote_attestation_message, remote_input_data) .await? .map(|verified| verified.measurements); + record.authenticate(VerifiedPeer { + measurements: measurements.clone(), + attestation_type: remote_attestation_type, + })?; Ok((tls_stream, measurements, remote_attestation_type)) } @@ -232,6 +273,9 @@ impl AttestedTlsServer { pub struct AttestedTlsClient { /// The connector for making TLS connections with out configuration connector: TlsConnector, + sessions: Arc, + #[cfg(test)] + counts: Arc, /// Quote generation type to use (including none) attestation_generator: AttestationGenerator, /// Verifier for remote attestation (including none) @@ -292,7 +336,8 @@ impl AttestedTlsClient { /// Create a new proxy client with given TLS configuration /// - /// This allows dangerous configuration but is used in tests + /// This allows dangerous configuration but is used in tests. The experiment + /// replaces session storage and disables early data. pub fn new_with_tls_config( mut client_config: ClientConfig, attestation_generator: AttestationGenerator, @@ -304,6 +349,7 @@ impl AttestedTlsClient { if client_config.client_auth_cert_resolver.has_certs() && cert_chain.is_none() { return Err(AttestedTlsError::ClientAuthWithoutClientCert); } + client_config.enable_early_data = false; // Ensure protocol version compatibility client_config.alpn_protocols = map_alpn_protocols(client_config.alpn_protocols); @@ -311,6 +357,9 @@ impl AttestedTlsClient { Ok(Self { connector, + sessions: Arc::default(), + #[cfg(test)] + counts: Arc::default(), attestation_generator, attestation_verifier, cert_chain, @@ -320,6 +369,7 @@ impl AttestedTlsClient { /// Given a connection to an attested TLS server, do a TLS handshake and attestation exchange, and return the TLS /// stream together with measurement details /// + /// Authenticated TLS resumptions reuse the previously verified peer metadata. /// This is transport agnostic and will work with any asynchronous stream pub async fn connect( &self, @@ -336,9 +386,14 @@ impl AttestedTlsClient { where IO: AsyncRead + AsyncWrite + Unpin, { - // Make a TLS handshake with the given connection - let mut tls_stream = self - .connector + let record = ConnectionRecord::default(); + let mut config = (**self.connector.config()).clone(); + config.resumption = rustls::client::Resumption::store(Arc::new(ClientStore { + cache: self.sessions.clone(), + target: target.to_owned(), + record: record.clone(), + })); + let mut tls_stream = TlsConnector::from(Arc::new(config)) .connect(server_name_from_host(target)?, outbound) .await?; @@ -350,9 +405,15 @@ impl AttestedTlsClient { } // Ensure that we agreed a protocol - let _negotiated_protocol = server_connection - .alpn_protocol() - .ok_or(AttestedTlsError::AlpnFailed)?; + validate_alpn(server_connection.alpn_protocol())?; + + if server_connection.handshake_kind() == Some(rustls::HandshakeKind::Resumed) { + let peer = record + .verified_peer()? + .ok_or(AttestedTlsError::MissingResumptionAttestation)?; + record.authenticate(peer.clone())?; + return Ok((tls_stream, peer.measurements, peer.attestation_type)); + } // Compute an exporter unique to the channel let mut exporter = [0u8; 32]; @@ -377,6 +438,10 @@ impl AttestedTlsClient { let remote_attestation_type = remote_attestation_message.attestation_type(); // Verify the remote attestation against our accepted measurements + #[cfg(test)] + self.counts + .verified + .fetch_add(1, std::sync::atomic::Ordering::Relaxed); let measurements = self .attestation_verifier .verify_attestation(remote_attestation_message, remote_input_data) @@ -385,6 +450,10 @@ impl AttestedTlsClient { // If we are in a CVM, provide an attestation let attestation = if self.attestation_generator.attestation_type != AttestationType::None { + #[cfg(test)] + self.counts + .generated + .fetch_add(1, std::sync::atomic::Ordering::Relaxed); let local_input_data = compute_report_input(self.cert_chain.as_deref(), exporter)?; let attestation_generator = self.attestation_generator.clone(); tokio::task::spawn_blocking(move || { @@ -401,6 +470,11 @@ impl AttestedTlsClient { let attestation_length_prefix = checked_length_prefix(&attestation)?; tls_stream.write_all(&attestation_length_prefix).await?; tls_stream.write_all(&attestation).await?; + tls_stream.flush().await?; + record.authenticate(VerifiedPeer { + measurements: measurements.clone(), + attestation_type: remote_attestation_type, + })?; Ok((tls_stream, measurements, remote_attestation_type)) } @@ -534,12 +608,33 @@ pub enum AttestedTlsError { NoCryptoProvider, #[error("Only TLS 1.3 is supported")] NotTls13, + #[error("Experimental attestation resumption requires stateful TLS tickets")] + StatelessResumptionUnsupported, + #[error("TLS resumed without verified attestation metadata")] + MissingResumptionAttestation, + #[error("Attestation resumption state was poisoned by a prior panic")] + PoisonedResumptionState, #[error("Attestation length {length} exceeds maximum {max}")] AttestationTooLarge { length: usize, max: usize }, #[error("Blocking task failed: {0}")] Join(#[from] JoinError), } +/// Requires the experimental protocol prefix, allowing an application-protocol suffix. +fn validate_alpn(protocol: Option<&[u8]>) -> Result<(), AttestedTlsError> { + let protocol = protocol.ok_or(AttestedTlsError::AlpnFailed)?; + if SUPPORTED_ALPN_PROTOCOL_VERSIONS.iter().any(|prefix| { + protocol == *prefix + || protocol + .strip_prefix(*prefix) + .is_some_and(|suffix| suffix.starts_with(b"+")) + }) { + Ok(()) + } else { + Err(AttestedTlsError::AlpnFailed) + } +} + /// Given a byte array, encode its length as a 4 byte big endian u32 fn length_prefix(input: &[u8]) -> [u8; 4] { let len = input.len() as u32; @@ -599,7 +694,7 @@ fn host_to_host_with_port(host: &str) -> String { } } -/// Ensure protocol compatibility with the other party by adding 'flashbots-ratls/' to the +/// Ensure protocol compatibility by adding the experimental attested TLS prefix to the /// protocol names of all supported protocols fn map_alpn_protocols(existing_protocols: Vec>) -> Vec> { let mut mapped_protocols = Vec::new(); @@ -700,7 +795,7 @@ mod tests { }) .await .unwrap(); - // The client finishes sending before the server checks its evidence. + // The client returns after sending; only the server knows whether it accepted the evidence. let _client_connection = client_result.unwrap(); if client_type == AttestationType::None { let (_stream, measurements, attestation_type) = server_result.unwrap(); diff --git a/attested-tls/src/resumption/mod.rs b/attested-tls/src/resumption/mod.rs new file mode 100644 index 0000000..53989e4 --- /dev/null +++ b/attested-tls/src/resumption/mod.rs @@ -0,0 +1,235 @@ +//! Experimental, process-local attestation-aware TLS session stores. +//! +//! Tickets are issued before attestation finishes. Each entry therefore references +//! a connection record that is promoted only after authentication succeeds. +//! These POC caches deliberately have no capacity or attestation-age limit. + +#[cfg(test)] +mod tests; + +use std::{ + collections::HashMap, + fmt, + sync::{Arc, Mutex}, +}; + +use crate::AttestedTlsError; +use attestation::{AttestationType, measurements::MultiMeasurements}; +use tokio_rustls::rustls::{ + NamedGroup, + client::{ClientSessionStore, Tls12ClientSessionValue, Tls13ClientSessionValue}, + pki_types::ServerName, + server::StoresServerSessions, +}; + +/// Peer attestation results retained for authenticated TLS resumption. +#[derive(Clone, Debug)] +pub(crate) struct VerifiedPeer { + pub measurements: Option, + pub attestation_type: AttestationType, +} + +/// Shared verification state; `None` means tickets are not yet eligible for reuse. +type Authentication = Arc>>; + +/// Tracks authentication for tickets issued by and selected for one connection. +#[derive(Clone, Debug, Default)] +pub(crate) struct ConnectionRecord { + issued: Authentication, + selected: Authentication, +} + +impl ConnectionRecord { + /// Returns the verified peer metadata associated with the selected ticket. + pub fn verified_peer(&self) -> Result, AttestedTlsError> { + Ok(self + .selected + .lock() + .map_err(|_| AttestedTlsError::PoisonedResumptionState)? + .clone()) + } + + /// Authorizes this connection's existing and future tickets using verified peer metadata. + pub fn authenticate(&self, peer: VerifiedPeer) -> Result<(), AttestedTlsError> { + *self + .issued + .lock() + .map_err(|_| AttestedTlsError::PoisonedResumptionState)? = Some(peer); + Ok(()) + } +} + +/// Pairs opaque TLS session state with its originating connection's authentication state. +struct Ticket { + value: T, + authentication: Authentication, +} + +/// Holds client tickets across connections, partitioned by the complete target string. +#[derive(Default)] +pub(crate) struct ClientCache { + // Scope by the caller's complete target (including port), not just SNI. + tickets: Mutex>>>, +} + +impl fmt::Debug for ClientCache { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("ClientCache").finish_non_exhaustive() + } +} + +/// Adapts the shared client cache to one connection's ticket callbacks. +#[derive(Debug)] +pub(crate) struct ClientStore { + pub cache: Arc, + pub target: String, + pub record: ConnectionRecord, +} + +impl ClientSessionStore for ClientStore { + fn set_kx_hint(&self, _: ServerName<'static>, _: NamedGroup) {} + fn kx_hint(&self, _: &ServerName<'_>) -> Option { + None + } + fn set_tls12_session(&self, _: ServerName<'static>, _: Tls12ClientSessionValue) {} + fn tls12_session(&self, _: &ServerName<'_>) -> Option { + None + } + fn remove_tls12_session(&self, _: &ServerName<'static>) {} + + /// Associates an arriving ticket with this connection's authentication state. + fn insert_tls13_ticket(&self, _: ServerName<'static>, value: Tls13ClientSessionValue) { + let Ok(mut cache) = self.cache.tickets.lock() else { + return; + }; + cache.entry(self.target.clone()).or_default().push(Ticket { + value, + authentication: self.record.issued.clone(), + }); + } + + /// Consumes an authenticated ticket and records its peer metadata for this connection. + fn take_tls13_ticket(&self, _: &ServerName<'static>) -> Option { + let mut cache = self.cache.tickets.lock().ok()?; + let tickets = cache.get_mut(&self.target)?; + let (index, peer) = tickets + .iter() + .enumerate() + .rev() + .find_map(|(index, ticket)| { + Some((index, ticket.authentication.lock().ok()?.clone()?)) + })?; + let mut selected = self.record.selected.lock().ok()?; + let ticket = tickets.remove(index); + *selected = Some(peer); + Some(ticket.value) + } +} + +/// Holds server session state and authentication records indexed by ticket identity. +#[derive(Default)] +pub(crate) struct ServerCache { + tickets: Mutex, Ticket>>>, + #[cfg(test)] + rejected_tickets: std::sync::atomic::AtomicUsize, +} + +impl fmt::Debug for ServerCache { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("ServerCache").finish_non_exhaustive() + } +} + +/// Adapts the shared server cache to one connection's ticket callbacks. +#[derive(Debug)] +pub(crate) struct ServerStore { + pub cache: Arc, + pub record: ConnectionRecord, +} + +impl StoresServerSessions for ServerStore { + /// Associates newly issued session state with this connection's authentication state. + fn put(&self, key: Vec, value: Vec) -> bool { + let Ok(mut cache) = self.cache.tickets.lock() else { + return false; + }; + cache.insert( + key, + Ticket { + value, + authentication: self.record.issued.clone(), + }, + ); + true + } + + /// Disables non-consuming lookups; TLS 1.3 resumption uses `take` instead. + fn get(&self, _: &[u8]) -> Option> { + None + } + + /// Consumes authenticated session state and records its peer metadata, rejecting pending tickets. + fn take(&self, key: &[u8]) -> Option> { + let mut cache = self.cache.tickets.lock().ok()?; + let peer = cache.get(key)?.authentication.lock().ok()?.clone(); + let Some(peer) = peer else { + #[cfg(test)] + self.cache + .rejected_tickets + .fetch_add(1, std::sync::atomic::Ordering::Relaxed); + return None; + }; + let mut selected = self.record.selected.lock().ok()?; + let ticket = cache.remove(key)?; + *selected = Some(peer); + Some(ticket.value) + } + + fn can_cache(&self) -> bool { + !self.cache.tickets.is_poisoned() + } +} + +/// Counts calls to attestation generation and verification across cloned endpoints. +#[cfg(test)] +#[derive(Debug, Default)] +pub(crate) struct AttestationCounts { + pub generated: std::sync::atomic::AtomicUsize, + pub verified: std::sync::atomic::AtomicUsize, +} + +#[cfg(test)] +impl ClientCache { + pub fn clear(&self) { + self.tickets.lock().unwrap().clear(); + } + pub fn pending_count(&self) -> usize { + self.tickets + .lock() + .unwrap() + .values() + .flatten() + .filter(|ticket| ticket.authentication.lock().unwrap().is_none()) + .count() + } +} + +#[cfg(test)] +impl ServerCache { + /// Returns how many offered tickets were rejected because authentication was still pending. + pub fn rejected_count(&self) -> usize { + self.rejected_tickets + .load(std::sync::atomic::Ordering::Relaxed) + } + pub fn clear(&self) { + self.tickets.lock().unwrap().clear(); + } + pub fn pending_count(&self) -> usize { + self.tickets + .lock() + .unwrap() + .values() + .filter(|ticket| ticket.authentication.lock().unwrap().is_none()) + .count() + } +} diff --git a/attested-tls/src/resumption/tests.rs b/attested-tls/src/resumption/tests.rs new file mode 100644 index 0000000..e6e714a --- /dev/null +++ b/attested-tls/src/resumption/tests.rs @@ -0,0 +1,426 @@ +//! Exercise the production connection paths; only the evidence is mocked. + +use crate::*; +use std::{sync::atomic::Ordering, time::Duration}; +use test_helpers::{generate_certificate_chain, generate_tls_config_with_client_auth}; +use tokio::io::DuplexStream; + +type ServerResult = Result< + ( + tokio_rustls::server::TlsStream, + Option, + AttestationType, + ), + AttestedTlsError, +>; +type ClientResult = Result< + ( + tokio_rustls::client::TlsStream, + Option, + AttestationType, + ), + AttestedTlsError, +>; + +/// Creates mutually trusted TLS credentials and configurations for the test peers. +fn configs() -> ( + Vec>, + Vec>, + ServerConfig, + ClientConfig, +) { + let (server_certs, server_key) = generate_certificate_chain("127.0.0.1".parse().unwrap()); + let (client_certs, client_key) = generate_certificate_chain("127.0.0.1".parse().unwrap()); + let ((server_config, _), (_, client_config)) = generate_tls_config_with_client_auth( + server_certs.clone(), + server_key, + client_certs.clone(), + client_key, + ); + (server_certs, client_certs, server_config, client_config) +} + +/// Builds endpoints requiring mock peer attestation, with a selectable server evidence type. +fn pair(server_type: AttestationType) -> (AttestedTlsServer, AttestedTlsClient) { + let (server_certs, client_certs, server_config, client_config) = configs(); + let server = AttestedTlsServer::new_with_tls_config( + server_certs, + server_config, + AttestationGenerator::new(server_type, None).unwrap(), + AttestationVerifier::mock(), + ) + .unwrap(); + let client = AttestedTlsClient::new_with_tls_config( + client_config, + AttestationGenerator::new(AttestationType::DcapTdx, None).unwrap(), + AttestationVerifier::mock(), + Some(client_certs), + ) + .unwrap(); + (server, client) +} + +/// Runs both production connection paths over an in-memory transport with a timeout. +async fn connect( + server: &AttestedTlsServer, + client: &AttestedTlsClient, +) -> (ServerResult, ClientResult) { + let (server_io, client_io) = tokio::io::duplex(128 * 1024); + tokio::time::timeout(Duration::from_secs(5), async { + tokio::join!( + server.handle_connection(server_io), + client.connect("127.0.0.1", client_io) + ) + }) + .await + .expect("connection must not deadlock") +} + +fn counts(counts: &resumption::AttestationCounts) -> (usize, usize) { + ( + counts.generated.load(Ordering::Relaxed), + counts.verified.load(Ordering::Relaxed), + ) +} + +#[tokio::test] +async fn mutual_attestation_is_reused_on_repeated_resumption() { + let (server, client) = pair(AttestationType::DcapTdx); + let mut previous = None; + for expected in [ + rustls::HandshakeKind::Full, + rustls::HandshakeKind::Resumed, + rustls::HandshakeKind::Resumed, + ] { + // Clones must share the authenticated session caches. + let (server_result, client_result) = connect(&server.clone(), &client.clone()).await; + let (mut server_stream, server_measurements, server_type) = server_result.unwrap(); + let (mut client_stream, client_measurements, client_type) = client_result.unwrap(); + assert_eq!(server_stream.get_ref().1.handshake_kind(), Some(expected)); + assert_eq!(client_stream.get_ref().1.handshake_kind(), Some(expected)); + assert_eq!(server_type, AttestationType::DcapTdx); + assert_eq!(client_type, AttestationType::DcapTdx); + assert!(server_measurements.is_some() && client_measurements.is_some()); + let metadata = ( + server_measurements, + client_measurements, + server_type, + client_type, + ); + if let Some(previous) = &previous { + assert_eq!(&metadata, previous); + } + previous = Some(metadata); + tokio::time::timeout(Duration::from_secs(5), async { + client_stream.write_all(b"ping").await.unwrap(); + client_stream.flush().await.unwrap(); + let mut buffer = [0; 4]; + server_stream.read_exact(&mut buffer).await.unwrap(); + assert_eq!(&buffer, b"ping"); + server_stream.write_all(b"pong").await.unwrap(); + server_stream.flush().await.unwrap(); + client_stream.read_exact(&mut buffer).await.unwrap(); + assert_eq!(&buffer, b"pong"); + }) + .await + .unwrap(); + assert_eq!(counts(&server.counts), (1, 1)); + assert_eq!(counts(&client.counts), (1, 1)); + println!( + "{expected:?}: server and client each generated 1 quote and verified 1 quote overall" + ); + } +} + +#[tokio::test] +async fn rejected_server_attestation_never_authorizes_client_tickets() { + let (server, client) = pair(AttestationType::None); + for attempt in 1..=2 { + let (server_result, client_result) = connect(&server, &client).await; + assert!(server_result.is_err()); + assert!(matches!( + client_result, + Err(AttestedTlsError::Attestation(_)) + )); + // The TLS tickets actually arrived before the rejected evidence. + assert!(client.sessions.pending_count() > 0); + assert_eq!(counts(&client.counts), (0, attempt)); + assert_eq!(counts(&server.counts).0, attempt); + let store = ClientStore { + cache: client.sessions.clone(), + target: "127.0.0.1".into(), + record: ConnectionRecord::default(), + }; + use rustls::client::ClientSessionStore; + assert!( + store + .take_tls13_ticket(&server_name_from_host("127.0.0.1").unwrap()) + .is_none() + ); + } +} + +/// Offers tickets from rejected or abandoned exchanges using an ungated attacker cache. +async fn unverified_client_cannot_resume(abandon: bool) { + let (server_certs, _, server_config, client_config) = configs(); + let server = AttestedTlsServer::new_with_tls_config( + server_certs, + server_config, + AttestationGenerator::new(AttestationType::DcapTdx, None).unwrap(), + AttestationVerifier::mock(), + ) + .unwrap(); + // An attacker deliberately uses rustls's ordinary, ungated client cache. + let connector = TlsConnector::from(Arc::new(client_config)); + for attempt in 1..=2 { + let (server_io, client_io) = tokio::io::duplex(128 * 1024); + let (server_result, ()) = tokio::time::timeout(Duration::from_secs(5), async { + tokio::join!(server.handle_connection(server_io), async { + let mut stream = connector + .connect(server_name_from_host("127.0.0.1").unwrap(), client_io) + .await + .unwrap(); + // Even on retry, the server must decline the attacker's ticket. + assert_eq!( + stream.get_ref().1.handshake_kind(), + Some(rustls::HandshakeKind::Full) + ); + // Reading the quote also processes the preceding TLS tickets. + read_length_prefixed_attestation(&mut stream).await.unwrap(); + if !abandon { + let message = AttestationExchangeMessage::without_attestation().encode(); + stream + .write_all(&checked_length_prefix(&message).unwrap()) + .await + .unwrap(); + stream.write_all(&message).await.unwrap(); + stream.flush().await.unwrap(); + // Rejection closes TLS without exposing any application data. + assert!(stream.read_u8().await.is_err()); + } + }) + }) + .await + .unwrap(); + if abandon { + assert!(matches!(server_result, Err(AttestedTlsError::Io(_)))); + } else { + assert!(matches!( + server_result, + Err(AttestedTlsError::Attestation(_)) + )); + assert_eq!(counts(&server.counts).1, attempt); + } + assert_eq!(counts(&server.counts).0, attempt); + assert!(server.sessions.pending_count() > 0); + // Prove retry offered a real ticket, rather than silently doing full TLS. + assert_eq!(server.sessions.rejected_count(), attempt - 1); + } +} + +#[tokio::test] +async fn rejected_client_ticket_cannot_bypass_attestation() { + unverified_client_cannot_resume(false).await; +} + +#[tokio::test] +async fn abandoned_exchange_ticket_cannot_bypass_attestation() { + unverified_client_cannot_resume(true).await; +} + +#[tokio::test] +async fn losing_either_cache_requires_fresh_attestation() { + for clear_server in [false, true] { + let (server, client) = pair(AttestationType::DcapTdx); + let (s, c) = connect(&server, &client).await; + s.unwrap(); + c.unwrap(); + if clear_server { + server.sessions.clear(); + } else { + client.sessions.clear(); + } + let (s, c) = connect(&server, &client).await; + assert_eq!( + s.unwrap().0.get_ref().1.handshake_kind(), + Some(rustls::HandshakeKind::Full) + ); + assert_eq!( + c.unwrap().0.get_ref().1.handshake_kind(), + Some(rustls::HandshakeKind::Full) + ); + assert_eq!(counts(&server.counts), (2, 2)); + assert_eq!(counts(&client.counts), (2, 2)); + } +} + +#[tokio::test] +async fn accepted_no_client_attestation_is_reused() { + let (mut server, mut client) = pair(AttestationType::DcapTdx); + server.attestation_verifier = AttestationVerifier::expect_none(); + client.attestation_generator = AttestationGenerator::with_no_attestation(); + for expected in [rustls::HandshakeKind::Full, rustls::HandshakeKind::Resumed] { + let (s, c) = connect(&server, &client).await; + let (server_stream, measurements, attestation_type) = s.unwrap(); + let (client_stream, client_measurements, _) = c.unwrap(); + assert_eq!(server_stream.get_ref().1.handshake_kind(), Some(expected)); + assert_eq!(client_stream.get_ref().1.handshake_kind(), Some(expected)); + assert!(measurements.is_none()); + assert_eq!(attestation_type, AttestationType::None); + assert!(client_measurements.is_some()); + assert_eq!(counts(&server.counts), (1, 1)); + assert_eq!(counts(&client.counts), (0, 1)); + } +} + +#[test] +fn experimental_configuration_disables_early_data_and_rejects_stateless_tickets() { + let (server_certs, client_certs, mut server_config, mut client_config) = configs(); + server_config.max_early_data_size = 1024; + client_config.enable_early_data = true; + let server = AttestedTlsServer::new_with_tls_config( + server_certs.clone(), + server_config.clone(), + AttestationGenerator::with_no_attestation(), + AttestationVerifier::expect_none(), + ) + .unwrap(); + let client = AttestedTlsClient::new_with_tls_config( + client_config, + AttestationGenerator::with_no_attestation(), + AttestationVerifier::expect_none(), + Some(client_certs), + ) + .unwrap(); + assert_eq!(server.acceptor.config().max_early_data_size, 0); + assert!(!client.connector.config().enable_early_data); + server_config.ticketer = rustls::crypto::aws_lc_rs::Ticketer::new().unwrap(); + assert!(matches!( + AttestedTlsServer::new_with_tls_config( + server_certs, + server_config, + AttestationGenerator::with_no_attestation(), + AttestationVerifier::expect_none(), + ), + Err(AttestedTlsError::StatelessResumptionUnsupported) + )); +} + +/// Poisons a mutex without letting the deliberate test panic escape. +fn poison(mutex: &std::sync::Mutex) { + let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| { + let _guard = mutex.lock().unwrap(); + panic!("deliberately poison test state"); + })); + assert!(result.is_err()); + assert!(mutex.is_poisoned()); +} + +#[tokio::test] +async fn poisoned_caches_and_ticket_metadata_require_fresh_attestation() { + for server_side in [false, true] { + for poison_metadata in [false, true] { + let (server, client) = pair(AttestationType::DcapTdx); + let (s, c) = connect(&server, &client).await; + s.unwrap(); + c.unwrap(); + if poison_metadata { + // All tickets from the first connection share this verified record. + let authentication = if server_side { + server + .sessions + .tickets + .lock() + .unwrap() + .values() + .next() + .unwrap() + .authentication + .clone() + } else { + client + .sessions + .tickets + .lock() + .unwrap() + .values() + .next() + .unwrap()[0] + .authentication + .clone() + }; + poison(&authentication); + } else if server_side { + poison(&server.sessions.tickets); + } else { + poison(&client.sessions.tickets); + } + let (s, c) = connect(&server, &client).await; + assert_eq!( + s.unwrap().0.get_ref().1.handshake_kind(), + Some(rustls::HandshakeKind::Full) + ); + assert_eq!( + c.unwrap().0.get_ref().1.handshake_kind(), + Some(rustls::HandshakeKind::Full) + ); + assert_eq!(counts(&server.counts), (2, 2)); + assert_eq!(counts(&client.counts), (2, 2)); + } + } +} + +#[test] +fn poisoned_connection_records_return_errors() { + let record = ConnectionRecord::default(); + poison(&record.selected); + assert!(matches!( + record.verified_peer(), + Err(AttestedTlsError::PoisonedResumptionState) + )); + poison(&record.issued); + assert!(matches!( + record.authenticate(VerifiedPeer { + measurements: None, + attestation_type: AttestationType::None, + }), + Err(AttestedTlsError::PoisonedResumptionState) + )); +} + +#[tokio::test] +async fn poisoned_selected_records_decline_tickets_without_poisoning_caches() { + use rustls::{client::ClientSessionStore, server::StoresServerSessions}; + let (server, client) = pair(AttestationType::DcapTdx); + let (s, c) = connect(&server, &client).await; + s.unwrap(); + c.unwrap(); + let record = ConnectionRecord::default(); + poison(&record.selected); + let client_store = ClientStore { + cache: client.sessions.clone(), + target: "127.0.0.1".into(), + record: record.clone(), + }; + assert!( + client_store + .take_tls13_ticket(&server_name_from_host("127.0.0.1").unwrap()) + .is_none() + ); + let key = server + .sessions + .tickets + .lock() + .unwrap() + .keys() + .next() + .unwrap() + .clone(); + let server_store = ServerStore { + cache: server.sessions.clone(), + record, + }; + assert!(server_store.take(&key).is_none()); + assert!(!server.sessions.tickets.is_poisoned()); + assert!(!client.sessions.tickets.is_poisoned()); +}