use std::collections::HashMap;
use std::sync::Arc;
use std::sync::Mutex;
use std::time::{Duration, Instant};
#[derive(Debug)]
pub struct TokenBucket {
capacity: f64,
refill_per_sec: f64,
state: Mutex<BucketState>,
}
#[derive(Debug, Clone, Copy)]
struct BucketState {
tokens: f64,
last_refill: Instant,
}
impl TokenBucket {
#[must_use]
pub fn new(capacity: u32) -> Self {
let cap = f64::from(capacity);
Self {
capacity: cap,
refill_per_sec: cap / 60.0,
state: Mutex::new(BucketState {
tokens: cap,
last_refill: Instant::now(),
}),
}
}
pub fn take_at(&self, at: Instant) -> Result<(), Duration> {
let mut state = self
.state
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let refill = elapsed_refill(&mut state, at, self.capacity, self.refill_per_sec);
if refill >= 1.0 {
state.tokens = refill - 1.0;
Ok(())
} else if self.refill_per_sec <= 0.0 {
Err(Duration::MAX)
} else {
let wait_secs = (1.0 - refill) / self.refill_per_sec;
Err(Duration::from_secs_f64(wait_secs))
}
}
pub fn take(&self) -> Result<(), Duration> {
self.take_at(Instant::now())
}
pub fn available_at(&self, at: Instant) -> f64 {
let mut state = self
.state
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
elapsed_refill(&mut state, at, self.capacity, self.refill_per_sec)
}
pub fn available(&self) -> f64 {
self.available_at(Instant::now())
}
}
fn elapsed_refill(state: &mut BucketState, at: Instant, capacity: f64, refill_per_sec: f64) -> f64 {
let elapsed = at.saturating_duration_since(state.last_refill);
if elapsed.is_zero() {
return state.tokens;
}
if refill_per_sec <= 0.0 {
return state.tokens;
}
let added = elapsed.as_secs_f64() * refill_per_sec;
let raw = state.tokens + added;
if raw >= capacity {
let needed = capacity - state.tokens;
let secs_to_fill = needed / refill_per_sec;
state.last_refill = state
.last_refill
.checked_add(Duration::from_secs_f64(secs_to_fill))
.unwrap_or(at);
state.tokens = capacity;
} else {
state.tokens = raw;
state.last_refill = at;
}
state.tokens
}
#[derive(Debug)]
pub struct RateLimiter {
buckets: Mutex<HashMap<String, Arc<TokenBucket>>>,
requests_per_minute: u32,
}
impl RateLimiter {
#[must_use]
pub fn new(requests_per_minute: u32) -> Self {
Self {
buckets: Mutex::new(HashMap::new()),
requests_per_minute,
}
}
pub fn acquire(&self, base_url: &str) -> Result<(), Duration> {
if self.requests_per_minute == 0 {
return Ok(());
}
let bucket = {
let mut map = self
.buckets
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
Arc::clone(
map.entry(base_url.to_owned())
.or_insert_with(|| Arc::new(TokenBucket::new(self.requests_per_minute))),
)
};
bucket.take()
}
#[must_use]
pub fn is_enabled(&self) -> bool {
self.requests_per_minute > 0
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn burst_then_throttle() {
let bucket = TokenBucket::new(5);
let now = Instant::now();
for i in 0..5 {
assert!(bucket.take_at(now).is_ok(), "burst slot {i} should succeed");
}
let wait = bucket
.take_at(now)
.expect_err("sixth take should wait for a refill");
assert!(
wait >= Duration::from_secs(10) && wait <= Duration::from_secs(14),
"expected ~12s wait, got {wait:?}"
);
}
#[test]
fn refill_after_idle() {
let bucket = TokenBucket::new(5);
let t0 = Instant::now();
for _ in 0..5 {
bucket.take_at(t0).expect("drain should succeed while full");
}
assert!(bucket.available_at(t0) < 1.0, "bucket should be drained");
let t1 = t0 + Duration::from_mins(1);
let available = bucket.available_at(t1);
assert!(
(available - 5.0).abs() < 0.5,
"expected ~5.0 after 60s idle, got {available}"
);
}
#[test]
fn subsecond_refill_precision() {
let bucket = TokenBucket::new(10);
let t0 = Instant::now();
for _ in 0..10 {
bucket.take_at(t0).expect("drain should succeed while full");
}
let t1 = t0 + Duration::from_secs(6);
let available = bucket.available_at(t1);
assert!(
(available - 1.0).abs() < 0.1,
"expected ~1.0 after 6s, got {available}"
);
}
#[test]
fn available_is_non_consuming() {
let bucket = TokenBucket::new(5);
let now = Instant::now();
let a = bucket.available_at(now);
let b = bucket.available_at(now);
assert!(
(a - b).abs() < f64::EPSILON,
"available must be non-consuming"
);
bucket
.take_at(now)
.expect("take should succeed on a full bucket");
let c = bucket.available_at(now);
assert!(c < a, "take must decrement available");
}
#[test]
fn disabled_short_circuits() {
let limiter = RateLimiter::new(0);
assert!(!limiter.is_enabled());
assert!(limiter.acquire("anywhere").is_ok());
assert!(limiter.buckets.lock().unwrap().is_empty());
}
#[test]
fn per_provider_isolation() {
let limiter = RateLimiter::new(1);
assert!(limiter.acquire("openai").is_ok());
assert!(
limiter.acquire("openai").is_err(),
"openai bucket should be empty"
);
assert!(
limiter.acquire("ollama").is_ok(),
"ollama bucket must be independent"
);
}
#[test]
fn zero_capacity_bucket_does_not_panic() {
let bucket = TokenBucket::new(0);
let result = bucket.take();
assert!(result.is_err(), "empty bucket should return Err");
}
#[test]
fn future_instant_breaks_rate_limit() {
let bucket = TokenBucket::new(10);
let now = Instant::now();
let one_hour_later = now + Duration::from_hours(1);
let one_min_later = now + Duration::from_mins(1);
for _ in 0..10 {
bucket.take_at(now).expect("burst tokens");
}
let _poison = bucket.take_at(one_hour_later);
for _ in 0..9 {
let _drain = bucket.take_at(one_hour_later);
}
let result = bucket.take_at(one_min_later);
assert!(
result.is_ok(),
"take_at must succeed — 1 min of refill should restore a token; the future-instant call must not freeze the bucket"
);
}
#[test]
fn past_instant_does_not_refund_or_corrupt() {
let bucket = TokenBucket::new(5);
let t0 = Instant::now();
let t1 = t0 + Duration::from_secs(6);
for _ in 0..5 {
bucket.take_at(t0).expect("drain while full");
}
assert!(bucket.take_at(t0).is_err(), "drained at t0");
assert!(
bucket.take_at(t1).is_err(),
"partial refill under one token"
);
assert!(
bucket.take_at(t0).is_err(),
"past-instant take must not refund tokens from negative elapsed time"
);
assert!(
bucket.take_at(t1).is_err(),
"last_refill must not rewind — bucket state unchanged by the past-instant probe"
);
}
}