use std::sync::Arc;
use std::time::Duration;
use axum::{
body::Body,
extract::State,
http::{Request, header},
middleware::Next,
response::{IntoResponse, Response},
};
use tokio::sync::Semaphore;
use tracing::warn;
use crate::error::Problem;
#[derive(Clone)]
pub struct Admission {
permits: Arc<Semaphore>,
wait: Duration,
deadline: Duration,
}
impl Admission {
#[must_use]
pub fn new(max_concurrent: usize, wait_ms: u64, deadline_ms: u64) -> Self {
let max_concurrent = max_concurrent.max(1);
Self {
permits: Arc::new(Semaphore::new(max_concurrent)),
wait: Duration::from_millis(wait_ms),
deadline: Duration::from_millis(deadline_ms),
}
}
#[must_use]
pub fn available(&self) -> usize {
self.permits.available_permits()
}
}
const RETRY_AFTER_SECONDS: &str = "5";
pub async fn admission_middleware(
State(admission): State<Admission>,
request: Request<Body>,
next: Next,
) -> Response {
let method = request.method().clone();
let path = request.uri().path().to_string();
let Ok(Ok(_permit)) = tokio::time::timeout(
admission.wait,
Arc::clone(&admission.permits).acquire_owned(),
)
.await
else {
warn!(event = "request_shed", outcome = "failure", method = %method, path = %path);
let mut response = Problem::service_unavailable(
"The server is at capacity and did not process this request; retry shortly",
)
.into_response();
response.headers_mut().insert(
header::RETRY_AFTER,
header::HeaderValue::from_static(RETRY_AFTER_SECONDS),
);
return response;
};
match tokio::time::timeout(admission.deadline, next.run(request)).await {
Ok(response) => response,
Err(_) => {
warn!(
event = "request_deadline_exceeded",
outcome = "failure",
method = %method,
path = %path,
deadline_ms = crate::millis(admission.deadline),
);
Problem::server_internal("The request took too long to process").into_response()
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use axum::http::StatusCode;
use axum::{Router, body::to_bytes, middleware, routing::get};
use tower::ServiceExt;
fn app(admission: Admission, handler_delay_ms: u64) -> Router {
Router::new()
.route(
"/slow",
get(move || async move {
tokio::time::sleep(Duration::from_millis(handler_delay_ms)).await;
"done"
}),
)
.layer(middleware::from_fn_with_state(
admission,
admission_middleware,
))
}
async fn get_slow(app: &Router) -> Response {
app.clone()
.oneshot(Request::get("/slow").body(Body::empty()).unwrap())
.await
.unwrap()
}
#[tokio::test]
async fn a_request_that_fits_is_served_normally() {
let response = get_slow(&app(Admission::new(1, 50, 10_000), 0)).await;
assert_eq!(response.status(), StatusCode::OK);
}
#[tokio::test]
async fn a_request_past_the_limit_is_shed_rather_than_queued() {
let app = app(Admission::new(1, 0, 10_000), 300);
let first = tokio::spawn({
let app = app.clone();
async move { get_slow(&app).await }
});
tokio::time::sleep(Duration::from_millis(50)).await;
let shed = get_slow(&app).await;
assert_eq!(shed.status(), StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(
shed.headers()
.get(header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok()),
Some("application/problem+json"),
"a refusal must still be a problem document"
);
assert_eq!(
shed.headers()
.get(header::RETRY_AFTER)
.and_then(|v| v.to_str().ok()),
Some(RETRY_AFTER_SECONDS)
);
let body = to_bytes(shed.into_body(), 64 * 1024).await.unwrap();
let problem: serde_json::Value = serde_json::from_slice(&body).unwrap();
assert_eq!(problem["status"], 503);
assert_eq!(problem["type"], "urn:ietf:params:acme:error:serverInternal");
assert_eq!(first.await.unwrap().status(), StatusCode::OK);
}
#[tokio::test]
async fn a_short_burst_waits_for_a_slot_instead_of_being_refused() {
let app = app(Admission::new(1, 1_000, 10_000), 100);
let first = tokio::spawn({
let app = app.clone();
async move { get_slow(&app).await }
});
tokio::time::sleep(Duration::from_millis(20)).await;
assert_eq!(get_slow(&app).await.status(), StatusCode::OK);
assert_eq!(first.await.unwrap().status(), StatusCode::OK);
}
#[tokio::test]
async fn a_request_past_its_deadline_answers_a_problem_document() {
let response = get_slow(&app(Admission::new(4, 50, 30), 5_000)).await;
assert_eq!(response.status(), StatusCode::INTERNAL_SERVER_ERROR);
assert_eq!(
response
.headers()
.get(header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok()),
Some("application/problem+json"),
);
}
#[tokio::test]
async fn a_timed_out_request_releases_its_slot() {
let admission = Admission::new(1, 50, 30);
let app = app(admission.clone(), 5_000);
assert_eq!(
get_slow(&app).await.status(),
StatusCode::INTERNAL_SERVER_ERROR
);
assert_eq!(admission.available(), 1);
}
#[tokio::test]
async fn a_limit_of_zero_still_serves_requests() {
assert_eq!(Admission::new(0, 50, 10_000).available(), 1);
}
}