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#[doc(hidden)]
51#[derive(Debug, Clone)]
52pub struct RateLimitContext {
53 limiter: RateLimiter,
54 ip: std::net::IpAddr,
55}
56
57impl RateLimitContext {
58 pub(crate) fn charge(&self, units: u32) -> Result<(), u64> {
60 self.limiter.check_n(self.ip, units)
61 }
62}
63
64impl RateLimiter {
65 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 pub fn disabled() -> Self {
75 Self::new(RateLimitConfig {
76 max_requests: 0,
77 window: std::time::Duration::from_secs(60),
78 })
79 }
80
81 pub fn is_enabled(&self) -> bool {
83 self.config.max_requests > 0
84 }
85
86 fn check(&self, ip: std::net::IpAddr) -> Result<(), u64> {
89 self.check_n(ip, 1)
90 }
91
92 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
124pub 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}