use std::collections::HashMap;
use std::num::NonZeroU32;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration;
use axum::{
body::Body,
extract::State,
http::{Request, StatusCode},
middleware::Next,
response::{IntoResponse, Response},
};
use governor::{
Quota, RateLimiter as GovLimiter,
clock::DefaultClock,
state::{InMemoryState, NotKeyed},
};
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct RateLimitConfig {
#[serde(default = "default_true")]
pub enabled: bool,
#[serde(default = "default_burst")]
pub burst: u32,
#[serde(default = "default_per_second")]
pub per_second: f64,
#[serde(default)]
pub endpoints: HashMap<String, EndpointLimit>,
#[serde(default)]
pub api_keys: HashMap<String, ApiKeyLimit>,
}
impl Default for RateLimitConfig {
fn default() -> Self {
Self {
enabled: true,
burst: default_burst(),
per_second: default_per_second(),
endpoints: HashMap::new(),
api_keys: HashMap::new(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct EndpointLimit {
pub burst: u32,
pub per_second: f64,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct ApiKeyLimit {
pub burst: u32,
pub per_second: f64,
}
fn default_true() -> bool {
true
}
fn default_burst() -> u32 {
30
}
fn default_per_second() -> f64 {
10.0
}
type Limiter = GovLimiter<NotKeyed, InMemoryState, DefaultClock>;
#[derive(Clone)]
pub struct RateLimiter {
inner: Arc<RateLimiterInner>,
}
struct RateLimiterInner {
config: RateLimitConfig,
default_limiter: Limiter,
key_limiters: HashMap<String, Limiter>,
endpoint_limiters: HashMap<String, Limiter>,
denied_count: AtomicU64,
}
impl RateLimiter {
pub fn new(config: RateLimitConfig) -> Self {
let default_quota = quota(config.burst, config.per_second);
let default_limiter = GovLimiter::direct(default_quota);
let key_limiters: HashMap<String, Limiter> = config
.api_keys
.iter()
.map(|(key, limit)| {
(
key.clone(),
GovLimiter::direct(quota(limit.burst, limit.per_second)),
)
})
.collect();
let endpoint_limiters: HashMap<String, Limiter> = config
.endpoints
.iter()
.map(|(path, limit)| {
(
path.clone(),
GovLimiter::direct(quota(limit.burst, limit.per_second)),
)
})
.collect();
Self {
inner: Arc::new(RateLimiterInner {
config,
default_limiter,
key_limiters,
endpoint_limiters,
denied_count: AtomicU64::new(0),
}),
}
}
pub fn check(&self, key: &str, path: &str) -> Result<(), Duration> {
let inner = &self.inner;
if !inner.config.enabled {
return Ok(());
}
let limiter = inner
.key_limiters
.get(key)
.or_else(|| {
inner
.endpoint_limiters
.iter()
.find(|(prefix, _)| path.starts_with(prefix.as_str()))
.map(|(_, l)| l)
})
.unwrap_or(&inner.default_limiter);
match limiter.check() {
Ok(_) => Ok(()),
Err(negative) => {
inner.denied_count.fetch_add(1, Ordering::Relaxed);
let wait = negative.wait_time_from(governor::clock::Clock::now(
&governor::clock::DefaultClock::default(),
));
Err(wait)
}
}
}
pub fn denied_count(&self) -> u64 {
self.inner.denied_count.load(Ordering::Relaxed)
}
}
fn quota(burst: u32, per_second: f64) -> Quota {
let replenish_interval = Duration::from_secs_f64(1.0 / per_second);
Quota::with_period(replenish_interval)
.unwrap()
.allow_burst(NonZeroU32::new(burst.max(1)).unwrap())
}
pub async fn rate_limit_middleware(
State(limiter): State<RateLimiter>,
req: Request<Body>,
next: Next,
) -> Result<Response, StatusCode> {
let path = req.uri().path().to_string();
let key = req
.headers()
.get("authorization")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.strip_prefix("Bearer "))
.unwrap_or("anonymous");
match limiter.check(key, &path) {
Ok(()) => Ok(next.run(req).await),
Err(wait) => {
let retry_after = wait.as_secs().max(1);
let body = serde_json::json!({
"error": "rate_limited",
"message": format!("Too many requests. Retry in {}s", retry_after),
"retry_after_seconds": retry_after,
});
Ok((
StatusCode::TOO_MANY_REQUESTS,
[("retry-after", retry_after.to_string())],
axum::Json(body),
)
.into_response())
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn disabled_always_allows() {
let config = RateLimitConfig {
enabled: false,
..Default::default()
};
let limiter = RateLimiter::new(config);
assert!(limiter.check("key1", "/v1/chat/completions").is_ok());
for _ in 0..1000 {
assert!(limiter.check("key1", "/v1/chat/completions").is_ok());
}
}
#[test]
fn enabled_limits_by_default() {
let config = RateLimitConfig {
enabled: true,
burst: 5,
per_second: 100.0,
..Default::default()
};
let limiter = RateLimiter::new(config);
for _ in 0..5 {
assert!(limiter.check("key1", "/v1/chat/completions").is_ok());
}
assert!(limiter.check("key1", "/v1/chat/completions").is_err());
}
#[test]
fn per_api_key_limits_override_default() {
let mut api_keys = HashMap::new();
api_keys.insert(
"premium".into(),
ApiKeyLimit {
burst: 100,
per_second: 100.0,
},
);
let config = RateLimitConfig {
enabled: true,
burst: 3,
per_second: 100.0,
api_keys,
..Default::default()
};
let limiter = RateLimiter::new(config);
for _ in 0..50 {
assert!(limiter.check("premium", "/v1/chat/completions").is_ok());
}
assert!(limiter.check("normal", "/v1/chat/completions").is_ok());
assert!(limiter.check("normal", "/v1/chat/completions").is_ok());
assert!(limiter.check("normal", "/v1/chat/completions").is_ok());
assert!(limiter.check("normal", "/v1/chat/completions").is_err());
}
}