use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::Mutex;
#[derive(Clone)]
pub struct RateLimiter {
state: Arc<Mutex<RateLimiterState>>,
config: RateLimiterConfig,
}
#[derive(Clone, Copy)]
struct RateLimiterConfig {
capacity: u32,
refill_interval: Duration,
}
struct RateLimiterState {
tokens: f64,
last_update: Instant,
}
impl RateLimiter {
#[must_use]
pub fn new(requests: u32, per_interval: Duration) -> Self {
Self {
state: Arc::new(Mutex::new(RateLimiterState {
tokens: requests as f64,
last_update: Instant::now(),
})),
config: RateLimiterConfig {
capacity: requests,
refill_interval: per_interval,
},
}
}
#[must_use]
pub fn default_booru() -> Self {
Self::new(2, Duration::from_secs(1))
}
pub async fn acquire(&self) {
loop {
let wait_time = {
let mut state = self.state.lock().await;
self.refill_tokens(&mut state);
if state.tokens >= 1.0 {
state.tokens -= 1.0;
return;
}
let tokens_needed = 1.0 - state.tokens;
let refill_rate =
self.config.capacity as f64 / self.config.refill_interval.as_secs_f64();
Duration::from_secs_f64(tokens_needed / refill_rate)
};
tokio::time::sleep(wait_time).await;
}
}
pub async fn try_acquire(&self) -> bool {
let mut state = self.state.lock().await;
self.refill_tokens(&mut state);
if state.tokens >= 1.0 {
state.tokens -= 1.0;
true
} else {
false
}
}
pub async fn available(&self) -> u32 {
let mut state = self.state.lock().await;
self.refill_tokens(&mut state);
state.tokens as u32
}
fn refill_tokens(&self, state: &mut RateLimiterState) {
let now = Instant::now();
let elapsed = now.duration_since(state.last_update);
if elapsed > Duration::ZERO {
let refill_rate =
self.config.capacity as f64 / self.config.refill_interval.as_secs_f64();
let new_tokens = elapsed.as_secs_f64() * refill_rate;
state.tokens = (state.tokens + new_tokens).min(self.config.capacity as f64);
state.last_update = now;
}
}
}
impl std::fmt::Debug for RateLimiter {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RateLimiter")
.field("capacity", &self.config.capacity)
.field("refill_interval", &self.config.refill_interval)
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_try_acquire() {
let limiter = RateLimiter::new(2, Duration::from_secs(1));
assert!(limiter.try_acquire().await);
assert!(limiter.try_acquire().await);
assert!(!limiter.try_acquire().await);
}
#[tokio::test]
async fn test_available() {
let limiter = RateLimiter::new(5, Duration::from_secs(1));
assert_eq!(limiter.available().await, 5);
limiter.acquire().await;
assert_eq!(limiter.available().await, 4);
}
#[tokio::test]
async fn test_refill() {
let limiter = RateLimiter::new(10, Duration::from_millis(100));
for _ in 0..10 {
limiter.try_acquire().await;
}
assert_eq!(limiter.available().await, 0);
tokio::time::sleep(Duration::from_millis(150)).await;
assert!(limiter.available().await >= 10);
}
}