use std::{
collections::HashMap,
sync::Mutex,
time::{Duration, Instant},
};
#[derive(Debug, Clone, Copy)]
pub struct RateLimit {
pub max_hits: u32,
pub window: Duration,
}
impl RateLimit {
pub fn new(max_hits: u32, window: Duration) -> Self {
Self { max_hits, window }
}
pub fn relaxed() -> Self {
Self::new(5, Duration::from_secs(3))
}
pub fn strict() -> Self {
Self::new(1, Duration::from_secs(3))
}
}
impl Default for RateLimit {
fn default() -> Self {
Self::relaxed()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Decision {
Allow,
Deny {
retry_after: Duration,
first: bool,
},
}
impl Decision {
pub fn allowed(&self) -> bool {
matches!(self, Decision::Allow)
}
}
#[derive(Debug)]
pub struct RateLimiter {
policy: RateLimit,
hits: Mutex<HashMap<String, KeyState>>,
calls: std::sync::atomic::AtomicU64,
}
#[derive(Debug, Default)]
struct KeyState {
hits: Vec<Instant>,
denied: bool,
}
const SWEEP_EVERY: u64 = 1024;
const SWEEP_MAP_SIZE: usize = 4096;
impl RateLimiter {
pub fn new(policy: RateLimit) -> Self {
Self {
policy,
hits: Mutex::new(HashMap::new()),
calls: std::sync::atomic::AtomicU64::new(0),
}
}
pub fn policy(&self) -> RateLimit {
self.policy
}
pub fn check(&self, key: &str) -> Decision {
let n = self
.calls
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
if n % SWEEP_EVERY == SWEEP_EVERY - 1 {
self.sweep();
}
let now = Instant::now();
let window = self.policy.window;
let mut map = match self.hits.lock() {
Ok(m) => m,
Err(poisoned) => poisoned.into_inner(),
};
if map.len() > SWEEP_MAP_SIZE {
map.retain(|_, state| {
state.hits.retain(|t| now.duration_since(*t) < window);
!state.hits.is_empty()
});
}
let entry = map.entry(key.to_owned()).or_default();
entry.hits.retain(|t| now.duration_since(*t) < window);
if entry.hits.len() as u32 >= self.policy.max_hits {
let first = !entry.denied;
entry.denied = true;
let oldest = entry.hits.first().copied().unwrap_or(now);
let retry_after = window.saturating_sub(now.duration_since(oldest));
return Decision::Deny { retry_after, first };
}
entry.denied = false;
entry.hits.push(now);
Decision::Allow
}
pub fn sweep(&self) {
let now = Instant::now();
let window = self.policy.window;
if let Ok(mut map) = self.hits.lock() {
map.retain(|_, state| {
state.hits.retain(|t| now.duration_since(*t) < window);
!state.hits.is_empty()
});
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn allows_up_to_the_limit_then_denies() {
let rl = RateLimiter::new(RateLimit::new(3, Duration::from_secs(10)));
assert!(rl.check("u").allowed());
assert!(rl.check("u").allowed());
assert!(rl.check("u").allowed());
assert!(!rl.check("u").allowed()); }
#[test]
fn keys_are_independent() {
let rl = RateLimiter::new(RateLimit::new(1, Duration::from_secs(10)));
assert!(rl.check("a").allowed());
assert!(rl.check("b").allowed());
assert!(!rl.check("a").allowed());
}
#[test]
fn window_frees_up_over_time() {
let rl = RateLimiter::new(RateLimit::new(1, Duration::from_millis(30)));
assert!(rl.check("u").allowed());
assert!(!rl.check("u").allowed());
std::thread::sleep(Duration::from_millis(40));
assert!(rl.check("u").allowed());
}
#[test]
fn deny_reports_retry_after() {
let rl = RateLimiter::new(RateLimit::new(1, Duration::from_secs(10)));
rl.check("u");
match rl.check("u") {
Decision::Deny { retry_after, .. } => assert!(retry_after <= Duration::from_secs(10)),
Decision::Allow => panic!("expected a deny"),
}
}
#[test]
fn only_first_deny_in_a_window_is_flagged() {
let rl = RateLimiter::new(RateLimit::new(1, Duration::from_millis(30)));
assert!(rl.check("u").allowed());
match rl.check("u") {
Decision::Deny { first, .. } => assert!(first),
Decision::Allow => panic!("expected a deny"),
}
match rl.check("u") {
Decision::Deny { first, .. } => assert!(!first),
Decision::Allow => panic!("expected a deny"),
}
std::thread::sleep(Duration::from_millis(40));
assert!(rl.check("u").allowed());
match rl.check("u") {
Decision::Deny { first, .. } => assert!(first),
Decision::Allow => panic!("expected a deny"),
}
}
#[test]
fn oversized_map_is_swept() {
let rl = RateLimiter::new(RateLimit::new(1, Duration::from_millis(1)));
{
let now = Instant::now();
let stale = now.checked_sub(Duration::from_secs(60)).unwrap_or(now);
let mut map = rl.hits.lock().unwrap();
for i in 0..(SWEEP_MAP_SIZE + 10) {
map.insert(
format!("k{i}"),
KeyState {
hits: vec![stale],
denied: false,
},
);
}
}
rl.check("fresh");
let len = rl.hits.lock().unwrap().len();
assert!(len <= 2, "map should have been swept, len={len}");
}
}