1use std::pin::Pin;
34use std::sync::atomic::{AtomicBool, Ordering};
35use std::sync::{Arc, Mutex};
36use std::task::{Context, Poll};
37
38use axum::body::Body;
39use axum::extract::{Request, State};
40use axum::middleware::Next;
41use axum::response::{IntoResponse, Response};
42use bytes::Bytes;
43use common_memory_manager::MemoryGuard;
44use http::StatusCode;
45use http_body::{Body as HttpBody, Frame};
46
47use crate::error::Result;
48use crate::request_memory_limiter::ServerMemoryLimiter;
49use crate::request_memory_metrics::RequestMemoryMetrics;
50
51#[derive(Clone, Copy)]
56pub(crate) struct ContentEncoded;
57
58pub async fn memory_limit_middleware(
59 State(limiter): State<ServerMemoryLimiter>,
60 req: Request,
61 next: Next,
62) -> Response {
63 let content_length = req
64 .headers()
65 .get(http::header::CONTENT_LENGTH)
66 .and_then(|v| v.to_str().ok())
67 .and_then(|v| v.parse::<u64>().ok())
68 .unwrap_or(0);
69
70 let _guard = match limiter.acquire(content_length).await {
71 Ok(guard) => guard,
72 Err(e) => {
73 return (
74 StatusCode::TOO_MANY_REQUESTS,
75 format!("Request body memory limit exceeded: {}", e),
76 )
77 .into_response();
78 }
79 };
80
81 let content_encoded = req
82 .headers()
83 .get(http::header::CONTENT_ENCODING)
84 .and_then(|v| v.to_str().ok())
85 .is_some_and(|v| !v.eq_ignore_ascii_case("identity"));
86
87 let accounting = BodyMemoryAccounting::default();
90 let retained_accounting = accounting.clone();
95 let (mut parts, body) = req.into_parts();
96 parts.extensions.insert(limiter.clone());
99 parts.extensions.insert(accounting.clone());
100 if content_encoded {
101 parts.extensions.insert(ContentEncoded);
102 }
103 let accounted = AccountedBody::new(body, limiter, content_length, accounting);
104 let req = Request::from_parts(parts, Body::new(accounted));
105
106 let response = next.run(req).await;
107 quota_exceeded_response(&retained_accounting, response)
110}
111
112fn quota_exceeded_response(accounting: &BodyMemoryAccounting, response: Response) -> Response {
118 if accounting.take_quota_exceeded() {
121 return (
122 StatusCode::TOO_MANY_REQUESTS,
123 "Request body memory limit exceeded",
124 )
125 .into_response();
126 }
127 response
128}
129
130pub(crate) async fn decoded_body_accounting_middleware(
140 State(limiter): State<ServerMemoryLimiter>,
141 req: Request,
142 next: Next,
143) -> Response {
144 if req.extensions().get::<ContentEncoded>().is_none() {
145 return next.run(req).await;
146 }
147
148 let content_length = req
149 .headers()
150 .get(http::header::CONTENT_LENGTH)
151 .and_then(|v| v.to_str().ok())
152 .and_then(|v| v.parse::<u64>().ok())
153 .unwrap_or(0);
154
155 let accounting = BodyMemoryAccounting::default();
156 let retained_accounting = accounting.clone();
158 let (mut parts, body) = req.into_parts();
159 parts.extensions.insert(accounting.clone());
160 let accounted = AccountedBody::new(body, limiter, content_length, accounting);
161 let req = Request::from_parts(parts, Body::new(accounted));
162
163 let response = next.run(req).await;
164 quota_exceeded_response(&retained_accounting, response)
167}
168
169#[derive(Clone, Default)]
175struct BodyMemoryAccounting {
176 guards: Arc<Mutex<Vec<MemoryGuard<RequestMemoryMetrics>>>>,
177 quota_exceeded: Arc<AtomicBool>,
180}
181
182impl BodyMemoryAccounting {
183 fn hold(&self, guard: MemoryGuard<RequestMemoryMetrics>) {
184 self.guards.lock().unwrap().push(guard);
185 }
186
187 fn mark_quota_exceeded(&self) {
188 self.quota_exceeded.store(true, Ordering::Release);
189 }
190
191 fn take_quota_exceeded(&self) -> bool {
192 self.quota_exceeded.swap(false, Ordering::AcqRel)
193 }
194}
195
196type AcquireFuture =
197 Pin<Box<dyn std::future::Future<Output = Result<MemoryGuard<RequestMemoryMetrics>>> + Send>>;
198
199enum ChargeOutcome {
200 Charged,
201 Pending,
202 Failed,
203}
204
205struct AccountedBody {
208 inner: Body,
209 limiter: ServerMemoryLimiter,
210 accounting: BodyMemoryAccounting,
211 streamed: u64,
213 incrementally_charged: u64,
215 pre_charged: u64,
217 pending: Option<(Frame<Bytes>, u64)>,
220 pending_acquire: Option<AcquireFuture>,
221}
222
223impl AccountedBody {
224 fn new(
225 inner: Body,
226 limiter: ServerMemoryLimiter,
227 pre_charged: u64,
228 accounting: BodyMemoryAccounting,
229 ) -> Self {
230 Self {
231 inner,
232 limiter,
233 accounting,
234 streamed: 0,
235 incrementally_charged: 0,
236 pre_charged,
237 pending: None,
238 pending_acquire: None,
239 }
240 }
241
242 fn uncharged(&self) -> u64 {
244 self.streamed
245 .saturating_sub(self.pre_charged)
246 .saturating_sub(self.incrementally_charged)
247 }
248
249 fn charge(&mut self, bytes: u64, cx: &mut Context<'_>) -> ChargeOutcome {
257 debug_assert!(bytes > 0);
258 if let Some(fut) = self.pending_acquire.as_mut() {
259 return match fut.as_mut().poll(cx) {
260 Poll::Ready(Ok(guard)) => {
261 self.accounting.hold(guard);
262 self.pending_acquire = None;
263 self.incrementally_charged += bytes;
264 ChargeOutcome::Charged
265 }
266 Poll::Ready(Err(_)) => {
267 self.pending_acquire = None;
268 ChargeOutcome::Failed
269 }
270 Poll::Pending => ChargeOutcome::Pending,
271 };
272 }
273 if let Some(guard) = self.limiter.try_acquire(bytes) {
275 self.accounting.hold(guard);
276 self.incrementally_charged += bytes;
277 return ChargeOutcome::Charged;
278 }
279 let limiter = self.limiter.clone();
280 let mut fut = Box::pin(async move { limiter.acquire(bytes).await });
281 match fut.as_mut().poll(cx) {
284 Poll::Ready(Ok(guard)) => {
285 self.accounting.hold(guard);
286 self.incrementally_charged += bytes;
287 ChargeOutcome::Charged
288 }
289 Poll::Ready(Err(_)) => ChargeOutcome::Failed,
290 Poll::Pending => {
291 self.pending_acquire = Some(fut);
292 ChargeOutcome::Pending
293 }
294 }
295 }
296
297 fn hand_out(
300 &mut self,
301 frame: Frame<Bytes>,
302 bytes: u64,
303 cx: &mut Context<'_>,
304 ) -> Poll<Option<Result<Frame<Bytes>, axum::Error>>> {
305 if bytes == 0 {
306 return Poll::Ready(Some(Ok(frame)));
307 }
308 match self.charge(bytes, cx) {
309 ChargeOutcome::Charged => Poll::Ready(Some(Ok(frame))),
310 ChargeOutcome::Pending => {
311 self.pending = Some((frame, bytes));
312 Poll::Pending
313 }
314 ChargeOutcome::Failed => {
315 self.accounting.mark_quota_exceeded();
316 Poll::Ready(Some(Err(limit_exceeded_error())))
317 }
318 }
319 }
320}
321
322impl HttpBody for AccountedBody {
323 type Data = Bytes;
324 type Error = axum::Error;
325
326 fn poll_frame(
327 self: Pin<&mut Self>,
328 cx: &mut Context<'_>,
329 ) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
330 let this = self.get_mut();
331
332 if let Some((frame, bytes)) = this.pending.take() {
334 return this.hand_out(frame, bytes, cx);
335 }
336
337 match Pin::new(&mut this.inner).poll_frame(cx) {
338 Poll::Pending => Poll::Pending,
339 Poll::Ready(None) => Poll::Ready(None),
340 Poll::Ready(Some(Err(e))) => Poll::Ready(Some(Err(e))),
341 Poll::Ready(Some(Ok(frame))) => {
342 if let Some(data) = frame.data_ref() {
343 this.streamed += data.len() as u64;
344 let uncharged = this.uncharged();
345 this.hand_out(frame, uncharged, cx)
346 } else {
347 Poll::Ready(Some(Ok(frame)))
349 }
350 }
351 }
352 }
353
354 fn is_end_stream(&self) -> bool {
355 self.pending.is_none() && self.inner.is_end_stream()
356 }
357
358 fn size_hint(&self) -> http_body::SizeHint {
359 self.inner.size_hint()
360 }
361}
362
363fn limit_exceeded_error() -> axum::Error {
364 axum::Error::new(std::io::Error::other(
365 "request body exceeded the aggregate memory limit",
366 ))
367}
368
369#[cfg(test)]
370mod tests {
371 use axum::body::{Body, Bytes};
372 use axum::http::{Request, StatusCode, header};
373 use axum::routing::post;
374 use axum::{Router, middleware};
375 use common_memory_manager::OnExhaustedPolicy;
376 use tower::ServiceExt;
377
378 use super::memory_limit_middleware;
379 use crate::request_memory_limiter::ServerMemoryLimiter;
380
381 fn chunked_request(uri: &str, body: Vec<u8>, chunk_size: usize) -> Request<Body> {
384 let stream = futures_util::stream::iter(
385 body.chunks(chunk_size)
386 .map(|c| Ok::<_, std::io::Error>(Bytes::copy_from_slice(c)))
387 .collect::<Vec<_>>(),
388 );
389 Request::builder()
390 .uri(uri)
391 .method("POST")
392 .body(Body::from_stream(stream))
393 .unwrap()
394 }
395
396 fn counted_body_app(limiter: ServerMemoryLimiter) -> Router {
397 Router::new()
398 .route(
399 "/echo",
400 post(|body: Bytes| async move { (StatusCode::OK, body.len().to_string()) }),
401 )
402 .layer(middleware::from_fn_with_state(
403 limiter,
404 memory_limit_middleware,
405 ))
406 }
407
408 #[tokio::test]
409 async fn test_chunked_body_larger_than_quota_is_rejected() {
410 let limiter = ServerMemoryLimiter::new(8 * 1024, OnExhaustedPolicy::Fail);
412 let app = counted_body_app(limiter.clone());
413
414 let res = app
415 .oneshot(chunked_request("/echo", vec![b'x'; 64 * 1024], 1024))
416 .await
417 .unwrap();
418 assert_eq!(
419 res.status(),
420 StatusCode::TOO_MANY_REQUESTS,
421 "a chunked body larger than the quota must be rejected with 429, not 400"
422 );
423 assert_eq!(0, limiter.used_bytes(), "guards must be released");
424 }
425
426 #[tokio::test]
427 async fn test_small_chunked_body_fits() {
428 let limiter = ServerMemoryLimiter::new(8 * 1024, OnExhaustedPolicy::Fail);
429 let app = counted_body_app(limiter.clone());
430
431 let res = app
432 .oneshot(chunked_request("/echo", vec![b'x'; 1024], 512))
433 .await
434 .unwrap();
435 assert_eq!(res.status(), StatusCode::OK);
436 let body = axum::body::to_bytes(res.into_body(), 1024).await.unwrap();
437 assert_eq!(&body[..], b"1024");
438 assert_eq!(0, limiter.used_bytes(), "guards must be released");
439 }
440
441 #[tokio::test]
442 async fn test_content_length_still_admitted_upfront() {
443 let limiter = ServerMemoryLimiter::new(1024, OnExhaustedPolicy::Fail);
446 let app = counted_body_app(limiter);
447
448 let req = Request::builder()
449 .uri("/echo")
450 .method("POST")
451 .header(header::CONTENT_LENGTH, "4096")
452 .body(Body::from(vec![b'x'; 4096]))
453 .unwrap();
454 let res = app.oneshot(req).await.unwrap();
455 assert_eq!(res.status(), StatusCode::TOO_MANY_REQUESTS);
456 }
457
458 #[tokio::test]
459 async fn test_understated_content_length_is_charged_the_difference() {
460 let limiter = ServerMemoryLimiter::new(2 * 1024, OnExhaustedPolicy::Fail);
463 let app = counted_body_app(limiter.clone());
464
465 let stream = futures_util::stream::iter(vec![Ok::<_, std::io::Error>(
466 Bytes::copy_from_slice(&vec![b'x'; 16 * 1024]),
467 )]);
468 let req = Request::builder()
469 .uri("/echo")
470 .method("POST")
471 .header(header::CONTENT_LENGTH, "1024")
472 .body(Body::from_stream(stream))
473 .unwrap();
474 let res = app.oneshot(req).await.unwrap();
475 assert_eq!(res.status(), StatusCode::TOO_MANY_REQUESTS);
476 assert_eq!(0, limiter.used_bytes());
477 }
478
479 #[tokio::test]
480 async fn test_honest_content_length_is_not_double_charged() {
481 let limiter = ServerMemoryLimiter::new(8 * 1024, OnExhaustedPolicy::Fail);
484 let app = counted_body_app(limiter.clone());
485
486 let req = Request::builder()
487 .uri("/echo")
488 .method("POST")
489 .header(header::CONTENT_LENGTH, "4096")
490 .body(Body::from(vec![b'x'; 4096]))
491 .unwrap();
492 let res = app.oneshot(req).await.unwrap();
493 assert_eq!(res.status(), StatusCode::OK);
494 assert_eq!(0, limiter.used_bytes());
495 }
496
497 fn decompressing_app(limiter: ServerMemoryLimiter) -> Router {
500 use tower_http::decompression::RequestDecompressionLayer;
501
502 Router::new()
503 .route(
504 "/echo",
505 post(|body: Bytes| async move { (StatusCode::OK, body.len().to_string()) }),
506 )
507 .layer(middleware::from_fn_with_state(
508 limiter,
509 super::decoded_body_accounting_middleware,
510 ))
511 .layer(RequestDecompressionLayer::new().pass_through_unaccepted(true))
512 }
513
514 fn zstd_body(decoded: &[u8]) -> Body {
515 let compressed = zstd::stream::encode_all(decoded, 3).unwrap();
516 Body::from(compressed)
517 }
518
519 #[tokio::test]
520 async fn test_decompressed_body_is_charged_for_content_encoded_requests() {
521 let limiter = ServerMemoryLimiter::new(16 * 1024, OnExhaustedPolicy::Fail);
524 let app = decompressing_app(limiter.clone());
525
526 let req = Request::builder()
527 .uri("/echo")
528 .method("POST")
529 .header("content-encoding", "zstd")
530 .extension(super::ContentEncoded)
531 .body(zstd_body(&vec![b'x'; 64 * 1024]))
532 .unwrap();
533 let res = app.oneshot(req).await.unwrap();
534 assert_eq!(
535 res.status(),
536 StatusCode::TOO_MANY_REQUESTS,
537 "decoded body larger than the quota must be rejected with 429"
538 );
539 assert_eq!(0, limiter.used_bytes(), "guards must be released");
540 }
541
542 #[tokio::test]
543 async fn test_small_decompressed_body_fits() {
544 let limiter = ServerMemoryLimiter::new(16 * 1024, OnExhaustedPolicy::Fail);
545 let app = decompressing_app(limiter.clone());
546
547 let req = Request::builder()
548 .uri("/echo")
549 .method("POST")
550 .header("content-encoding", "zstd")
551 .extension(super::ContentEncoded)
552 .body(zstd_body(&vec![b'x'; 1024]))
553 .unwrap();
554 let res = app.oneshot(req).await.unwrap();
555 assert_eq!(res.status(), StatusCode::OK);
556 let body = axum::body::to_bytes(res.into_body(), 1024).await.unwrap();
557 assert_eq!(&body[..], b"1024", "the handler must see the decoded body");
558 assert_eq!(0, limiter.used_bytes(), "guards must be released");
559 }
560
561 #[tokio::test]
562 async fn test_pending_acquisition_is_attributed_to_its_own_bytes() {
563 use std::time::Duration;
564
565 use tokio_stream::wrappers::ReceiverStream;
566
567 let limiter = ServerMemoryLimiter::new(
578 8 * 1024,
579 OnExhaustedPolicy::Wait {
580 timeout: Duration::from_millis(500),
581 },
582 );
583 let app = counted_body_app(limiter.clone());
584
585 let (tx, rx) = tokio::sync::mpsc::channel::<Result<Bytes, std::io::Error>>(2);
586 let driver_limiter = limiter.clone();
587 let driver = tokio::spawn(async move {
588 let external = driver_limiter.acquire(8 * 1024).await.unwrap();
590 tx.send(Ok(Bytes::from(vec![b'a'; 4096]))).await.unwrap();
591 tokio::time::sleep(Duration::from_millis(300)).await;
592 drop(external);
594 tokio::time::sleep(Duration::from_millis(300)).await;
595 tx.send(Ok(Bytes::from(vec![b'b'; 8192]))).await.unwrap();
596 tokio::time::sleep(Duration::from_millis(1500)).await;
598 });
599
600 let res = app
601 .oneshot(
602 Request::builder()
603 .uri("/echo")
604 .method("POST")
605 .body(Body::from_stream(ReceiverStream::new(rx)))
606 .unwrap(),
607 )
608 .await
609 .unwrap();
610 driver.await.unwrap();
611
612 assert_eq!(
613 res.status(),
614 StatusCode::TOO_MANY_REQUESTS,
615 "frame B must wait for its own 8 KiB and time out, not borrow A's reservation"
616 );
617 }
618
619 #[tokio::test]
620 async fn test_plain_requests_are_not_charged_by_the_decoded_layer() {
621 let limiter = ServerMemoryLimiter::new(16 * 1024, OnExhaustedPolicy::Fail);
624 let app = decompressing_app(limiter.clone());
625
626 let req = Request::builder()
628 .uri("/echo")
629 .method("POST")
630 .body(Body::from(vec![b'x'; 1024]))
631 .unwrap();
632 let res = app.clone().oneshot(req).await.unwrap();
633 assert_eq!(res.status(), StatusCode::OK);
634
635 let req = Request::builder()
638 .uri("/echo")
639 .method("POST")
640 .header("content-encoding", "zstd")
641 .body(zstd_body(&vec![b'x'; 1024]))
642 .unwrap();
643 let res = app.oneshot(req).await.unwrap();
644 assert_eq!(res.status(), StatusCode::OK);
645 }
646}