use std::time::{Duration, Instant};
use super::config::RateLimitConfig;
#[derive(Debug)]
pub struct RateLimiter {
tokens: f64,
capacity: f64,
refill_rate: f64,
last_update: Instant,
enabled: bool,
}
impl RateLimiter {
#[must_use]
pub fn new(config: &RateLimitConfig) -> Self {
Self {
tokens: f64::from(config.burst_size),
capacity: f64::from(config.burst_size),
refill_rate: f64::from(config.requests_per_second),
last_update: Instant::now(),
enabled: config.enabled,
}
}
pub fn try_acquire(&mut self) -> bool {
if !self.enabled {
return true;
}
self.refill();
if self.tokens >= 1.0 {
self.tokens -= 1.0;
true
} else {
false
}
}
#[must_use]
pub fn time_until_available(&mut self) -> Duration {
if !self.enabled {
return Duration::ZERO;
}
self.refill();
if self.tokens >= 1.0 {
Duration::ZERO
} else {
let tokens_needed = 1.0 - self.tokens;
let seconds_needed = tokens_needed / self.refill_rate;
Duration::from_secs_f64(seconds_needed)
}
}
#[must_use]
pub const fn is_enabled(&self) -> bool {
self.enabled
}
#[must_use]
pub fn available_tokens(&mut self) -> f64 {
self.refill();
self.tokens
}
fn refill(&mut self) {
let now = Instant::now();
let elapsed = now.duration_since(self.last_update).as_secs_f64();
self.last_update = now;
self.tokens = elapsed.mul_add(self.refill_rate, self.tokens).min(self.capacity);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_rate_limiter_initial_burst() {
let config = RateLimitConfig {
enabled: true,
requests_per_second: 10,
burst_size: 5,
};
let mut limiter = RateLimiter::new(&config);
for _ in 0..5 {
assert!(limiter.try_acquire(), "Should allow initial burst");
}
assert!(!limiter.try_acquire(), "Should rate limit after burst");
}
#[test]
fn test_rate_limiter_disabled() {
let config = RateLimitConfig {
enabled: false,
requests_per_second: 10,
burst_size: 5,
};
let mut limiter = RateLimiter::new(&config);
for _ in 0..100 {
assert!(limiter.try_acquire(), "Should allow all requests when disabled");
}
}
#[test]
fn test_rate_limiter_disabled_via_config() {
let config = RateLimitConfig {
enabled: false,
requests_per_second: 10,
burst_size: 5,
};
let mut limiter = RateLimiter::new(&config);
assert!(!limiter.is_enabled());
assert!(limiter.try_acquire());
assert_eq!(limiter.time_until_available(), Duration::ZERO);
}
#[test]
fn test_rate_limiter_time_until_available() {
let config = RateLimitConfig {
enabled: true,
requests_per_second: 10,
burst_size: 1,
};
let mut limiter = RateLimiter::new(&config);
assert!(limiter.try_acquire());
let wait_time = limiter.time_until_available();
assert!(wait_time > Duration::ZERO);
assert!(wait_time <= Duration::from_millis(150));
}
#[test]
fn test_rate_limiter_refill() {
let config = RateLimitConfig {
enabled: true,
requests_per_second: 1000, burst_size: 1,
};
let mut limiter = RateLimiter::new(&config);
assert!(limiter.try_acquire());
assert!(!limiter.try_acquire());
std::thread::sleep(Duration::from_millis(2));
assert!(limiter.try_acquire());
}
#[test]
fn test_rate_limiter_capacity_limit() {
let config = RateLimitConfig {
enabled: true,
requests_per_second: 1000,
burst_size: 3,
};
let mut limiter = RateLimiter::new(&config);
std::thread::sleep(Duration::from_millis(10));
let tokens = limiter.available_tokens();
assert!(tokens <= 3.0, "Tokens should not exceed capacity");
}
#[test]
fn test_default_config_rate_limiter() {
let config = RateLimitConfig::default();
let mut limiter = RateLimiter::new(&config);
assert!(limiter.is_enabled());
for _ in 0..50 {
assert!(limiter.try_acquire());
}
}
}