use std::sync::Mutex;
use std::time::Instant;
use tracing::Metadata;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub struct RateLimit {
pub max_events: u64,
pub per_secs: u64,
}
impl RateLimit {
pub fn new(max_events: u64, per_secs: u64) -> Self {
Self {
max_events,
per_secs,
}
}
pub fn per_second(max_events: u64) -> Self {
Self::new(max_events, 1)
}
pub fn per_minute(max_events: u64) -> Self {
Self::new(max_events, 60)
}
}
#[derive(Debug)]
pub struct RateLimiter {
config: RateLimit,
state: Mutex<Bucket>,
}
#[derive(Debug)]
struct Bucket {
tokens: f64,
last_refill: Instant,
}
impl RateLimiter {
pub fn new(config: RateLimit) -> Self {
Self {
config,
state: Mutex::new(Bucket {
tokens: config.max_events as f64,
last_refill: Instant::now(),
}),
}
}
pub fn allow(&self) -> bool {
let mut bucket = match self.state.lock() {
Ok(b) => b,
Err(_) => return false, };
let now = Instant::now();
let elapsed = now.duration_since(bucket.last_refill).as_secs_f64();
if elapsed > 0.0 {
let refill_rate = self.config.max_events as f64 / self.config.per_secs as f64;
bucket.tokens =
(bucket.tokens + elapsed * refill_rate).min(self.config.max_events as f64);
bucket.last_refill = now;
}
if bucket.tokens >= 1.0 {
bucket.tokens -= 1.0;
true
} else {
false
}
}
}
impl super::gate::EventGate for RateLimiter {
fn allows(&self, _meta: &Metadata<'_>) -> bool {
self.allow()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn per_second_defaults() {
let rl = RateLimit::per_second(10);
assert_eq!(rl.max_events, 10);
assert_eq!(rl.per_secs, 1);
}
#[test]
fn per_minute_defaults() {
let rl = RateLimit::per_minute(60);
assert_eq!(rl.max_events, 60);
assert_eq!(rl.per_secs, 60);
}
#[test]
fn allows_events_up_to_limit() {
let limiter = RateLimiter::new(RateLimit::per_second(5));
for _ in 0..5 {
assert!(limiter.allow(), "first 5 events should be allowed");
}
}
#[test]
fn throttles_after_limit() {
let limiter = RateLimiter::new(RateLimit::per_second(3));
for _ in 0..3 {
assert!(limiter.allow());
}
assert!(!limiter.allow(), "4th event should be throttled");
}
#[test]
fn refills_over_time() {
let limiter = RateLimiter::new(RateLimit::per_second(100));
for _ in 0..100 {
assert!(limiter.allow());
}
assert!(!limiter.allow());
std::thread::sleep(std::time::Duration::from_millis(150));
for _ in 0..10 {
assert!(limiter.allow(), "should refill ~10 tokens after 150ms");
}
}
#[test]
fn token_cap_respected() {
let limiter = RateLimiter::new(RateLimit::per_second(5));
std::thread::sleep(std::time::Duration::from_secs(2));
let mut allowed = 0;
while limiter.allow() {
allowed += 1;
}
assert_eq!(allowed, 5, "token bucket should cap at max_events");
}
}