Skip to main content

id_effect_opentelemetry/
axum.rs

1//! Axum middleware: W3C trace context extraction and server spans.
2
3use axum::extract::Request;
4use axum::middleware::Next;
5use axum::response::Response;
6use opentelemetry::Context;
7use tracing_opentelemetry::OpenTelemetrySpanExt;
8
9use crate::propagation::extract_trace_context_from_headers;
10
11/// Create a server span from incoming W3C headers and run the rest of the stack under it.
12pub async fn trace_request(request: Request, next: Next) -> Response {
13  let pairs: Vec<(String, String)> = request
14    .headers()
15    .iter()
16    .filter_map(|(k, v)| Some((k.as_str().to_string(), v.to_str().ok()?.to_string())))
17    .collect();
18  let parent = extract_trace_context_from_headers(&Context::new(), &pairs);
19  let span = tracing::info_span!("http.request", otel.name = %request.uri().path());
20  let _ = span.set_parent(parent);
21  let _guard = span.enter();
22  next.run(request).await
23}
24
25#[cfg(test)]
26mod tests {
27  use super::*;
28  use axum::Router;
29  use axum::body::Body;
30  use axum::routing::get;
31  use tower::ServiceExt;
32
33  #[tokio::test]
34  async fn trace_request_middleware_runs() {
35    let app = Router::new()
36      .route("/", get(|| async { "ok" }))
37      .layer(axum::middleware::from_fn(trace_request));
38    let res = app
39      .oneshot(Request::builder().uri("/").body(Body::empty()).unwrap())
40      .await
41      .unwrap();
42    assert_eq!(res.status(), axum::http::StatusCode::OK);
43  }
44}