1use crate::error::refusal;
35use crate::metrics::route_and_surface;
36use axum::body::Body;
37use axum::extract::{Request, State};
38use axum::http::StatusCode;
39use axum::middleware::Next;
40use axum::response::Response;
41use futures::StreamExt as _;
42use notedthat_core::metrics::{label, name, refused_reason};
43use std::sync::Arc;
44use std::time::Duration;
45use tokio::sync::{Semaphore, oneshot};
46use tower_http::request_id::RequestId;
47
48#[derive(Debug, Clone)]
54pub struct RequestBounds {
55 timeout: Duration,
56 client_idle: Duration,
57 in_flight: Arc<Semaphore>,
58}
59
60impl RequestBounds {
61 #[must_use]
68 pub fn new(timeout: Duration, client_idle: Duration, in_flight: Arc<Semaphore>) -> Self {
69 Self {
70 timeout,
71 client_idle,
72 in_flight,
73 }
74 }
75
76 #[must_use]
78 pub fn with_timeout(&self, timeout: Duration) -> Self {
79 Self {
80 timeout,
81 client_idle: self.client_idle,
82 in_flight: Arc::clone(&self.in_flight),
83 }
84 }
85
86 #[must_use]
90 pub fn unbounded() -> Self {
91 Self::new(
92 Duration::from_hours(24),
93 Duration::from_hours(24),
94 Arc::new(Semaphore::new(Semaphore::MAX_PERMITS)),
95 )
96 }
97}
98
99#[derive(Debug, Clone, Copy, PartialEq, Eq)]
101enum BodyEnd {
102 Complete,
104 Idle,
106}
107
108fn watch_body(body: Body, idle: Duration, signal: oneshot::Sender<BodyEnd>) -> Body {
116 let stream = futures::stream::unfold(
117 (body.into_data_stream(), Some(signal), false),
118 move |(mut frames, mut signal, done)| async move {
119 if done {
120 return None;
121 }
122 match tokio::time::timeout(idle, frames.next()).await {
123 Ok(Some(Ok(chunk))) => Some((Ok(chunk), (frames, signal, false))),
124 Ok(Some(Err(e))) => {
125 if let Some(signal) = signal.take() {
128 let _ = signal.send(BodyEnd::Complete);
129 }
130 Some((Err(e), (frames, signal, true)))
131 }
132 Ok(None) => {
133 if let Some(signal) = signal.take() {
134 let _ = signal.send(BodyEnd::Complete);
135 }
136 None
137 }
138 Err(_) => {
139 if let Some(signal) = signal.take() {
140 let _ = signal.send(BodyEnd::Idle);
141 }
142 Some((
147 Err(axum::Error::new(std::io::Error::new(
148 std::io::ErrorKind::TimedOut,
149 "the request body stopped arriving",
150 ))),
151 (frames, signal, true),
152 ))
153 }
154 }
155 },
156 );
157 Body::from_stream(stream)
158}
159
160pub async fn bound(State(bounds): State<RequestBounds>, req: Request, next: Next) -> Response {
166 let Ok(_permit) = Arc::clone(&bounds.in_flight).try_acquire_owned() else {
167 count_refusal(&req, refused_reason::IN_FLIGHT);
168 return refusal(
169 StatusCode::SERVICE_UNAVAILABLE,
170 "backend_unavailable",
171 "the server is at its limit of requests in flight; retry shortly".to_string(),
172 request_id(&req),
173 );
174 };
175 let refused = RefusedRequest::of(&req);
178
179 let (req, body) = if has_declared_body(&req) {
187 let (signal, ended) = oneshot::channel();
188 let req = req.map(|body| watch_body(body, bounds.client_idle, signal));
189 (req, Some(ended))
190 } else {
191 (req, None)
192 };
193
194 let mut handler = std::pin::pin!(next.run(req));
195 if let Some(mut ended) = body {
196 tokio::select! {
197 biased;
203 end = &mut ended => match end {
204 Ok(BodyEnd::Idle) => {
205 refused.count(refused_reason::TIMEOUT);
206 return refusal(
207 StatusCode::REQUEST_TIMEOUT,
208 "request_timeout",
209 format!(
210 "the request body stopped arriving for more than {} ms",
211 bounds.client_idle.as_millis()
212 ),
213 refused.request_id,
214 );
215 }
216 Ok(BodyEnd::Complete) | Err(_) => {}
219 },
220 response = &mut handler => {
223 if matches!(ended.try_recv(), Ok(BodyEnd::Idle)) {
230 refused.count(refused_reason::TIMEOUT);
231 return refusal(
232 StatusCode::REQUEST_TIMEOUT,
233 "request_timeout",
234 format!(
235 "the request body stopped arriving for more than {} ms",
236 bounds.client_idle.as_millis()
237 ),
238 refused.request_id,
239 );
240 }
241 return response;
242 }
243 }
244 }
245
246 if let Ok(response) = tokio::time::timeout(bounds.timeout, handler).await {
247 response
248 } else {
249 refused.count(refused_reason::TIMEOUT);
250 refusal(
251 StatusCode::GATEWAY_TIMEOUT,
252 "request_timeout",
253 format!(
254 "the request did not complete within {} ms",
255 bounds.timeout.as_millis()
256 ),
257 refused.request_id,
258 )
259 }
260}
261
262fn has_declared_body(req: &Request) -> bool {
276 let declared = req
277 .headers()
278 .get(http::header::CONTENT_LENGTH)
279 .and_then(|value| value.to_str().ok())
280 .and_then(|value| value.parse::<u64>().ok());
281 let chunked = req.headers().contains_key(http::header::TRANSFER_ENCODING);
282 chunked || declared.is_some_and(|len| len > 0)
283}
284
285fn request_id(req: &Request) -> Option<String> {
291 req.extensions()
292 .get::<RequestId>()
293 .and_then(|id| id.header_value().to_str().ok())
294 .map(str::to_owned)
295}
296
297struct RefusedRequest {
299 route: String,
300 surface: &'static str,
301 request_id: Option<String>,
302}
303
304impl RefusedRequest {
305 fn of(req: &Request) -> Self {
306 let (route, surface) = route_and_surface(req);
307 Self {
308 route,
309 surface,
310 request_id: request_id(req),
311 }
312 }
313
314 fn count(&self, reason: &'static str) {
315 metrics::counter!(
316 name::HTTP_REQUESTS_REFUSED,
317 label::SURFACE => self.surface,
318 label::ROUTE => self.route.clone(),
319 label::REASON => reason,
320 )
321 .increment(1);
322 }
323}
324
325fn count_refusal(req: &Request, reason: &'static str) {
326 RefusedRequest::of(req).count(reason);
327}
328
329#[cfg(test)]
330mod tests {
331 use super::{RequestBounds, bound};
332 use axum::Router;
333 use axum::body::{Body, to_bytes};
334 use axum::http::{Request, StatusCode, header::RETRY_AFTER};
335 use axum::middleware::from_fn_with_state;
336 use axum::response::Response;
337 use axum::routing::get;
338 use metrics_util::debugging::DebuggingRecorder;
339 use std::sync::Arc;
340 use std::time::Duration;
341 use tokio::sync::{Notify, Semaphore};
342 use tower::ServiceExt;
343 use tower_http::request_id::{MakeRequestUuid, SetRequestIdLayer};
344
345 const TIMEOUT: Duration = Duration::from_secs(5);
346 const IDLE: Duration = Duration::from_secs(300);
348
349 fn app(bounds: &RequestBounds, started: Arc<Notify>, release: Arc<Notify>) -> Router {
353 let bounded = Router::new()
354 .route("/fast", get(|| async { "fast" }))
355 .route(
356 "/slow",
357 get(move || {
358 let started = Arc::clone(&started);
359 let release = Arc::clone(&release);
360 async move {
361 started.notify_one();
362 release.notified().await;
363 "slow"
364 }
365 }),
366 )
367 .route_layer(from_fn_with_state(bounds.clone(), bound));
368 Router::new()
369 .route(
370 "/stream",
371 get(|| async {
372 tokio::time::sleep(Duration::from_secs(60)).await;
373 "stream"
374 }),
375 )
376 .merge(bounded)
377 .layer(SetRequestIdLayer::x_request_id(MakeRequestUuid))
378 }
379
380 fn get_req(path: &str) -> Request<Body> {
381 Request::get(path).body(Body::empty()).expect("request")
382 }
383
384 async fn json(response: Response) -> serde_json::Value {
385 let bytes = to_bytes(response.into_body(), usize::MAX)
386 .await
387 .expect("body");
388 serde_json::from_slice(&bytes).expect("a JSON body")
389 }
390
391 #[tokio::test(start_paused = true)]
392 async fn a_request_past_its_timeout_is_504_in_the_error_envelope() {
393 let bounds = RequestBounds::new(TIMEOUT, IDLE, Arc::new(Semaphore::new(4)));
394 let app = app(&bounds, Arc::new(Notify::new()), Arc::new(Notify::new()));
395
396 let response = app.oneshot(get_req("/slow")).await.expect("infallible");
397
398 assert_eq!(response.status(), StatusCode::GATEWAY_TIMEOUT);
399 assert!(
400 response.headers().get(RETRY_AFTER).is_none(),
401 "retrying the same slow request will not make it faster"
402 );
403 let body = json(response).await;
404 assert_eq!(body["error"], "request_timeout");
405 assert!(
406 body["message"]
407 .as_str()
408 .is_some_and(|m| m.contains("5000 ms")),
409 "{body}"
410 );
411 assert!(
412 body["request_id"].as_str().is_some_and(|id| !id.is_empty()),
413 "the id the request-id layer assigned belongs in the body: {body}"
414 );
415 }
416
417 #[tokio::test(start_paused = true)]
418 async fn a_route_outside_the_layer_outlives_the_timeout() {
419 let bounds = RequestBounds::new(TIMEOUT, IDLE, Arc::new(Semaphore::new(1)));
420 let app = app(&bounds, Arc::new(Notify::new()), Arc::new(Notify::new()));
421
422 let response = app.oneshot(get_req("/stream")).await.expect("infallible");
423
424 assert_eq!(response.status(), StatusCode::OK);
425 }
426
427 #[tokio::test(start_paused = true)]
428 async fn past_the_cap_a_request_is_refused_503_with_retry_after_at_once() {
429 let bounds = RequestBounds::new(TIMEOUT, IDLE, Arc::new(Semaphore::new(1)));
430 let started = Arc::new(Notify::new());
431 let app = app(&bounds, Arc::clone(&started), Arc::new(Notify::new()));
432
433 let holder = tokio::spawn(app.clone().oneshot(get_req("/slow")));
435 started.notified().await;
436
437 let refused = app
439 .clone()
440 .oneshot(get_req("/fast"))
441 .await
442 .expect("infallible");
443
444 assert_eq!(refused.status(), StatusCode::SERVICE_UNAVAILABLE);
446 assert_eq!(
447 refused
448 .headers()
449 .get(RETRY_AFTER)
450 .map(axum::http::HeaderValue::as_bytes),
451 Some(&b"5"[..])
452 );
453 assert_eq!(json(refused).await["error"], "backend_unavailable");
454
455 let stream = tokio::spawn(app.clone().oneshot(get_req("/stream")));
457 tokio::time::advance(Duration::from_secs(61)).await;
458 assert_eq!(
459 stream.await.expect("join").expect("infallible").status(),
460 StatusCode::OK
461 );
462
463 let held = holder.await.expect("join").expect("infallible");
465 assert_eq!(held.status(), StatusCode::GATEWAY_TIMEOUT);
466 }
467
468 #[tokio::test(start_paused = true)]
469 async fn a_timed_out_request_gives_its_permit_back() {
470 let bounds = RequestBounds::new(TIMEOUT, IDLE, Arc::new(Semaphore::new(1)));
471 let app = app(&bounds, Arc::new(Notify::new()), Arc::new(Notify::new()));
472
473 let first = app
474 .clone()
475 .oneshot(get_req("/slow"))
476 .await
477 .expect("infallible");
478 assert_eq!(first.status(), StatusCode::GATEWAY_TIMEOUT);
479
480 let second = app.oneshot(get_req("/fast")).await.expect("infallible");
481 assert_eq!(second.status(), StatusCode::OK);
482 }
483
484 #[tokio::test(start_paused = true)]
485 async fn a_client_that_disconnects_gives_its_permit_back() {
486 let bounds = RequestBounds::new(TIMEOUT, IDLE, Arc::new(Semaphore::new(1)));
487 let started = Arc::new(Notify::new());
488 let app = app(&bounds, Arc::clone(&started), Arc::new(Notify::new()));
489
490 let abandoned = tokio::spawn(app.clone().oneshot(get_req("/slow")));
492 started.notified().await;
493 abandoned.abort();
494 let _ = abandoned.await;
495
496 let next = app.oneshot(get_req("/fast")).await.expect("infallible");
497 assert_eq!(next.status(), StatusCode::OK);
498 }
499
500 #[test]
501 fn each_refusal_is_counted_by_route_and_reason() {
502 let recorder = DebuggingRecorder::new();
503 let snapshotter = recorder.snapshotter();
504 metrics::with_local_recorder(&recorder, || {
505 let runtime = tokio::runtime::Builder::new_current_thread()
506 .enable_time()
507 .start_paused(true)
508 .build()
509 .expect("runtime");
510 runtime.block_on(async {
511 let bounds = RequestBounds::new(TIMEOUT, IDLE, Arc::new(Semaphore::new(1)));
512 let started = Arc::new(Notify::new());
513 let app = app(&bounds, Arc::clone(&started), Arc::new(Notify::new()));
514 let holder = tokio::spawn(app.clone().oneshot(get_req("/slow")));
515 started.notified().await;
516 let refused = app.clone().oneshot(get_req("/fast")).await;
517 assert_eq!(
518 refused.expect("infallible").status(),
519 StatusCode::SERVICE_UNAVAILABLE
520 );
521 let timed_out = holder.await.expect("join").expect("infallible");
522 assert_eq!(timed_out.status(), StatusCode::GATEWAY_TIMEOUT);
523 });
524 });
525
526 let mut series: Vec<String> = snapshotter
527 .snapshot()
528 .into_vec()
529 .into_iter()
530 .map(|(key, _, _, _)| {
531 let key = key.key();
532 let labels = key
533 .labels()
534 .map(|l| format!("{}={}", l.key(), l.value()))
535 .collect::<Vec<_>>()
536 .join(",");
537 format!("{}{{{labels}}}", key.name())
538 })
539 .collect();
540 series.sort();
541 assert_eq!(
542 series,
543 [
544 "notedthat_http_requests_refused_total{surface=root,route=/fast,reason=in_flight}",
545 "notedthat_http_requests_refused_total{surface=root,route=/slow,reason=timeout}",
546 ]
547 );
548 }
549}