Skip to main content

rest/
middleware.rs

1use crate::{metrics::HttpMetrics, route::RequestPolicy};
2use actix_web::{
3    body::{EitherBody, MessageBody},
4    dev::{Service, ServiceRequest, ServiceResponse, Transform},
5    http::{header, StatusCode},
6    web::BytesMut,
7    Error, HttpMessage, HttpResponse,
8};
9use futures::future::{ok, LocalBoxFuture, Ready};
10use futures::{Stream, StreamExt};
11use rust_zero_core::{AdaptiveShedder, ShedPermit};
12use std::{
13    io::Read,
14    pin::Pin,
15    rc::Rc,
16    sync::{Arc, Mutex},
17    task::{Context, Poll},
18    time::{Duration, Instant},
19};
20use tokio::sync::Semaphore;
21
22/// Rejects requests that exceed the configured maximum execution time.
23pub struct Timeout {
24    duration: Duration,
25    metrics: Option<HttpMetrics>,
26}
27
28impl Clone for Timeout {
29    fn clone(&self) -> Self {
30        Self {
31            duration: self.duration,
32            metrics: self.metrics.clone(),
33        }
34    }
35}
36
37impl Timeout {
38    pub fn new(duration: Duration) -> Self {
39        assert!(
40            !duration.is_zero(),
41            "timeout duration must be greater than zero"
42        );
43        Self {
44            duration,
45            metrics: None,
46        }
47    }
48
49    pub fn with_metrics(mut self, metrics: HttpMetrics) -> Self {
50        self.metrics = Some(metrics);
51        self
52    }
53}
54
55impl<S, B> Transform<S, ServiceRequest> for Timeout
56where
57    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
58    S::Future: 'static,
59    B: MessageBody + 'static,
60{
61    type Response = ServiceResponse<B>;
62    type Error = Error;
63    type Transform = TimeoutMiddleware<S>;
64    type InitError = ();
65    type Future = Ready<Result<Self::Transform, Self::InitError>>;
66
67    fn new_transform(&self, service: S) -> Self::Future {
68        ok(TimeoutMiddleware {
69            service,
70            duration: self.duration,
71            metrics: self.metrics.clone(),
72        })
73    }
74}
75
76pub struct TimeoutMiddleware<S> {
77    service: S,
78    duration: Duration,
79    metrics: Option<HttpMetrics>,
80}
81
82impl<S, B> Service<ServiceRequest> for TimeoutMiddleware<S>
83where
84    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
85    S::Future: 'static,
86    B: MessageBody + 'static,
87{
88    type Response = ServiceResponse<B>;
89    type Error = Error;
90    type Future = LocalBoxFuture<'static, Result<Self::Response, Self::Error>>;
91
92    fn poll_ready(&self, context: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
93        self.service.poll_ready(context)
94    }
95
96    fn call(&self, request: ServiceRequest) -> Self::Future {
97        let (duration, sse) = request
98            .extensions()
99            .get::<RequestPolicy>()
100            .map(|policy| (policy.timeout.unwrap_or(self.duration), policy.sse))
101            .unwrap_or((self.duration, false));
102        let future = self.service.call(request);
103        let metrics = self.metrics.clone();
104
105        Box::pin(async move {
106            if sse {
107                return future.await;
108            }
109            match actix_rt::time::timeout(duration, future).await {
110                Ok(response) => response,
111                Err(_) => {
112                    if let Some(metrics) = metrics {
113                        metrics.record_protection("timeout", "rejected");
114                    }
115                    Err(actix_web::error::ErrorGatewayTimeout("request timed out"))
116                }
117            }
118        })
119    }
120}
121
122/// Sheds excess load instead of queueing requests when all execution slots are busy.
123pub struct ConcurrencyLimit {
124    semaphore: Arc<Semaphore>,
125    priority_reserve: Arc<Semaphore>,
126    metrics: Option<HttpMetrics>,
127}
128
129impl Clone for ConcurrencyLimit {
130    fn clone(&self) -> Self {
131        Self {
132            semaphore: Arc::clone(&self.semaphore),
133            priority_reserve: Arc::clone(&self.priority_reserve),
134            metrics: self.metrics.clone(),
135        }
136    }
137}
138
139impl ConcurrencyLimit {
140    pub fn new(max_concurrent_requests: usize) -> Self {
141        assert!(
142            max_concurrent_requests > 0,
143            "maximum concurrent requests must be greater than zero"
144        );
145        Self {
146            semaphore: Arc::new(Semaphore::new(max_concurrent_requests)),
147            priority_reserve: Arc::new(Semaphore::new(max_concurrent_requests.div_ceil(4))),
148            metrics: None,
149        }
150    }
151
152    /// Sets the additional capacity reserved exclusively for priority routes.
153    pub fn with_priority_reserve(mut self, capacity: usize) -> Self {
154        assert!(capacity > 0, "priority reserve must be greater than zero");
155        self.priority_reserve = Arc::new(Semaphore::new(capacity));
156        self
157    }
158
159    pub fn with_metrics(mut self, metrics: HttpMetrics) -> Self {
160        self.metrics = Some(metrics);
161        self
162    }
163}
164
165impl<S, B> Transform<S, ServiceRequest> for ConcurrencyLimit
166where
167    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
168    S::Future: 'static,
169    B: MessageBody + 'static,
170{
171    type Response = ServiceResponse<EitherBody<B>>;
172    type Error = Error;
173    type Transform = ConcurrencyLimitMiddleware<S>;
174    type InitError = ();
175    type Future = Ready<Result<Self::Transform, Self::InitError>>;
176
177    fn new_transform(&self, service: S) -> Self::Future {
178        ok(ConcurrencyLimitMiddleware {
179            service,
180            semaphore: Arc::clone(&self.semaphore),
181            priority_reserve: Arc::clone(&self.priority_reserve),
182            metrics: self.metrics.clone(),
183        })
184    }
185}
186
187pub struct ConcurrencyLimitMiddleware<S> {
188    service: S,
189    semaphore: Arc<Semaphore>,
190    priority_reserve: Arc<Semaphore>,
191    metrics: Option<HttpMetrics>,
192}
193
194/// Applies rust-zero's CPU- and throughput-aware admission control to HTTP requests.
195#[derive(Clone)]
196pub struct AdaptiveLoadShed {
197    shedder: AdaptiveShedder,
198    metrics: Option<HttpMetrics>,
199}
200
201impl AdaptiveLoadShed {
202    pub fn new(shedder: AdaptiveShedder) -> Self {
203        Self {
204            shedder,
205            metrics: None,
206        }
207    }
208
209    pub fn with_metrics(mut self, metrics: HttpMetrics) -> Self {
210        self.metrics = Some(metrics);
211        self
212    }
213}
214
215impl<S, B> Transform<S, ServiceRequest> for AdaptiveLoadShed
216where
217    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
218    S::Future: 'static,
219    B: MessageBody + 'static,
220{
221    type Response = ServiceResponse<EitherBody<PermitBody<B>>>;
222    type Error = Error;
223    type Transform = AdaptiveLoadShedMiddleware<S>;
224    type InitError = ();
225    type Future = Ready<Result<Self::Transform, Self::InitError>>;
226
227    fn new_transform(&self, service: S) -> Self::Future {
228        ok(AdaptiveLoadShedMiddleware {
229            service,
230            shedder: self.shedder.clone(),
231            metrics: self.metrics.clone(),
232        })
233    }
234}
235
236pub struct AdaptiveLoadShedMiddleware<S> {
237    service: S,
238    shedder: AdaptiveShedder,
239    metrics: Option<HttpMetrics>,
240}
241
242impl<S, B> Service<ServiceRequest> for AdaptiveLoadShedMiddleware<S>
243where
244    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
245    S::Future: 'static,
246    B: MessageBody + 'static,
247{
248    type Response = ServiceResponse<EitherBody<PermitBody<B>>>;
249    type Error = Error;
250    type Future = LocalBoxFuture<'static, Result<Self::Response, Self::Error>>;
251
252    fn poll_ready(&self, context: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
253        self.service.poll_ready(context)
254    }
255
256    fn call(&self, request: ServiceRequest) -> Self::Future {
257        let Some(permit) = self.shedder.try_acquire() else {
258            if let Some(metrics) = &self.metrics {
259                metrics.record_protection("load_shedder", "rejected");
260            }
261            return Box::pin(async move {
262                Ok(request.into_response(
263                    HttpResponse::build(StatusCode::SERVICE_UNAVAILABLE)
264                        .body("server is overloaded")
265                        .map_into_right_body(),
266                ))
267            });
268        };
269        let future = self.service.call(request);
270        Box::pin(async move {
271            let response = future
272                .await?
273                .map_body(move |_, body| PermitBody::new(body, permit))
274                .map_into_left_body();
275            Ok(response)
276        })
277    }
278}
279
280pub struct PermitBody<B> {
281    inner: Pin<Box<B>>,
282    _permit: ShedPermit,
283}
284
285impl<B> PermitBody<B> {
286    fn new(body: B, permit: ShedPermit) -> Self {
287        Self {
288            inner: Box::pin(body),
289            _permit: permit,
290        }
291    }
292}
293
294impl<B: MessageBody> MessageBody for PermitBody<B> {
295    type Error = B::Error;
296
297    fn size(&self) -> actix_web::body::BodySize {
298        self.inner.as_ref().get_ref().size()
299    }
300
301    fn poll_next(
302        mut self: Pin<&mut Self>,
303        context: &mut Context<'_>,
304    ) -> Poll<Option<Result<actix_web::web::Bytes, Self::Error>>> {
305        self.inner.as_mut().poll_next(context)
306    }
307}
308
309impl<S, B> Service<ServiceRequest> for ConcurrencyLimitMiddleware<S>
310where
311    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
312    S::Future: 'static,
313    B: MessageBody + 'static,
314{
315    type Response = ServiceResponse<EitherBody<B>>;
316    type Error = Error;
317    type Future = LocalBoxFuture<'static, Result<Self::Response, Self::Error>>;
318
319    fn poll_ready(&self, context: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
320        self.service.poll_ready(context)
321    }
322
323    fn call(&self, request: ServiceRequest) -> Self::Future {
324        let priority = request
325            .extensions()
326            .get::<RequestPolicy>()
327            .is_some_and(|policy| policy.priority);
328        let permit = match Arc::clone(&self.semaphore)
329            .try_acquire_owned()
330            .or_else(|error| {
331                if priority {
332                    Arc::clone(&self.priority_reserve).try_acquire_owned()
333                } else {
334                    Err(error)
335                }
336            }) {
337            Ok(permit) => permit,
338            Err(_) => {
339                if let Some(metrics) = &self.metrics {
340                    metrics.record_protection("concurrency", "rejected");
341                }
342                return Box::pin(async move {
343                    Ok(request.into_response(
344                        HttpResponse::build(StatusCode::SERVICE_UNAVAILABLE)
345                            .body("server is overloaded")
346                            .map_into_right_body(),
347                    ))
348                });
349            }
350        };
351
352        let future = self.service.call(request);
353        Box::pin(async move {
354            let response = future.await?.map_into_left_body();
355            drop(permit);
356            Ok(response)
357        })
358    }
359}
360
361/// A token-bucket limiter that can be shared by cloning it into Actix workers.
362pub struct RateLimit {
363    state: Arc<Mutex<TokenBucket>>,
364    permits_per_second: f64,
365    metrics: Option<HttpMetrics>,
366}
367
368impl Clone for RateLimit {
369    fn clone(&self) -> Self {
370        Self {
371            state: Arc::clone(&self.state),
372            permits_per_second: self.permits_per_second,
373            metrics: self.metrics.clone(),
374        }
375    }
376}
377
378impl RateLimit {
379    pub fn new(permits_per_second: u32, burst: u32) -> Self {
380        assert!(
381            permits_per_second > 0,
382            "permits per second must be greater than zero"
383        );
384        assert!(burst > 0, "burst capacity must be greater than zero");
385
386        Self {
387            state: Arc::new(Mutex::new(TokenBucket {
388                available: f64::from(burst),
389                capacity: f64::from(burst),
390                last_refill: Instant::now(),
391            })),
392            permits_per_second: f64::from(permits_per_second),
393            metrics: None,
394        }
395    }
396
397    pub fn with_metrics(mut self, metrics: HttpMetrics) -> Self {
398        self.metrics = Some(metrics);
399        self
400    }
401}
402
403impl<S, B> Transform<S, ServiceRequest> for RateLimit
404where
405    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
406    S::Future: 'static,
407    B: MessageBody + 'static,
408{
409    type Response = ServiceResponse<EitherBody<B>>;
410    type Error = Error;
411    type Transform = RateLimitMiddleware<S>;
412    type InitError = ();
413    type Future = Ready<Result<Self::Transform, Self::InitError>>;
414
415    fn new_transform(&self, service: S) -> Self::Future {
416        ok(RateLimitMiddleware {
417            service,
418            state: Arc::clone(&self.state),
419            permits_per_second: self.permits_per_second,
420            metrics: self.metrics.clone(),
421        })
422    }
423}
424
425pub struct RateLimitMiddleware<S> {
426    service: S,
427    state: Arc<Mutex<TokenBucket>>,
428    permits_per_second: f64,
429    metrics: Option<HttpMetrics>,
430}
431
432impl<S, B> Service<ServiceRequest> for RateLimitMiddleware<S>
433where
434    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
435    S::Future: 'static,
436    B: MessageBody + 'static,
437{
438    type Response = ServiceResponse<EitherBody<B>>;
439    type Error = Error;
440    type Future = LocalBoxFuture<'static, Result<Self::Response, Self::Error>>;
441
442    fn poll_ready(&self, context: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
443        self.service.poll_ready(context)
444    }
445
446    fn call(&self, request: ServiceRequest) -> Self::Future {
447        let retry_after = self
448            .state
449            .lock()
450            .expect("rate limiter state lock poisoned")
451            .try_acquire(self.permits_per_second);
452
453        if let Some(retry_after) = retry_after {
454            if let Some(metrics) = &self.metrics {
455                metrics.record_protection("rate_limit", "rejected");
456            }
457            let retry_after_seconds = retry_after.as_secs_f64().ceil().max(1.0) as u64;
458            return Box::pin(async move {
459                Ok(request.into_response(
460                    HttpResponse::build(StatusCode::TOO_MANY_REQUESTS)
461                        .insert_header((header::RETRY_AFTER, retry_after_seconds.to_string()))
462                        .body("rate limit exceeded")
463                        .map_into_right_body(),
464                ))
465            });
466        }
467
468        let future = self.service.call(request);
469        Box::pin(async move { Ok(future.await?.map_into_left_body()) })
470    }
471}
472
473struct TokenBucket {
474    available: f64,
475    capacity: f64,
476    last_refill: Instant,
477}
478
479/// Buffers request bodies, optionally expands gzip input, and rejects oversized payloads.
480///
481/// Both compressed and expanded data are checked, preventing a small gzip payload from expanding
482/// beyond the configured application limit.
483#[derive(Debug, Clone)]
484pub struct RequestBodyLimit {
485    max_bytes: usize,
486    decompress_gzip: bool,
487}
488
489impl RequestBodyLimit {
490    pub fn new(max_bytes: usize) -> Self {
491        assert!(
492            max_bytes > 0,
493            "request body limit must be greater than zero"
494        );
495        Self {
496            max_bytes,
497            decompress_gzip: true,
498        }
499    }
500
501    pub fn decompress_gzip(mut self, enabled: bool) -> Self {
502        self.decompress_gzip = enabled;
503        self
504    }
505}
506
507impl<S, B> Transform<S, ServiceRequest> for RequestBodyLimit
508where
509    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
510    S::Future: 'static,
511    B: MessageBody + 'static,
512{
513    type Response = ServiceResponse<EitherBody<B>>;
514    type Error = Error;
515    type Transform = RequestBodyLimitMiddleware<S>;
516    type InitError = ();
517    type Future = Ready<Result<Self::Transform, Self::InitError>>;
518
519    fn new_transform(&self, service: S) -> Self::Future {
520        ok(RequestBodyLimitMiddleware {
521            service: Rc::new(service),
522            max_bytes: self.max_bytes,
523            decompress_gzip: self.decompress_gzip,
524        })
525    }
526}
527
528pub struct RequestBodyLimitMiddleware<S> {
529    service: Rc<S>,
530    max_bytes: usize,
531    decompress_gzip: bool,
532}
533
534impl<S, B> Service<ServiceRequest> for RequestBodyLimitMiddleware<S>
535where
536    S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
537    S::Future: 'static,
538    B: MessageBody + 'static,
539{
540    type Response = ServiceResponse<EitherBody<B>>;
541    type Error = Error;
542    type Future = LocalBoxFuture<'static, Result<Self::Response, Self::Error>>;
543
544    fn poll_ready(&self, context: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
545        self.service.poll_ready(context)
546    }
547
548    fn call(&self, mut request: ServiceRequest) -> Self::Future {
549        let max_bytes = request
550            .extensions()
551            .get::<RequestPolicy>()
552            .and_then(|policy| policy.max_body_bytes)
553            .unwrap_or(self.max_bytes);
554        let decompress_gzip = self.decompress_gzip;
555        let service = Rc::clone(&self.service);
556
557        Box::pin(async move {
558            if request
559                .headers()
560                .get(header::CONTENT_LENGTH)
561                .and_then(|value| value.to_str().ok())
562                .and_then(|value| value.parse::<usize>().ok())
563                .is_some_and(|length| length > max_bytes)
564            {
565                return Ok(payload_too_large(request, max_bytes));
566            }
567
568            let gzip = decompress_gzip
569                && request
570                    .headers()
571                    .get(header::CONTENT_ENCODING)
572                    .and_then(|value| value.to_str().ok())
573                    .is_some_and(|value| {
574                        value
575                            .split(',')
576                            .any(|encoding| encoding.trim().eq_ignore_ascii_case("gzip"))
577                    });
578
579            let mut payload = request.take_payload();
580            let mut body = BytesMut::new();
581            while let Some(chunk) = payload.next().await {
582                let chunk = match chunk {
583                    Ok(chunk) => chunk,
584                    Err(_) => {
585                        return Ok(request.into_response(
586                            HttpResponse::BadRequest()
587                                .body("invalid request body")
588                                .map_into_right_body(),
589                        ));
590                    }
591                };
592                if body.len().saturating_add(chunk.len()) > max_bytes {
593                    return Ok(payload_too_large(request, max_bytes));
594                }
595                body.extend_from_slice(&chunk);
596            }
597
598            if gzip {
599                let decoder = flate2::read::GzDecoder::new(body.as_ref());
600                let mut expanded = Vec::new();
601                if decoder
602                    .take(max_bytes as u64 + 1)
603                    .read_to_end(&mut expanded)
604                    .is_err()
605                {
606                    return Ok(request.into_response(
607                        HttpResponse::BadRequest()
608                            .body("invalid gzip request body")
609                            .map_into_right_body(),
610                    ));
611                }
612                if expanded.len() > max_bytes {
613                    return Ok(payload_too_large(request, max_bytes));
614                }
615                body = BytesMut::from(expanded.as_slice());
616                request.headers_mut().remove(header::CONTENT_ENCODING);
617                request.headers_mut().remove(header::CONTENT_LENGTH);
618            }
619
620            let body = body.freeze();
621            let payload =
622                futures::stream::once(async move { Ok::<_, actix_web::error::PayloadError>(body) });
623            let payload: Pin<
624                Box<
625                    dyn Stream<
626                        Item = Result<actix_web::web::Bytes, actix_web::error::PayloadError>,
627                    >,
628                >,
629            > = Box::pin(payload);
630            request.set_payload(payload.into());
631            Ok(service.call(request).await?.map_into_left_body())
632        })
633    }
634}
635
636fn payload_too_large<B>(
637    request: ServiceRequest,
638    max_bytes: usize,
639) -> ServiceResponse<EitherBody<B>> {
640    request.into_response(
641        HttpResponse::build(StatusCode::PAYLOAD_TOO_LARGE)
642            .body(format!("request body exceeds {max_bytes} bytes"))
643            .map_into_right_body(),
644    )
645}
646
647impl TokenBucket {
648    /// Returns the time until the next permit when the bucket is empty.
649    fn try_acquire(&mut self, permits_per_second: f64) -> Option<Duration> {
650        let now = Instant::now();
651        let elapsed = now.duration_since(self.last_refill).as_secs_f64();
652        self.available = (self.available + elapsed * permits_per_second).min(self.capacity);
653        self.last_refill = now;
654
655        if self.available >= 1.0 {
656            self.available -= 1.0;
657            None
658        } else {
659            Some(Duration::from_secs_f64(
660                (1.0 - self.available) / permits_per_second,
661            ))
662        }
663    }
664}
665
666#[cfg(test)]
667mod tests {
668    use super::*;
669    use actix_web::{
670        test,
671        web::{self, Data},
672        App, HttpResponse,
673    };
674    use rust_zero_core::Metrics;
675    use std::{
676        future::{poll_fn, Future},
677        sync::Arc,
678        task::Poll,
679    };
680    use tokio::sync::Notify;
681
682    #[actix_rt::test]
683    async fn timeout_returns_gateway_timeout() {
684        let metrics = Metrics::new();
685        let http_metrics = HttpMetrics::new(&metrics, "test").unwrap();
686        let app = test::init_service(
687            App::new()
688                .wrap(Timeout::new(Duration::from_millis(5)).with_metrics(http_metrics))
689                .route(
690                    "/",
691                    web::get().to(|| async {
692                        actix_rt::time::sleep(Duration::from_millis(50)).await;
693                        HttpResponse::Ok().finish()
694                    }),
695                ),
696        )
697        .await;
698
699        let error = test::try_call_service(&app, test::TestRequest::get().uri("/").to_request())
700            .await
701            .expect_err("slow request should time out");
702
703        assert_eq!(
704            actix_web::error::ResponseError::status_code(error.as_response_error()),
705            StatusCode::GATEWAY_TIMEOUT
706        );
707        assert!(metrics.render().contains(
708            "test_http_protection_decisions_total{mechanism=\"timeout\",decision=\"rejected\"} 1"
709        ));
710    }
711
712    #[actix_rt::test]
713    async fn concurrency_limit_sheds_busy_requests() {
714        let release = Arc::new(Notify::new());
715        let metrics = Metrics::new();
716        let http_metrics = HttpMetrics::new(&metrics, "test").unwrap();
717        let app = test::init_service(
718            App::new()
719                .app_data(Data::from(Arc::clone(&release)))
720                .wrap(ConcurrencyLimit::new(1).with_metrics(http_metrics))
721                .route(
722                    "/",
723                    web::get().to(|release: Data<Notify>| async move {
724                        release.notified().await;
725                        HttpResponse::Ok().finish()
726                    }),
727                ),
728        )
729        .await;
730
731        let first = test::call_service(&app, test::TestRequest::get().uri("/").to_request());
732        futures::pin_mut!(first);
733        poll_fn(|context| {
734            assert!(first.as_mut().poll(context).is_pending());
735            Poll::Ready(())
736        })
737        .await;
738
739        let second = test::call_service(&app, test::TestRequest::get().uri("/").to_request()).await;
740        assert_eq!(second.status(), StatusCode::SERVICE_UNAVAILABLE);
741        assert!(metrics.render().contains(
742            "test_http_protection_decisions_total{mechanism=\"concurrency\",decision=\"rejected\"} 1"
743        ));
744
745        release.notify_waiters();
746        assert_eq!(first.await.status(), StatusCode::OK);
747    }
748
749    #[actix_rt::test]
750    async fn adaptive_load_shed_is_installed_as_http_middleware() {
751        let release = Arc::new(Notify::new());
752        let metrics = Metrics::new();
753        let http_metrics = HttpMetrics::new(&metrics, "adaptive").unwrap();
754        let shedder = AdaptiveShedder::new(rust_zero_core::LoadShedderConfig::new(
755            1,
756            Duration::from_secs(1),
757        ));
758        let app = test::init_service(
759            App::new()
760                .app_data(Data::from(Arc::clone(&release)))
761                .wrap(AdaptiveLoadShed::new(shedder).with_metrics(http_metrics))
762                .route(
763                    "/",
764                    web::get().to(|release: Data<Notify>| async move {
765                        release.notified().await;
766                        HttpResponse::Ok().finish()
767                    }),
768                ),
769        )
770        .await;
771
772        let first = test::call_service(&app, test::TestRequest::get().uri("/").to_request());
773        futures::pin_mut!(first);
774        poll_fn(|context| {
775            assert!(first.as_mut().poll(context).is_pending());
776            Poll::Ready(())
777        })
778        .await;
779
780        let rejected =
781            test::call_service(&app, test::TestRequest::get().uri("/").to_request()).await;
782        assert_eq!(rejected.status(), StatusCode::SERVICE_UNAVAILABLE);
783        assert!(metrics.render().contains(
784            "adaptive_http_protection_decisions_total{mechanism=\"load_shedder\",decision=\"rejected\"} 1"
785        ));
786
787        release.notify_waiters();
788        assert_eq!(first.await.status(), StatusCode::OK);
789    }
790
791    #[actix_rt::test]
792    async fn adaptive_permit_lives_until_the_response_body_is_dropped() {
793        let shedder = AdaptiveShedder::new(rust_zero_core::LoadShedderConfig::new(
794            1,
795            Duration::from_secs(1),
796        ));
797        let body = PermitBody::new((), shedder.try_acquire().unwrap());
798        assert!(shedder.try_acquire().is_none());
799        drop(body);
800        assert!(shedder.try_acquire().is_some());
801    }
802
803    #[actix_rt::test]
804    async fn rate_limit_returns_retry_after() {
805        let metrics = Metrics::new();
806        let http_metrics = HttpMetrics::new(&metrics, "test").unwrap();
807        let app = test::init_service(
808            App::new()
809                .wrap(RateLimit::new(1, 1).with_metrics(http_metrics))
810                .route("/", web::get().to(|| async { HttpResponse::Ok().finish() })),
811        )
812        .await;
813
814        let first = test::call_service(&app, test::TestRequest::get().uri("/").to_request()).await;
815        let second = test::call_service(&app, test::TestRequest::get().uri("/").to_request()).await;
816
817        assert_eq!(first.status(), StatusCode::OK);
818        assert_eq!(second.status(), StatusCode::TOO_MANY_REQUESTS);
819        assert_eq!(second.headers().get(header::RETRY_AFTER).unwrap(), "1");
820        assert!(metrics.render().contains(
821            "test_http_protection_decisions_total{mechanism=\"rate_limit\",decision=\"rejected\"} 1"
822        ));
823    }
824
825    #[actix_rt::test]
826    async fn request_body_limit_rejects_oversized_streams() {
827        let app = test::init_service(
828            App::new()
829                .wrap(RequestBodyLimit::new(4))
830                .route("/", web::post().to(|body: String| async move { body })),
831        )
832        .await;
833
834        let response = test::call_service(
835            &app,
836            test::TestRequest::post()
837                .uri("/")
838                .set_payload("12345")
839                .to_request(),
840        )
841        .await;
842        assert_eq!(response.status(), StatusCode::PAYLOAD_TOO_LARGE);
843    }
844
845    #[actix_rt::test]
846    async fn request_body_limit_expands_gzip_input() {
847        use flate2::{write::GzEncoder, Compression};
848        use std::io::Write;
849
850        let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
851        encoder.write_all(b"hello").unwrap();
852        let compressed = encoder.finish().unwrap();
853        let app = test::init_service(
854            App::new()
855                .wrap(RequestBodyLimit::new(64))
856                .route("/", web::post().to(|body: String| async move { body })),
857        )
858        .await;
859
860        let response = test::call_service(
861            &app,
862            test::TestRequest::post()
863                .uri("/")
864                .insert_header((header::CONTENT_ENCODING, "gzip"))
865                .set_payload(compressed)
866                .to_request(),
867        )
868        .await;
869        assert_eq!(response.status(), StatusCode::OK);
870        assert_eq!(test::read_body(response).await, "hello");
871    }
872}