use axum::{
extract::Request,
http::{header::HeaderValue, HeaderName},
middleware::{from_fn, FromFnLayer, Next},
response::Response,
};
use std::future::Future;
use std::pin::Pin;
use tracing::Instrument;
use uuid::Uuid;
pub const TRACEPARENT_HEADER: &str = "traceparent";
tokio::task_local! {
static CURRENT_TRACE_ID: String;
}
#[derive(Clone, Debug)]
pub struct TraceId(pub String);
type BoxFut = Pin<Box<dyn Future<Output = Response> + Send>>;
type TraceFn = fn(Request, Next) -> BoxFut;
pub fn trace_id_layer() -> FromFnLayer<TraceFn, (), (Request,)> {
from_fn(trace_id_mw as TraceFn)
}
pub fn current_trace_id() -> Option<String> {
CURRENT_TRACE_ID.try_with(|t| t.clone()).ok()
}
pub fn outbound_traceparent() -> Option<String> {
current_trace_id().map(|tid| format_traceparent(&tid))
}
fn format_traceparent(trace_id: &str) -> String {
let span_id = &Uuid::new_v4().simple().to_string()[..16];
format!("00-{trace_id}-{span_id}-01")
}
fn parse_trace_id(traceparent: &str) -> Option<String> {
let parts: [&str; 4] = {
let mut it = traceparent.trim().split('-');
let a = it.next()?;
let b = it.next()?;
let c = it.next()?;
let d = it.next()?;
if it.next().is_some() {
return None; }
[a, b, c, d]
};
let [version, trace_id, parent_id, flags] = parts;
let is_hex = |s: &str, n: usize| s.len() == n && s.bytes().all(|b| b.is_ascii_hexdigit());
if !is_hex(version, 2) || !is_hex(trace_id, 32) || !is_hex(parent_id, 16) || !is_hex(flags, 2) {
return None;
}
if version.eq_ignore_ascii_case("ff") {
return None;
}
let tid = trace_id.to_ascii_lowercase();
if tid.bytes().all(|b| b == b'0') {
return None;
}
Some(tid)
}
fn new_trace_id() -> String {
Uuid::new_v4().simple().to_string()
}
fn trace_id_mw(mut req: Request, next: Next) -> BoxFut {
let trace_id = req
.headers()
.get(TRACEPARENT_HEADER)
.and_then(|v| v.to_str().ok())
.and_then(parse_trace_id)
.unwrap_or_else(new_trace_id);
let method = req.method().clone();
let path = req.uri().path().to_string();
req.extensions_mut().insert(TraceId(trace_id.clone()));
let span = tracing::info_span!(
"request",
trace_id = %trace_id,
method = %method,
path = %path,
);
let response_tp = format_traceparent(&trace_id);
Box::pin(
CURRENT_TRACE_ID
.scope(trace_id, async move {
let mut response = next.run(req).await;
if let Ok(value) = HeaderValue::from_str(&response_tp) {
response
.headers_mut()
.insert(HeaderName::from_static(TRACEPARENT_HEADER), value);
}
response
})
.instrument(span),
)
}
#[cfg(test)]
mod tests {
use super::*;
use axum::{body::Body, http::Request as HttpRequest, routing::get, Router};
use std::sync::{Arc, Mutex};
use tower::ServiceExt;
#[test]
fn parse_trace_id_accepts_valid_traceparent() {
let tp = "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01";
assert_eq!(
parse_trace_id(tp).as_deref(),
Some("4bf92f3577b34da6a3ce929d0e0e4736")
);
}
#[test]
fn parse_trace_id_rejects_malformed_and_zero() {
assert_eq!(parse_trace_id("garbage"), None);
assert_eq!(parse_trace_id("00-abc-00f067aa0ba902b7-01"), None); assert_eq!(
parse_trace_id("ff-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01"),
None
); assert_eq!(
parse_trace_id("00-00000000000000000000000000000000-00f067aa0ba902b7-01"),
None
); }
#[derive(Clone)]
struct BufWriter(Arc<Mutex<Vec<u8>>>);
impl std::io::Write for BufWriter {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.0.lock().unwrap().extend_from_slice(buf);
Ok(buf.len())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
impl<'a> tracing_subscriber::fmt::MakeWriter<'a> for BufWriter {
type Writer = BufWriter;
fn make_writer(&'a self) -> Self::Writer {
self.clone()
}
}
async fn handler() -> &'static str {
tracing::info!("handling request");
"ok"
}
fn app() -> Router {
Router::new()
.route("/x", get(handler))
.layer(trace_id_layer())
}
fn resp_trace_id(header: Option<&str>) -> Option<String> {
header.and_then(parse_trace_id)
}
#[test]
fn incoming_traceparent_flows_into_logs_and_response() {
let buf = Arc::new(Mutex::new(Vec::new()));
let subscriber = tracing_subscriber::fmt()
.json()
.flatten_event(true)
.with_current_span(true)
.with_span_list(true)
.with_max_level(tracing::Level::TRACE)
.with_writer(BufWriter(buf.clone()))
.finish();
tracing::subscriber::set_global_default(subscriber)
.expect("no other test should set a global subscriber");
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
let incoming = "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01";
let echoed = rt.block_on(async {
let req = HttpRequest::builder()
.uri("/x")
.header(TRACEPARENT_HEADER, incoming)
.body(Body::empty())
.unwrap();
let resp = app().oneshot(req).await.unwrap();
resp.headers()
.get(TRACEPARENT_HEADER)
.and_then(|v| v.to_str().ok())
.map(String::from)
});
assert_eq!(
resp_trace_id(echoed.as_deref()).as_deref(),
Some("4bf92f3577b34da6a3ce929d0e0e4736")
);
let logs = String::from_utf8(buf.lock().unwrap().clone()).unwrap();
assert!(
logs.contains("\"trace_id\":\"4bf92f3577b34da6a3ce929d0e0e4736\""),
"expected trace_id in JSON logs, got: {logs}"
);
assert!(logs.contains("handling request"));
}
#[test]
fn missing_header_generates_root_trace() {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap();
let echoed = rt.block_on(async {
let req = HttpRequest::builder()
.uri("/x")
.body(Body::empty())
.unwrap();
let resp = app().oneshot(req).await.unwrap();
resp.headers()
.get(TRACEPARENT_HEADER)
.and_then(|v| v.to_str().ok())
.map(String::from)
});
let tid = resp_trace_id(echoed.as_deref());
assert!(
tid.as_deref().map(|t| t.len() == 32).unwrap_or(false),
"expected a generated 32-hex trace id, got: {echoed:?}"
);
}
}