use std::collections::HashMap;
use chia_protocol::Bytes32;
use crate::constants::{FRESHNESS_WINDOW_MS, MAX_TRACKED_SENDERS, REPLAY_WINDOW};
const WINDOW_WORDS: usize = REPLAY_WINDOW / 64;
type SenderKey = (Bytes32, u32);
#[derive(Clone)]
struct SenderWindow {
highest: u64,
bits: [u64; WINDOW_WORDS],
last_seen_ms: u64,
started: bool,
}
impl SenderWindow {
fn new() -> Self {
Self {
highest: 0,
bits: [0; WINDOW_WORDS],
last_seen_ms: 0,
started: false,
}
}
fn get_bit(&self, offset: u64) -> bool {
let i = offset as usize;
(self.bits[i / 64] >> (i % 64)) & 1 == 1
}
fn set_bit(&mut self, offset: u64) {
let i = offset as usize;
self.bits[i / 64] |= 1 << (i % 64);
}
fn shift_left(&mut self, diff: u64) {
if diff as usize >= REPLAY_WINDOW {
self.bits = [0; WINDOW_WORDS];
return;
}
let shift = diff as usize;
let word_shift = shift / 64;
let bit_shift = shift % 64;
let mut out = [0u64; WINDOW_WORDS];
for i in (0..WINDOW_WORDS).rev() {
let mut v = 0u64;
if i >= word_shift {
v = self.bits[i - word_shift] << bit_shift;
if bit_shift > 0 && i > word_shift {
v |= self.bits[i - word_shift - 1] >> (64 - bit_shift);
}
}
out[i] = v;
}
self.bits = out;
}
fn admit(&mut self, counter: u64) -> bool {
if !self.started {
self.started = true;
self.highest = counter;
self.set_bit(0);
return true;
}
if counter > self.highest {
let diff = counter - self.highest;
self.shift_left(diff);
self.highest = counter;
self.set_bit(0);
true
} else {
let offset = self.highest - counter;
if offset as usize >= REPLAY_WINDOW || self.get_bit(offset) {
false
} else {
self.set_bit(offset);
true
}
}
}
}
#[derive(Default)]
pub struct ReplayGuard {
senders: HashMap<SenderKey, SenderWindow>,
}
impl ReplayGuard {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn tracked_senders(&self) -> usize {
self.senders.len()
}
#[must_use]
pub fn check_and_admit(
&mut self,
sender: Bytes32,
sender_epoch: u32,
counter: u64,
timestamp_ms: u64,
now_ms: u64,
) -> bool {
let lower = now_ms.saturating_sub(FRESHNESS_WINDOW_MS);
let upper = now_ms.saturating_add(FRESHNESS_WINDOW_MS);
if timestamp_ms < lower || timestamp_ms > upper {
return false;
}
let key = (sender, sender_epoch);
if !self.senders.contains_key(&key) {
self.evict_if_full();
}
let window = self.senders.entry(key).or_insert_with(SenderWindow::new);
if window.admit(counter) {
window.last_seen_ms = window.last_seen_ms.max(timestamp_ms);
true
} else {
false
}
}
fn evict_if_full(&mut self) {
if self.senders.len() < MAX_TRACKED_SENDERS {
return;
}
if let Some(victim) = self
.senders
.iter()
.min_by_key(|(_, w)| w.last_seen_ms)
.map(|(k, _)| *k)
{
self.senders.remove(&victim);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn sender(n: u8) -> Bytes32 {
Bytes32::new([n; 32])
}
const NOW: u64 = 1_700_000_000_000;
#[test]
fn first_message_is_accepted() {
let mut g = ReplayGuard::new();
assert!(g.check_and_admit(sender(1), 0, 0, NOW, NOW));
}
#[test]
fn duplicate_counter_is_rejected() {
let mut g = ReplayGuard::new();
assert!(g.check_and_admit(sender(1), 0, 5, NOW, NOW));
assert!(!g.check_and_admit(sender(1), 0, 5, NOW, NOW));
}
#[test]
fn monotonic_advance_accepts() {
let mut g = ReplayGuard::new();
for c in 0..100 {
assert!(g.check_and_admit(sender(1), 0, c, NOW, NOW), "counter {c}");
}
}
#[test]
fn in_window_reorder_accepts_then_rejects_replay() {
let mut g = ReplayGuard::new();
assert!(g.check_and_admit(sender(1), 0, 10, NOW, NOW));
assert!(g.check_and_admit(sender(1), 0, 7, NOW, NOW));
assert!(g.check_and_admit(sender(1), 0, 3, NOW, NOW));
assert!(!g.check_and_admit(sender(1), 0, 7, NOW, NOW));
assert!(!g.check_and_admit(sender(1), 0, 10, NOW, NOW));
}
#[test]
fn counter_below_the_window_is_rejected() {
let mut g = ReplayGuard::new();
assert!(g.check_and_admit(sender(1), 0, REPLAY_WINDOW as u64 + 100, NOW, NOW));
assert!(!g.check_and_admit(sender(1), 0, 1, NOW, NOW));
}
#[test]
fn far_future_jump_clears_and_accepts() {
let mut g = ReplayGuard::new();
assert!(g.check_and_admit(sender(1), 0, 5, NOW, NOW));
let far = 5 + REPLAY_WINDOW as u64 * 3;
assert!(g.check_and_admit(sender(1), 0, far, NOW, NOW));
assert!(!g.check_and_admit(sender(1), 0, 5, NOW, NOW));
}
#[test]
fn stale_timestamp_rejected() {
let mut g = ReplayGuard::new();
assert!(!g.check_and_admit(sender(1), 0, 0, NOW - FRESHNESS_WINDOW_MS - 1, NOW));
}
#[test]
fn future_timestamp_beyond_window_rejected() {
let mut g = ReplayGuard::new();
assert!(!g.check_and_admit(sender(1), 0, 0, NOW + FRESHNESS_WINDOW_MS + 1, NOW));
}
#[test]
fn timestamp_at_window_edges_accepted() {
let mut g = ReplayGuard::new();
assert!(g.check_and_admit(sender(1), 0, 0, NOW - FRESHNESS_WINDOW_MS, NOW));
assert!(g.check_and_admit(sender(2), 0, 0, NOW + FRESHNESS_WINDOW_MS, NOW));
}
#[test]
fn distinct_senders_have_independent_counters() {
let mut g = ReplayGuard::new();
assert!(g.check_and_admit(sender(1), 0, 0, NOW, NOW));
assert!(g.check_and_admit(sender(2), 0, 0, NOW, NOW));
}
#[test]
fn distinct_epochs_are_independent() {
let mut g = ReplayGuard::new();
assert!(g.check_and_admit(sender(1), 0, 0, NOW, NOW));
assert!(g.check_and_admit(sender(1), 1, 0, NOW, NOW));
}
#[test]
fn self_pair_is_a_valid_independent_stream() {
let mut g = ReplayGuard::new();
let me = sender(7);
assert!(g.check_and_admit(me, 0, 0, NOW, NOW));
assert!(g.check_and_admit(me, 0, 1, NOW, NOW));
assert!(!g.check_and_admit(me, 0, 0, NOW, NOW));
}
#[test]
fn per_sender_state_is_bounded_under_a_counter_flood() {
let mut g = ReplayGuard::new();
for c in (0..10_000).step_by(1) {
let _ = g.check_and_admit(sender(1), 0, c, NOW, NOW);
}
assert_eq!(g.tracked_senders(), 1, "one sender = one bounded entry");
}
}