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,
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];
+131 -4
View File
@@ -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<ReplayKey, ReplayEntry>,
pending: HashMap<ReplayKey, u64>,
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<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),
additions: 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)
}
/// 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 {
self.check_and_add_handshake(data)
}
@@ -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");
}