use super::StoreError;
use std::collections::HashMap;
use std::sync::Mutex;
pub trait ThrottleStore: Send + Sync {
fn record_failure(&self, key: &str, now: u64, window_secs: u64) -> Result<u32, StoreError>;
fn failure_count(&self, key: &str, now: u64, window_secs: u64) -> Result<u32, StoreError>;
fn is_banned(&self, key: &str, now: u64) -> Result<Option<u64>, StoreError>;
fn ban(&self, key: &str, until: u64) -> Result<(), StoreError>;
fn clear_failures(&self, key: &str) -> Result<(), StoreError>;
fn reset(&self, key: &str) -> Result<(), StoreError>;
fn purge_expired(&self, now: u64) -> Result<usize, StoreError>;
}
#[derive(Debug, Default)]
struct Entry {
failures: Vec<u64>,
banned_until: Option<u64>,
window_secs: u64,
}
#[derive(Debug)]
pub struct MemoryThrottleStore {
entries: Mutex<HashMap<String, Entry>>,
}
impl Default for MemoryThrottleStore {
fn default() -> Self {
Self::new()
}
}
impl MemoryThrottleStore {
pub fn new() -> Self {
Self {
entries: Mutex::new(HashMap::new()),
}
}
fn lock<T>(m: &Mutex<T>) -> std::sync::MutexGuard<'_, T> {
m.lock().unwrap_or_else(|e| e.into_inner())
}
fn cutoff(now: u64, window_secs: u64) -> u64 {
now.saturating_sub(window_secs)
}
}
impl ThrottleStore for MemoryThrottleStore {
fn record_failure(&self, key: &str, now: u64, window_secs: u64) -> Result<u32, StoreError> {
let cutoff = Self::cutoff(now, window_secs);
let mut g = Self::lock(&self.entries);
let e = g.entry(key.to_string()).or_default();
e.window_secs = window_secs;
e.failures.retain(|t| *t > cutoff);
e.failures.push(now);
Ok(e.failures.len() as u32)
}
fn failure_count(&self, key: &str, now: u64, window_secs: u64) -> Result<u32, StoreError> {
let cutoff = Self::cutoff(now, window_secs);
let g = Self::lock(&self.entries);
Ok(g.get(key)
.map(|e| e.failures.iter().filter(|t| **t > cutoff).count() as u32)
.unwrap_or(0))
}
fn is_banned(&self, key: &str, now: u64) -> Result<Option<u64>, StoreError> {
Ok(Self::lock(&self.entries)
.get(key)
.and_then(|e| e.banned_until)
.filter(|until| *until > now))
}
fn ban(&self, key: &str, until: u64) -> Result<(), StoreError> {
let mut g = Self::lock(&self.entries);
let e = g.entry(key.to_string()).or_default();
e.banned_until = Some(e.banned_until.map_or(until, |existing| existing.max(until)));
Ok(())
}
fn clear_failures(&self, key: &str) -> Result<(), StoreError> {
if let Some(e) = Self::lock(&self.entries).get_mut(key) {
e.failures.clear();
}
Ok(())
}
fn reset(&self, key: &str) -> Result<(), StoreError> {
Self::lock(&self.entries).remove(key);
Ok(())
}
fn purge_expired(&self, now: u64) -> Result<usize, StoreError> {
let mut g = Self::lock(&self.entries);
let before = g.len();
g.retain(|_, e| {
let banned = e.banned_until.is_some_and(|until| until > now);
let fresh = e
.failures
.iter()
.any(|t| *t > Self::cutoff(now, e.window_secs));
banned || fresh
});
Ok(before - g.len())
}
}
#[cfg(test)]
mod tests {
use super::*;
const NOW: u64 = 1_000_000;
fn store() -> MemoryThrottleStore {
MemoryThrottleStore::new()
}
#[test]
fn record_failure_counts_within_window() {
let s = store();
assert_eq!(s.record_failure("k", NOW, 60).unwrap(), 1);
assert_eq!(s.record_failure("k", NOW + 1, 60).unwrap(), 2);
}
#[test]
fn failures_at_window_edge_are_dropped() {
let s = store();
s.record_failure("k", NOW, 60).unwrap();
assert_eq!(s.record_failure("k", NOW + 60, 60).unwrap(), 1);
assert_eq!(s.record_failure("k", NOW + 61, 60).unwrap(), 2);
}
#[test]
fn failure_count_is_read_only() {
let s = store();
assert_eq!(s.failure_count("k", NOW, 60).unwrap(), 0);
s.record_failure("k", NOW, 60).unwrap();
assert_eq!(s.failure_count("k", NOW, 60).unwrap(), 1);
for _ in 0..10 {
assert_eq!(s.failure_count("k", NOW, 60).unwrap(), 1);
}
assert_eq!(s.record_failure("k", NOW, 60).unwrap(), 2);
}
#[test]
fn failure_count_of_unknown_key_is_zero() {
let s = store();
assert_eq!(s.failure_count("ghost", NOW, 60).unwrap(), 0);
}
#[test]
fn failures_vec_stays_bounded_to_window() {
let s = store();
for i in 0..100 {
s.record_failure("k", NOW + i, 60).unwrap();
}
assert_eq!(len(&s, "k"), 60, "窗口外的失败必须真被丢弃");
s.record_failure("k", NOW + 1_000, 60).unwrap();
assert_eq!(len(&s, "k"), 1);
}
#[test]
fn banned_expires_at_until_exclusive() {
let s = store();
s.ban("k", NOW + 100).unwrap();
assert_eq!(s.is_banned("k", NOW).unwrap(), Some(NOW + 100));
assert_eq!(s.is_banned("k", NOW + 99).unwrap(), Some(NOW + 100));
assert_eq!(s.is_banned("k", NOW + 100).unwrap(), None);
}
#[test]
fn is_banned_unknown_key_is_none() {
let s = store();
assert_eq!(s.is_banned("ghost", NOW).unwrap(), None);
}
#[test]
fn reset_clears_failures_and_ban() {
let s = store();
s.record_failure("k", NOW, 60).unwrap();
s.ban("k", NOW + 100).unwrap();
s.reset("k").unwrap();
assert_eq!(s.failure_count("k", NOW, 60).unwrap(), 0);
assert_eq!(s.is_banned("k", NOW).unwrap(), None);
}
#[test]
fn reset_unknown_key_is_ok() {
let s = store();
assert!(s.reset("ghost").is_ok());
}
#[test]
fn purge_keeps_banned_and_fresh_entries() {
let s = store();
s.record_failure("stale", NOW, 60).unwrap();
s.record_failure("fresh", NOW + 990, 60).unwrap();
s.ban("banned", NOW + 5_000).unwrap();
assert_eq!(s.purge_expired(NOW + 1_000).unwrap(), 1);
assert_eq!(s.failure_count("stale", NOW + 1_000, 60).unwrap(), 0);
assert_eq!(s.failure_count("fresh", NOW + 1_000, 60).unwrap(), 1);
assert!(s.is_banned("banned", NOW + 1_000).unwrap().is_some());
}
#[test]
fn purge_with_expired_ban_drops_entry() {
let s = store();
s.record_failure("k", NOW, 60).unwrap();
s.ban("k", NOW + 10).unwrap();
assert_eq!(s.purge_expired(NOW + 100).unwrap(), 1);
assert_eq!(s.is_banned("k", NOW + 100).unwrap(), None);
}
#[test]
fn purge_empty_store_is_zero() {
let s = store();
assert_eq!(s.purge_expired(NOW).unwrap(), 0);
}
#[test]
fn purge_keeps_in_window_failures_after_clock_rollback() {
let s = store();
s.record_failure("k", NOW, 60).unwrap();
s.record_failure("k", NOW - 200, 60).unwrap();
assert_eq!(s.failure_count("k", NOW, 60).unwrap(), 1);
assert_eq!(
s.purge_expired(NOW).unwrap(),
0,
"窗口内仍有有效失败,不该清"
);
assert_eq!(
s.failure_count("k", NOW, 60).unwrap(),
1,
"计数不能被 purge 免费重置"
);
}
#[test]
fn ban_cannot_be_shortened_by_clock_rollback() {
let s = store();
s.ban("k", NOW + 900).unwrap();
s.ban("k", NOW - 5_000 + 900).unwrap();
assert_eq!(
s.is_banned("k", NOW).unwrap(),
Some(NOW + 900),
"封禁只能延长,回拨不能提前解封"
);
}
#[test]
fn empty_key_is_a_normal_key() {
let s = store();
assert_eq!(s.record_failure("", NOW, 60).unwrap(), 1);
assert_eq!(s.record_failure("", NOW, 60).unwrap(), 2);
assert_eq!(s.record_failure("k", NOW, 60).unwrap(), 1);
}
#[test]
fn lock_recovers_from_poisoned_mutex() {
let m = Mutex::new(Entry {
failures: vec![NOW],
banned_until: Some(NOW + 1),
window_secs: 60,
});
std::panic::catch_unwind(|| {
let _guard = m.lock().unwrap();
panic!("poison");
})
.unwrap_err();
assert!(m.is_poisoned());
let g = MemoryThrottleStore::lock(&m);
assert_eq!(g.failures, vec![NOW]);
assert_eq!(g.banned_until, Some(NOW + 1));
}
fn len(s: &MemoryThrottleStore, key: &str) -> usize {
MemoryThrottleStore::lock(&s.entries)
.get(key)
.map(|e| e.failures.len())
.unwrap_or(0)
}
}