use std::collections::{HashMap, VecDeque};
use std::net::{IpAddr, SocketAddr};
use std::time::Instant;
use crate::constants;
use crate::identity::Identity;
use super::guard::{ChainPin, GuardUndo};
use super::staged::{ChainState, IntroId};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub(crate) enum SourceKey {
V4([u8; 4]),
V6Prefix64([u8; 8]),
}
impl SourceKey {
pub(crate) fn of(addr: SocketAddr) -> Self {
match addr.ip() {
IpAddr::V4(v4) => SourceKey::V4(v4.octets()),
IpAddr::V6(v6) => {
let octets = v6.octets();
let mut prefix = [0u8; 8];
prefix.copy_from_slice(&octets[..8]);
SourceKey::V6Prefix64(prefix)
}
}
}
}
pub(crate) struct IntroEntry<I: Identity> {
pub(crate) id: IntroId,
pub(crate) src: SocketAddr,
pub(crate) source_key: SourceKey,
pub(crate) sender_index: u32,
pub(crate) msg1: Vec<u8>,
pub(crate) state: ChainState<I>,
pub(crate) consumed: bool,
pub(crate) guard_undo: Option<GuardUndo>,
pub(crate) guard_pin: Option<ChainPin>,
refreshed_at: Instant,
}
impl<I: Identity> IntroEntry<I> {
pub(crate) fn age_key(&self) -> Instant {
self.refreshed_at
}
pub(crate) fn deadline(&self) -> Instant {
self.refreshed_at + constants::INTRO_TTL
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum Arrival {
Parked(IntroId),
Refreshed(IntroId),
Dropped,
}
pub(crate) struct ArrivalOutcome<I: Identity> {
pub(crate) arrival: Arrival,
pub(crate) evicted: Option<IntroEntry<I>>,
}
pub(crate) struct IntroQueue<I: Identity> {
entries: HashMap<IntroId, IntroEntry<I>>,
by_addr: HashMap<SocketAddr, IntroId>,
per_source: HashMap<SourceKey, u32>,
evicted: VecDeque<IntroId>,
next_id: u64,
cap: usize,
max_per_source: usize,
}
impl<I: Identity> IntroQueue<I> {
pub(crate) fn new(cap: usize, max_per_source: usize) -> Self {
Self {
entries: HashMap::new(),
by_addr: HashMap::new(),
per_source: HashMap::new(),
evicted: VecDeque::new(),
next_id: 0,
cap,
max_per_source,
}
}
fn note_eviction(&mut self, id: IntroId) {
if self.cap == 0 {
return;
}
while self.evicted.len() >= self.cap {
self.evicted.pop_front();
}
self.evicted.push_back(id);
}
pub(crate) fn was_evicted(&self, id: IntroId) -> bool {
self.evicted.contains(&id)
}
pub(crate) fn arrive(
&mut self,
now: Instant,
src: SocketAddr,
sender_index: u32,
msg1: &[u8],
) -> ArrivalOutcome<I> {
if let Some(&id) = self.by_addr.get(&src)
&& let Some(entry) = self.entries.get_mut(&id)
{
debug_assert!(!entry.consumed, "by_addr holds unconsumed entries only");
entry.msg1.clear();
entry.msg1.extend_from_slice(msg1);
entry.sender_index = sender_index;
entry.refreshed_at = now;
return ArrivalOutcome {
arrival: Arrival::Refreshed(id),
evicted: None,
};
}
let source_key = SourceKey::of(src);
let mut evicted = None;
if u64::from(self.count_for(source_key)) >= self.max_per_source as u64 {
match self.oldest_unconsumed(Some(source_key)) {
Some(victim) => {
self.note_eviction(victim);
evicted = self.remove(victim);
}
None => {
return ArrivalOutcome {
arrival: Arrival::Dropped,
evicted: None,
};
}
}
}
if self.entries.len() >= self.cap {
match self.oldest_unconsumed(None) {
Some(victim) => {
self.note_eviction(victim);
evicted = self.remove(victim);
}
None => {
return ArrivalOutcome {
arrival: Arrival::Dropped,
evicted: None,
};
}
}
}
let id = IntroId::from_raw(self.next_id);
self.next_id += 1;
self.entries.insert(
id,
IntroEntry {
id,
src,
source_key,
sender_index,
msg1: msg1.to_vec(),
state: ChainState::Parked,
consumed: false,
guard_undo: None,
guard_pin: None,
refreshed_at: now,
},
);
self.by_addr.insert(src, id);
*self.per_source.entry(source_key).or_insert(0) += 1;
ArrivalOutcome {
arrival: Arrival::Parked(id),
evicted,
}
}
pub(crate) fn consume(&mut self, id: IntroId) {
if let Some(entry) = self.entries.get_mut(&id)
&& !entry.consumed
{
entry.consumed = true;
let src = entry.src;
if self.by_addr.get(&src) == Some(&id) {
self.by_addr.remove(&src);
}
}
}
pub(crate) fn get(&self, id: IntroId) -> Option<&IntroEntry<I>> {
self.entries.get(&id)
}
pub(crate) fn get_mut(&mut self, id: IntroId) -> Option<&mut IntroEntry<I>> {
self.entries.get_mut(&id)
}
pub(crate) fn remove(&mut self, id: IntroId) -> Option<IntroEntry<I>> {
let entry = self.entries.remove(&id)?;
if self.by_addr.get(&entry.src) == Some(&id) {
self.by_addr.remove(&entry.src);
}
match self.per_source.get_mut(&entry.source_key) {
Some(count) if *count > 1 => *count -= 1,
Some(_) => {
self.per_source.remove(&entry.source_key);
}
None => debug_assert!(false, "per-source counter underflow"),
}
Some(entry)
}
pub(crate) fn expire(&mut self, now: Instant) -> Vec<IntroEntry<I>> {
let due: Vec<IntroId> = self
.entries
.values()
.filter(|entry| entry.deadline() <= now)
.map(|entry| entry.id)
.collect();
due.into_iter().filter_map(|id| self.remove(id)).collect()
}
pub(crate) fn next_deadline(&self) -> Option<Instant> {
self.entries.values().map(IntroEntry::deadline).min()
}
pub(crate) fn count_for(&self, key: SourceKey) -> u32 {
self.per_source.get(&key).copied().unwrap_or(0)
}
fn oldest_unconsumed(&self, within: Option<SourceKey>) -> Option<IntroId> {
self.entries
.values()
.filter(|entry| !entry.consumed)
.filter(|entry| within.is_none_or(|key| entry.source_key == key))
.min_by(|a, b| a.age_key().cmp(&b.age_key()).then_with(|| a.id.cmp(&b.id)))
.map(|entry| entry.id)
}
}