use crate::credentials::{Credentials, KeyId};
use bitvec::BitArr;
use std::sync::{
atomic::{AtomicU64, Ordering},
Mutex,
};
const WINDOW: usize = 896;
type Seen = BitArr!(for WINDOW);
#[derive(Debug)]
pub struct State {
max_seen_key_id: AtomicU64,
seen: Mutex<Seen>,
}
impl super::map::SizeOf for Mutex<Seen> {
fn size(&self) -> usize {
if cfg!(target_os = "linux") {
assert!(
!std::mem::needs_drop::<Self>(),
"{:?} requires custom SizeOf impl",
std::any::type_name::<Self>()
);
}
std::mem::size_of::<Self>()
}
}
impl super::map::SizeOf for State {
fn size(&self) -> usize {
let State {
max_seen_key_id,
seen,
} = self;
max_seen_key_id.size() + seen.size()
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
pub enum Error {
#[error("packet definitely already seen before")]
AlreadyExists,
#[error("packet may have been seen before")]
Unknown,
}
impl State {
pub fn new() -> State {
State {
max_seen_key_id: AtomicU64::new(u64::MAX),
seen: Default::default(),
}
}
pub fn pre_authentication(&self, _identity: &Credentials) -> Result<(), Error> {
Ok(())
}
pub fn minimum_unseen_key_id(&self) -> KeyId {
KeyId::try_from(
self.max_seen_key_id
.load(Ordering::Relaxed)
.wrapping_add(1),
)
.unwrap()
}
pub fn post_authentication(&self, identity: &Credentials) -> Result<(), Error> {
let mut seen = self.seen.lock().unwrap();
let key_id = *identity.key_id;
let mut previous_max = self.max_seen_key_id.load(Ordering::Relaxed);
let new_max = if previous_max == u64::MAX {
previous_max = 0;
key_id
} else {
previous_max.max(key_id)
};
self.max_seen_key_id.store(new_max, Ordering::Relaxed);
let delta = new_max - previous_max;
if delta > seen.len() as u64 {
seen.fill(false);
} else {
seen.shift_right(delta as usize);
}
let Ok(idx) = usize::try_from(new_max - key_id) else {
return Err(Error::Unknown);
};
let ret = if let Some(mut entry) = seen.get_mut(idx) {
if *entry {
return Err(Error::AlreadyExists);
}
entry.set(true);
Ok(())
} else {
return Err(Error::Unknown);
};
ret
}
}
impl Default for State {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests;