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/// Shared transport budgets with an independent failed-authentication allowance.
49#[derive(Debug, Clone)]
50pub struct RequestLimits {
51    /// Budget for authenticated or anonymous API work.
52    pub evaluation: RateLimiter,
53    /// Budget charged only when credentials fail verification.
54    pub failed_auth: RateLimiter,
55}
56
57impl RequestLimits {
58    /// Construct independent budgets with the same configured limit and window.
59    pub fn new(config: RateLimitConfig) -> Self {
60        Self {
61            evaluation: RateLimiter::new(config.clone()),
62            failed_auth: RateLimiter::new(config),
63        }
64    }
65}
66
67impl Default for RequestLimits {
68    fn default() -> Self {
69        Self::new(RateLimitConfig::default())
70    }
71}
72
73impl From<RateLimiter> for RequestLimits {
74    fn from(evaluation: RateLimiter) -> Self {
75        Self {
76            failed_auth: RateLimiter::new(evaluation.config.clone()),
77            evaluation,
78        }
79    }
80}
81
82/// Per-request handle used by bulk handlers to charge work beyond the one
83/// unit already recorded by [`rate_limit_layer`].
84#[doc(hidden)]
85#[derive(Debug, Clone)]
86pub struct RateLimitContext {
87    limiter: RateLimiter,
88    ip: std::net::IpAddr,
89}
90
91impl RateLimitContext {
92    /// Atomically charge additional units to the current client's window.
93    pub(crate) fn charge(&self, units: u32) -> Result<(), u64> {
94        self.limiter.check_n(self.ip, units)
95    }
96}
97
98impl RateLimiter {
99    /// Construct a `RateLimiter` with the given configuration.
100    pub fn new(config: RateLimitConfig) -> Self {
101        Self {
102            config,
103            buckets: Arc::new(std::sync::Mutex::new(std::collections::HashMap::default())),
104        }
105    }
106
107    /// Construct a disabled `RateLimiter` (all requests pass).
108    pub fn disabled() -> Self {
109        Self::new(RateLimitConfig {
110            max_requests: 0,
111            window: std::time::Duration::from_secs(60),
112        })
113    }
114
115    /// Whether this limiter enforces anything.
116    pub fn is_enabled(&self) -> bool {
117        self.config.max_requests > 0
118    }
119
120    /// Record one request for `ip`. Returns `Ok(())` when under the limit,
121    /// or `Err(retry_after_ms)` when the client has exhausted its window.
122    pub(crate) fn check(&self, ip: std::net::IpAddr) -> Result<(), u64> {
123        self.check_n(ip, 1)
124    }
125
126    /// Atomically record `units` for `ip` without partially spending a
127    /// rejected charge.
128    fn check_n(&self, ip: std::net::IpAddr, units: u32) -> Result<(), u64> {
129        if !self.is_enabled() {
130            return Ok(());
131        }
132        let mut buckets = self
133            .buckets
134            .lock()
135            .unwrap_or_else(|poisoned| poisoned.into_inner());
136        let now = std::time::Instant::now();
137        if buckets.len() >= RATE_LIMIT_SWEEP_THRESHOLD {
138            buckets.retain(|_, (_, start)| now.duration_since(*start) < self.config.window);
139        }
140        let window = self.config.window;
141        let entry = buckets.entry(ip).or_insert((0, now));
142        if now.duration_since(entry.1) >= window {
143            *entry = (0, now);
144        }
145        if units > self.config.max_requests.saturating_sub(entry.0) {
146            let elapsed = now.duration_since(entry.1);
147            let remaining_ms = window
148                .saturating_sub(elapsed)
149                .as_millis()
150                .min(u64::MAX as u128) as u64;
151            return Err(remaining_ms.max(1));
152        }
153        entry.0 += units;
154        Ok(())
155    }
156}
157
158/// Stackable middleware function: fixed-window rate limit on `/v1/*` per
159/// client IP (taken from the `ConnectInfo` extension, which `axum::serve`
160/// provides when the router is served via
161/// `into_make_service_with_connect_info`). Requests without connect info
162/// (unit tests, or any future non-TCP transport such as a Unix socket or a
163/// Windows named pipe) are passed through — per-IP limiting is only
164/// enforceable when the peer address is known.
165pub async fn rate_limit_layer(
166    State(limiter): State<RateLimiter>,
167    req: Request<Body>,
168    next: Next,
169) -> Response {
170    if !limiter.is_enabled()
171        || !req.uri().path().starts_with("/v1/")
172        || req.method() == axum::http::Method::OPTIONS
173    {
174        return next.run(req).await;
175    }
176    let peer_ip = req
177        .extensions()
178        .get::<axum::extract::ConnectInfo<std::net::SocketAddr>>()
179        .map(|c| c.0.ip());
180    match peer_ip {
181        Some(ip) => match limiter.check(ip) {
182            Ok(()) => {
183                let mut req = req;
184                req.extensions_mut().insert(RateLimitContext {
185                    limiter: limiter.clone(),
186                    ip,
187                });
188                next.run(req).await
189            }
190            Err(retry_after_ms) => {
191                tracing::debug!(%ip, retry_after_ms, "rate limited");
192                crate::error::ApiError::RateLimited { retry_after_ms }.into_response()
193            }
194        },
195        None => next.run(req).await,
196    }
197}