TLS Replay Digest In-flight Ownership

This commit is contained in:
Alexey
2026-09-13 21:56:20 +03:00
parent 1ea5f7a1d6
commit b37f1ebdeb
3 changed files with 171 additions and 9 deletions
+10 -5
View File
@@ -257,14 +257,15 @@ where
secret: validated_secret, secret: validated_secret,
user_id: validated_user_id, user_id: validated_user_id,
} = validation; } = 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]; 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()); auth_probe_record_failure_in(shared, peer.ip(), Instant::now());
maybe_apply_server_hello_delay(config).await; maybe_apply_server_hello_delay(config).await;
warn!(peer = %peer, "TLS replay attack detected (duplicate digest)"); warn!(peer = %peer, "TLS replay attack detected (duplicate digest)");
return HandshakeResult::BadClient { reader, writer }; return HandshakeResult::BadClient { reader, writer };
} };
let selected_tls_domain = matched_tls_domain.unwrap_or(config.censorship.tls_domain.as_str()); let selected_tls_domain = matched_tls_domain.unwrap_or(config.censorship.tls_domain.as_str());
let cached_entry = if config.censorship.tls_emulation { let cached_entry = if config.censorship.tls_emulation {
@@ -337,8 +338,12 @@ where
None None
}; };
// Add replay digest only for policy-valid handshakes. // Commit only policy-valid handshakes; early returns release the pending claim.
replay_checker.add_tls_digest(digest_half); 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]; let validation_session_id_slice = &validation_session_id[..validation_session_id_len];
+131 -4
View File
@@ -1,5 +1,5 @@
use std::borrow::Borrow; use std::borrow::Borrow;
use std::collections::VecDeque; use std::collections::{HashMap, VecDeque};
use std::collections::hash_map::DefaultHasher; use std::collections::hash_map::DefaultHasher;
use std::hash::{Hash, Hasher}; use std::hash::{Hash, Hasher};
use std::num::NonZeroUsize; use std::num::NonZeroUsize;
@@ -74,6 +74,7 @@ pub struct ReplayChecker {
hits: AtomicU64, hits: AtomicU64,
additions: AtomicU64, additions: AtomicU64,
cleanups: AtomicU64, cleanups: AtomicU64,
next_claim_token: AtomicU64,
} }
struct ReplayEntry { struct ReplayEntry {
@@ -82,6 +83,7 @@ struct ReplayEntry {
struct ReplayShard { struct ReplayShard {
cache: LruCache<ReplayKey, ReplayEntry>, cache: LruCache<ReplayKey, ReplayEntry>,
pending: HashMap<ReplayKey, u64>,
queue: VecDeque<(Instant, ReplayKey, u64)>, queue: VecDeque<(Instant, ReplayKey, u64)>,
seq_counter: u64, seq_counter: u64,
capacity: usize, capacity: usize,
@@ -91,6 +93,7 @@ impl ReplayShard {
fn new(cap: NonZeroUsize) -> Self { fn new(cap: NonZeroUsize) -> Self {
Self { Self {
cache: LruCache::new(cap), cache: LruCache::new(cap),
pending: HashMap::new(),
queue: VecDeque::with_capacity(cap.get()), queue: VecDeque::with_capacity(cap.get()),
seq_counter: 0, seq_counter: 0,
capacity: cap.get(), capacity: cap.get(),
@@ -135,7 +138,7 @@ impl ReplayShard {
return false; return false;
} }
self.cleanup(now, window); 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) { fn add_owned(&mut self, key: ReplayKey, now: Instant, window: Duration) {
@@ -143,7 +146,7 @@ impl ReplayShard {
return; return;
} }
self.cleanup(now, window); 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; return;
} }
while self.queue.len() >= self.capacity { while self.queue.len() >= self.capacity {
@@ -155,8 +158,83 @@ impl ReplayShard {
self.queue.push_back((now, key, seq)); 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 { 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<ReplayKey>,
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), hits: AtomicU64::new(0),
additions: AtomicU64::new(0), additions: AtomicU64::new(0),
cleanups: AtomicU64::new(0), cleanups: AtomicU64::new(0),
next_claim_token: AtomicU64::new(1),
}
}
fn reserve_claim_token(&self) -> Option<u64> {
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) 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<TlsReplayClaim<'_>> {
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 { pub fn check_handshake(&self, data: &[u8]) -> bool {
self.check_and_add_handshake(data) self.check_and_add_handshake(data)
} }
@@ -1,4 +1,5 @@
use super::*; use super::*;
use std::sync::{Arc, Barrier};
use std::time::Duration; use std::time::Duration;
#[test] #[test]
@@ -78,3 +79,32 @@ fn replay_checker_stats_reflect_dual_shard_domains() {
"stats should expose both shard domains (handshake + TLS)" "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");
}