use std::sync::Arc;
use std::time::Duration;
use axum::body::Body;
use axum::http::{Request, StatusCode};
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
use majra::ratelimit::RateLimiter;
use tracing::warn;
pub struct RateLimitState {
limiter: RateLimiter,
eviction_interval: Duration,
}
impl RateLimitState {
pub fn new(rate: f64, burst: usize) -> Self {
Self {
limiter: RateLimiter::new(rate, burst),
eviction_interval: Duration::from_secs(300),
}
}
#[must_use]
pub fn check(&self, key: &str) -> bool {
self.limiter.check(key)
}
#[must_use]
pub fn evict_stale(&self) -> usize {
self.limiter.evict_stale(self.eviction_interval)
}
pub fn stats(&self) -> majra::ratelimit::RateLimitStats {
self.limiter.stats()
}
}
fn extract_client_key(req: &Request<Body>) -> String {
if let Some(xff) = req.headers().get("x-forwarded-for")
&& let Ok(s) = xff.to_str()
&& let Some(first_ip) = s.split(',').next()
{
return first_ip.trim().to_string();
}
if let Some(xri) = req.headers().get("x-real-ip")
&& let Ok(s) = xri.to_str()
{
return s.trim().to_string();
}
"unknown".to_string()
}
pub async fn rate_limit_middleware(req: Request<Body>, next: Next) -> Response {
let limiter = req.extensions().get::<Arc<RateLimitState>>();
if let Some(limiter) = limiter {
let key = extract_client_key(&req);
if !limiter.check(&key) {
warn!(client = %key, "rate limit exceeded");
return StatusCode::TOO_MANY_REQUESTS.into_response();
}
}
next.run(req).await
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rate_limit_state_allows_within_burst() {
let state = RateLimitState::new(10.0, 5);
for _ in 0..5 {
assert!(state.check("client-1"));
}
assert!(!state.check("client-1"));
}
#[test]
fn rate_limit_separate_keys() {
let state = RateLimitState::new(10.0, 2);
assert!(state.check("client-a"));
assert!(state.check("client-a"));
assert!(!state.check("client-a")); assert!(state.check("client-b")); }
#[test]
fn rate_limit_stats() {
let state = RateLimitState::new(10.0, 1);
let _ = state.check("key");
let _ = state.check("key"); let stats = state.stats();
assert_eq!(stats.total_allowed, 1);
assert_eq!(stats.total_rejected, 1);
}
#[test]
fn evict_stale_returns_count() {
let state = RateLimitState::new(10.0, 5);
let _ = state.check("key-a");
assert_eq!(state.evict_stale(), 0);
}
#[test]
fn extract_key_from_xff_header() {
let req = Request::builder()
.header("x-forwarded-for", "1.2.3.4, 5.6.7.8")
.body(Body::empty())
.unwrap();
assert_eq!(extract_client_key(&req), "1.2.3.4");
}
#[test]
fn extract_key_from_xri_header() {
let req = Request::builder()
.header("x-real-ip", "10.0.0.1")
.body(Body::empty())
.unwrap();
assert_eq!(extract_client_key(&req), "10.0.0.1");
}
#[test]
fn extract_key_fallback() {
let req = Request::builder().body(Body::empty()).unwrap();
assert_eq!(extract_client_key(&req), "unknown");
}
}