use std::sync::Arc;
use axum::{
body::Body,
extract::State,
http::Request,
middleware::Next,
response::{IntoResponse, Response},
};
#[derive(Debug, Clone)]
pub struct RateLimitConfig {
pub max_requests: u32,
pub window: std::time::Duration,
}
impl Default for RateLimitConfig {
fn default() -> Self {
Self {
max_requests: 120,
window: std::time::Duration::from_secs(60),
}
}
}
pub(crate) const RATE_LIMIT_SWEEP_THRESHOLD: usize = 4096;
#[derive(Debug, Clone)]
pub struct RateLimiter {
config: RateLimitConfig,
buckets: Arc<
std::sync::Mutex<std::collections::HashMap<std::net::IpAddr, (u32, std::time::Instant)>>,
>,
}
#[doc(hidden)]
#[derive(Debug, Clone)]
pub struct RateLimitContext {
limiter: RateLimiter,
ip: std::net::IpAddr,
}
impl RateLimitContext {
pub(crate) fn charge(&self, units: u32) -> Result<(), u64> {
self.limiter.check_n(self.ip, units)
}
}
impl RateLimiter {
pub fn new(config: RateLimitConfig) -> Self {
Self {
config,
buckets: Arc::new(std::sync::Mutex::new(std::collections::HashMap::default())),
}
}
pub fn disabled() -> Self {
Self::new(RateLimitConfig {
max_requests: 0,
window: std::time::Duration::from_secs(60),
})
}
pub fn is_enabled(&self) -> bool {
self.config.max_requests > 0
}
fn check(&self, ip: std::net::IpAddr) -> Result<(), u64> {
self.check_n(ip, 1)
}
fn check_n(&self, ip: std::net::IpAddr, units: u32) -> Result<(), u64> {
if !self.is_enabled() {
return Ok(());
}
let mut buckets = self
.buckets
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
let now = std::time::Instant::now();
if buckets.len() >= RATE_LIMIT_SWEEP_THRESHOLD {
buckets.retain(|_, (_, start)| now.duration_since(*start) < self.config.window);
}
let window = self.config.window;
let entry = buckets.entry(ip).or_insert((0, now));
if now.duration_since(entry.1) >= window {
*entry = (0, now);
}
if units > self.config.max_requests.saturating_sub(entry.0) {
let elapsed = now.duration_since(entry.1);
let remaining_ms = window
.saturating_sub(elapsed)
.as_millis()
.min(u64::MAX as u128) as u64;
return Err(remaining_ms.max(1));
}
entry.0 += units;
Ok(())
}
}
pub async fn rate_limit_layer(
State(limiter): State<RateLimiter>,
req: Request<Body>,
next: Next,
) -> Response {
if !limiter.is_enabled()
|| !req.uri().path().starts_with("/v1/")
|| req.method() == axum::http::Method::OPTIONS
{
return next.run(req).await;
}
let peer_ip = req
.extensions()
.get::<axum::extract::ConnectInfo<std::net::SocketAddr>>()
.map(|c| c.0.ip());
match peer_ip {
Some(ip) => match limiter.check(ip) {
Ok(()) => {
let mut req = req;
req.extensions_mut().insert(RateLimitContext {
limiter: limiter.clone(),
ip,
});
next.run(req).await
}
Err(retry_after_ms) => {
tracing::debug!(%ip, retry_after_ms, "rate limited");
crate::error::ApiError::RateLimited { retry_after_ms }.into_response()
}
},
None => next.run(req).await,
}
}