From b37f1ebdeb13bac5acdaa50e50dd26939dd9ff8a Mon Sep 17 00:00:00 2001 From: Alexey <247128645+axkurcom@users.noreply.github.com> Date: Sun, 13 Sep 2026 21:56:20 +0300 Subject: [PATCH] TLS Replay Digest In-flight Ownership --- src/proxy/handshake/tls_handshake.rs | 15 +- src/stats/replay.rs | 135 +++++++++++++++++- .../tests/replay_checker_security_tests.rs | 30 ++++ 3 files changed, 171 insertions(+), 9 deletions(-) diff --git a/src/proxy/handshake/tls_handshake.rs b/src/proxy/handshake/tls_handshake.rs index c98dc17..f1838ea 100644 --- a/src/proxy/handshake/tls_handshake.rs +++ b/src/proxy/handshake/tls_handshake.rs @@ -257,14 +257,15 @@ where secret: validated_secret, user_id: validated_user_id, } = validation; - // Reject known replay digests before expensive cache/domain/ALPN policy work. + // Reserve the replay digest before any asynchronous policy work so a concurrent + // duplicate cannot pass the check-to-commit window. let digest_half = &validation_digest[..tls::TLS_DIGEST_HALF_LEN]; - if replay_checker.check_tls_digest(digest_half) { + let Some(replay_claim) = replay_checker.claim_tls_digest(digest_half) else { auth_probe_record_failure_in(shared, peer.ip(), Instant::now()); maybe_apply_server_hello_delay(config).await; warn!(peer = %peer, "TLS replay attack detected (duplicate digest)"); return HandshakeResult::BadClient { reader, writer }; - } + }; let selected_tls_domain = matched_tls_domain.unwrap_or(config.censorship.tls_domain.as_str()); let cached_entry = if config.censorship.tls_emulation { @@ -337,8 +338,12 @@ where None }; - // Add replay digest only for policy-valid handshakes. - replay_checker.add_tls_digest(digest_half); + // Commit only policy-valid handshakes; early returns release the pending claim. + if !replay_claim.commit() { + auth_probe_record_failure_in(shared, peer.ip(), Instant::now()); + warn!(peer = %peer, "TLS replay claim lost before commit"); + return HandshakeResult::BadClient { reader, writer }; + } let validation_session_id_slice = &validation_session_id[..validation_session_id_len]; diff --git a/src/stats/replay.rs b/src/stats/replay.rs index c0e2257..13b0d8f 100644 --- a/src/stats/replay.rs +++ b/src/stats/replay.rs @@ -1,5 +1,5 @@ use std::borrow::Borrow; -use std::collections::VecDeque; +use std::collections::{HashMap, VecDeque}; use std::collections::hash_map::DefaultHasher; use std::hash::{Hash, Hasher}; use std::num::NonZeroUsize; @@ -74,6 +74,7 @@ pub struct ReplayChecker { hits: AtomicU64, additions: AtomicU64, cleanups: AtomicU64, + next_claim_token: AtomicU64, } struct ReplayEntry { @@ -82,6 +83,7 @@ struct ReplayEntry { struct ReplayShard { cache: LruCache, + pending: HashMap, queue: VecDeque<(Instant, ReplayKey, u64)>, seq_counter: u64, capacity: usize, @@ -91,6 +93,7 @@ impl ReplayShard { fn new(cap: NonZeroUsize) -> Self { Self { cache: LruCache::new(cap), + pending: HashMap::new(), queue: VecDeque::with_capacity(cap.get()), seq_counter: 0, capacity: cap.get(), @@ -135,7 +138,7 @@ impl ReplayShard { return false; } self.cleanup(now, window); - self.cache.get(key).is_some() + self.cache.get(key).is_some() || self.pending.contains_key(key) } fn add_owned(&mut self, key: ReplayKey, now: Instant, window: Duration) { @@ -143,7 +146,7 @@ impl ReplayShard { return; } self.cleanup(now, window); - if self.cache.peek(key.as_slice()).is_some() { + if self.cache.peek(key.as_slice()).is_some() || self.pending.contains_key(key.as_slice()) { return; } while self.queue.len() >= self.capacity { @@ -155,8 +158,83 @@ impl ReplayShard { self.queue.push_back((now, key, seq)); } + fn claim_owned( + &mut self, + key: ReplayKey, + now: Instant, + window: Duration, + token: u64, + ) -> bool { + if window.is_zero() { + return true; + } + self.cleanup(now, window); + if self.cache.peek(key.as_slice()).is_some() || self.pending.contains_key(key.as_slice()) { + return false; + } + while self.cache.len().saturating_add(self.pending.len()) >= self.capacity { + if self.queue.is_empty() { + return false; + } + self.evict_queue_front(); + } + self.pending.insert(key, token); + true + } + + fn remove_pending(&mut self, key: &[u8], token: u64) -> bool { + if self.pending.get(key).copied() != Some(token) { + return false; + } + self.pending.remove(key); + true + } + fn len(&self) -> usize { - self.cache.len() + self.cache.len().saturating_add(self.pending.len()) + } +} + +/// Exclusive in-flight ownership of one TLS replay digest. +#[must_use = "the claim must be committed after all TLS policy checks"] +pub(crate) struct TlsReplayClaim<'a> { + checker: &'a ReplayChecker, + shard_idx: usize, + key: Option, + token: u64, + reserved: bool, +} + +impl TlsReplayClaim<'_> { + /// Commits the claimed digest into the replay window. + pub(crate) fn commit(mut self) -> bool { + if !self.reserved { + return true; + } + let Some(key) = self.key.take() else { + return false; + }; + let mut shard = self.checker.tls_shards[self.shard_idx].lock(); + if !shard.remove_pending(key.as_slice(), self.token) { + return false; + } + shard.add_owned(key, Instant::now(), self.checker.tls_window); + self.checker.additions.fetch_add(1, Ordering::Relaxed); + self.reserved = false; + true + } +} + +impl Drop for TlsReplayClaim<'_> { + fn drop(&mut self) { + if !self.reserved { + return; + } + if let Some(key) = self.key.as_ref() { + self.checker.tls_shards[self.shard_idx] + .lock() + .remove_pending(key.as_slice(), self.token); + } } } @@ -184,6 +262,25 @@ impl ReplayChecker { hits: AtomicU64::new(0), additions: AtomicU64::new(0), cleanups: AtomicU64::new(0), + next_claim_token: AtomicU64::new(1), + } + } + + fn reserve_claim_token(&self) -> Option { + let mut current = self.next_claim_token.load(Ordering::Relaxed); + loop { + if current == u64::MAX { + return None; + } + match self.next_claim_token.compare_exchange_weak( + current, + current + 1, + Ordering::Relaxed, + Ordering::Relaxed, + ) { + Ok(_) => return Some(current), + Err(actual) => current = actual, + } } } @@ -246,6 +343,36 @@ impl ReplayChecker { self.check_and_add_internal(data, &self.tls_shards, self.tls_window) } + /// Claims one TLS digest until the policy-valid handshake commits or aborts. + pub(crate) fn claim_tls_digest(&self, data: &[u8]) -> Option> { + self.checks.fetch_add(1, Ordering::Relaxed); + let shard_idx = self.get_shard_idx(data); + let key = ReplayKey::from_slice(data); + if self.tls_window.is_zero() { + return Some(TlsReplayClaim { + checker: self, + shard_idx, + key: None, + token: 0, + reserved: false, + }); + } + let token = self.reserve_claim_token()?; + let mut shard = self.tls_shards[shard_idx].lock(); + if !shard.claim_owned(key.clone(), Instant::now(), self.tls_window, token) { + self.hits.fetch_add(1, Ordering::Relaxed); + return None; + } + drop(shard); + Some(TlsReplayClaim { + checker: self, + shard_idx, + key: Some(key), + token, + reserved: true, + }) + } + pub fn check_handshake(&self, data: &[u8]) -> bool { self.check_and_add_handshake(data) } diff --git a/src/stats/tests/replay_checker_security_tests.rs b/src/stats/tests/replay_checker_security_tests.rs index 8e73204..e20f39f 100644 --- a/src/stats/tests/replay_checker_security_tests.rs +++ b/src/stats/tests/replay_checker_security_tests.rs @@ -1,4 +1,5 @@ use super::*; +use std::sync::{Arc, Barrier}; use std::time::Duration; #[test] @@ -78,3 +79,32 @@ fn replay_checker_stats_reflect_dual_shard_domains() { "stats should expose both shard domains (handshake + TLS)" ); } + +#[test] +fn concurrent_tls_claim_allows_exactly_one_pending_handshake() { + const WORKERS: usize = 16; + + let checker = Arc::new(ReplayChecker::new(128, Duration::from_secs(1))); + let start = Arc::new(Barrier::new(WORKERS)); + let finish = Arc::new(Barrier::new(WORKERS)); + let mut handles = Vec::with_capacity(WORKERS); + + for _ in 0..WORKERS { + let checker = Arc::clone(&checker); + let start = Arc::clone(&start); + let finish = Arc::clone(&finish); + handles.push(std::thread::spawn(move || { + start.wait(); + let claim = checker.claim_tls_digest(b"parallel-client-hello"); + finish.wait(); + claim.map(|claim| claim.commit()).unwrap_or(false) + })); + } + + let accepted = handles + .into_iter() + .map(|handle| handle.join().expect("TLS claim worker must not panic")) + .filter(|accepted| *accepted) + .count(); + assert_eq!(accepted, 1, "only one concurrent TLS claim may commit"); +}