use axum::{ extract::{MatchedPath, Request}, http::{HeaderName, HeaderValue}, middleware::Next, response::Response, }; use crank_core::{CorrelationContext, RequestId, TraceContext}; use crank_metrics::ExemplarTraceId; use crank_observability::with_request_correlation; use tracing::{Instrument, info, info_span}; pub const REQUEST_ID_HEADER: HeaderName = HeaderName::from_static("x-request-id"); pub const TRACE_ID_HEADER: HeaderName = HeaderName::from_static("x-trace-id"); #[derive(Clone, Debug)] pub struct RequestContext { pub correlation: CorrelationContext, } impl RequestContext { pub fn request_id(&self) -> &str { self.correlation.request_id().as_str() } pub fn trace_id(&self) -> &str { self.correlation.trace_id().as_str() } } pub async fn apply_request_context(mut request: Request, next: Next) -> Response { let (request_id, remote_parent) = resolve_correlation(request.headers()); let method = request.method().clone(); let route = request .extensions() .get::() .map_or("unmatched", MatchedPath::as_str) .to_owned(); let span = info_span!( target: "crank::trace", "http.request", request_id = %request_id, trace_id = tracing::field::Empty, ); if let Some(remote_parent) = remote_parent.as_ref() { set_canonical_parent(&span, remote_parent); } let trace_context = crank_trace::trace_context_for_span(&span).unwrap_or_else(|| { remote_parent .as_ref() .map_or_else(TraceContext::generate, TraceContext::continue_local) }); span.record("trace_id", trace_context.trace_id().as_str()); let context = RequestContext { correlation: CorrelationContext::new(request_id, trace_context), }; request.extensions_mut().insert(context.clone()); with_request_correlation( context.correlation.request_id().to_string(), context.correlation.trace_id().to_string(), async move { let mut response = next.run(request).instrument(span).await; info!( name: "admin.request.completed", request_id = %context.correlation.request_id(), trace_id = %context.correlation.trace_id(), method = %method, route, status = response.status().as_u16(), "admin request completed" ); if let Ok(value) = HeaderValue::from_str(context.correlation.request_id().as_str()) { response.headers_mut().insert(REQUEST_ID_HEADER, value); } if let Ok(value) = HeaderValue::from_str(context.correlation.trace_id().as_str()) { response.headers_mut().insert(TRACE_ID_HEADER, value); } if context.correlation.trace_context().is_sampled() && let Some(exemplar) = ExemplarTraceId::parse(context.correlation.trace_id().as_str()) { response.extensions_mut().insert(exemplar); } response }, ) .await } fn resolve_correlation(headers: &axum::http::HeaderMap) -> (RequestId, Option) { let _tracestate_accepted = one_auxiliary_header_within_budget( headers, "tracestate", TraceContext::tracestate_within_budget, ); let _baggage_accepted = one_auxiliary_header_within_budget(headers, "baggage", TraceContext::baggage_within_budget); let mut request_ids = headers.get_all(REQUEST_ID_HEADER).iter(); let request_id = request_ids.next().and_then(|value| value.to_str().ok()); let request_id = if request_ids.next().is_some() { RequestId::generate() } else { RequestId::resolve(request_id) }; let mut traceparents = headers.get_all("traceparent").iter(); let traceparent = traceparents.next().and_then(|value| value.to_str().ok()); let remote_parent = if traceparents.next().is_some() { None } else { traceparent.and_then(|value| TraceContext::parse(value).ok()) }; (request_id, remote_parent) } fn one_auxiliary_header_within_budget( headers: &axum::http::HeaderMap, name: &'static str, validate: fn(&str) -> bool, ) -> bool { let mut values = headers.get_all(name).iter(); let value = values.next().and_then(|value| value.to_str().ok()); values.next().is_none() && value.is_some_and(validate) } fn set_canonical_parent(span: &tracing::Span, context: &TraceContext) { crank_trace::set_parent_from_trace_context(span, context); } #[cfg(test)] mod tests { #[test] fn accepts_visible_ascii_request_ids() { assert!(crank_core::RequestId::is_valid("req_test_123")); assert!(crank_core::RequestId::is_valid("trace-123/abc")); } #[test] fn rejects_empty_or_control_request_ids() { assert!(!crank_core::RequestId::is_valid("")); assert!(!crank_core::RequestId::is_valid("bad value")); assert!(!crank_core::RequestId::is_valid("bad\nvalue")); } }