openkind_api/middleware/
rate_limit.rs1use 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#[derive(Debug, Clone)]
18pub struct RateLimitConfig {
19 pub max_requests: u32,
21 pub window: std::time::Duration,
23}
24
25impl Default for RateLimitConfig {
26 fn default() -> Self {
27 Self {
29 max_requests: 120,
30 window: std::time::Duration::from_secs(60),
31 }
32 }
33}
34
35pub(crate) const RATE_LIMIT_SWEEP_THRESHOLD: usize = 4096;
38
39#[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#[derive(Debug, Clone)]
50pub struct RequestLimits {
51 pub evaluation: RateLimiter,
53 pub failed_auth: RateLimiter,
55}
56
57impl RequestLimits {
58 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#[doc(hidden)]
85#[derive(Debug, Clone)]
86pub struct RateLimitContext {
87 limiter: RateLimiter,
88 ip: std::net::IpAddr,
89}
90
91impl RateLimitContext {
92 pub(crate) fn charge(&self, units: u32) -> Result<(), u64> {
94 self.limiter.check_n(self.ip, units)
95 }
96}
97
98impl RateLimiter {
99 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 pub fn disabled() -> Self {
109 Self::new(RateLimitConfig {
110 max_requests: 0,
111 window: std::time::Duration::from_secs(60),
112 })
113 }
114
115 pub fn is_enabled(&self) -> bool {
117 self.config.max_requests > 0
118 }
119
120 pub(crate) fn check(&self, ip: std::net::IpAddr) -> Result<(), u64> {
123 self.check_n(ip, 1)
124 }
125
126 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
158pub 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}