use std::collections::HashMap;
use std::net::IpAddr;
use std::sync::Arc;
use std::sync::Mutex;
use std::time::{Instant, SystemTime, UNIX_EPOCH};
use async_trait::async_trait;
use boatramp_core::access::RateLimit;
use boatramp_core::kv::KvStore;
use serde::{Deserialize, Serialize};
const MAX_BUCKETS: usize = 100_000;
struct Bucket {
tokens: f64,
last: Instant,
}
#[derive(Default)]
pub struct RateLimiter {
buckets: Mutex<HashMap<(String, IpAddr), Bucket>>,
}
impl RateLimiter {
pub fn new() -> Self {
Self::default()
}
pub fn check(&self, site: &str, ip: IpAddr, limit: &RateLimit) -> bool {
let capacity = limit.burst_capacity() as f64;
let rate = limit.rps.max(1) as f64;
let now = Instant::now();
let mut buckets = self.buckets.lock().unwrap();
if buckets.len() >= MAX_BUCKETS && !buckets.contains_key(&(site.to_string(), ip)) {
buckets.retain(|_, b| {
let refilled = b.tokens + now.duration_since(b.last).as_secs_f64() * rate;
refilled < capacity
});
}
let bucket = buckets.entry((site.to_string(), ip)).or_insert(Bucket {
tokens: capacity,
last: now,
});
let elapsed = now.duration_since(bucket.last).as_secs_f64();
bucket.last = now;
bucket.tokens = (bucket.tokens + elapsed * rate).min(capacity);
if bucket.tokens >= 1.0 {
bucket.tokens -= 1.0;
true
} else {
false
}
}
}
#[async_trait]
pub trait RateLimitStore: Send + Sync {
async fn check(&self, site: &str, ip: IpAddr, limit: &RateLimit) -> bool;
}
#[async_trait]
impl RateLimitStore for RateLimiter {
async fn check(&self, site: &str, ip: IpAddr, limit: &RateLimit) -> bool {
Self::check(self, site, ip, limit)
}
}
const WINDOW_SECS: u64 = 1;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
struct RateLimitWindow {
window_start: u64,
count: u64,
}
fn fixed_window_decision(
window: Option<RateLimitWindow>,
now: u64,
max: u64,
) -> (bool, RateLimitWindow) {
match window {
Some(w) if now < w.window_start.saturating_add(WINDOW_SECS) => {
if w.count >= max {
(false, w)
} else {
(
true,
RateLimitWindow {
window_start: w.window_start,
count: w.count + 1,
},
)
}
}
_ => (
true,
RateLimitWindow {
window_start: now,
count: 1,
},
),
}
}
pub struct KvRateLimiter {
kv: Arc<dyn KvStore>,
fail_open: bool,
}
impl KvRateLimiter {
pub fn new(kv: Arc<dyn KvStore>, fail_open: bool) -> Self {
Self { kv, fail_open }
}
}
#[async_trait]
impl RateLimitStore for KvRateLimiter {
async fn check(&self, site: &str, ip: IpAddr, limit: &RateLimit) -> bool {
let key = format!("ratelimit/{site}/{ip}");
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_secs())
.unwrap_or(0);
let current = match self.kv.get(&key).await {
Ok(Some(bytes)) => serde_json::from_slice::<RateLimitWindow>(&bytes).ok(),
Ok(None) => None,
Err(_) => return self.fail_open,
};
let (allowed, window) = fixed_window_decision(current, now, limit.burst_capacity() as u64);
if allowed {
if let Ok(bytes) = serde_json::to_vec(&window) {
let _ = self.kv.put(&key, bytes).await;
}
}
allowed
}
}
#[cfg(test)]
mod tests {
use super::*;
fn ip() -> IpAddr {
"203.0.113.1".parse().unwrap()
}
#[test]
fn fixed_window_allows_up_to_max_then_blocks_then_rolls() {
let max = 2;
let (ok, w) = fixed_window_decision(None, 100, max);
assert!(ok && w.count == 1 && w.window_start == 100);
let (ok, w) = fixed_window_decision(Some(w), 100, max);
assert!(ok && w.count == 2);
let (ok, w2) = fixed_window_decision(Some(w), 100, max);
assert!(!ok && w2 == w);
let (ok, w3) = fixed_window_decision(Some(w), 101, max);
assert!(ok && w3.count == 1 && w3.window_start == 101);
}
#[tokio::test]
async fn kv_limiter_shares_count_across_checks() {
use boatramp_core::kv::MemoryKv;
let limiter = KvRateLimiter::new(Arc::new(MemoryKv::new()), false);
let limit = RateLimit { rps: 1, burst: 2 }; let cap = limit.burst_capacity();
for _ in 0..cap {
assert!(limiter.check("s", ip(), &limit).await);
}
assert!(!limiter.check("s", ip(), &limit).await);
assert!(limiter.check("other", ip(), &limit).await);
}
#[tokio::test]
async fn kv_limiter_fails_closed_on_kv_error_unless_opted_in() {
use async_trait::async_trait;
use boatramp_core::error::KvError;
struct FailingKv;
#[async_trait]
impl KvStore for FailingKv {
async fn get(&self, _key: &str) -> Result<Option<Vec<u8>>, KvError> {
Err(KvError::Backend("kv down".into()))
}
async fn put(&self, _key: &str, _value: Vec<u8>) -> Result<(), KvError> {
Ok(())
}
async fn delete(&self, _key: &str) -> Result<(), KvError> {
Ok(())
}
async fn list_prefix(&self, _prefix: &str) -> Result<Vec<String>, KvError> {
Ok(Vec::new())
}
}
let limit = RateLimit { rps: 1, burst: 2 };
let closed = KvRateLimiter::new(Arc::new(FailingKv), false);
assert!(!closed.check("s", ip(), &limit).await);
let open = KvRateLimiter::new(Arc::new(FailingKv), true);
assert!(open.check("s", ip(), &limit).await);
}
#[test]
fn allows_burst_then_blocks() {
let limiter = RateLimiter::new();
let limit = RateLimit { rps: 1, burst: 3 };
assert!(limiter.check("s", ip(), &limit));
assert!(limiter.check("s", ip(), &limit));
assert!(limiter.check("s", ip(), &limit));
assert!(!limiter.check("s", ip(), &limit));
}
#[test]
fn separate_keys_have_separate_buckets() {
let limiter = RateLimiter::new();
let limit = RateLimit { rps: 1, burst: 1 };
assert!(limiter.check("s", ip(), &limit));
assert!(!limiter.check("s", ip(), &limit));
assert!(limiter.check("other", ip(), &limit));
assert!(limiter.check("s", "203.0.113.2".parse().unwrap(), &limit));
}
}