use std::collections::HashMap;
use std::sync::Mutex;
use std::time::Instant;
pub trait RateLimiter: Send + Sync {
fn check(&self, key: &str) -> bool;
fn peek(&self, key: &str) -> u32;
fn reset(&self, key: &str);
}
#[derive(Debug, Clone)]
pub struct TokenBucketConfig {
pub capacity: u32,
pub refill_per_second: f64,
}
pub struct TokenBucketRateLimiter {
config: TokenBucketConfig,
buckets: Mutex<HashMap<String, Bucket>>,
}
struct Bucket {
tokens: f64,
last_refill: Instant,
}
impl TokenBucketRateLimiter {
pub fn new(config: TokenBucketConfig) -> Self {
Self {
config,
buckets: Mutex::new(HashMap::new()),
}
}
pub fn with_rate(capacity: u32, refill_per_second: f64) -> Self {
Self::new(TokenBucketConfig {
capacity,
refill_per_second,
})
}
fn refill_bucket(&self, bucket: &mut Bucket) {
let now = Instant::now();
let elapsed = now.duration_since(bucket.last_refill).as_secs_f64();
let refilled = elapsed * self.config.refill_per_second;
bucket.tokens = (bucket.tokens + refilled).min(self.config.capacity as f64);
bucket.last_refill = now;
}
fn get_or_create_bucket(&self, _key: &str) -> Bucket {
Bucket {
tokens: self.config.capacity as f64,
last_refill: Instant::now(),
}
}
}
impl RateLimiter for TokenBucketRateLimiter {
fn check(&self, key: &str) -> bool {
let mut buckets = self.buckets.lock().unwrap();
let bucket = buckets
.entry(key.to_string())
.or_insert_with(|| self.get_or_create_bucket(key));
self.refill_bucket(bucket);
if bucket.tokens >= 1.0 {
bucket.tokens -= 1.0;
true
} else {
false
}
}
fn peek(&self, key: &str) -> u32 {
let mut buckets = self.buckets.lock().unwrap();
let bucket = buckets
.entry(key.to_string())
.or_insert_with(|| self.get_or_create_bucket(key));
self.refill_bucket(bucket);
bucket.tokens as u32
}
fn reset(&self, key: &str) {
self.buckets.lock().unwrap().remove(key);
}
}
pub struct NoopRateLimiter;
impl RateLimiter for NoopRateLimiter {
fn check(&self, _key: &str) -> bool {
true
}
fn peek(&self, _key: &str) -> u32 {
u32::MAX
}
fn reset(&self, _key: &str) {}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
#[test]
fn allows_until_capacity() {
let limiter = TokenBucketRateLimiter::with_rate(5, 0.0);
for _ in 0..5 {
assert!(limiter.check("client-1"));
}
assert!(!limiter.check("client-1"));
}
#[test]
fn separate_keys_have_separate_buckets() {
let limiter = TokenBucketRateLimiter::with_rate(2, 0.0);
assert!(limiter.check("a"));
assert!(limiter.check("a"));
assert!(limiter.check("b"));
assert!(limiter.check("b"));
assert!(!limiter.check("a"));
assert!(!limiter.check("b"));
}
#[test]
fn noop_always_allows() {
let limiter = NoopRateLimiter;
for _ in 0..100 {
assert!(limiter.check("anyone"));
}
}
#[test]
fn peek_does_not_consume() {
let limiter = TokenBucketRateLimiter::with_rate(3, 0.0);
assert_eq!(limiter.peek("k"), 3);
assert_eq!(limiter.peek("k"), 3);
limiter.check("k");
assert_eq!(limiter.peek("k"), 2);
}
#[test]
fn reset_clears_bucket() {
let limiter = TokenBucketRateLimiter::with_rate(1, 0.0);
assert!(limiter.check("k"));
assert!(!limiter.check("k"));
limiter.reset("k");
assert!(limiter.check("k"));
}
#[test]
fn refill_restores_tokens_over_time() {
let limiter = TokenBucketRateLimiter::with_rate(1, 1000.0);
assert!(limiter.check("k"));
assert!(!limiter.check("k"));
std::thread::sleep(Duration::from_millis(10));
assert!(limiter.check("k"));
}
#[test]
fn capacity_ceiling_prevents_overfill() {
let limiter = TokenBucketRateLimiter::with_rate(3, 1000.0);
std::thread::sleep(Duration::from_millis(50));
assert_eq!(limiter.peek("k"), 3);
}
#[test]
fn new_key_starts_at_full_capacity() {
let limiter = TokenBucketRateLimiter::with_rate(7, 1.0);
assert_eq!(limiter.peek("fresh"), 7);
}
#[test]
fn unknown_key_peek_creates_full_bucket() {
let limiter = TokenBucketRateLimiter::with_rate(5, 0.0);
assert_eq!(limiter.peek("unknown"), 5);
}
#[test]
fn config_values_preserved() {
let config = TokenBucketConfig {
capacity: 10,
refill_per_second: 2.5,
};
let limiter = TokenBucketRateLimiter::new(config);
assert_eq!(limiter.peek("k"), 10);
}
}