runique 2.2.0

A Django-inspired web framework for Rust with ORM, templates, and comprehensive security middleware
Documentation
//! Rate limiter by key (IP or other) with sliding window and 429 response.
use crate::utils::trad::t;
use axum::{
    body::Body,
    extract::State,
    http::{Method, Request, StatusCode, header},
    middleware::Next,
    response::{IntoResponse, Response},
};
use std::{
    collections::HashMap,
    sync::{Arc, Mutex},
    time::{Duration, Instant},
};
use tokio::time::interval;

/// Entry per key: (request count in window, window start)
type Store = Arc<Mutex<HashMap<String, (u32, Instant)>>>;

/// Configurable rate limiter — sliding window per key (IP or other)
#[derive(Clone)]
pub struct RateLimiter {
    store: Store,
    /// Maximum number of requests allowed in the window
    pub max_requests: u32,
    /// Window duration
    pub window: Duration,
    /// If set, only these HTTP methods are counted (others pass through freely)
    pub methods: Option<Vec<Method>>,
}

impl RateLimiter {
    /// Creates a rate limiter with default values (60 Req / 60 s).
    ///
    /// # Example
    /// ```rust,ignore
    /// RateLimiter::new()
    ///     .max_requests(100)
    ///     .retry_after(60)
    /// ```
    pub fn new() -> Self {
        Self {
            store: Arc::new(Mutex::new(HashMap::new())),
            max_requests: 60,
            window: Duration::from_secs(60),
            methods: None,
        }
    }

    /// Restricts rate limiting to the given HTTP methods.
    /// GET requests are never counted if POST is the only method listed.
    #[must_use]
    pub fn only_methods(mut self, methods: Vec<Method>) -> Self {
        self.methods = Some(methods);
        self
    }

    /// Maximum number of requests allowed in the window
    #[must_use]
    pub fn max_requests(mut self, max: u32) -> Self {
        self.max_requests = max;
        self
    }

    /// Window duration in seconds
    #[must_use]
    pub fn retry_after(mut self, secs: u64) -> Self {
        self.window = Duration::from_secs(secs);
        self
    }

    /// Spawns a Tokio task that periodically purges expired entries.
    /// Should be called once at application startup.
    pub fn spawn_cleanup(&self, period: tokio::time::Duration) {
        let store = self.store.clone();
        let window = self.window;
        tokio::spawn(async move {
            let mut ticker = interval(period);
            loop {
                ticker.tick().await;
                let mut guard = match store.lock() {
                    Ok(g) => g,
                    Err(p) => p.into_inner(),
                };
                let now = Instant::now();
                guard.retain(|_, (_, start)| now.duration_since(*start) < window);
            }
        });
    }

    /// Seconds remaining before window reset for this key.
    /// Returns `0` if the key is unknown or if the window is already expired.
    #[must_use]
    pub fn retry_after_secs(&self, key: &str) -> u64 {
        let store = match self.store.lock() {
            Ok(s) => s,
            Err(p) => p.into_inner(),
        };
        match store.get(key) {
            Some((_, start)) => {
                let interval = Instant::now().duration_since(*start);
                self.window.saturating_sub(interval).as_secs()
            }
            None => 0,
        }
    }

    /// Returns `true` if the key is under the limit, `false` if exceeded
    #[must_use]
    pub fn is_allowed(&self, key: &str) -> bool {
        let mut store = match self.store.lock() {
            Ok(s) => s,
            Err(p) => p.into_inner(),
        };
        let now = Instant::now();
        let entry = store.entry(key.to_string()).or_insert((0, now));

        if now.duration_since(entry.1) >= self.window {
            // New window
            *entry = (1, now);
            true
        } else if entry.0 < self.max_requests {
            entry.0 = entry.0.saturating_add(1);
            true
        } else {
            false
        }
    }
}

impl Default for RateLimiter {
    fn default() -> Self {
        Self::new()
    }
}

impl std::fmt::Debug for RateLimiter {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        f.debug_struct("RateLimiter")
            .field("max_requests", &self.max_requests)
            .field("window", &self.window)
            .field("methods", &self.methods)
            .finish()
    }
}

/// Extracts the IP key from headers (`X-Forwarded-For`, `X-Real-IP`, fallback `"unknown"`).
///
/// Fallback client key, used only when the `ClientIp` extension is absent — i.e.
/// the `trusted_proxies` middleware is not in the stack (standalone use of this
/// middleware). Keys on the **real socket peer** from `ConnectInfo`, never on the
/// client-controlled `X-Forwarded-For`, so the fallback can't be spoofed.
///
/// Behind a proxy without `trusted_proxies`, every request shares the proxy's IP
/// (one bucket) — add `trusted_proxies` (which injects `ClientIp`) for per-client
/// limiting. For login brute-force, prefer [`LoginGuard`] (limits by username,
/// never by IP).
fn extract_ip(req: &Request<Body>) -> String {
    req.extensions()
        .get::<axum::extract::ConnectInfo<std::net::SocketAddr>>()
        .map(|ci| ci.0.ip().to_string())
        .unwrap_or_else(|| "unknown".to_string())
}

/// Rate limiting middleware — to be applied on sensitive routes (login, etc.)
///
/// # Example
/// ```rust,ignore
/// use runique::prelude::*;
/// use std::sync::Arc;
///
/// let limiter = Arc::new(RateLimiter::new().max_requests(5).retry_after(60));
///
/// Router::new()
///     .route("/login", post(login_view))
///     .layer(axum::middleware::from_fn_with_state(
///         limiter,
///         rate_limit_middleware,
///     ))
/// ```
pub async fn rate_limit_middleware(
    State(limiter): State<Arc<RateLimiter>>,
    req: Request<Body>,
    next: Next,
) -> Response {
    if let Some(ref methods) = limiter.methods
        && !methods.contains(req.method())
    {
        return next.run(req).await;
    }
    let ip = req
        .extensions()
        .get::<crate::middleware::security::trusted_proxies::ClientIp>()
        .map(|c| c.0.to_string())
        .unwrap_or_else(|| extract_ip(&req));
    if limiter.is_allowed(&ip) {
        next.run(req).await
    } else {
        let retry_after_secs = limiter.retry_after_secs(&ip);
        if let Some(level) = crate::utils::runique_log::get_log()
            .middleware
            .as_ref()
            .and_then(|m| m.rate_limit)
        {
            crate::runique_log!(level, %ip, retry_after = retry_after_secs, "rate limited");
        }
        (
            StatusCode::TOO_MANY_REQUESTS,
            [(header::RETRY_AFTER, retry_after_secs.to_string())],
            t("html.429_text").into_owned(),
        )
            .into_response()
    }
}

#[cfg(test)]
mod tests {
    use super::*;

    #[test]
    fn test_is_allowed_basic() {
        let limiter = RateLimiter::new().max_requests(3).retry_after(60);
        assert!(limiter.is_allowed("ip1"));
        assert!(limiter.is_allowed("ip1"));
        assert!(limiter.is_allowed("ip1"));
        assert!(!limiter.is_allowed("ip1"));
        // different key unaffected
        assert!(limiter.is_allowed("ip2"));
    }

    #[test]
    fn test_only_methods_skips_unlisted() {
        let limiter = RateLimiter::new()
            .max_requests(2)
            .retry_after(60)
            .only_methods(vec![Method::POST]);

        // GET is not counted — can call indefinitely
        for _ in 0..10 {
            // simulate method check as the middleware does
            let method = Method::GET;
            let methods = limiter.methods.as_ref().unwrap();
            assert!(!methods.contains(&method), "GET should not be in the list");
        }

        // POST IS counted
        assert!(limiter.is_allowed("ip1"));
        assert!(limiter.is_allowed("ip1"));
        assert!(!limiter.is_allowed("ip1"));
    }

    #[test]
    fn test_only_methods_counts_listed() {
        let limiter = RateLimiter::new()
            .max_requests(1)
            .retry_after(60)
            .only_methods(vec![Method::POST, Method::PUT]);

        let methods = limiter.methods.as_ref().unwrap();
        assert!(methods.contains(&Method::POST));
        assert!(methods.contains(&Method::PUT));
        assert!(!methods.contains(&Method::GET));
        assert!(!methods.contains(&Method::DELETE));
    }

    #[test]
    fn test_no_methods_filter_counts_all() {
        let limiter = RateLimiter::new().max_requests(1).retry_after(60);
        assert!(limiter.methods.is_none());
        assert!(limiter.is_allowed("ip1"));
        assert!(!limiter.is_allowed("ip1"));
    }

    #[test]
    fn test_retry_after_secs_unknown_key() {
        let limiter = RateLimiter::new().max_requests(5).retry_after(60);
        assert_eq!(limiter.retry_after_secs("unknown"), 0);
    }

    #[test]
    fn test_retry_after_secs_known_key() {
        let limiter = RateLimiter::new().max_requests(1).retry_after(60);
        let _ = limiter.is_allowed("ip1");
        let _ = limiter.is_allowed("ip1"); // exceeds limit
        let remaining = limiter.retry_after_secs("ip1");
        assert!(remaining > 0 && remaining <= 60);
    }
}