mirror of
https://github.com/telemt/telemt.git
synced 2026-09-28 13:36:00 +03:00
TLS Replay Digest In-flight Ownership
This commit is contained in:
@@ -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
@@ -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");
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user