diff --git a/Cargo.lock b/Cargo.lock index 5c05d5da..fdbb65a5 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2546,6 +2546,7 @@ dependencies = [ "x25519-dalek", "x509-parser", "yasna", + "zeroize", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index 3f4ffe80..61b51701 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -73,6 +73,7 @@ validator = { version = "0.20.0", features = ["derive"] } regex = "1.9.0" lazy_static = "1.4.0" subtle = "2.6.1" +zeroize = { version = "1.8", features = ["derive"] } tiktoken-rs = "0.5" once_cell = "1.19" diff --git a/src/bounded_ttl_map.rs b/src/bounded_ttl_map.rs new file mode 100644 index 00000000..e18652ae --- /dev/null +++ b/src/bounded_ttl_map.rs @@ -0,0 +1,235 @@ +use std::{borrow::Borrow, collections::HashMap, hash::Hash, time::Duration}; +use tokio::time::Instant; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) struct CapacityError; + +struct Entry { + value: V, + expires_at: Instant, +} + +/// A fixed-capacity map whose entries expire after a period of inactivity. +/// +/// Callers are expected to put this behind a lock. Capacity checks and inserts +/// then happen atomically, so concurrent requests cannot exceed the configured +/// memory bound. +pub(crate) struct BoundedTtlMap { + entries: HashMap>, + capacity: usize, + idle_ttl: Duration, +} + +impl BoundedTtlMap +where + K: Eq + Hash, +{ + pub(crate) fn new(capacity: usize, idle_ttl: Duration) -> Self { + assert!(capacity > 0, "bounded map capacity must be non-zero"); + assert!(!idle_ttl.is_zero(), "bounded map TTL must be non-zero"); + + Self { + entries: HashMap::new(), + capacity, + idle_ttl, + } + } + + pub(crate) fn try_insert(&mut self, key: K, value: V) -> Result, CapacityError> { + self.try_insert_at(key, value, Instant::now()) + } + + pub(crate) fn try_insert_at( + &mut self, + key: K, + value: V, + now: Instant, + ) -> Result, CapacityError> { + self.prune_expired_at(now); + + if self.entries.len() >= self.capacity && !self.entries.contains_key(&key) { + return Err(CapacityError); + } + + Ok(self + .entries + .insert( + key, + Entry { + value, + expires_at: now + self.idle_ttl, + }, + ) + .map(|entry| entry.value)) + } + + pub(crate) fn remove_live(&mut self, key: &Q) -> Option + where + K: Borrow, + Q: Eq + Hash + ?Sized, + { + self.remove_live_at(key, Instant::now()) + } + + fn remove_live_at(&mut self, key: &Q, now: Instant) -> Option + where + K: Borrow, + Q: Eq + Hash + ?Sized, + { + let entry = self.entries.remove(key)?; + (entry.expires_at > now).then_some(entry.value) + } + + pub(crate) fn get_cloned_and_touch_at(&mut self, key: &K, now: Instant) -> Option + where + V: Clone, + { + if self.entries.get(key)?.expires_at <= now { + self.entries.remove(key); + return None; + } + + let entry = self.entries.get_mut(key)?; + entry.expires_at = now + self.idle_ttl; + Some(entry.value.clone()) + } + + fn prune_expired_at(&mut self, now: Instant) { + self.entries.retain(|_, entry| entry.expires_at > now); + } + + #[cfg(test)] + fn len(&self) -> usize { + self.entries.len() + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::{ + atomic::{AtomicUsize, Ordering}, + Arc, + }; + use tokio::sync::{Barrier, RwLock}; + + const TTL: Duration = Duration::from_secs(10); + + #[test] + fn capacity_rejection_preserves_live_entries() { + let now = Instant::now(); + let mut map = BoundedTtlMap::new(2, TTL); + + assert_eq!(map.try_insert_at("one", 1, now), Ok(None)); + assert_eq!(map.try_insert_at("two", 2, now), Ok(None)); + assert_eq!(map.try_insert_at("three", 3, now), Err(CapacityError)); + + assert_eq!(map.len(), 2); + assert_eq!(map.get_cloned_and_touch_at(&"one", now), Some(1)); + assert_eq!(map.get_cloned_and_touch_at(&"two", now), Some(2)); + } + + #[test] + fn insert_prunes_expired_entries_without_removing_live_entries() { + let now = Instant::now(); + let mut map = BoundedTtlMap::new(2, TTL); + + map.try_insert_at("expired", 1, now).unwrap(); + map.try_insert_at("live", 2, now + Duration::from_secs(5)) + .unwrap(); + + assert_eq!( + map.try_insert_at("new", 3, now + TTL), + Ok(None), + "the expired slot should be reclaimed" + ); + assert_eq!(map.len(), 2); + assert_eq!(map.get_cloned_and_touch_at(&"live", now + TTL), Some(2)); + assert_eq!(map.get_cloned_and_touch_at(&"expired", now + TTL), None); + } + + #[test] + fn removal_is_one_time_and_rejects_expired_entries() { + let now = Instant::now(); + let mut map = BoundedTtlMap::new(2, TTL); + + map.try_insert_at("once", 1, now).unwrap(); + assert_eq!(map.remove_live_at(&"once", now), Some(1)); + assert_eq!(map.remove_live_at(&"once", now), None); + + map.try_insert_at("expired", 2, now).unwrap(); + assert_eq!(map.remove_live_at(&"expired", now + TTL), None); + assert_eq!(map.len(), 0); + } + + #[test] + fn successful_access_extends_idle_expiry() { + let now = Instant::now(); + let mut map = BoundedTtlMap::new(1, TTL); + map.try_insert_at("session", 7, now).unwrap(); + + assert_eq!( + map.get_cloned_and_touch_at(&"session", now + Duration::from_secs(5)), + Some(7) + ); + assert_eq!( + map.get_cloned_and_touch_at(&"session", now + Duration::from_secs(14)), + Some(7), + "touching the session should keep an active request alive" + ); + assert_eq!( + map.get_cloned_and_touch_at(&"session", now + Duration::from_secs(24)), + None + ); + assert_eq!(map.len(), 0); + } + + #[test] + fn expired_lookup_drops_the_value_immediately() { + struct DropSpy(Arc); + + impl Drop for DropSpy { + fn drop(&mut self) { + self.0.fetch_add(1, Ordering::SeqCst); + } + } + + let now = Instant::now(); + let drops = Arc::new(AtomicUsize::new(0)); + let mut map = BoundedTtlMap::new(1, TTL); + map.try_insert_at("session", Arc::new(DropSpy(Arc::clone(&drops))), now) + .unwrap(); + + assert!(map.get_cloned_and_touch_at(&"session", now + TTL).is_none()); + assert_eq!(map.len(), 0); + assert_eq!(drops.load(Ordering::SeqCst), 1); + } + + #[tokio::test] + async fn concurrent_inserts_cannot_exceed_capacity() { + const CAPACITY: usize = 8; + const REQUESTS: usize = 64; + + let map = Arc::new(RwLock::new(BoundedTtlMap::new(CAPACITY, TTL))); + let barrier = Arc::new(Barrier::new(REQUESTS)); + let now = Instant::now(); + let mut tasks = Vec::with_capacity(REQUESTS); + + for key in 0..REQUESTS { + let map = Arc::clone(&map); + let barrier = Arc::clone(&barrier); + tasks.push(tokio::spawn(async move { + barrier.wait().await; + map.write().await.try_insert_at(key, key, now).is_ok() + })); + } + + let mut successful_inserts = 0; + for task in tasks { + successful_inserts += usize::from(task.await.unwrap()); + } + + assert_eq!(successful_inserts, CAPACITY); + assert_eq!(map.read().await.len(), CAPACITY); + } +} diff --git a/src/main.rs b/src/main.rs index 86cc5ff6..e76a0980 100644 --- a/src/main.rs +++ b/src/main.rs @@ -60,7 +60,6 @@ use secp256k1::SecretKey; use serde::{Deserialize, Serialize}; use serde_json::Value; use sha2::{Digest, Sha256}; -use std::collections::HashMap; use std::env; use std::fmt; use std::io::Write; @@ -83,6 +82,7 @@ use x25519_dalek::{EphemeralSecret, PublicKey}; mod apple_signin; mod aws_credentials; mod billing; +mod bounded_ttl_map; mod brave; mod db; mod email; @@ -109,6 +109,7 @@ mod web; mod aead_db_tamper_tests; use apple_signin::AppleJwtVerifier; +use bounded_ttl_map::BoundedTtlMap; use oauth::{AppleProvider, GithubProvider, GoogleProvider, OAuthManager}; use provider_routing::{ProviderName, ProviderPreference, ProviderRouter}; use proxy_config::ProxyRouter; @@ -117,6 +118,21 @@ const ENCLAVE_KEY_NAME: &str = "enclave_key"; const OPENAI_API_KEY_NAME: &str = "openai_api_key"; const JWT_SECRET_KEY_NAME: &str = "jwt_secret"; +// Attestation nonces are only needed between the attestation request and key +// exchange. Nitro attestation documents permit at most 512 nonce bytes. That +// size limit, a short TTL, and a fixed entry cap bound abandoned handshakes. +const MAX_ATTESTATION_NONCE_BYTES: usize = 512; +const EPHEMERAL_KEY_TTL: Duration = Duration::from_secs(5 * 60); +const MAX_EPHEMERAL_KEYS: usize = 4_096; + +// JWT refresh does not rotate encryption sessions. This 65-minute idle TTL is +// an independent in-memory lifecycle choice: official clients re-attest when a +// session is missing, while successful use extends the TTL so long-running +// streams are not interrupted. The fixed cap, rather than the TTL, is the hard +// memory bound. +const SESSION_STATE_IDLE_TTL: Duration = Duration::from_secs(65 * 60); +const MAX_SESSION_STATES: usize = 65_536; + // General secret key names const GITHUB_CLIENT_ID_NAME: &str = "github_client_id"; const GITHUB_CLIENT_SECRET_NAME: &str = "github_client_secret"; @@ -327,6 +343,9 @@ pub enum ApiError { #[error("Payload too large")] PayloadTooLarge, + + #[error("Too many requests")] + TooManyRequests, } impl IntoResponse for ApiError { @@ -352,6 +371,7 @@ impl IntoResponse for ApiError { ApiError::NotFound => StatusCode::NOT_FOUND, ApiError::UnprocessableEntity => StatusCode::UNPROCESSABLE_ENTITY, ApiError::PayloadTooLarge => StatusCode::PAYLOAD_TOO_LARGE, + ApiError::TooManyRequests => StatusCode::TOO_MANY_REQUESTS, }; ( status, @@ -364,6 +384,63 @@ impl IntoResponse for ApiError { } } +fn live_session_state_at( + session_states: &mut BoundedTtlMap>, + session_id: &Uuid, + now: tokio::time::Instant, +) -> Result, ApiError> { + session_states + .get_cloned_and_touch_at(session_id, now) + // The official SDKs use 400 as the signal to establish a fresh + // attestation session. This also covers GET/DELETE routes, where the + // first session lookup may happen while encrypting the response. + .ok_or(ApiError::BadRequest) +} + +fn validate_attestation_nonce(nonce: &str) -> Result<(), ApiError> { + if nonce.len() > MAX_ATTESTATION_NONCE_BYTES { + return Err(ApiError::BadRequest); + } + Ok(()) +} + +#[cfg(test)] +mod attestation_session_state_tests { + use super::*; + + #[test] + fn missing_and_expired_sessions_use_the_reattest_response() { + let now = tokio::time::Instant::now(); + let session_id = Uuid::new_v4(); + let mut sessions = BoundedTtlMap::new(1, Duration::from_secs(10)); + + let missing = live_session_state_at(&mut sessions, &session_id, now) + .err() + .expect("missing session should fail"); + assert_eq!(missing.into_response().status(), StatusCode::BAD_REQUEST); + + sessions + .try_insert_at(session_id, Arc::new(SessionState::new([7_u8; 32])), now) + .unwrap(); + assert!( + live_session_state_at(&mut sessions, &session_id, now + Duration::from_secs(9)).is_ok() + ); + assert!( + live_session_state_at(&mut sessions, &session_id, now + Duration::from_secs(19)) + .is_err() + ); + } + + #[test] + fn attestation_nonce_size_is_bounded_in_bytes() { + assert!(validate_attestation_nonce(&"a".repeat(MAX_ATTESTATION_NONCE_BYTES)).is_ok()); + + let oversized = validate_attestation_nonce(&"a".repeat(MAX_ATTESTATION_NONCE_BYTES + 1)) + .expect_err("oversized nonce should fail"); + assert_eq!(oversized.into_response().status(), StatusCode::BAD_REQUEST); + } +} + impl From for ApiError { fn from(err: DBError) -> Self { error!("Database error: {:?}", err); @@ -466,8 +543,8 @@ pub struct AppState { proxy_router: Arc, provider_router: Arc, resend_api_key: Option, - ephemeral_keys: Arc>>, - session_states: Arc>>, + ephemeral_keys: Arc>>, + session_states: Arc>>>, oauth_manager: Arc, sqs_publisher: Option>, billing_client: Option, @@ -769,8 +846,14 @@ impl AppStateBuilder { proxy_router, provider_router, resend_api_key: self.resend_api_key, - ephemeral_keys: Arc::new(RwLock::new(HashMap::new())), - session_states: Arc::new(tokio::sync::RwLock::new(HashMap::new())), + ephemeral_keys: Arc::new(RwLock::new(BoundedTtlMap::new( + MAX_EPHEMERAL_KEYS, + EPHEMERAL_KEY_TTL, + ))), + session_states: Arc::new(RwLock::new(BoundedTtlMap::new( + MAX_SESSION_STATES, + SESSION_STATE_IDLE_TTL, + ))), oauth_manager, sqs_publisher, billing_client, @@ -1486,7 +1569,9 @@ impl AppState { self.enclave_key.clone() } - pub async fn create_ephemeral_key(&self, nonce: String) -> PublicKey { + pub async fn create_ephemeral_key(&self, nonce: String) -> Result { + validate_attestation_nonce(&nonce)?; + let custom_rng = CustomRng::new(); // Use a wrapper that implements RngCore and CryptoRng @@ -1499,13 +1584,39 @@ impl AppState { self.ephemeral_keys .write() .await - .insert(nonce, ephemeral_secret); + .try_insert(nonce, ephemeral_secret) + .map_err(|_| { + warn!( + max_entries = MAX_EPHEMERAL_KEYS, + "Attestation ephemeral-key capacity reached" + ); + ApiError::TooManyRequests + })?; - public_key + Ok(public_key) } pub async fn get_and_remove_ephemeral_secret(&self, nonce: &str) -> Option { - self.ephemeral_keys.write().await.remove(nonce) + self.ephemeral_keys.write().await.remove_live(nonce) + } + + pub async fn store_session_state( + &self, + session_id: Uuid, + session_state: Arc, + ) -> Result<(), ApiError> { + self.session_states + .write() + .await + .try_insert(session_id, session_state) + .map(|_| ()) + .map_err(|_| { + warn!( + max_entries = MAX_SESSION_STATES, + "Attestation session capacity reached" + ); + ApiError::TooManyRequests + }) } pub async fn decrypt_session_data( @@ -1538,19 +1649,19 @@ impl AppState { tracing::trace!("nonce: {:?}", nonce_array); tracing::trace!("ciphertext length: {}", ciphertext.len()); - self.session_states - .read() - .await - .get(session_id) - .ok_or_else(|| { - tracing::error!("Session not found: {}", session_id); - ApiError::Unauthorized - }) - .and_then(|state| { - state.decrypt(ciphertext, &nonce_array).map_err(|e| { - tracing::error!("Decryption failed: {:?}", e); - e - }) + let session_state = { + let mut session_states = self.session_states.write().await; + live_session_state_at(&mut session_states, session_id, tokio::time::Instant::now()) + } + .inspect_err(|_| { + tracing::warn!("Session missing or expired: {}", session_id); + })?; + + session_state + .decrypt(ciphertext, &nonce_array) + .map_err(|e| { + tracing::error!("Decryption failed: {:?}", e); + e }) } @@ -1559,13 +1670,13 @@ impl AppState { session_id: &Uuid, data: &[u8], ) -> Result, ApiError> { - let session_states = self.session_states.read().await; - let session_state = session_states - .get(session_id) - .ok_or(ApiError::Unauthorized)?; + let session_state = { + let mut session_states = self.session_states.write().await; + live_session_state_at(&mut session_states, session_id, tokio::time::Instant::now()) + }?; let session_key = session_state.get_session_key(); - let key = Key::from_slice(session_key.as_ref()); + let key = Key::from_slice(session_key); let nonce_bytes: [u8; 12] = crate::encrypt::generate_random(); let nonce = Nonce::from_slice(&nonce_bytes); diff --git a/src/web/attestation_routes.rs b/src/web/attestation_routes.rs index a623632f..9019cfdd 100644 --- a/src/web/attestation_routes.rs +++ b/src/web/attestation_routes.rs @@ -20,7 +20,9 @@ use tracing::{error, trace}; use uuid::Uuid; use yasna::models::ObjectIdentifier; use yasna::{construct_der, Tag}; +use zeroize::{Zeroize, ZeroizeOnDrop}; +#[derive(Zeroize, ZeroizeOnDrop)] pub struct SessionState { session_key: [u8; 32], } @@ -30,13 +32,12 @@ impl SessionState { Self { session_key } } - pub fn get_session_key(&self) -> [u8; 32] { - self.session_key + pub fn get_session_key(&self) -> &[u8; 32] { + &self.session_key } pub fn decrypt(&self, encrypted_data: &[u8], nonce: &[u8; 12]) -> Result, ApiError> { tracing::trace!("decrypting encrypted data"); - tracing::trace!("session key: {:?}", self.session_key); tracing::trace!("nonce: {:?}", nonce); tracing::trace!("encrypted data length: {}", encrypted_data.len()); @@ -81,7 +82,7 @@ async fn get_attestation( ) -> Result<(StatusCode, Json), ApiError> { // Create an ephemeral key pair for this request trace!("Creating ephemeral key"); - let enclave_public_key = data.create_ephemeral_key(nonce.clone()).await; + let enclave_public_key = data.create_ephemeral_key(nonce.clone()).await?; trace!("Ephemeral key created"); // Create a request for the attestation document @@ -402,7 +403,7 @@ async fn key_exchange( let shared_secret = ephemeral_secret.diffie_hellman(&client_public_key); // Generate a random session key using your secure random function - let session_key: [u8; 32] = crate::encrypt::generate_random(); + let session_state = Arc::new(SessionState::new(crate::encrypt::generate_random())); // Encrypt the session key using the shared secret let nonce_bytes: [u8; 12] = crate::encrypt::generate_random(); @@ -411,24 +412,15 @@ async fn key_exchange( let mut encrypted_session_key = nonce_bytes.to_vec(); encrypted_session_key.extend_from_slice( &cipher - .encrypt(nonce, session_key.as_ref()) + .encrypt(nonce, session_state.get_session_key().as_ref()) .map_err(|_| ApiError::InternalServerError)?, ); // Generate a new UUID for the session let session_id = Uuid::new_v4(); - trace!( - "Generated session key {:?} for nonce {:?}", - session_key, - nonce - ); - // Store the session state - data.session_states - .write() - .await - .insert(session_id, SessionState::new(session_key)); + data.store_session_state(session_id, session_state).await?; Ok(Json(KeyExchangeResponse { session_id, encrypted_session_key: general_purpose::STANDARD.encode(&encrypted_session_key),