use axum::{
Json,
extract::{Request, State},
http::{HeaderMap, HeaderValue, StatusCode},
middleware::Next,
response::{IntoResponse, Response},
};
use governor::{
Quota, RateLimiter as GovernorRateLimiter,
clock::{Clock, DefaultClock},
state::{InMemoryState, NotKeyed},
};
use serde_json::json;
use std::num::NonZeroU32;
use std::sync::Arc;
use std::time::Duration;
use crate::server::metrics;
pub type RateLimiter = Arc<GovernorRateLimiter<NotKeyed, InMemoryState, DefaultClock>>;
#[derive(Clone, Debug)]
pub struct RateLimitConfig {
pub requests_per_second: u32,
pub burst_size: u32,
pub enabled: bool,
}
impl Default for RateLimitConfig {
fn default() -> Self {
Self {
requests_per_second: 100,
burst_size: 200,
enabled: false,
}
}
}
impl RateLimitConfig {
pub fn build_limiter(&self) -> Option<RateLimiter> {
if !self.enabled {
return None;
}
let quota = Quota::per_second(
NonZeroU32::new(self.requests_per_second).expect("Invalid requests_per_second"),
)
.allow_burst(NonZeroU32::new(self.burst_size).expect("Invalid burst_size"));
Some(Arc::new(GovernorRateLimiter::new(
quota,
InMemoryState::default(),
DefaultClock::default(),
)))
}
}
pub async fn rate_limit_middleware(
State(limiter): State<Option<RateLimiter>>,
request: Request,
next: Next,
) -> Response {
let Some(limiter) = limiter else {
return next.run(request).await;
};
match limiter.check() {
Ok(()) => {
next.run(request).await
}
Err(not_until) => {
let retry_after = not_until.wait_time_from(DefaultClock::default().now());
let retry_after_secs = retry_after.as_secs();
metrics::record_rate_limited("global", "token_exhausted");
tracing::warn!(
"Rate limit exceeded, retry after {} seconds",
retry_after_secs
);
let mut response = (
StatusCode::TOO_MANY_REQUESTS,
Json(json!({
"error": {
"message": "Rate limit exceeded. Please retry later.",
"type": "rate_limit_exceeded",
"code": "rate_limit",
"retry_after_seconds": retry_after_secs
}
})),
)
.into_response();
response.headers_mut().insert(
"Retry-After",
HeaderValue::from_str(&retry_after_secs.to_string()).unwrap(),
);
response
}
}
}
pub struct ClientRateLimiter {
limiters: dashmap::DashMap<String, (RateLimiter, std::time::Instant)>,
config: RateLimitConfig,
}
impl ClientRateLimiter {
pub fn new(config: RateLimitConfig) -> Self {
Self {
limiters: dashmap::DashMap::new(),
config,
}
}
pub fn get_limiter(&self, client_id: &str) -> RateLimiter {
let now = std::time::Instant::now();
self.limiters
.entry(client_id.to_string())
.and_modify(|(_, last_access)| *last_access = now)
.or_insert_with(|| {
(
self.config
.build_limiter()
.expect("Rate limiting should be enabled"),
now,
)
})
.0
.clone()
}
pub fn cleanup(&self, max_age: Duration) {
let now = std::time::Instant::now();
let before_size = self.limiters.len();
self.limiters
.retain(|_key, (_limiter, last_access)| now.duration_since(*last_access) < max_age);
let after_size = self.limiters.len();
if before_size > after_size {
tracing::debug!(
"Cleaned {} expired entries from client rate limiter cache ({} remaining)",
before_size - after_size,
after_size
);
}
if after_size > 10000 {
tracing::warn!(
"Client rate limiter cache has {} entries after cleanup, consider decreasing max_age",
after_size
);
}
}
}
pub fn extract_client_id(headers: &HeaderMap) -> String {
headers
.get("X-Forwarded-For")
.and_then(|v| v.to_str().ok())
.and_then(|s| s.split(',').next())
.map(|s| s.trim().to_string())
.or_else(|| {
headers
.get("X-Real-IP")
.and_then(|v| v.to_str().ok())
.map(str::to_string)
})
.unwrap_or_else(|| "unknown".to_string())
}