use std::collections::HashMap;
use std::hash::{DefaultHasher, Hash, Hasher};
use std::net::IpAddr;
use std::sync::Mutex;
use std::time::Instant;
const NUM_SHARDS: usize = 16;
const CLEANUP_THRESHOLD_PER_SHARD: usize = 1_000;
const CLEANUP_IDLE_SECS: f64 = 300.0;
pub struct RateLimiter {
burst: f64,
rate: f64,
shards: [Mutex<HashMap<IpAddr, Bucket>>; NUM_SHARDS],
}
struct Bucket {
tokens: f64,
last_refill: Instant,
}
fn shard_index(ip: IpAddr) -> usize {
let mut hasher = DefaultHasher::new();
ip.hash(&mut hasher);
hasher.finish() as usize & (NUM_SHARDS - 1)
}
impl RateLimiter {
pub fn new(rate: f64, burst: f64) -> Self {
Self {
burst,
rate,
shards: std::array::from_fn(|_| Mutex::new(HashMap::new())),
}
}
pub fn check(&self, ip: IpAddr) -> bool {
let now = Instant::now();
let idx = shard_index(ip);
let mut shard = self.shards[idx]
.lock()
.expect("rate limiter shard lock poisoned");
if shard.len() > CLEANUP_THRESHOLD_PER_SHARD {
shard.retain(|_, bucket| {
now.saturating_duration_since(bucket.last_refill)
.as_secs_f64()
< CLEANUP_IDLE_SECS
});
}
let bucket = shard.entry(ip).or_insert_with(|| Bucket {
tokens: self.burst,
last_refill: now,
});
let elapsed = now
.saturating_duration_since(bucket.last_refill)
.as_secs_f64();
bucket.tokens = (bucket.tokens + elapsed * self.rate).min(self.burst);
bucket.last_refill = now;
if bucket.tokens >= 1.0 {
bucket.tokens -= 1.0;
true
} else {
false
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::net::{Ipv4Addr, Ipv6Addr};
use std::thread;
use std::time::Duration;
const NO_REFILL: f64 = 0.001;
#[test]
fn allows_requests_within_burst() {
let limiter = RateLimiter::new(NO_REFILL, 5.0);
let ip = IpAddr::V4(Ipv4Addr::LOCALHOST);
for i in 0..5 {
assert!(limiter.check(ip), "request {i} should be allowed");
}
assert!(!limiter.check(ip), "request 6 should be rejected");
}
#[test]
fn refills_tokens_over_time() {
let limiter = RateLimiter::new(10.0, 2.0);
let ip = IpAddr::V4(Ipv4Addr::LOCALHOST);
assert!(limiter.check(ip));
assert!(limiter.check(ip));
assert!(!limiter.check(ip));
thread::sleep(Duration::from_millis(120));
assert!(limiter.check(ip), "should have refilled at least 1 token");
}
#[test]
fn independent_per_ip() {
let limiter = RateLimiter::new(NO_REFILL, 1.0);
let ip_a = IpAddr::V4(Ipv4Addr::new(10, 0, 0, 1));
let ip_b = IpAddr::V4(Ipv4Addr::new(10, 0, 0, 2));
assert!(limiter.check(ip_a));
assert!(!limiter.check(ip_a));
assert!(limiter.check(ip_b));
}
#[test]
fn ipv6_addresses_work() {
let limiter = RateLimiter::new(NO_REFILL, 1.0);
let ip = IpAddr::V6(Ipv6Addr::LOCALHOST);
assert!(limiter.check(ip));
assert!(!limiter.check(ip));
}
#[test]
fn tokens_cap_at_burst() {
let burst = 3.0;
let limiter = RateLimiter::new(1000.0, burst);
let ip = IpAddr::V4(Ipv4Addr::LOCALHOST);
assert!(limiter.check(ip));
thread::sleep(Duration::from_millis(50));
assert!(limiter.check(ip));
let shard = limiter.shards[shard_index(ip)]
.lock()
.expect("shard lock is not poisoned");
let tokens = shard
.get(&ip)
.expect("the checks above created the bucket")
.tokens;
assert!(
tokens <= burst - 1.0,
"the idle period banked {tokens} tokens over a burst of {burst}"
);
}
#[test]
fn cleanup_removes_idle_entries() {
let limiter = RateLimiter::new(10.0, 1.0);
let Some(old) = Instant::now().checked_sub(Duration::from_secs(600)) else {
return; };
let target_ip = IpAddr::V4(Ipv4Addr::LOCALHOST);
let target_shard = shard_index(target_ip);
{
let mut shard = limiter.shards[target_shard].lock().unwrap();
for i in 0..(CLEANUP_THRESHOLD_PER_SHARD + 100) {
let ip = IpAddr::V4(Ipv4Addr::from((i as u32).to_be_bytes()));
shard.insert(
ip,
Bucket {
tokens: 0.0,
last_refill: old,
},
);
}
assert!(shard.len() > CLEANUP_THRESHOLD_PER_SHARD);
}
limiter.check(target_ip);
let shard = limiter.shards[target_shard].lock().unwrap();
assert_eq!(shard.len(), 1);
}
#[test]
fn shard_index_distributes_across_shards() {
let mut seen = std::collections::HashSet::new();
for i in 0..256u32 {
let ip = IpAddr::V4(Ipv4Addr::from(i.to_be_bytes()));
seen.insert(shard_index(ip));
}
assert!(seen.len() >= NUM_SHARDS / 2, "poor distribution: {seen:?}");
}
#[test]
fn concurrent_access_does_not_panic() {
use std::sync::Arc;
let limiter = Arc::new(RateLimiter::new(100.0, 10.0));
let handles: Vec<_> = (0..8)
.map(|t| {
let limiter = Arc::clone(&limiter);
thread::spawn(move || {
for i in 0..100u32 {
let ip = IpAddr::V4(Ipv4Addr::new(10, 0, t, i as u8));
limiter.check(ip);
}
})
})
.collect();
for h in handles {
h.join().unwrap();
}
}
}