use std::net::SocketAddr;
use std::time::Instant;
use axum::{
body::Body,
extract::ConnectInfo,
http::{HeaderName, HeaderValue, Method, Request},
middleware::Next,
response::IntoResponse,
};
use tracing::{Instrument, debug, field, info, info_span, warn};
use uuid::Uuid;
use acme_proxy_core::client::RequestId;
pub const X_REQUEST_ID: HeaderName = HeaderName::from_static("x-request-id");
fn is_probe(method: &Method, path: &str) -> bool {
matches!(*method, Method::GET | Method::HEAD) && matches!(path, "/health" | "/")
}
const REQUEST_ID_MAX: usize = 128;
fn usable_request_id(value: &str) -> bool {
!value.is_empty()
&& value.len() <= REQUEST_ID_MAX
&& value
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.' | b':'))
}
pub async fn add_access_middleware(mut request: Request<Body>, next: Next) -> impl IntoResponse {
let id_str = match request.headers().get(&X_REQUEST_ID) {
Some(value) => match value.to_str() {
Ok(value) => {
let value = value.trim();
if usable_request_id(value) {
value.to_string()
} else {
debug!(event = "request_id_header_invalid", outcome = "failure");
String::new()
}
}
Err(_) => {
debug!(event = "request_id_header_invalid", outcome = "failure");
String::new()
}
},
None => String::new(),
};
let id_str = if id_str.is_empty() {
Uuid::now_v7().to_string()
} else {
id_str
};
request.extensions_mut().insert(RequestId(id_str.clone()));
let probe = is_probe(request.method(), request.uri().path());
let peer = request
.extensions()
.get::<ConnectInfo<SocketAddr>>()
.map(|ConnectInfo(addr)| acme_proxy_core::client::canonical(addr.ip()));
if peer.is_some() {
request
.extensions_mut()
.insert(acme_proxy_core::client::ClientIp(peer));
}
let span = info_span!(
"request",
method = %request.method(),
uri = %request.uri(),
version = ?request.version(),
request_id = %id_str,
profile = field::Empty,
client_ip = field::Empty,
);
if let Some(peer) = peer {
span.record("client_ip", field::display(peer));
}
let started = Instant::now();
let mut response = next.run(request).instrument(span.clone()).await;
let latency_ms = acme_proxy_core::logfields::millis(started.elapsed());
let status = response.status().as_u16();
span.in_scope(|| {
if response.status().is_server_error() {
warn!(
event = "request_completed",
outcome = "failure",
status,
latency_ms
);
} else if probe {
debug!(
event = "request_completed",
outcome = "success",
status,
latency_ms
);
} else {
info!(
event = "request_completed",
outcome = "success",
status,
latency_ms
);
}
});
if let Ok(header_val) = HeaderValue::from_str(&id_str) {
response.headers_mut().insert(X_REQUEST_ID, header_val);
} else {
debug!(event = "request_id_header_not_returned", outcome = "success", request_id = %id_str);
}
response
}
#[cfg(test)]
mod tests {
use super::*;
use axum::{Router, http::StatusCode, middleware, routing::get};
use tower::ServiceExt;
fn app() -> Router {
Router::new()
.route("/test", get(|| async { "ok" }))
.route("/health", get(|| async { "ok" }))
.route("/boom", get(|| async { StatusCode::INTERNAL_SERVER_ERROR }))
.layer(middleware::from_fn(add_access_middleware))
}
async fn fields_for(request: Request<Body>) -> acme_proxy_core::testutil::SpanFields {
acme_proxy_core::testutil::capture_request_span(app().oneshot(request)).await
}
#[tokio::test]
async fn the_peer_address_is_recorded_as_the_client() {
let mut req = Request::builder().uri("/test").body(Body::empty()).unwrap();
req.extensions_mut()
.insert(ConnectInfo(SocketAddr::from(([192, 0, 2, 7], 4711))));
let fields = fields_for(req).await;
assert_eq!(fields.get("client_ip").as_deref(), Some("192.0.2.7"));
}
#[tokio::test]
async fn a_v4_mapped_peer_is_canonicalized() {
let mapped: std::net::IpAddr = "::ffff:192.0.2.7".parse().unwrap();
let mut req = Request::builder().uri("/test").body(Body::empty()).unwrap();
req.extensions_mut()
.insert(ConnectInfo(SocketAddr::from((mapped, 4711))));
let fields = fields_for(req).await;
assert_eq!(fields.get("client_ip").as_deref(), Some("192.0.2.7"));
}
#[tokio::test]
async fn no_connect_info_leaves_the_client_unnamed() {
let req = Request::builder().uri("/test").body(Body::empty()).unwrap();
let fields = fields_for(req).await;
assert_eq!(fields.get("client_ip"), None);
assert!(fields.get("request_id").is_some());
}
#[tokio::test]
async fn request_id_header_generated_when_absent() {
let req = Request::builder().uri("/test").body(Body::empty()).unwrap();
let res = app().oneshot(req).await.unwrap();
assert!(res.headers().contains_key(&X_REQUEST_ID));
}
#[tokio::test]
async fn request_id_header_preserved_when_present() {
let req = Request::builder()
.uri("/test")
.header("x-request-id", "custom-req-id-123")
.body(Body::empty())
.unwrap();
let res = app().oneshot(req).await.unwrap();
let id = res.headers().get(&X_REQUEST_ID).unwrap().to_str().unwrap();
assert_eq!(id, "custom-req-id-123");
}
#[tokio::test]
async fn a_non_ascii_request_id_is_replaced() {
let req = Request::builder()
.uri("/test")
.header(
"x-request-id",
HeaderValue::from_bytes(&[0xff, 0xfe]).unwrap(),
)
.body(Body::empty())
.unwrap();
let res = app().oneshot(req).await.unwrap();
let id = res.headers().get(&X_REQUEST_ID).unwrap().to_str().unwrap();
assert_eq!(id.len(), Uuid::now_v7().to_string().len());
}
#[tokio::test]
async fn a_blank_request_id_is_replaced() {
let req = Request::builder()
.uri("/test")
.header("x-request-id", " ")
.body(Body::empty())
.unwrap();
let res = app().oneshot(req).await.unwrap();
let id = res.headers().get(&X_REQUEST_ID).unwrap().to_str().unwrap();
assert!(!id.trim().is_empty());
}
#[tokio::test]
async fn an_oversized_request_id_is_replaced() {
let req = Request::builder()
.uri("/test")
.header("x-request-id", "a".repeat(REQUEST_ID_MAX + 1))
.body(Body::empty())
.unwrap();
let res = app().oneshot(req).await.unwrap();
let id = res.headers().get(&X_REQUEST_ID).unwrap().to_str().unwrap();
assert_eq!(id.len(), Uuid::now_v7().to_string().len());
let req = Request::builder()
.uri("/test")
.header("x-request-id", "b".repeat(REQUEST_ID_MAX))
.body(Body::empty())
.unwrap();
let res = app().oneshot(req).await.unwrap();
assert_eq!(
res.headers().get(&X_REQUEST_ID).unwrap().to_str().unwrap(),
"b".repeat(REQUEST_ID_MAX)
);
}
#[tokio::test]
async fn a_request_id_that_could_forge_log_structure_is_replaced() {
let generated = Uuid::now_v7().to_string().len();
for forged in [
"abc outcome=success",
"abc=def",
"abc\tstatus=200",
"abc\"quoted\"",
] {
let Ok(header) = HeaderValue::from_str(forged) else {
continue;
};
let req = Request::builder()
.uri("/test")
.header("x-request-id", header)
.body(Body::empty())
.unwrap();
let res = app().oneshot(req).await.unwrap();
let id = res.headers().get(&X_REQUEST_ID).unwrap().to_str().unwrap();
assert_eq!(
id.len(),
generated,
"`{forged}` must not be adopted verbatim"
);
}
}
#[tokio::test]
async fn ordinary_correlation_ids_are_still_honoured() {
for ordinary in [
"550e8400-e29b-41d4-a716-446655440000",
"custom-req-id-123",
"trace.4bf92f:span_1",
] {
let req = Request::builder()
.uri("/test")
.header("x-request-id", ordinary)
.body(Body::empty())
.unwrap();
let res = app().oneshot(req).await.unwrap();
assert_eq!(
res.headers().get(&X_REQUEST_ID).unwrap().to_str().unwrap(),
ordinary
);
}
}
#[tokio::test]
async fn every_access_line_level_is_reachable() {
for (uri, expected) in [
("/test", StatusCode::OK),
("/health", StatusCode::OK),
("/boom", StatusCode::INTERNAL_SERVER_ERROR),
] {
let req = Request::builder().uri(uri).body(Body::empty()).unwrap();
let res = app().oneshot(req).await.unwrap();
assert_eq!(res.status(), expected);
}
}
#[test]
fn probe_routes_are_recognized() {
assert!(is_probe(&Method::GET, "/health"));
assert!(is_probe(&Method::HEAD, "/health"));
assert!(is_probe(&Method::GET, "/"));
assert!(!is_probe(&Method::POST, "/health"));
assert!(!is_probe(&Method::GET, "/directory"));
}
}