use std::time::Duration;
use axum::{
body::{Body, to_bytes},
extract::Request,
http::{HeaderMap, HeaderName, HeaderValue, StatusCode, header, response::Parts},
middleware::Next,
response::Response,
};
use toolkit_canonical_errors::{CanonicalError, ForeignPassthrough, Http, Problem};
const PROBLEM_JSON: &str = "application/problem+json";
const MAX_FOREIGN_BODY_LOG_BYTES: usize = 8 * 1024;
const MAX_FOREIGN_BODY_LOG_CHARS: usize = 256;
const FOREIGN_BODY_READ_TIMEOUT: Duration = Duration::from_secs(2);
const MAX_PROBLEM_BODY_BYTES: usize = 64 * 1024;
pub async fn canonical_error_middleware(request: Request, next: Next) -> Response {
let uri_path = request.uri().path().to_owned();
let request_headers = request.headers().clone();
let response = next.run(request).await;
if response.extensions().get::<ForeignPassthrough>().is_some() {
return response;
}
if is_problem_response(&response) {
return enrich_problem_response(response, &uri_path, &request_headers).await;
}
let status = response.status();
let is_error_status = status.is_client_error() || status.is_server_error();
if is_error_status && is_unstructured_error_body(&response) {
let recovered = response.extensions().get::<CanonicalError>().cloned();
return wrap_foreign_response(response, &uri_path, &request_headers, recovered).await;
}
response
}
async fn enrich_problem_response(
response: Response,
uri_path: &str,
request_headers: &HeaderMap,
) -> Response {
let (parts, body) = response.into_parts();
let canonical_err = parts.extensions.get::<CanonicalError>().cloned();
let bytes = match tokio::time::timeout(
FOREIGN_BODY_READ_TIMEOUT,
to_bytes(body, MAX_PROBLEM_BODY_BYTES),
)
.await
{
Ok(Ok(b)) => b,
Ok(Err(e)) => {
tracing::error!(error = %e, "canonical error middleware: failed to read response body");
return finish_unreadable_problem_body(parts, uri_path, request_headers, canonical_err)
.await;
}
Err(_) => {
tracing::error!(
instance = uri_path,
"canonical error middleware: timed out reading problem+json response body"
);
return finish_unreadable_problem_body(parts, uri_path, request_headers, canonical_err)
.await;
}
};
let mut problem: Problem = match serde_json::from_slice(&bytes) {
Ok(p) => p,
Err(e) => {
tracing::error!(error = %e, "canonical error middleware: failed to deserialize problem+json body");
let status = parts.status;
if !(status.is_client_error() || status.is_server_error()) {
tracing::warn!(
status = status.as_u16(),
instance = uri_path,
"canonical error middleware: non-error response claimed application/problem+json but its body is not valid Problem JSON; passing through unchanged"
);
return Response::from_parts(parts, Body::from(bytes));
}
let foreign = Response::from_parts(parts, Body::from(bytes));
return wrap_foreign_response(foreign, uri_path, request_headers, canonical_err).await;
}
};
if problem.instance.is_none() {
problem.instance = Some(uri_path.to_owned());
}
if problem.trace_id.is_none() {
problem.trace_id = extract_trace_id(request_headers);
}
problem.status.get_or_insert(parts.status.as_u16());
log_problem(&problem, canonical_err.as_ref());
let new_bytes = match serde_json::to_vec(&problem) {
Ok(b) => b,
Err(e) => {
tracing::error!(
error = %e,
"canonical error middleware: failed to re-serialize problem+json body"
);
return Response::from_parts(parts, Body::from(bytes));
}
};
let len = new_bytes.len();
let mut response = Response::from_parts(parts, Body::from(new_bytes));
response
.headers_mut()
.insert(header::CONTENT_LENGTH, HeaderValue::from(len));
response
}
async fn finish_unreadable_problem_body(
parts: Parts,
uri_path: &str,
request_headers: &HeaderMap,
canonical_err: Option<CanonicalError>,
) -> Response {
let status = parts.status;
if !(status.is_client_error() || status.is_server_error()) {
return Response::from_parts(parts, Body::empty());
}
let foreign = Response::from_parts(parts, Body::empty());
wrap_foreign_response(foreign, uri_path, request_headers, canonical_err).await
}
async fn wrap_foreign_response(
response: Response,
uri_path: &str,
request_headers: &HeaderMap,
recovered: Option<CanonicalError>,
) -> Response {
if let Some(canonical) = recovered {
return wrap_recovered_canonical_error(response, uri_path, request_headers, canonical)
.await;
}
if response.status().is_server_error() {
wrap_as_internal_problem(response, uri_path, request_headers).await
} else {
wrap_as_generic_problem(response, uri_path, request_headers).await
}
}
async fn wrap_recovered_canonical_error(
response: Response,
uri_path: &str,
request_headers: &HeaderMap,
canonical: CanonicalError,
) -> Response {
let status = response.status();
let (parts, body) = response.into_parts();
log_foreign_body(body, status, uri_path).await;
let mut problem: Problem = canonical.clone().into();
problem.instance = Some(uri_path.to_owned());
problem.trace_id = extract_trace_id(request_headers);
log_problem(&problem, Some(&canonical));
finish_wrapped_response(parts, &problem)
}
async fn log_foreign_body(body: Body, status: StatusCode, uri_path: &str) {
tracing::warn!(
status = status.as_u16(),
instance = uri_path,
"canonical error middleware: wrapping a foreign error response"
);
match tokio::time::timeout(
FOREIGN_BODY_READ_TIMEOUT,
to_bytes(body, MAX_FOREIGN_BODY_LOG_BYTES),
)
.await
{
Ok(Ok(bytes)) => {
let text = String::from_utf8_lossy(&bytes);
let text = text.trim();
if !text.is_empty() {
let truncated: String = text.chars().take(MAX_FOREIGN_BODY_LOG_CHARS).collect();
tracing::debug!(
status = status.as_u16(),
instance = uri_path,
body = %truncated.escape_debug(),
"canonical error middleware: foreign response body (diagnostic only, never sent to client)"
);
}
}
Ok(Err(e)) => {
tracing::warn!(error = %e, "canonical error middleware: failed to read foreign response body while wrapping");
}
Err(_) => {
tracing::warn!(
status = status.as_u16(),
instance = uri_path,
"canonical error middleware: timed out reading foreign response body while wrapping"
);
}
}
}
const PRESERVED_FOREIGN_HEADERS: &[axum::http::HeaderName] = &[
header::RETRY_AFTER,
header::ACCESS_CONTROL_ALLOW_ORIGIN,
header::ACCESS_CONTROL_ALLOW_CREDENTIALS,
header::ACCESS_CONTROL_EXPOSE_HEADERS,
header::VARY,
];
const PRESERVED_RATE_LIMIT_HEADERS: &[&str] = &[
"ratelimit-policy",
"ratelimit-limit",
"x-ratelimit-limit",
"x-ratelimit-remaining",
"x-ratelimit-reset",
];
const PRESERVED_MULTI_VALUE_FOREIGN_HEADERS: &[axum::http::HeaderName] =
&[header::SET_COOKIE, header::WWW_AUTHENTICATE];
fn finish_wrapped_response(parts: axum::http::response::Parts, problem: &Problem) -> Response {
let bytes = match serde_json::to_vec(problem) {
Ok(b) => b,
Err(e) => {
tracing::error!(
error = %e,
"canonical error middleware: failed to serialize generic-wrap problem body"
);
return Response::from_parts(parts, Body::empty());
}
};
let len = bytes.len();
let mut headers = HeaderMap::new();
for name in PRESERVED_FOREIGN_HEADERS {
if let Some(value) = parts.headers.get(name) {
headers.insert(name.clone(), value.clone());
}
}
for name in PRESERVED_RATE_LIMIT_HEADERS {
let name = HeaderName::from_static(name);
if let Some(value) = parts.headers.get(&name) {
headers.insert(name, value.clone());
}
}
for name in PRESERVED_MULTI_VALUE_FOREIGN_HEADERS {
for value in parts.headers.get_all(name) {
headers.append(name.clone(), value.clone());
}
}
headers.insert(header::CONTENT_TYPE, HeaderValue::from_static(PROBLEM_JSON));
headers.insert(header::CONTENT_LENGTH, HeaderValue::from(len));
let mut response = Response::from_parts(parts, Body::from(bytes));
*response.headers_mut() = headers;
response
}
async fn wrap_as_generic_problem(
response: Response,
uri_path: &str,
request_headers: &HeaderMap,
) -> Response {
let status = response.status();
let (parts, body) = response.into_parts();
let reason = status.canonical_reason().unwrap_or("Error");
log_foreign_body(body, status, uri_path).await;
let problem = Problem {
problem_type: "about:blank".to_owned(),
title: reason.to_owned(),
status: Some(status.as_u16()),
detail: reason.to_owned(),
instance: Some(uri_path.to_owned()),
trace_id: extract_trace_id(request_headers),
context: serde_json::Value::Object(serde_json::Map::new()),
error_code: None,
error_domain: None,
};
log_problem(&problem, None);
finish_wrapped_response(parts, &problem)
}
async fn wrap_as_internal_problem(
response: Response,
uri_path: &str,
request_headers: &HeaderMap,
) -> Response {
let status = response.status();
let (parts, body) = response.into_parts();
let reason = status.canonical_reason().unwrap_or("Error");
log_foreign_body(body, status, uri_path).await;
let canonical = if status == StatusCode::INTERNAL_SERVER_ERROR {
CanonicalError::internal(reason).create()
} else {
CanonicalError::internal(reason)
.with_override(Http::status_code(status.as_u16()))
.create()
};
let mut problem: Problem = canonical.clone().into();
problem.instance = Some(uri_path.to_owned());
problem.trace_id = extract_trace_id(request_headers);
log_problem(&problem, Some(&canonical));
finish_wrapped_response(parts, &problem)
}
fn content_type_media(response: &Response) -> Option<&str> {
response
.headers()
.get(header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.map(|ct| {
ct.split_once(';')
.map_or(ct, |(media_type, _)| media_type)
.trim()
})
}
fn is_problem_response(response: &Response) -> bool {
content_type_media(response).is_some_and(|mt| mt.eq_ignore_ascii_case(PROBLEM_JSON))
}
fn is_unstructured_error_body(response: &Response) -> bool {
match content_type_media(response) {
None => true,
Some(mt) => mt.eq_ignore_ascii_case("text/plain"),
}
}
fn extract_trace_id(headers: &HeaderMap) -> Option<String> {
if let Some(tp) = headers.get("traceparent").and_then(|v| v.to_str().ok())
&& let Some(trace_id) = parse_w3c_trace_id(tp)
{
return Some(trace_id);
}
for name in ["x-trace-id", "x-request-id"] {
if let Some(v) = headers.get(name).and_then(|v| v.to_str().ok()) {
return Some(v.to_owned());
}
}
tracing::Span::current()
.id()
.map(|id| id.into_u64().to_string())
}
fn parse_w3c_trace_id(traceparent: &str) -> Option<String> {
let parts: Vec<&str> = traceparent.split('-').collect();
if parts.len() >= 4 && parts[0] == "00" {
Some(parts[1].to_owned())
} else {
None
}
}
fn log_problem(problem: &Problem, canonical: Option<&CanonicalError>) {
let status = problem.status.unwrap_or(0);
let problem_type = problem.problem_type.as_str();
let instance = problem.instance.as_deref().unwrap_or("");
let trace_id = problem.trace_id.as_deref().unwrap_or("");
let description = canonical.and_then(CanonicalError::diagnostic).unwrap_or("");
if (400..500).contains(&status) {
tracing::warn!(
status,
problem_type,
instance,
trace_id,
"canonical error response (client)"
);
} else if (500..600).contains(&status) {
tracing::error!(
status,
problem_type,
instance,
trace_id,
description,
"canonical error response (server)"
);
}
}
#[cfg(test)]
#[cfg_attr(coverage_nightly, coverage(off))]
#[path = "canonical_error_layer_tests.rs"]
mod tests;