use std::collections::{HashMap, VecDeque};
use std::hash::Hash;
pub struct FailureLimiter<K: Eq + Hash + Clone> {
window_ms: u64,
max_failures: usize,
fails: HashMap<K, VecDeque<u64>>,
}
impl<K: Eq + Hash + Clone> FailureLimiter<K> {
pub fn new(window_ms: u64, max_failures: usize) -> Self {
Self {
window_ms,
max_failures,
fails: HashMap::new(),
}
}
pub fn check(&mut self, key: &K, now: u64) -> bool {
let Some(q) = self.fails.get_mut(key) else {
return true;
};
prune(q, now, self.window_ms);
let count = q.len();
if count == 0 {
self.fails.remove(key);
return true;
}
count < self.max_failures
}
pub fn record_failure(&mut self, key: K, now: u64) {
let q = self.fails.entry(key).or_default();
prune(q, now, self.window_ms);
q.push_back(now);
}
pub fn record_success(&mut self, key: &K) {
self.fails.remove(key);
}
pub fn gc(&mut self, now: u64) {
let window = self.window_ms;
self.fails.retain(|_, q| {
prune(q, now, window);
!q.is_empty()
});
}
pub fn tracked_keys(&self) -> usize {
self.fails.len()
}
}
fn prune(q: &mut VecDeque<u64>, now: u64, window_ms: u64) {
while let Some(&front) = q.front() {
if front.saturating_add(window_ms) <= now {
q.pop_front();
} else {
break;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn blocks_after_max_failures_within_window_then_recovers() {
let mut lim = FailureLimiter::new(1000, 3);
lim.record_failure("peer", 0);
lim.record_failure("peer", 10);
lim.record_failure("peer", 20);
assert!(
!lim.check(&"peer", 30),
"3 failures in the window must block"
);
assert!(
lim.check(&"peer", 1001),
"an expired failure must free up a slot"
);
}
#[test]
fn record_success_resets_the_peer() {
let mut lim = FailureLimiter::new(1000, 3);
lim.record_failure("peer", 0);
lim.record_failure("peer", 10);
lim.record_failure("peer", 20);
assert!(!lim.check(&"peer", 30), "blocked before success");
lim.record_success(&"peer");
assert!(
lim.check(&"peer", 30),
"a successful handshake clears the failure history"
);
assert_eq!(lim.tracked_keys(), 0, "no residual entry after success");
}
#[test]
fn gc_evicts_keys_whose_failures_all_expired() {
let mut lim = FailureLimiter::new(1000, 3);
lim.record_failure("a", 0);
lim.record_failure("a", 10);
lim.record_failure("b", 20);
assert_eq!(lim.tracked_keys(), 2, "two peers tracked");
lim.gc(2000);
assert_eq!(lim.tracked_keys(), 0, "gc must evict fully-expired keys");
assert!(lim.check(&"c", 2000));
}
#[test]
fn unknown_key_is_always_allowed() {
let mut lim = FailureLimiter::new(1000, 3);
assert!(lim.check(&"never-seen", 0));
assert_eq!(
lim.tracked_keys(),
0,
"checking an unknown key must not allocate an entry"
);
}
}