use std::{sync::Arc, time::Duration};
use tokio::sync::Mutex;
#[non_exhaustive]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RateLimit {
Unlimited,
Limited {
max_requests: u32,
per: Duration,
},
}
impl RateLimit {
pub const fn per(max_requests: u32, per: Duration) -> Self {
RateLimit::Limited { max_requests, per }
}
pub(crate) fn validate(&self) -> Result<(), String> {
if let RateLimit::Limited { max_requests, per } = self {
if *max_requests == 0 {
return Err("rate limit `max_requests` must be greater than 0".to_string());
}
if per.is_zero() {
return Err("rate limit `per` window must be greater than 0".to_string());
}
}
Ok(())
}
}
#[derive(Debug, Clone)]
pub(crate) struct RateLimiter {
bucket: Option<Arc<Mutex<Bucket>>>,
interval: Duration,
capacity: f64,
}
#[derive(Debug)]
struct Bucket {
tokens: f64,
last: Option<tokio::time::Instant>,
}
impl RateLimiter {
pub(crate) fn new(limit: RateLimit) -> Self {
match limit {
RateLimit::Unlimited => Self {
bucket: None,
interval: Duration::ZERO,
capacity: 0.0,
},
RateLimit::Limited { max_requests, per } => {
let capacity = f64::from(max_requests.max(1));
Self {
bucket: Some(Arc::new(Mutex::new(Bucket {
tokens: capacity, last: None,
}))),
interval: per / max_requests.max(1),
capacity,
}
}
}
}
pub(crate) async fn acquire(&self) {
let Some(bucket) = &self.bucket else {
return;
};
let refill_per_sec = 1.0 / self.interval.as_secs_f64();
let mut b = bucket.lock().await;
let now = tokio::time::Instant::now();
if let Some(last) = b.last {
let elapsed = now.saturating_duration_since(last).as_secs_f64();
b.tokens = (b.tokens + elapsed * refill_per_sec).min(self.capacity);
}
b.last = Some(now);
if b.tokens >= 1.0 {
b.tokens -= 1.0;
return;
}
let deficit = 1.0 - b.tokens;
let until = now + self.interval.mul_f64(deficit);
b.tokens = 0.0;
b.last = Some(until);
tokio::time::sleep_until(until).await;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn validate_rejects_zero_values() {
assert!(
RateLimit::Limited {
max_requests: 0,
per: Duration::from_secs(1)
}
.validate()
.is_err()
);
assert!(
RateLimit::Limited {
max_requests: 5,
per: Duration::ZERO
}
.validate()
.is_err()
);
assert!(RateLimit::per(5, Duration::from_secs(1)).validate().is_ok());
assert!(RateLimit::Unlimited.validate().is_ok());
}
#[tokio::test(start_paused = true)]
async fn sustained_rate_is_bounded() {
let limiter = RateLimiter::new(RateLimit::per(5, Duration::from_secs(1)));
let start = tokio::time::Instant::now();
for _ in 0..15 {
limiter.acquire().await;
}
let elapsed = start.elapsed();
assert!(
elapsed >= Duration::from_millis(1900),
"15 requests at 5/s took only {elapsed:?}; rate is not bounded"
);
}
#[tokio::test(start_paused = true)]
async fn unlimited_never_waits() {
let limiter = RateLimiter::new(RateLimit::Unlimited);
let start = tokio::time::Instant::now();
for _ in 0..1000 {
limiter.acquire().await;
}
assert_eq!(start.elapsed(), Duration::ZERO);
}
#[tokio::test(start_paused = true)]
async fn initial_burst_capped_at_capacity() {
let limiter = RateLimiter::new(RateLimit::per(3, Duration::from_secs(1)));
let start = tokio::time::Instant::now();
for _ in 0..3 {
limiter.acquire().await;
}
assert_eq!(start.elapsed(), Duration::ZERO);
limiter.acquire().await;
assert!(start.elapsed() >= Duration::from_millis(300));
}
}