use std::{
collections::HashMap,
sync::{Arc, Mutex},
time::{Duration, Instant},
};
const IDLE_EVICTION: Duration = Duration::from_secs(600);
const SWEEP_EVERY: usize = 512;
struct Bucket {
tokens: f64,
last: Instant,
}
struct Inner {
buckets: HashMap<String, Bucket>,
since_sweep: usize,
}
#[derive(Clone)]
pub struct RateLimiter {
inner: Arc<Mutex<Inner>>,
refill_per_sec: f64,
burst: f64,
}
pub struct Throttled {
pub retry_after_secs: u64,
}
impl RateLimiter {
pub fn new(per_minute: u32) -> Option<Self> {
if per_minute == 0 {
return None;
}
let burst = ((per_minute as f64) / 10.0).max(10.0);
Some(Self {
inner: Arc::new(Mutex::new(Inner {
buckets: HashMap::new(),
since_sweep: 0,
})),
refill_per_sec: per_minute as f64 / 60.0,
burst,
})
}
pub fn check(&self, key: &str) -> Result<(), Throttled> {
let now = Instant::now();
let mut inner = match self.inner.lock() {
Ok(guard) => guard,
Err(poisoned) => poisoned.into_inner(),
};
inner.since_sweep += 1;
if inner.since_sweep >= SWEEP_EVERY {
inner.since_sweep = 0;
inner
.buckets
.retain(|_, b| now.duration_since(b.last) < IDLE_EVICTION);
}
let burst = self.burst;
let refill = self.refill_per_sec;
let bucket = inner.buckets.entry(key.to_owned()).or_insert(Bucket {
tokens: burst,
last: now,
});
let elapsed = now.duration_since(bucket.last).as_secs_f64();
bucket.tokens = (bucket.tokens + elapsed * refill).min(burst);
bucket.last = now;
if bucket.tokens >= 1.0 {
bucket.tokens -= 1.0;
return Ok(());
}
let deficit = 1.0 - bucket.tokens;
Err(Throttled {
retry_after_secs: (deficit / refill).ceil().max(1.0) as u64,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn disabled_when_zero() {
assert!(RateLimiter::new(0).is_none());
}
#[test]
fn allows_a_burst_then_throttles() {
let rl = RateLimiter::new(60).expect("enabled");
for i in 0..10 {
assert!(rl.check("k").is_ok(), "burst request {i} should pass");
}
let throttled = rl.check("k");
assert!(
throttled.is_err(),
"the 11th immediate request is throttled"
);
assert!(throttled.expect_err("throttled").retry_after_secs >= 1);
}
#[test]
fn keys_are_independent() {
let rl = RateLimiter::new(60).expect("enabled");
for _ in 0..10 {
let _ = rl.check("a");
}
assert!(
rl.check("b").is_ok(),
"one caller's budget must not affect another's"
);
}
#[test]
fn refills_over_time() {
let rl = RateLimiter::new(600).expect("enabled"); for _ in 0..60 {
let _ = rl.check("k");
}
assert!(rl.check("k").is_err(), "bucket is drained");
std::thread::sleep(Duration::from_millis(250));
assert!(
rl.check("k").is_ok(),
"a quarter second refills 2.5 tokens at 10/s"
);
}
}