use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::{Mutex, OwnedSemaphorePermit, Semaphore};
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum RateLimitError {
#[error("request rate limit exceeded")]
TooManyRequests,
#[error("too many concurrent requests")]
ConcurrencyLimitExceeded,
}
struct WindowState {
window_start: Instant,
count: usize,
}
pub struct RateLimiter {
semaphore: Arc<Semaphore>,
window: Duration,
max_requests: usize,
state: Mutex<WindowState>,
}
impl RateLimiter {
pub fn new(max_concurrent: usize, max_requests_per_minute: usize) -> Self {
let permits = if max_concurrent == 0 {
Semaphore::MAX_PERMITS
} else {
max_concurrent
};
Self {
semaphore: Arc::new(Semaphore::new(permits)),
window: Duration::from_secs(60),
max_requests: max_requests_per_minute,
state: Mutex::new(WindowState {
window_start: Instant::now(),
count: 0,
}),
}
}
pub async fn try_acquire(&self) -> Result<RateLimitPermit, RateLimitError> {
if self.max_requests > 0 {
let mut state = self.state.lock().await;
let now = Instant::now();
if now.duration_since(state.window_start) >= self.window {
state.window_start = now;
state.count = 0;
}
if state.count >= self.max_requests {
return Err(RateLimitError::TooManyRequests);
}
state.count += 1;
}
let permit = self
.semaphore
.clone()
.try_acquire_owned()
.map_err(|_| RateLimitError::ConcurrencyLimitExceeded)?;
Ok(RateLimitPermit { _permit: permit })
}
}
pub struct RateLimitPermit {
_permit: OwnedSemaphorePermit,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn unlimited_acquires() {
let limiter = RateLimiter::new(0, 0);
let rt = tokio::runtime::Runtime::new().unwrap();
rt.block_on(async {
let permit = limiter.try_acquire().await;
assert!(permit.is_ok());
drop(permit);
assert!(limiter.try_acquire().await.is_ok());
});
}
#[test]
fn enforces_per_minute_limit() {
let limiter = RateLimiter::new(0, 2);
let rt = tokio::runtime::Runtime::new().unwrap();
rt.block_on(async {
assert!(limiter.try_acquire().await.is_ok());
assert!(limiter.try_acquire().await.is_ok());
let err = limiter.try_acquire().await;
assert!(matches!(err, Err(RateLimitError::TooManyRequests)));
});
}
#[test]
fn enforces_concurrency_limit() {
let limiter = RateLimiter::new(1, 0);
let rt = tokio::runtime::Runtime::new().unwrap();
rt.block_on(async {
let p1 = limiter.try_acquire().await.unwrap();
let err = limiter.try_acquire().await;
assert!(matches!(err, Err(RateLimitError::ConcurrencyLimitExceeded)));
drop(p1);
assert!(limiter.try_acquire().await.is_ok());
});
}
}