use std::{ sync::{Arc, Mutex}, time::Duration, }; use axum::{ body::Body, http::{HeaderValue, Request, StatusCode}, }; use opentelemetry::{ global, trace::{SpanId, TraceId, TracerProvider as _}, }; use opentelemetry_sdk::{ error::OTelSdkResult, propagation::TraceContextPropagator, trace::{SdkTracerProvider, SpanData, SpanExporter}, }; use tower::ServiceExt; use tracing::instrument::WithSubscriber; use tracing_subscriber::layer::SubscriberExt; use super::common::{build_test_app, test_registry}; const REMOTE_TRACE_ID: &str = "0af7651916cd43dd8448eb211c80319c"; static TRACING_TEST_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(()); #[tokio::test(flavor = "current_thread")] async fn covers_valid_invalid_and_absent_traceparent_on_mcp_boundary() { let _tracing_test_guard = TRACING_TEST_LOCK.lock().await; global::set_text_map_propagator(TraceContextPropagator::new()); let exported = Arc::new(Mutex::new(Vec::new())); let provider = SdkTracerProvider::builder() .with_simple_exporter(CapturingExporter(Arc::clone(&exported))) .build(); let tracer = provider.tracer("mcp-request-context-test"); let subscriber = tracing_subscriber::registry().with(tracing_opentelemetry::layer().with_tracer(tracer)); let dispatch = tracing::Dispatch::new(subscriber); let app = build_test_app(test_registry().await, Duration::ZERO, None); let (valid, invalid, absent) = async { let valid = send_health( app.clone(), Some("00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01"), Some("request-id-is-separate"), ) .await; let invalid = send_health( app.clone(), Some("canary-invalid-traceparent"), Some("bad,value"), ) .await; let absent = send_health(app, None, None).await; (valid, invalid, absent) } .with_subscriber(dispatch) .await; provider.force_flush().unwrap(); assert_eq!(valid.status, StatusCode::OK); assert_eq!(valid.request_id.as_deref(), Some("request-id-is-separate")); assert_eq!(valid.trace_id.as_deref(), Some(REMOTE_TRACE_ID)); assert_eq!(invalid.status, StatusCode::OK); assert_eq!(absent.status, StatusCode::OK); assert!(valid.traceparent_response.is_none()); assert!(invalid.traceparent_response.is_none()); assert!(absent.traceparent_response.is_none()); let trace_ids: Vec<_> = exported .lock() .unwrap() .iter() .filter(|span| span.name.as_ref() == "mcp.request") .map(|span| span.span_context.trace_id()) .collect(); assert_eq!(trace_ids.len(), 3); assert_eq!(trace_ids[0].to_string(), REMOTE_TRACE_ID); assert_ne!(trace_ids[1], trace_ids[0]); assert_ne!(trace_ids[2], trace_ids[0]); assert_ne!(trace_ids[1], trace_ids[2]); assert!(!trace_ids.contains(&TraceId::INVALID)); let request_spans = exported .lock() .unwrap() .iter() .filter(|span| span.name.as_ref() == "mcp.request") .cloned() .collect::>(); assert_eq!( request_spans[0].parent_span_id.to_string(), "b7ad6b7169203331" ); assert_eq!(request_spans[1].parent_span_id, SpanId::INVALID); assert_eq!(request_spans[2].parent_span_id, SpanId::INVALID); provider.shutdown().unwrap(); } #[tokio::test(flavor = "current_thread")] async fn sampling_off_still_returns_a_local_trace_identity() { let _tracing_test_guard = TRACING_TEST_LOCK.lock().await; global::set_text_map_propagator(TraceContextPropagator::new()); let provider = SdkTracerProvider::builder() .with_sampler(opentelemetry_sdk::trace::Sampler::AlwaysOff) .build(); let tracer = provider.tracer("mcp-request-context-sampling-off-test"); let subscriber = tracing_subscriber::registry().with(tracing_opentelemetry::layer().with_tracer(tracer)); let app = build_test_app(test_registry().await, Duration::ZERO, None); let response = send_health(app, None, None) .with_subscriber(subscriber) .await; let trace_id = response.trace_id.expect("local trace id"); assert_eq!(trace_id.len(), 32); assert_ne!(trace_id, "00000000000000000000000000000000"); provider.shutdown().unwrap(); } #[tokio::test] async fn replaces_multiple_request_id_headers_with_one_uuid_v7() { let _tracing_test_guard = TRACING_TEST_LOCK.lock().await; let app = build_test_app(test_registry().await, Duration::ZERO, None); let mut request = Request::builder() .uri("/health") .body(Body::empty()) .unwrap(); request .headers_mut() .append("x-request-id", HeaderValue::from_static("first-request-id")); request.headers_mut().append( "x-request-id", HeaderValue::from_static("second-request-id"), ); request.headers_mut().append( "traceparent", HeaderValue::from_static("00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01"), ); request.headers_mut().append( "traceparent", HeaderValue::from_static("00-1af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01"), ); let response = app .oneshot(request) .with_subscriber(tracing_subscriber::registry()) .await .unwrap(); let generated = response.headers()["x-request-id"].to_str().unwrap(); assert_ne!(generated, "first-request-id"); assert_ne!(generated, "second-request-id"); assert_eq!( uuid::Uuid::parse_str(generated).unwrap().get_version(), Some(uuid::Version::SortRand) ); let trace_id = response.headers()["x-trace-id"].to_str().unwrap(); assert_ne!(trace_id, "0af7651916cd43dd8448eb211c80319c"); assert_ne!(trace_id, "1af7651916cd43dd8448eb211c80319c"); } async fn send_health( app: axum::Router, traceparent: Option<&str>, request_id: Option<&str>, ) -> ProbeResponse { let mut request = Request::builder().uri("/health"); if let Some(traceparent) = traceparent { request = request.header("traceparent", traceparent); } if let Some(request_id) = request_id { request = request.header("x-request-id", request_id); } let response = app .oneshot(request.body(Body::empty()).unwrap()) .await .unwrap(); ProbeResponse { status: response.status(), request_id: response .headers() .get("x-request-id") .and_then(|value| value.to_str().ok()) .map(str::to_owned), traceparent_response: response .headers() .get("traceparent") .and_then(|value| value.to_str().ok()) .map(str::to_owned), trace_id: response .headers() .get("x-trace-id") .and_then(|value| value.to_str().ok()) .map(str::to_owned), } } struct ProbeResponse { status: StatusCode, request_id: Option, traceparent_response: Option, trace_id: Option, } #[derive(Clone, Debug)] struct CapturingExporter(Arc>>); impl SpanExporter for CapturingExporter { async fn export(&self, batch: Vec) -> OTelSdkResult { self.0.lock().unwrap().extend(batch); Ok(()) } }