Skip to main content

headless_lms_server/domain/
rate_limit_middleware_builder.rs

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/// Whose requests share one quota.
35#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
36pub enum RateLimitKey {
37    #[default]
38    ClientIp,
39    /// The signed-in user, falling back to the client IP for anonymous requests. For endpoints
40    /// that act on the caller's own data, where users behind one NAT address must not share a quota.
41    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    /// Global `/api/v0` limits aligned with nginx ingress `limit-rps` and `limit-rpm`; relaxed when `TEST_MODE` is set.
116    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    /// Shares each quota by `key` instead of by client IP.
142    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        // hour window should dominate minute window; be tolerant but meaningful
485        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        // Trigger retain_recent() at 1024 calls; keep limit huge to avoid 429.
529        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}