use axum::{ extract::Request, http::{HeaderMap, HeaderName, HeaderValue}, middleware::Next, response::Response, }; use tracing::info; use uuid::Uuid; pub const REQUEST_ID_HEADER: HeaderName = HeaderName::from_static("x-request-id"); const MAX_REQUEST_ID_LEN: usize = 128; #[derive(Clone, Debug)] pub struct RequestContext { pub request_id: String, } pub async fn apply_request_context(mut request: Request, next: Next) -> Response { let context = RequestContext { request_id: resolve_request_id(request.headers()), }; let method = request.method().clone(); let path = request.uri().path().to_owned(); request.extensions_mut().insert(context.clone()); let mut response = next.run(request).await; info!( request_id = %context.request_id, method = %method, path, status = response.status().as_u16(), "admin request completed" ); if let Ok(value) = HeaderValue::from_str(&context.request_id) { response.headers_mut().insert(REQUEST_ID_HEADER, value); } response } fn resolve_request_id(headers: &HeaderMap) -> String { headers .get(&REQUEST_ID_HEADER) .and_then(|value| value.to_str().ok()) .map(str::trim) .filter(|value| is_valid_request_id(value)) .map(ToOwned::to_owned) .unwrap_or_else(|| Uuid::now_v7().to_string()) } fn is_valid_request_id(value: &str) -> bool { !value.is_empty() && value.len() <= MAX_REQUEST_ID_LEN && value .bytes() .all(|byte| matches!(byte, 0x21..=0x7e) && byte != b',' && byte != b';') } #[cfg(test)] mod tests { use std::io; use std::sync::{Arc, Mutex}; use axum::{Router, routing::get}; use reqwest::Client; use tokio::net::TcpListener; use tracing_subscriber::{filter::LevelFilter, fmt::MakeWriter, prelude::*}; use super::{REQUEST_ID_HEADER, apply_request_context, is_valid_request_id}; #[test] fn accepts_visible_ascii_request_ids() { assert!(is_valid_request_id("req_test_123")); assert!(is_valid_request_id("trace-123/abc")); } #[test] fn rejects_empty_or_control_request_ids() { assert!(!is_valid_request_id("")); assert!(!is_valid_request_id("bad value")); assert!(!is_valid_request_id("bad\nvalue")); } #[derive(Clone, Default)] struct SharedLogWriter { buffer: Arc>>, } impl SharedLogWriter { fn output(&self) -> String { String::from_utf8(self.buffer.lock().unwrap().clone()).unwrap() } } impl<'a> MakeWriter<'a> for SharedLogWriter { type Writer = SharedLogGuard; fn make_writer(&'a self) -> Self::Writer { SharedLogGuard { buffer: Arc::clone(&self.buffer), } } } struct SharedLogGuard { buffer: Arc>>, } impl io::Write for SharedLogGuard { fn write(&mut self, bytes: &[u8]) -> io::Result { self.buffer.lock().unwrap().extend_from_slice(bytes); Ok(bytes.len()) } fn flush(&mut self) -> io::Result<()> { Ok(()) } } #[tokio::test] async fn logs_request_completion_with_request_id() { let writer = SharedLogWriter::default(); let subscriber = tracing_subscriber::registry().with( tracing_subscriber::fmt::layer() .with_writer(writer.clone()) .without_time() .with_ansi(false) .with_target(false) .compact() .with_filter(LevelFilter::INFO), ); let dispatch = tracing::Dispatch::new(subscriber); let app = Router::new() .route("/probe", get(|| async { "ok" })) .layer(axum::middleware::from_fn(apply_request_context)); let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); let address = listener.local_addr().unwrap(); let _guard = tracing::dispatcher::set_default(&dispatch); tokio::spawn(async move { axum::serve(listener, app).await.unwrap(); }); let response = Client::new() .get(format!("http://{address}/probe")) .header(REQUEST_ID_HEADER.as_str(), "req_admin_trace_123") .send() .await .unwrap(); assert_eq!(response.status(), reqwest::StatusCode::OK); assert_eq!( response.headers()[REQUEST_ID_HEADER.as_str()] .to_str() .unwrap(), "req_admin_trace_123" ); let logs = writer.output(); assert!(logs.contains("admin request completed")); assert!(logs.contains("req_admin_trace_123")); assert!(logs.contains("GET")); assert!(logs.contains("/probe")); assert!(logs.contains("status=200")); } }