use parking_lot::Mutex;
use serde::{Deserialize, Serialize};
use std::collections::HashSet;
use std::hash::{BuildHasher, Hasher};
use std::marker::PhantomData;
use std::sync::atomic::{AtomicU64, Ordering};
pub const HISTORY_LEN: u64 = 1024;
#[derive(Serialize, Deserialize)]
pub struct AntiReplayAttackContainer {
history: Mutex<(u64, HashSet<u64, NoHashHasher<u64>>)>,
counter_out: AtomicU64,
}
const ORDERING: Ordering = Ordering::Relaxed;
impl AntiReplayAttackContainer {
#[inline]
pub fn get_next_pid(&self) -> u64 {
self.counter_out.fetch_add(1, ORDERING)
}
#[allow(unused_results)]
pub fn on_pid_received(&self, pid_received: u64) -> bool {
let mut queue = self.history.lock();
if queue.1.contains(&pid_received) {
log::error!(target: "citadel", "[ARA] packet {} already arrived!", pid_received);
false
} else {
let min = queue.0.saturating_sub(HISTORY_LEN);
if pid_received >= min {
if queue.1.len() >= HISTORY_LEN as _ {
let lowest = queue.0;
if queue.1.remove(&lowest) {
queue.0 += 1;
}
}
queue.1.insert(pid_received);
true
} else {
log::error!(target: "citadel", "[ARA] out of range! Recv: {}. Expected >= {}", pid_received, min);
false
}
}
}
pub fn has_tracked_packets(&self) -> bool {
(self.counter_out.load(ORDERING) != 0) || (self.history.lock().0 != 0)
}
pub fn reset(&self) {
self.counter_out.store(0, ORDERING);
let mut lock = self.history.lock();
lock.0 = 0;
lock.1 = HashSet::with_capacity_and_hasher(HISTORY_LEN as usize, Default::default());
}
}
impl Default for AntiReplayAttackContainer {
fn default() -> Self {
Self {
history: Mutex::new((
0,
HashSet::with_capacity_and_hasher(HISTORY_LEN as usize, Default::default()),
)),
counter_out: AtomicU64::new(0),
}
}
}
struct NoHashHasher<T>(u64, PhantomData<T>);
impl<T> Default for NoHashHasher<T> {
fn default() -> Self {
NoHashHasher(0, PhantomData)
}
}
trait IsEnabled {}
impl IsEnabled for u64 {}
impl<T: IsEnabled> Hasher for NoHashHasher<T> {
fn finish(&self) -> u64 {
self.0
}
fn write(&mut self, _: &[u8]) {
panic!("Invalid use of NoHashHasher")
}
fn write_u64(&mut self, n: u64) {
self.0 = n
}
}
impl<T: IsEnabled> BuildHasher for NoHashHasher<T> {
type Hasher = Self;
fn build_hasher(&self) -> Self::Hasher {
Self::default()
}
}