//! Simple in-memory rate limiter for Axum. //! //! Uses a sliding-window approach: each client has a rolling list of //! timestamps. Requests arriving within the window that exceed the //! configured max are rejected. use std::collections::HashMap; use std::future::Future; use std::pin::Pin; use std::sync::Mutex; use std::task::{Context, Poll}; use std::time::{SystemTime, UNIX_EPOCH}; use axum::http::{Request, StatusCode}; use axum::response::{IntoResponse, Response}; use tower::{Layer, Service}; /// In-memory sliding-window rate limiter. /// /// Thread-safe via interior mutability (`Mutex`). Each client (identified /// by a string key, e.g. IP address or session id) has a Vec of entry /// timestamps (in seconds). Old entries are cleaned on every check. #[derive(Debug)] pub struct RateLimiter { windows: Mutex>>, trust_proxy_headers: bool, } impl RateLimiter { /// Create a rate limiter that keys strictly on the real connection /// socket address (default, safe when not behind a trusted proxy). pub fn new() -> Self { Self::with_proxy_trust(false) } /// Create a rate limiter with an explicit proxy-header trust policy. /// /// When `trust_proxy_headers` is `true`, the `X-Forwarded-For` and /// `X-Real-IP` headers are used to derive the client bucket key. /// This must only be enabled when the middleware sits behind a /// reverse proxy known to overwrite (not merge) these headers. pub fn with_proxy_trust(trust_proxy_headers: bool) -> Self { Self { windows: Mutex::new(HashMap::new()), trust_proxy_headers, } } /// Check whether a request from `client_id` should be allowed. /// /// * `max_requests` — max number of requests permitted within the /// window. /// * `window_secs` — width of the sliding window in seconds. /// /// Returns `Ok(true)` if the request is allowed (and records it), /// or `Ok(false)` if the client has exceeded the limit. /// /// The window is **sliding**: only timestamps falling within /// `[now - window_secs, now]` are counted. pub fn check_rate_limit( &self, client_id: &str, max_requests: u32, window_secs: u64, ) -> anyhow::Result { let now = SystemTime::now() .duration_since(UNIX_EPOCH) .unwrap_or_default() .as_secs() as i64; let cutoff = now.saturating_sub(window_secs as i64); let mut windows = self .windows .lock() .map_err(|e| anyhow::anyhow!("rate limiter lock poisoned: {e}"))?; let timestamps = windows .entry(client_id.to_string()) .or_insert_with(Vec::new); // Discard entries older than the window. timestamps.retain(|&ts| ts >= cutoff); if timestamps.len() >= max_requests as usize { return Ok(false); } timestamps.push(now); Ok(true) } /// Convenience wrapper that returns an Axum [`Response`] on rejection /// so it can be used directly in middleware. pub fn check_or_429( &self, client_id: &str, max_requests: u32, window_secs: u64, ) -> Result<(), Box> { match self.check_rate_limit(client_id, max_requests, window_secs) { Ok(true) => Ok(()), Ok(false) => Err(Box::new( ( StatusCode::TOO_MANY_REQUESTS, "rate limit exceeded, try again later", ) .into_response(), )), Err(e) => Err(Box::new( (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()).into_response(), )), } } /// Remove all stored windows (for testing / reset). pub fn reset(&self) -> anyhow::Result<()> { let mut windows = self .windows .lock() .map_err(|e| anyhow::anyhow!("rate limiter lock poisoned: {e}"))?; windows.clear(); Ok(()) } } impl Default for RateLimiter { fn default() -> Self { Self::new() } } // --------------------------------------------------------------------------- // Tower Layer / Service // --------------------------------------------------------------------------- /// Configuration for the rate-limit middleware layer. #[derive(Debug, Clone)] pub struct RateLimitLayer { limiter: std::sync::Arc, max_requests: u32, window_secs: u64, } impl RateLimitLayer { /// Create a new layer with the given limits, keying strictly on the /// real connection socket address (default, safe when not behind a /// trusted proxy). /// /// * `max_requests` — max requests per window per client. /// * `window_secs` — sliding-window width in seconds. pub fn new(max_requests: u32, window_secs: u64) -> Self { Self::with_proxy_trust(max_requests, window_secs, false) } /// Create a new layer with an explicit proxy-header trust policy. /// /// When `trust_proxy_headers` is `true`, the `X-Forwarded-For` and /// `X-Real-IP` headers are used to derive the client bucket key. /// This must only be enabled when the middleware sits behind a /// reverse proxy known to overwrite (not merge) these headers. pub fn with_proxy_trust( max_requests: u32, window_secs: u64, trust_proxy_headers: bool, ) -> Self { Self { limiter: std::sync::Arc::new(RateLimiter::with_proxy_trust(trust_proxy_headers)), max_requests, window_secs, } } /// Return a reference to the shared [`RateLimiter`] so callers can /// reset it or perform manual checks. pub fn limiter(&self) -> &std::sync::Arc { &self.limiter } } impl Layer for RateLimitLayer { type Service = RateLimitMiddleware; fn layer(&self, inner: S) -> Self::Service { RateLimitMiddleware { inner, limiter: std::sync::Arc::clone(&self.limiter), max_requests: self.max_requests, window_secs: self.window_secs, trust_proxy_headers: self.limiter.trust_proxy_headers, } } } /// Derive the per-client rate-limit bucket key for a request. /// /// Flow: if `trust_proxy_headers` is true, use `X-Forwarded-For` (first /// hop) then `X-Real-IP`; otherwise always use the real connection /// socket address, ignoring any client-supplied headers. /// /// Why: without a trusted reverse proxy stripping/overwriting these /// headers, they are attacker-controlled — trusting them by default lets /// any direct caller reset their own rate-limit bucket on every request. /// `trust_proxy_headers` must only be set to `true` when this middleware /// sits behind a proxy that is known to overwrite (not merge) these headers. fn client_id( headers: &axum::http::HeaderMap, socket_addr: std::net::SocketAddr, trust_proxy_headers: bool, ) -> String { if trust_proxy_headers { if let Some(fwd) = headers .get("x-forwarded-for") .and_then(|v| v.to_str().ok()) .and_then(|v| v.split(',').next()) .map(str::trim) { if !fwd.is_empty() { return fwd.to_string(); } } if let Some(real_ip) = headers.get("x-real-ip").and_then(|v| v.to_str().ok()) { if !real_ip.is_empty() { return real_ip.to_string(); } } } socket_addr.ip().to_string() } /// Tower [`Service`] wrapping each request with a rate-limit check. /// /// Client identity is extracted from the real connection socket address /// by default (safe). When `trust_proxy_headers` is `true`, /// `X-Forwarded-For`/`X-Real-IP` headers are also considered — only /// enable this behind a trusted reverse proxy. #[derive(Debug, Clone)] pub struct RateLimitMiddleware { inner: S, limiter: std::sync::Arc, max_requests: u32, window_secs: u64, trust_proxy_headers: bool, } impl Service> for RateLimitMiddleware where S: Service, Response = Response> + Send + 'static, S::Future: Send + 'static, ReqBody: Send + 'static, { type Response = S::Response; type Error = S::Error; type Future = Pin> + Send + 'static>>; fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll> { self.inner.poll_ready(cx) } fn call(&mut self, req: Request) -> Self::Future { let client_id = req .extensions() .get::>() .map(|ci| client_id(req.headers(), ci.0, self.trust_proxy_headers)) .unwrap_or_else(|| "unknown".to_string()); let limiter = std::sync::Arc::clone(&self.limiter); let max_requests = self.max_requests; let window_secs = self.window_secs; match limiter.check_or_429(&client_id, max_requests, window_secs) { Ok(()) => {} Err(resp) => { return Box::pin(async move { Ok(*resp) }); } } let fut = self.inner.call(req); Box::pin(fut) } } // --------------------------------------------------------------------------- // Tests // --------------------------------------------------------------------------- #[cfg(test)] mod tests { use super::*; #[test] fn test_rate_limiter_allows_within_limit() { let limiter = RateLimiter::new(); assert!(limiter.check_rate_limit("client-1", 5, 60).unwrap()); assert!(limiter.check_rate_limit("client-1", 5, 60).unwrap()); assert!(limiter.check_rate_limit("client-1", 5, 60).unwrap()); } #[test] fn test_rate_limiter_rejects_excess() { let limiter = RateLimiter::new(); assert!(limiter.check_rate_limit("client-2", 3, 60).unwrap()); assert!(limiter.check_rate_limit("client-2", 3, 60).unwrap()); assert!(limiter.check_rate_limit("client-2", 3, 60).unwrap()); assert!(!limiter.check_rate_limit("client-2", 3, 60).unwrap()); } #[test] fn test_rate_limiter_independent_clients() { let limiter = RateLimiter::new(); assert!(limiter.check_rate_limit("alice", 2, 60).unwrap()); assert!(limiter.check_rate_limit("alice", 2, 60).unwrap()); assert!(!limiter.check_rate_limit("alice", 2, 60).unwrap()); assert!(limiter.check_rate_limit("bob", 2, 60).unwrap()); } #[test] fn test_rate_limiter_reset() { let limiter = RateLimiter::new(); assert!(limiter.check_rate_limit("client-3", 1, 60).unwrap()); assert!(!limiter.check_rate_limit("client-3", 1, 60).unwrap()); limiter.reset().unwrap(); assert!(limiter.check_rate_limit("client-3", 1, 60).unwrap()); } #[test] fn client_id_ignores_spoofed_forwarded_headers_by_default() { // A request carrying a spoofed X-Forwarded-For must NOT be treated // as a distinct client from one with a different spoofed value — // both should resolve to the same real socket address. let socket_addr: std::net::SocketAddr = "127.0.0.1:9999".parse().unwrap(); let mut headers_a = axum::http::HeaderMap::new(); headers_a.insert("x-forwarded-for", "1.2.3.4".parse().unwrap()); let mut headers_b = axum::http::HeaderMap::new(); headers_b.insert("x-forwarded-for", "5.6.7.8".parse().unwrap()); let id_a = client_id(&headers_a, socket_addr, false); let id_b = client_id(&headers_b, socket_addr, false); assert_eq!( id_a, id_b, "client_id must key on the real socket address when trust_proxy_headers is false, \ not on attacker-controlled X-Forwarded-For" ); } #[test] fn client_id_uses_forwarded_header_when_trust_enabled() { // When explicitly told to trust a fronting proxy, the header value // should be used (this is the opt-in, documented-risk path). let socket_addr: std::net::SocketAddr = "127.0.0.1:9999".parse().unwrap(); let mut headers = axum::http::HeaderMap::new(); headers.insert("x-forwarded-for", "1.2.3.4".parse().unwrap()); let id = client_id(&headers, socket_addr, true); assert_eq!(id, "1.2.3.4"); } }