Skip to main content

openkind_api/middleware/
rate_limit.rs

1//! Fixed-window per-client-IP rate limiting middleware and state tracking.
2
3use std::sync::Arc;
4
5use axum::{
6    body::Body,
7    extract::State,
8    http::Request,
9    middleware::Next,
10    response::{IntoResponse, Response},
11};
12
13/// Fixed-window per-client-IP rate limit configuration.
14///
15/// `max_requests` of `0` disables limiting. The window is a plain fixed
16/// window (not sliding): counters reset every `window` interval per IP.
17#[derive(Debug, Clone)]
18pub struct RateLimitConfig {
19    /// Maximum requests allowed per client IP within `window`. `0` disables rate limiting.
20    pub max_requests: u32,
21    /// Length of the counting window.
22    pub window: std::time::Duration,
23}
24
25impl Default for RateLimitConfig {
26    fn default() -> Self {
27        // Generous for SDK clients, but caps runaway loops and brute force.
28        Self {
29            max_requests: 120,
30            window: std::time::Duration::from_secs(60),
31        }
32    }
33}
34
35/// Sweep the bucket map once it grows past this many entries so a large,
36/// rotating client population cannot grow state without bound.
37pub(crate) const RATE_LIMIT_SWEEP_THRESHOLD: usize = 4096;
38
39/// Shared fixed-window counter state for [`rate_limit_layer`].
40#[derive(Debug, Clone)]
41pub struct RateLimiter {
42    config: RateLimitConfig,
43    buckets: Arc<
44        std::sync::Mutex<std::collections::HashMap<std::net::IpAddr, (u32, std::time::Instant)>>,
45    >,
46}
47
48/// Per-request handle used by bulk handlers to charge work beyond the one
49/// unit already recorded by [`rate_limit_layer`].
50#[doc(hidden)]
51#[derive(Debug, Clone)]
52pub struct RateLimitContext {
53    limiter: RateLimiter,
54    ip: std::net::IpAddr,
55}
56
57impl RateLimitContext {
58    /// Atomically charge additional units to the current client's window.
59    pub(crate) fn charge(&self, units: u32) -> Result<(), u64> {
60        self.limiter.check_n(self.ip, units)
61    }
62}
63
64impl RateLimiter {
65    /// Construct a `RateLimiter` with the given configuration.
66    pub fn new(config: RateLimitConfig) -> Self {
67        Self {
68            config,
69            buckets: Arc::new(std::sync::Mutex::new(std::collections::HashMap::default())),
70        }
71    }
72
73    /// Construct a disabled `RateLimiter` (all requests pass).
74    pub fn disabled() -> Self {
75        Self::new(RateLimitConfig {
76            max_requests: 0,
77            window: std::time::Duration::from_secs(60),
78        })
79    }
80
81    /// Whether this limiter enforces anything.
82    pub fn is_enabled(&self) -> bool {
83        self.config.max_requests > 0
84    }
85
86    /// Record one request for `ip`. Returns `Ok(())` when under the limit,
87    /// or `Err(retry_after_ms)` when the client has exhausted its window.
88    fn check(&self, ip: std::net::IpAddr) -> Result<(), u64> {
89        self.check_n(ip, 1)
90    }
91
92    /// Atomically record `units` for `ip` without partially spending a
93    /// rejected charge.
94    fn check_n(&self, ip: std::net::IpAddr, units: u32) -> Result<(), u64> {
95        if !self.is_enabled() {
96            return Ok(());
97        }
98        let mut buckets = self
99            .buckets
100            .lock()
101            .unwrap_or_else(|poisoned| poisoned.into_inner());
102        let now = std::time::Instant::now();
103        if buckets.len() >= RATE_LIMIT_SWEEP_THRESHOLD {
104            buckets.retain(|_, (_, start)| now.duration_since(*start) < self.config.window);
105        }
106        let window = self.config.window;
107        let entry = buckets.entry(ip).or_insert((0, now));
108        if now.duration_since(entry.1) >= window {
109            *entry = (0, now);
110        }
111        if units > self.config.max_requests.saturating_sub(entry.0) {
112            let elapsed = now.duration_since(entry.1);
113            let remaining_ms = window
114                .saturating_sub(elapsed)
115                .as_millis()
116                .min(u64::MAX as u128) as u64;
117            return Err(remaining_ms.max(1));
118        }
119        entry.0 += units;
120        Ok(())
121    }
122}
123
124/// Stackable middleware function: fixed-window rate limit on `/v1/*` per
125/// client IP (taken from the `ConnectInfo` extension, which `axum::serve`
126/// provides when the router is served via
127/// `into_make_service_with_connect_info`). Requests without connect info
128/// (unit tests, or any future non-TCP transport such as a Unix socket or a
129/// Windows named pipe) are passed through — per-IP limiting is only
130/// enforceable when the peer address is known.
131pub async fn rate_limit_layer(
132    State(limiter): State<RateLimiter>,
133    req: Request<Body>,
134    next: Next,
135) -> Response {
136    if !limiter.is_enabled()
137        || !req.uri().path().starts_with("/v1/")
138        || req.method() == axum::http::Method::OPTIONS
139    {
140        return next.run(req).await;
141    }
142    let peer_ip = req
143        .extensions()
144        .get::<axum::extract::ConnectInfo<std::net::SocketAddr>>()
145        .map(|c| c.0.ip());
146    match peer_ip {
147        Some(ip) => match limiter.check(ip) {
148            Ok(()) => {
149                let mut req = req;
150                req.extensions_mut().insert(RateLimitContext {
151                    limiter: limiter.clone(),
152                    ip,
153                });
154                next.run(req).await
155            }
156            Err(retry_after_ms) => {
157                tracing::debug!(%ip, retry_after_ms, "rate limited");
158                crate::error::ApiError::RateLimited { retry_after_ms }.into_response()
159            }
160        },
161        None => next.run(req).await,
162    }
163}