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
22pub 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
122pub 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 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#[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
361pub 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#[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 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}