1use crate::domain::authentication::session_user_id;
2use actix_session::SessionExt;
3use actix_web::{
4 Error, HttpResponse,
5 body::{EitherBody, MessageBody},
6 dev::{Service, ServiceRequest, ServiceResponse, Transform},
7 http::{StatusCode, header},
8};
9use futures_util::future::{LocalBoxFuture, Ready, ready};
10use governor::{
11 Quota, RateLimiter,
12 clock::{Clock, DefaultClock},
13 state::keyed::DefaultKeyedStateStore,
14};
15use std::{
16 num::NonZeroU32,
17 sync::{
18 Arc,
19 atomic::{AtomicU64, Ordering},
20 },
21 task::{Context, Poll},
22 time::Duration,
23};
24
25#[derive(Clone, Debug, Default)]
26pub struct RateLimitConfig {
27 pub per_second: Option<u64>,
28 pub per_minute: Option<u64>,
29 pub per_hour: Option<u64>,
30 pub per_day: Option<u64>,
31 pub per_month: Option<u64>,
32}
33
34#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
36pub enum RateLimitKey {
37 #[default]
38 ClientIp,
39 User,
42}
43
44type Key = String;
45type Limiter = RateLimiter<Key, DefaultKeyedStateStore<Key>, DefaultClock>;
46
47#[derive(Clone)]
48struct EndpointLimiters {
49 month: Option<Arc<Limiter>>,
50 day: Option<Arc<Limiter>>,
51 hour: Option<Arc<Limiter>>,
52 minute: Option<Arc<Limiter>>,
53 second: Option<Arc<Limiter>>,
54}
55
56impl EndpointLimiters {
57 fn from_config(cfg: &RateLimitConfig) -> Self {
58 Self {
59 month: cfg.per_month.and_then(|n| {
60 build_custom_period_limiter(n, Duration::from_secs(30 * 24 * 60 * 60))
61 }),
62 day: cfg
63 .per_day
64 .and_then(|n| build_custom_period_limiter(n, Duration::from_secs(24 * 60 * 60))),
65 hour: cfg.per_hour.and_then(|n| build_limiter(n, Quota::per_hour)),
66 minute: cfg
67 .per_minute
68 .and_then(|n| build_limiter(n, Quota::per_minute)),
69 second: cfg
70 .per_second
71 .and_then(|n| build_limiter(n, Quota::per_second)),
72 }
73 }
74
75 fn iter(&self) -> impl Iterator<Item = &Arc<Limiter>> {
76 self.month
77 .iter()
78 .chain(self.day.iter())
79 .chain(self.hour.iter())
80 .chain(self.minute.iter())
81 .chain(self.second.iter())
82 }
83
84 fn is_empty(&self) -> bool {
85 self.second.is_none()
86 && self.minute.is_none()
87 && self.hour.is_none()
88 && self.day.is_none()
89 && self.month.is_none()
90 }
91}
92
93fn build_limiter<F>(n: u64, quota_fn: F) -> Option<Arc<Limiter>>
94where
95 F: FnOnce(NonZeroU32) -> Quota,
96{
97 let n32 = NonZeroU32::new(u32::try_from(n).ok()?)?;
98 Some(Arc::new(RateLimiter::keyed(quota_fn(n32))))
99}
100
101fn build_custom_period_limiter(n: u64, period: Duration) -> Option<Arc<Limiter>> {
102 let n32 = NonZeroU32::new(u32::try_from(n).ok()?)?;
103 let quota = Quota::with_period(period)?.allow_burst(n32);
104 Some(Arc::new(RateLimiter::keyed(quota)))
105}
106
107#[derive(Clone)]
108pub struct RateLimit {
109 limiters: Arc<EndpointLimiters>,
110 key: RateLimitKey,
111 calls: Arc<AtomicU64>,
112}
113
114impl RateLimit {
115 pub fn global_api_rate_limit_config(test_mode: bool) -> RateLimitConfig {
117 if test_mode {
118 RateLimitConfig {
119 per_second: Some(10000),
120 per_minute: Some(200000),
121 ..Default::default()
122 }
123 } else {
124 RateLimitConfig {
125 per_second: Some(20),
126 per_minute: Some(1000),
127 per_hour: Some(10000),
128 ..Default::default()
129 }
130 }
131 }
132
133 pub fn new(cfg: RateLimitConfig) -> Self {
134 Self {
135 limiters: Arc::new(EndpointLimiters::from_config(&cfg)),
136 key: RateLimitKey::default(),
137 calls: Arc::new(AtomicU64::new(0)),
138 }
139 }
140
141 pub fn keyed_by(self, key: RateLimitKey) -> Self {
143 Self { key, ..self }
144 }
145}
146
147impl<S, B> Transform<S, ServiceRequest> for RateLimit
148where
149 S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
150 B: MessageBody + 'static,
151{
152 type Response = ServiceResponse<EitherBody<B>>;
153 type Error = Error;
154 type InitError = ();
155 type Transform = RateLimitInner<S>;
156 type Future = Ready<Result<Self::Transform, Self::InitError>>;
157
158 fn new_transform(&self, service: S) -> Self::Future {
159 ready(Ok(RateLimitInner {
160 service,
161 limiters: self.limiters.clone(),
162 key: self.key,
163 calls: self.calls.clone(),
164 }))
165 }
166}
167
168pub struct RateLimitInner<S> {
169 service: S,
170 limiters: Arc<EndpointLimiters>,
171 key: RateLimitKey,
172 calls: Arc<AtomicU64>,
173}
174
175impl<S, B> Service<ServiceRequest> for RateLimitInner<S>
176where
177 S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
178 B: MessageBody + 'static,
179{
180 type Response = ServiceResponse<EitherBody<B>>;
181 type Error = Error;
182 type Future = LocalBoxFuture<'static, Result<Self::Response, Self::Error>>;
183
184 fn poll_ready(&self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
185 self.service.poll_ready(cx)
186 }
187
188 fn call(&self, req: ServiceRequest) -> Self::Future {
189 if self.limiters.is_empty() {
190 let fut = self.service.call(req);
191 return Box::pin(async move { fut.await.map(|r| r.map_into_left_body()) });
192 }
193
194 const RETAIN_EVERY: u64 = 1024;
195 const SHRINK_EVERY: u64 = 65_536;
196
197 let n = self.calls.fetch_add(1, Ordering::Relaxed) + 1;
198 if n.is_multiple_of(RETAIN_EVERY) {
199 for limiter in self.limiters.iter() {
200 limiter.retain_recent();
201 }
202 if n.is_multiple_of(SHRINK_EVERY) {
203 for limiter in self.limiters.iter() {
204 limiter.shrink_to_fit();
205 }
206 }
207 }
208
209 let clock = DefaultClock::default();
210 let key = match self.key {
211 RateLimitKey::User => session_user_id(&req.get_session())
212 .map(|id| format!("user:{id}"))
213 .unwrap_or_else(|| extract_client_ip_key(&req)),
214 RateLimitKey::ClientIp => extract_client_ip_key(&req),
215 };
216
217 let mut retry_after: Option<Duration> = None;
218 for limiter in self.limiters.iter() {
219 if let Err(negative) = limiter.check_key(&key) {
220 let wait = negative.wait_time_from(clock.now());
221 retry_after = Some(retry_after.map_or(wait, |cur| cur.max(wait)));
222 }
223 }
224
225 if let Some(wait) = retry_after {
226 let secs = wait.as_secs().max(1);
227 let resp = HttpResponse::build(StatusCode::TOO_MANY_REQUESTS)
228 .insert_header((header::RETRY_AFTER, secs.to_string()))
229 .content_type("application/json")
230 .body(r#"{"type":"rate_limit","message_key":"rate_limited","message":"Too many requests. Please try again later."}"#.to_string());
231 return Box::pin(async move { Ok(req.into_response(resp).map_into_right_body()) });
232 }
233
234 let fut = self.service.call(req);
235 Box::pin(async move { fut.await.map(|r| r.map_into_left_body()) })
236 }
237}
238
239fn extract_client_ip_key(req: &ServiceRequest) -> String {
240 if let Some(s) = req.connection_info().realip_remote_addr() {
241 let s = s.trim();
242 if !s.is_empty() {
243 return s.to_string();
244 }
245 }
246
247 if let Some(sa) = req.peer_addr() {
248 return sa.ip().to_string();
249 }
250
251 format!("unknown:{}|{}", req.connection_info().host(), req.path())
252}
253
254#[cfg(test)]
255mod tests {
256 use super::*;
257 use actix_http::Request;
258 use actix_web::{
259 App, HttpResponse,
260 body::{BoxBody, EitherBody},
261 dev::{Service, ServiceResponse},
262 http::header,
263 test, web,
264 };
265 use std::net::{IpAddr, Ipv4Addr, SocketAddr};
266
267 fn mw(
268 per_minute: Option<u64>,
269 per_hour: Option<u64>,
270 per_day: Option<u64>,
271 per_month: Option<u64>,
272 ) -> RateLimit {
273 RateLimit::new(RateLimitConfig {
274 per_minute,
275 per_hour,
276 per_day,
277 per_month,
278 ..Default::default()
279 })
280 }
281
282 #[actix_web::test]
283 async fn global_api_rate_limit_config_uses_test_mode_argument() {
284 let test_cfg = RateLimit::global_api_rate_limit_config(true);
285 assert_eq!(test_cfg.per_second, Some(10000));
286 assert_eq!(test_cfg.per_minute, Some(200000));
287 assert_eq!(test_cfg.per_hour, None);
288
289 let production_cfg = RateLimit::global_api_rate_limit_config(false);
290 assert_eq!(production_cfg.per_second, Some(20));
291 assert_eq!(production_cfg.per_minute, Some(1000));
292 assert_eq!(production_cfg.per_hour, Some(10000));
293 }
294
295 async fn call_get<S>(
296 app: &S,
297 uri: &str,
298 xff: Option<&str>,
299 peer: Option<SocketAddr>,
300 ) -> ServiceResponse<EitherBody<BoxBody>>
301 where
302 S: Service<Request, Response = ServiceResponse<EitherBody<BoxBody>>, Error = Error>,
303 {
304 let mut tr = test::TestRequest::get().uri(uri);
305 if let Some(v) = xff {
306 tr = tr.insert_header(("x-forwarded-for", v));
307 }
308 if let Some(p) = peer {
309 tr = tr.peer_addr(p);
310 }
311 test::call_service(app, tr.to_request()).await
312 }
313
314 fn retry_after_secs(resp: &ServiceResponse<EitherBody<BoxBody>>) -> Option<u64> {
315 resp.headers()
316 .get(header::RETRY_AFTER)
317 .and_then(|v| v.to_str().ok())
318 .and_then(|s| s.parse::<u64>().ok())
319 }
320
321 async fn app(
322 mw: RateLimit,
323 ) -> impl Service<Request, Response = ServiceResponse<EitherBody<BoxBody>>, Error = Error> {
324 test::init_service(
325 App::new()
326 .wrap(mw)
327 .route("/", web::get().to(|| async { HttpResponse::Ok().finish() }))
328 .route(
329 "/other",
330 web::get().to(|| async { HttpResponse::Ok().finish() }),
331 ),
332 )
333 .await
334 }
335
336 #[actix_web::test]
337 async fn key_trims_realip_value() {
338 let req = test::TestRequest::get()
339 .uri("/x")
340 .insert_header(("x-forwarded-for", " 9.9.9.9 "))
341 .to_srv_request();
342
343 assert_eq!(super::extract_client_ip_key(&req), "9.9.9.9");
344 }
345
346 #[actix_web::test]
347 async fn key_empty_realip_falls_back_to_peer_ip() {
348 let peer = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(1, 2, 3, 4)), 5555);
349 let req = test::TestRequest::get()
350 .uri("/x")
351 .insert_header(("x-forwarded-for", " "))
352 .peer_addr(peer)
353 .to_srv_request();
354
355 assert_eq!(super::extract_client_ip_key(&req), "1.2.3.4");
356 }
357
358 #[actix_web::test]
359 async fn key_no_realip_uses_peer_ip() {
360 let peer = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(10, 0, 0, 9)), 1234);
361 let req = test::TestRequest::get()
362 .uri("/x")
363 .peer_addr(peer)
364 .to_srv_request();
365 assert_eq!(super::extract_client_ip_key(&req), "10.0.0.9");
366 }
367
368 #[actix_web::test]
369 async fn key_no_realip_and_no_peer_uses_unknown_host_and_path() {
370 let req = test::TestRequest::get().uri("/path123").to_srv_request();
371 let k = super::extract_client_ip_key(&req);
372 assert!(k.starts_with("unknown:"), "key={k}");
373 assert!(k.contains("|/path123"), "key={k}");
374 }
375
376 #[actix_web::test]
377 async fn passthrough_when_no_limits() {
378 let app = app(mw(None, None, None, None)).await;
379
380 let r1 = call_get(&app, "/", Some("1.2.3.4"), None).await;
381 let r2 = call_get(&app, "/", Some("1.2.3.4"), None).await;
382
383 assert_eq!(r1.status(), StatusCode::OK);
384 assert_eq!(r2.status(), StatusCode::OK);
385 assert!(r2.headers().get(header::RETRY_AFTER).is_none());
386 }
387
388 #[actix_web::test]
389 async fn per_minute_zero_disables_window() {
390 let app = app(mw(Some(0), None, None, None)).await;
391
392 let r1 = call_get(&app, "/", Some("1.2.3.4"), None).await;
393 let r2 = call_get(&app, "/", Some("1.2.3.4"), None).await;
394
395 assert_eq!(r1.status(), StatusCode::OK);
396 assert_eq!(r2.status(), StatusCode::OK);
397 }
398
399 #[actix_web::test]
400 async fn per_minute_over_u32_disables_window() {
401 let app = app(mw(Some(u64::from(u32::MAX) + 1), None, None, None)).await;
402
403 let r1 = call_get(&app, "/", Some("1.2.3.4"), None).await;
404 let r2 = call_get(&app, "/", Some("1.2.3.4"), None).await;
405
406 assert_eq!(r1.status(), StatusCode::OK);
407 assert_eq!(r2.status(), StatusCode::OK);
408 }
409
410 #[actix_web::test]
411 async fn blocks_second_request_same_key() {
412 let app = app(mw(Some(1), None, None, None)).await;
413
414 let ok = call_get(&app, "/", Some("1.2.3.4"), None).await;
415 assert_eq!(ok.status(), StatusCode::OK);
416
417 let blocked = call_get(&app, "/", Some("1.2.3.4"), None).await;
418 assert_eq!(blocked.status(), StatusCode::TOO_MANY_REQUESTS);
419
420 let ra = retry_after_secs(&blocked).expect("missing Retry-After");
421 assert!(ra >= 1);
422 }
423
424 #[actix_web::test]
425 async fn retry_after_is_integer_seconds_and_body_is_json() {
426 let app = app(mw(Some(1), None, None, None)).await;
427
428 let _ = call_get(&app, "/", Some("1.2.3.4"), None).await;
429 let blocked = call_get(&app, "/", Some("1.2.3.4"), None).await;
430
431 assert_eq!(blocked.status(), StatusCode::TOO_MANY_REQUESTS);
432
433 let ra_hdr = blocked.headers().get(header::RETRY_AFTER).unwrap();
434 let ra_str = ra_hdr.to_str().unwrap();
435 assert!(
436 ra_str.parse::<u64>().is_ok(),
437 "Retry-After not int: {ra_str}"
438 );
439
440 let bytes = test::read_body(blocked).await;
441 let body = std::str::from_utf8(&bytes).unwrap();
442 assert!(body.contains(r#""type":"rate_limit""#), "body={body}");
443 let v: serde_json::Value = serde_json::from_str(body).unwrap();
444 assert_eq!(v["type"], "rate_limit");
445 assert_eq!(v["message_key"], "rate_limited");
446 assert!(v.get("retry_after").is_none());
447 }
448
449 #[actix_web::test]
450 async fn different_keys_independent() {
451 let app = app(mw(Some(1), None, None, None)).await;
452
453 let a1 = call_get(&app, "/", Some("10.0.0.1"), None).await;
454 let b1 = call_get(&app, "/", Some("10.0.0.2"), None).await;
455 assert_eq!(a1.status(), StatusCode::OK);
456 assert_eq!(b1.status(), StatusCode::OK);
457
458 let a2 = call_get(&app, "/", Some("10.0.0.1"), None).await;
459 assert_eq!(a2.status(), StatusCode::TOO_MANY_REQUESTS);
460 }
461
462 #[actix_web::test]
463 async fn same_key_shared_across_routes_in_same_app() {
464 let app = app(mw(Some(1), None, None, None)).await;
465
466 let r1 = call_get(&app, "/", Some("1.2.3.4"), None).await;
467 assert_eq!(r1.status(), StatusCode::OK);
468
469 let r2 = call_get(&app, "/other", Some("1.2.3.4"), None).await;
470 assert_eq!(r2.status(), StatusCode::TOO_MANY_REQUESTS);
471 }
472
473 #[actix_web::test]
474 async fn max_retry_after_prefers_longer_window() {
475 let app = app(mw(Some(1), Some(1), None, None)).await;
476
477 let r1 = call_get(&app, "/", Some("1.2.3.4"), None).await;
478 assert_eq!(r1.status(), StatusCode::OK);
479
480 let r2 = call_get(&app, "/", Some("1.2.3.4"), None).await;
481 assert_eq!(r2.status(), StatusCode::TOO_MANY_REQUESTS);
482
483 let ra = retry_after_secs(&r2).expect("missing Retry-After");
484 assert!(ra >= 120, "expected hour-dominated Retry-After, got {ra}");
486 }
487
488 #[actix_web::test]
489 async fn peer_addr_used_when_no_forwarded_headers() {
490 let app = app(mw(Some(1), None, None, None)).await;
491
492 let peer = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(7, 7, 7, 7)), 9999);
493 let r1 = call_get(&app, "/", None, Some(peer)).await;
494 let r2 = call_get(&app, "/", None, Some(peer)).await;
495
496 assert_eq!(r1.status(), StatusCode::OK);
497 assert_eq!(r2.status(), StatusCode::TOO_MANY_REQUESTS);
498 }
499
500 #[actix_web::test]
501 async fn empty_forwarded_header_does_not_create_empty_key_bucket() {
502 let app = app(mw(Some(1), None, None, None)).await;
503
504 let peer = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8)), 1111);
505 let r1 = call_get(&app, "/", Some(" "), Some(peer)).await;
506 let r2 = call_get(&app, "/", Some(" "), Some(peer)).await;
507
508 assert_eq!(r1.status(), StatusCode::OK);
509 assert_eq!(r2.status(), StatusCode::TOO_MANY_REQUESTS);
510 }
511
512 #[actix_web::test]
513 async fn unknown_bucket_includes_path_to_reduce_collisions() {
514 let app = app(mw(Some(1), None, None, None)).await;
515
516 let r1 = call_get(&app, "/", None, None).await;
517 let r2 = call_get(&app, "/other", None, None).await;
518
519 assert_eq!(r1.status(), StatusCode::OK);
520 assert_eq!(r2.status(), StatusCode::OK);
521
522 let r1b = call_get(&app, "/", None, None).await;
523 assert_eq!(r1b.status(), StatusCode::TOO_MANY_REQUESTS);
524 }
525
526 #[actix_web::test]
527 async fn housekeeping_retain_recent_path_executes() {
528 let app = app(mw(Some(1_000_000), None, None, None)).await;
530
531 for i in 0..=1024 {
532 let ip = format!("192.0.2.{}", (i % 250) + 1);
533 let resp = call_get(&app, "/", Some(&ip), None).await;
534 assert_eq!(resp.status(), StatusCode::OK, "i={i} ip={ip}");
535 }
536 }
537
538 #[actix_web::test]
539 async fn keys_are_exact_strings_no_normalization_means_ports_are_distinct() {
540 let app = app(mw(Some(1), None, None, None)).await;
541
542 let a = call_get(&app, "/", Some("1.2.3.4"), None).await;
543 let b = call_get(&app, "/", Some("1.2.3.4:12345"), None).await;
544
545 assert_eq!(a.status(), StatusCode::OK);
546 assert_eq!(b.status(), StatusCode::OK);
547
548 let a2 = call_get(&app, "/", Some("1.2.3.4"), None).await;
549 let b2 = call_get(&app, "/", Some("1.2.3.4:12345"), None).await;
550
551 assert_eq!(a2.status(), StatusCode::TOO_MANY_REQUESTS);
552 assert_eq!(b2.status(), StatusCode::TOO_MANY_REQUESTS);
553 }
554}