#[cfg(feature = "server")]
use axum::{
body::Body,
extract::Request,
http::{HeaderMap, HeaderValue, StatusCode},
middleware::Next,
response::{IntoResponse, Response},
Json,
};
#[cfg(feature = "server")]
use serde_json::json;
#[cfg(feature = "server")]
use std::collections::HashMap;
#[cfg(feature = "server")]
use std::sync::{Arc, Mutex};
#[cfg(feature = "server")]
use std::time::{Duration, Instant};
#[cfg(feature = "server")]
use crate::server::config::RateLimitConfig;
#[cfg(feature = "server")]
#[derive(Clone)]
pub struct RateLimiter {
config: RateLimitConfig,
requests: Arc<Mutex<HashMap<String, (u32, Instant)>>>,
window_duration: Duration,
}
#[cfg(feature = "server")]
impl RateLimiter {
pub fn new(config: RateLimitConfig) -> Self {
Self {
config,
requests: Arc::new(Mutex::new(HashMap::new())),
window_duration: Duration::from_secs(60), }
}
fn check_rate_limit(&self, key: &str, limit: u32) -> (bool, u32, u64) {
let mut requests = self.requests.lock().unwrap();
let now = Instant::now();
let entry = requests.entry(key.to_string()).or_insert((0, now));
if now.duration_since(entry.1) > self.window_duration {
*entry = (1, now);
return (true, limit - 1, self.window_duration.as_secs());
}
entry.0 += 1;
let remaining = limit.saturating_sub(entry.0);
let reset_seconds = self.window_duration.as_secs() - now.duration_since(entry.1).as_secs();
(entry.0 <= limit, remaining, reset_seconds)
}
fn get_client_ip(headers: &HeaderMap) -> Option<String> {
if let Some(forwarded) = headers.get("x-forwarded-for") {
if let Ok(forwarded_str) = forwarded.to_str() {
return forwarded_str
.split(',')
.next()
.map(|s| s.trim().to_string());
}
}
if let Some(real_ip) = headers.get("x-real-ip") {
if let Ok(ip_str) = real_ip.to_str() {
return Some(ip_str.to_string());
}
}
None
}
fn get_api_key(headers: &HeaderMap) -> Option<String> {
headers
.get("x-api-key")
.and_then(|v| v.to_str().ok())
.map(|s| s.to_string())
}
fn get_user_id(request: &Request<Body>) -> Option<String> {
use crate::server::auth::jwt::Claims;
request
.extensions()
.get::<Claims>()
.map(|claims| claims.sub.clone())
}
}
#[cfg(feature = "server")]
pub async fn rate_limit_middleware(
limiter: RateLimiter,
request: Request<Body>,
next: Next,
) -> Result<Response, Response> {
if !limiter.config.enabled {
return Ok(next.run(request).await);
}
let headers = request.headers();
let mut limit: Option<u32> = None;
let mut remaining: Option<u32> = None;
let mut reset: Option<u64> = None;
if let Some(ip) = RateLimiter::get_client_ip(headers) {
let (allowed, rem, rst) = limiter.check_rate_limit(
&format!("global:{}", ip),
limiter.config.per_ip_requests_per_minute,
);
limit = Some(limiter.config.per_ip_requests_per_minute);
remaining = Some(rem);
reset = Some(rst);
if !allowed {
return Err((
StatusCode::TOO_MANY_REQUESTS,
[
(
"X-RateLimit-Limit",
limiter.config.per_ip_requests_per_minute.to_string(),
),
("X-RateLimit-Remaining", "0".to_string()),
("X-RateLimit-Reset", rst.to_string()),
("Retry-After", rst.to_string()),
],
Json(json!({
"error": "Rate limit exceeded",
"message": "Too many requests from your IP",
"retry_after": rst,
})),
)
.into_response());
}
}
if let Some(api_key) = RateLimiter::get_api_key(headers) {
let (allowed, rem, rst) = limiter.check_rate_limit(
&format!("api_key:{}", api_key),
limiter.config.per_api_key_requests_per_minute,
);
if let Some(current_remaining) = remaining {
if rem < current_remaining {
limit = Some(limiter.config.per_api_key_requests_per_minute);
remaining = Some(rem);
reset = Some(rst);
}
} else {
limit = Some(limiter.config.per_api_key_requests_per_minute);
remaining = Some(rem);
reset = Some(rst);
}
if !allowed {
return Err((
StatusCode::TOO_MANY_REQUESTS,
[
(
"X-RateLimit-Limit",
limiter.config.per_api_key_requests_per_minute.to_string(),
),
("X-RateLimit-Remaining", "0".to_string()),
("X-RateLimit-Reset", rst.to_string()),
("Retry-After", rst.to_string()),
],
Json(json!({
"error": "Rate limit exceeded",
"message": "Too many requests for this API key",
"retry_after": rst,
})),
)
.into_response());
}
}
if let Some(user_id) = RateLimiter::get_user_id(&request) {
let (allowed, rem, rst) = limiter.check_rate_limit(
&format!("user:{}", user_id),
limiter.config.per_user_requests_per_minute,
);
if let Some(current_remaining) = remaining {
if rem < current_remaining {
limit = Some(limiter.config.per_user_requests_per_minute);
remaining = Some(rem);
reset = Some(rst);
}
} else {
limit = Some(limiter.config.per_user_requests_per_minute);
remaining = Some(rem);
reset = Some(rst);
}
if !allowed {
return Err((
StatusCode::TOO_MANY_REQUESTS,
[
(
"X-RateLimit-Limit",
limiter.config.per_user_requests_per_minute.to_string(),
),
("X-RateLimit-Remaining", "0".to_string()),
("X-RateLimit-Reset", rst.to_string()),
("Retry-After", rst.to_string()),
],
Json(json!({
"error": "Rate limit exceeded",
"message": "Too many requests for your account",
"retry_after": rst,
})),
)
.into_response());
}
}
let mut response = next.run(request).await;
if let (Some(lim), Some(rem), Some(rst)) = (limit, remaining, reset) {
let headers = response.headers_mut();
if let Ok(value) = HeaderValue::from_str(&lim.to_string()) {
headers.insert("X-RateLimit-Limit", value);
}
if let Ok(value) = HeaderValue::from_str(&rem.to_string()) {
headers.insert("X-RateLimit-Remaining", value);
}
if let Ok(value) = HeaderValue::from_str(&rst.to_string()) {
headers.insert("X-RateLimit-Reset", value);
}
}
Ok(response)
}
#[cfg(all(test, feature = "server"))]
mod tests {
use super::*;
#[test]
fn test_rate_limiter() {
let config = RateLimitConfig {
enabled: true,
global_requests_per_minute: 100,
per_user_requests_per_minute: 50,
per_api_key_requests_per_minute: 200,
per_ip_requests_per_minute: 60,
};
let limiter = RateLimiter::new(config);
let (allowed, remaining, _) = limiter.check_rate_limit("test_key", 10);
assert!(allowed);
assert_eq!(remaining, 9);
for i in 0..9 {
let (allowed, remaining, _) = limiter.check_rate_limit("test_key", 10);
assert!(allowed);
assert_eq!(remaining, 9 - i - 1);
}
let (allowed, remaining, _) = limiter.check_rate_limit("test_key", 10);
assert!(!allowed);
assert_eq!(remaining, 0);
}
}