use std::time::{Duration, Instant};
use dashmap::DashMap;
use tracing::debug;
use crate::util::{Counter, evict_from_dashmap};
struct Bucket {
tokens: f64,
last_access: Instant,
}
#[derive(Debug, Clone, Default)]
#[must_use]
pub struct RateLimitStats {
pub total_allowed: u64,
pub total_rejected: u64,
pub active_keys: usize,
pub total_evicted: u64,
}
#[must_use]
pub struct RateLimiter {
rate: f64,
burst: usize,
buckets: DashMap<String, Bucket>,
total_allowed: Counter,
total_rejected: Counter,
total_evicted: Counter,
}
impl RateLimiter {
pub fn new(rate: f64, burst: usize) -> Self {
Self {
rate,
burst,
buckets: DashMap::new(),
total_allowed: Counter::new(),
total_rejected: Counter::new(),
total_evicted: Counter::new(),
}
}
#[must_use]
pub fn check(&self, key: &str) -> bool {
let now = Instant::now();
let burst = self.burst as f64;
let mut entry = self.buckets.entry(key.to_string()).or_insert(Bucket {
tokens: burst,
last_access: now,
});
let bucket = entry.value_mut();
let elapsed = now.duration_since(bucket.last_access).as_secs_f64();
bucket.tokens = (bucket.tokens + elapsed * self.rate).min(burst);
bucket.last_access = now;
if bucket.tokens >= 1.0 {
bucket.tokens -= 1.0;
self.total_allowed.inc();
true
} else {
self.total_rejected.inc();
false
}
}
#[inline]
#[must_use]
pub fn key_count(&self) -> usize {
self.buckets.len()
}
#[must_use]
pub fn evict_stale(&self, max_idle: Duration) -> usize {
let now = Instant::now();
let count = evict_from_dashmap(&self.buckets, |_key, bucket| {
now.duration_since(bucket.last_access) >= max_idle
});
self.total_evicted.add(count as u64);
if count > 0 {
debug!(count, "ratelimit: evicted stale keys");
}
count
}
pub fn stats(&self) -> RateLimitStats {
RateLimitStats {
total_allowed: self.total_allowed.get(),
total_rejected: self.total_rejected.get(),
active_keys: self.buckets.len(),
total_evicted: self.total_evicted.get(),
}
}
pub fn compact(&self) {
self.buckets.shrink_to_fit();
}
}
struct WindowCounter {
prev_count: u64,
curr_count: u64,
window_start: Instant,
}
pub struct SlidingWindowLimiter {
windows: DashMap<String, WindowCounter>,
max_requests: u64,
window: Duration,
total_allowed: Counter,
total_rejected: Counter,
}
impl SlidingWindowLimiter {
pub fn new(max_requests: u64, window: Duration) -> Self {
Self {
windows: DashMap::new(),
max_requests,
window,
total_allowed: Counter::new(),
total_rejected: Counter::new(),
}
}
pub fn check(&self, key: &str) -> bool {
let now = Instant::now();
let mut entry = self
.windows
.entry(key.to_string())
.or_insert_with(|| WindowCounter {
prev_count: 0,
curr_count: 0,
window_start: now,
});
let elapsed = now.duration_since(entry.window_start);
if elapsed >= self.window * 2 {
entry.prev_count = 0;
entry.curr_count = 0;
entry.window_start = now;
} else if elapsed >= self.window {
entry.prev_count = entry.curr_count;
entry.curr_count = 0;
entry.window_start += self.window;
}
let elapsed_in_current = now.duration_since(entry.window_start);
let weight = 1.0 - (elapsed_in_current.as_secs_f64() / self.window.as_secs_f64());
let estimate = (entry.prev_count as f64 * weight) + entry.curr_count as f64;
if estimate < self.max_requests as f64 {
entry.curr_count += 1;
drop(entry);
self.total_allowed.inc();
true
} else {
drop(entry);
self.total_rejected.inc();
false
}
}
#[inline]
pub fn key_count(&self) -> usize {
self.windows.len()
}
pub fn evict_stale(&self, max_idle: Duration) -> usize {
let now = Instant::now();
crate::util::evict_from_dashmap(&self.windows, |_, v| {
now.duration_since(v.window_start) > max_idle
})
}
#[inline]
pub fn total_allowed(&self) -> u64 {
self.total_allowed.get()
}
#[inline]
pub fn total_rejected(&self) -> u64 {
self.total_rejected.get()
}
pub fn compact(&self) {
self.windows.shrink_to_fit();
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn allows_up_to_burst() {
let limiter = RateLimiter::new(1.0, 3);
assert!(limiter.check("a"));
assert!(limiter.check("a"));
assert!(limiter.check("a"));
assert!(!limiter.check("a"));
}
#[test]
fn separate_keys_independent() {
let limiter = RateLimiter::new(1.0, 1);
assert!(limiter.check("a"));
assert!(limiter.check("b"));
assert!(!limiter.check("a"));
assert!(!limiter.check("b"));
}
#[test]
fn refills_over_time() {
let limiter = RateLimiter::new(100.0, 1);
assert!(limiter.check("a"));
assert!(!limiter.check("a"));
std::thread::sleep(Duration::from_millis(20));
assert!(limiter.check("a"));
}
#[test]
fn key_count() {
let limiter = RateLimiter::new(1.0, 1);
let _ = limiter.check("a");
let _ = limiter.check("b");
assert_eq!(limiter.key_count(), 2);
}
#[test]
fn concurrent_access() {
use std::sync::Arc;
use std::thread;
let limiter = Arc::new(RateLimiter::new(1000.0, 10));
let mut handles = Vec::new();
for i in 0..4 {
let l = limiter.clone();
handles.push(thread::spawn(move || {
let key = format!("key-{i}");
let mut allowed = 0;
for _ in 0..20 {
if l.check(&key) {
allowed += 1;
}
}
allowed
}));
}
let total: usize = handles.into_iter().map(|h| h.join().unwrap()).sum();
assert!(total >= 40, "expected at least 40 allowed, got {total}");
assert_eq!(limiter.key_count(), 4);
}
#[test]
fn stats_tracking() {
let limiter = RateLimiter::new(1.0, 2);
let _ = limiter.check("a"); let _ = limiter.check("a"); let _ = limiter.check("a");
let stats = limiter.stats();
assert_eq!(stats.total_allowed, 2);
assert_eq!(stats.total_rejected, 1);
assert_eq!(stats.active_keys, 1);
}
#[test]
fn evict_stale_keys() {
let limiter = RateLimiter::new(1.0, 1);
let _ = limiter.check("fresh");
let _ = limiter.check("stale");
std::thread::sleep(Duration::from_millis(20));
let _ = limiter.check("fresh");
let evicted = limiter.evict_stale(Duration::from_millis(15));
assert_eq!(evicted, 1);
assert_eq!(limiter.key_count(), 1);
let stats = limiter.stats();
assert_eq!(stats.total_evicted, 1);
}
#[test]
fn evict_stale_no_keys() {
let limiter = RateLimiter::new(1.0, 1);
let _ = limiter.check("a");
let evicted = limiter.evict_stale(Duration::from_secs(60));
assert_eq!(evicted, 0);
}
#[test]
fn evict_all_stale() {
let limiter = RateLimiter::new(1.0, 1);
let _ = limiter.check("a");
let _ = limiter.check("b");
let _ = limiter.check("c");
std::thread::sleep(Duration::from_millis(15));
let evicted = limiter.evict_stale(Duration::from_millis(10));
assert_eq!(evicted, 3);
assert_eq!(limiter.key_count(), 0);
}
#[test]
fn sliding_window_allows_up_to_max() {
let limiter = SlidingWindowLimiter::new(3, Duration::from_secs(1));
assert!(limiter.check("a"));
assert!(limiter.check("a"));
assert!(limiter.check("a"));
assert!(!limiter.check("a"));
}
#[test]
fn sliding_window_separate_keys() {
let limiter = SlidingWindowLimiter::new(1, Duration::from_secs(1));
assert!(limiter.check("a"));
assert!(limiter.check("b"));
assert!(!limiter.check("a"));
assert!(!limiter.check("b"));
}
#[test]
fn sliding_window_stats() {
let limiter = SlidingWindowLimiter::new(2, Duration::from_secs(1));
let _ = limiter.check("a");
let _ = limiter.check("a");
let _ = limiter.check("a"); assert_eq!(limiter.total_allowed(), 2);
assert_eq!(limiter.total_rejected(), 1);
assert_eq!(limiter.key_count(), 1);
}
#[test]
fn sliding_window_refills_after_window() {
let limiter = SlidingWindowLimiter::new(1, Duration::from_millis(20));
assert!(limiter.check("a"));
assert!(!limiter.check("a"));
std::thread::sleep(Duration::from_millis(45));
assert!(limiter.check("a"));
}
}