use crate::error::refusal;
use crate::metrics::route_and_surface;
use axum::body::Body;
use axum::extract::{Request, State};
use axum::http::StatusCode;
use axum::middleware::Next;
use axum::response::Response;
use futures::StreamExt as _;
use notedthat_core::metrics::{label, name, refused_reason};
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::{Semaphore, oneshot};
use tower_http::request_id::RequestId;
#[derive(Debug, Clone)]
pub struct RequestBounds {
timeout: Duration,
client_idle: Duration,
in_flight: Arc<Semaphore>,
}
impl RequestBounds {
#[must_use]
pub fn new(timeout: Duration, client_idle: Duration, in_flight: Arc<Semaphore>) -> Self {
Self {
timeout,
client_idle,
in_flight,
}
}
#[must_use]
pub fn with_timeout(&self, timeout: Duration) -> Self {
Self {
timeout,
client_idle: self.client_idle,
in_flight: Arc::clone(&self.in_flight),
}
}
#[must_use]
pub fn unbounded() -> Self {
Self::new(
Duration::from_hours(24),
Duration::from_hours(24),
Arc::new(Semaphore::new(Semaphore::MAX_PERMITS)),
)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum BodyEnd {
Complete,
Idle,
}
fn watch_body(body: Body, idle: Duration, signal: oneshot::Sender<BodyEnd>) -> Body {
let stream = futures::stream::unfold(
(body.into_data_stream(), Some(signal), false),
move |(mut frames, mut signal, done)| async move {
if done {
return None;
}
match tokio::time::timeout(idle, frames.next()).await {
Ok(Some(Ok(chunk))) => Some((Ok(chunk), (frames, signal, false))),
Ok(Some(Err(e))) => {
if let Some(signal) = signal.take() {
let _ = signal.send(BodyEnd::Complete);
}
Some((Err(e), (frames, signal, true)))
}
Ok(None) => {
if let Some(signal) = signal.take() {
let _ = signal.send(BodyEnd::Complete);
}
None
}
Err(_) => {
if let Some(signal) = signal.take() {
let _ = signal.send(BodyEnd::Idle);
}
Some((
Err(axum::Error::new(std::io::Error::new(
std::io::ErrorKind::TimedOut,
"the request body stopped arriving",
))),
(frames, signal, true),
))
}
}
},
);
Body::from_stream(stream)
}
pub async fn bound(State(bounds): State<RequestBounds>, req: Request, next: Next) -> Response {
let Ok(_permit) = Arc::clone(&bounds.in_flight).try_acquire_owned() else {
count_refusal(&req, refused_reason::IN_FLIGHT);
return refusal(
StatusCode::SERVICE_UNAVAILABLE,
"backend_unavailable",
"the server is at its limit of requests in flight; retry shortly".to_string(),
request_id(&req),
);
};
let refused = RefusedRequest::of(&req);
let (req, body) = if has_declared_body(&req) {
let (signal, ended) = oneshot::channel();
let req = req.map(|body| watch_body(body, bounds.client_idle, signal));
(req, Some(ended))
} else {
(req, None)
};
let mut handler = std::pin::pin!(next.run(req));
if let Some(mut ended) = body {
tokio::select! {
biased;
end = &mut ended => match end {
Ok(BodyEnd::Idle) => {
refused.count(refused_reason::TIMEOUT);
return refusal(
StatusCode::REQUEST_TIMEOUT,
"request_timeout",
format!(
"the request body stopped arriving for more than {} ms",
bounds.client_idle.as_millis()
),
refused.request_id,
);
}
Ok(BodyEnd::Complete) | Err(_) => {}
},
response = &mut handler => {
if matches!(ended.try_recv(), Ok(BodyEnd::Idle)) {
refused.count(refused_reason::TIMEOUT);
return refusal(
StatusCode::REQUEST_TIMEOUT,
"request_timeout",
format!(
"the request body stopped arriving for more than {} ms",
bounds.client_idle.as_millis()
),
refused.request_id,
);
}
return response;
}
}
}
if let Ok(response) = tokio::time::timeout(bounds.timeout, handler).await {
response
} else {
refused.count(refused_reason::TIMEOUT);
refusal(
StatusCode::GATEWAY_TIMEOUT,
"request_timeout",
format!(
"the request did not complete within {} ms",
bounds.timeout.as_millis()
),
refused.request_id,
)
}
}
fn has_declared_body(req: &Request) -> bool {
let declared = req
.headers()
.get(http::header::CONTENT_LENGTH)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.parse::<u64>().ok());
let chunked = req.headers().contains_key(http::header::TRANSFER_ENCODING);
chunked || declared.is_some_and(|len| len > 0)
}
fn request_id(req: &Request) -> Option<String> {
req.extensions()
.get::<RequestId>()
.and_then(|id| id.header_value().to_str().ok())
.map(str::to_owned)
}
struct RefusedRequest {
route: String,
surface: &'static str,
request_id: Option<String>,
}
impl RefusedRequest {
fn of(req: &Request) -> Self {
let (route, surface) = route_and_surface(req);
Self {
route,
surface,
request_id: request_id(req),
}
}
fn count(&self, reason: &'static str) {
metrics::counter!(
name::HTTP_REQUESTS_REFUSED,
label::SURFACE => self.surface,
label::ROUTE => self.route.clone(),
label::REASON => reason,
)
.increment(1);
}
}
fn count_refusal(req: &Request, reason: &'static str) {
RefusedRequest::of(req).count(reason);
}
#[cfg(test)]
mod tests {
use super::{RequestBounds, bound};
use axum::Router;
use axum::body::{Body, to_bytes};
use axum::http::{Request, StatusCode, header::RETRY_AFTER};
use axum::middleware::from_fn_with_state;
use axum::response::Response;
use axum::routing::get;
use metrics_util::debugging::DebuggingRecorder;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::{Notify, Semaphore};
use tower::ServiceExt;
use tower_http::request_id::{MakeRequestUuid, SetRequestIdLayer};
const TIMEOUT: Duration = Duration::from_secs(5);
const IDLE: Duration = Duration::from_secs(300);
fn app(bounds: &RequestBounds, started: Arc<Notify>, release: Arc<Notify>) -> Router {
let bounded = Router::new()
.route("/fast", get(|| async { "fast" }))
.route(
"/slow",
get(move || {
let started = Arc::clone(&started);
let release = Arc::clone(&release);
async move {
started.notify_one();
release.notified().await;
"slow"
}
}),
)
.route_layer(from_fn_with_state(bounds.clone(), bound));
Router::new()
.route(
"/stream",
get(|| async {
tokio::time::sleep(Duration::from_secs(60)).await;
"stream"
}),
)
.merge(bounded)
.layer(SetRequestIdLayer::x_request_id(MakeRequestUuid))
}
fn get_req(path: &str) -> Request<Body> {
Request::get(path).body(Body::empty()).expect("request")
}
async fn json(response: Response) -> serde_json::Value {
let bytes = to_bytes(response.into_body(), usize::MAX)
.await
.expect("body");
serde_json::from_slice(&bytes).expect("a JSON body")
}
#[tokio::test(start_paused = true)]
async fn a_request_past_its_timeout_is_504_in_the_error_envelope() {
let bounds = RequestBounds::new(TIMEOUT, IDLE, Arc::new(Semaphore::new(4)));
let app = app(&bounds, Arc::new(Notify::new()), Arc::new(Notify::new()));
let response = app.oneshot(get_req("/slow")).await.expect("infallible");
assert_eq!(response.status(), StatusCode::GATEWAY_TIMEOUT);
assert!(
response.headers().get(RETRY_AFTER).is_none(),
"retrying the same slow request will not make it faster"
);
let body = json(response).await;
assert_eq!(body["error"], "request_timeout");
assert!(
body["message"]
.as_str()
.is_some_and(|m| m.contains("5000 ms")),
"{body}"
);
assert!(
body["request_id"].as_str().is_some_and(|id| !id.is_empty()),
"the id the request-id layer assigned belongs in the body: {body}"
);
}
#[tokio::test(start_paused = true)]
async fn a_route_outside_the_layer_outlives_the_timeout() {
let bounds = RequestBounds::new(TIMEOUT, IDLE, Arc::new(Semaphore::new(1)));
let app = app(&bounds, Arc::new(Notify::new()), Arc::new(Notify::new()));
let response = app.oneshot(get_req("/stream")).await.expect("infallible");
assert_eq!(response.status(), StatusCode::OK);
}
#[tokio::test(start_paused = true)]
async fn past_the_cap_a_request_is_refused_503_with_retry_after_at_once() {
let bounds = RequestBounds::new(TIMEOUT, IDLE, Arc::new(Semaphore::new(1)));
let started = Arc::new(Notify::new());
let app = app(&bounds, Arc::clone(&started), Arc::new(Notify::new()));
let holder = tokio::spawn(app.clone().oneshot(get_req("/slow")));
started.notified().await;
let refused = app
.clone()
.oneshot(get_req("/fast"))
.await
.expect("infallible");
assert_eq!(refused.status(), StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(
refused
.headers()
.get(RETRY_AFTER)
.map(axum::http::HeaderValue::as_bytes),
Some(&b"5"[..])
);
assert_eq!(json(refused).await["error"], "backend_unavailable");
let stream = tokio::spawn(app.clone().oneshot(get_req("/stream")));
tokio::time::advance(Duration::from_secs(61)).await;
assert_eq!(
stream.await.expect("join").expect("infallible").status(),
StatusCode::OK
);
let held = holder.await.expect("join").expect("infallible");
assert_eq!(held.status(), StatusCode::GATEWAY_TIMEOUT);
}
#[tokio::test(start_paused = true)]
async fn a_timed_out_request_gives_its_permit_back() {
let bounds = RequestBounds::new(TIMEOUT, IDLE, Arc::new(Semaphore::new(1)));
let app = app(&bounds, Arc::new(Notify::new()), Arc::new(Notify::new()));
let first = app
.clone()
.oneshot(get_req("/slow"))
.await
.expect("infallible");
assert_eq!(first.status(), StatusCode::GATEWAY_TIMEOUT);
let second = app.oneshot(get_req("/fast")).await.expect("infallible");
assert_eq!(second.status(), StatusCode::OK);
}
#[tokio::test(start_paused = true)]
async fn a_client_that_disconnects_gives_its_permit_back() {
let bounds = RequestBounds::new(TIMEOUT, IDLE, Arc::new(Semaphore::new(1)));
let started = Arc::new(Notify::new());
let app = app(&bounds, Arc::clone(&started), Arc::new(Notify::new()));
let abandoned = tokio::spawn(app.clone().oneshot(get_req("/slow")));
started.notified().await;
abandoned.abort();
let _ = abandoned.await;
let next = app.oneshot(get_req("/fast")).await.expect("infallible");
assert_eq!(next.status(), StatusCode::OK);
}
#[test]
fn each_refusal_is_counted_by_route_and_reason() {
let recorder = DebuggingRecorder::new();
let snapshotter = recorder.snapshotter();
metrics::with_local_recorder(&recorder, || {
let runtime = tokio::runtime::Builder::new_current_thread()
.enable_time()
.start_paused(true)
.build()
.expect("runtime");
runtime.block_on(async {
let bounds = RequestBounds::new(TIMEOUT, IDLE, Arc::new(Semaphore::new(1)));
let started = Arc::new(Notify::new());
let app = app(&bounds, Arc::clone(&started), Arc::new(Notify::new()));
let holder = tokio::spawn(app.clone().oneshot(get_req("/slow")));
started.notified().await;
let refused = app.clone().oneshot(get_req("/fast")).await;
assert_eq!(
refused.expect("infallible").status(),
StatusCode::SERVICE_UNAVAILABLE
);
let timed_out = holder.await.expect("join").expect("infallible");
assert_eq!(timed_out.status(), StatusCode::GATEWAY_TIMEOUT);
});
});
let mut series: Vec<String> = snapshotter
.snapshot()
.into_vec()
.into_iter()
.map(|(key, _, _, _)| {
let key = key.key();
let labels = key
.labels()
.map(|l| format!("{}={}", l.key(), l.value()))
.collect::<Vec<_>>()
.join(",");
format!("{}{{{labels}}}", key.name())
})
.collect();
series.sort();
assert_eq!(
series,
[
"notedthat_http_requests_refused_total{surface=root,route=/fast,reason=in_flight}",
"notedthat_http_requests_refused_total{surface=root,route=/slow,reason=timeout}",
]
);
}
}