use std::{ io, sync::{Arc, Mutex}, }; use admin_api::request_context::{REQUEST_ID_HEADER, TRACE_ID_HEADER, apply_request_context}; use axum::{ Router, body::{Body, to_bytes}, http::{HeaderMap, HeaderValue, Request, StatusCode}, routing::get, }; use crank_observability::{ ObservabilityConfig, RedactionLimits, ServiceIdentity, inject_current_trace_context, }; use opentelemetry::{global, trace::TracerProvider as _}; use opentelemetry_sdk::{propagation::TraceContextPropagator, trace::SdkTracerProvider}; use tower::ServiceExt; use tracing_subscriber::{fmt::MakeWriter, layer::SubscriberExt}; use uuid::Version; static TRACING_TEST_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(()); #[tokio::test(flavor = "current_thread")] async fn logs_request_completion_and_rejects_untrusted_values() { let _tracing_test_guard = TRACING_TEST_LOCK.lock().await; let writer = SharedLogWriter::default(); let subscriber = crank_observability::build_subscriber( ObservabilityConfig::new( ServiceIdentity::try_new("admin-api", "test", "test").unwrap(), "info", RedactionLimits::default(), ), writer.clone(), ) .unwrap(); let dispatch = tracing::Dispatch::new(subscriber); let _dispatch_guard = tracing::dispatcher::set_default(&dispatch); let app = probe_app(); let response = app .clone() .oneshot( Request::builder() .uri("/probe") .header(REQUEST_ID_HEADER.as_str(), "req_admin_trace_123") .body(Body::empty()) .unwrap(), ) .await .unwrap(); assert_eq!(response.status(), StatusCode::OK); assert_eq!( response.headers()[REQUEST_ID_HEADER.as_str()] .to_str() .unwrap(), "req_admin_trace_123" ); let event: serde_json::Value = writer .output() .lines() .map(|line| serde_json::from_str(line).unwrap()) .find(|event: &serde_json::Value| event["event"] == "admin.request.completed") .unwrap(); assert_eq!(event["request_id"], "req_admin_trace_123"); let trace_id = response.headers()[TRACE_ID_HEADER.as_str()] .to_str() .unwrap(); assert_eq!(trace_id.len(), 32); assert_eq!(event["trace_id"], trace_id); assert_eq!(event["fields"]["status"], 200); assert_eq!(event["fields"]["route"], "/probe"); let invalid_response = app .oneshot( Request::builder() .uri("/probe") .header(REQUEST_ID_HEADER.as_str(), "bad,value") .header("traceparent", "canary-invalid-traceparent") .body(Body::empty()) .unwrap(), ) .await .unwrap(); let generated = invalid_response.headers()[REQUEST_ID_HEADER.as_str()] .to_str() .unwrap(); assert_eq!( uuid::Uuid::parse_str(generated).unwrap().get_version(), Some(Version::SortRand) ); let generated_trace = invalid_response.headers()[TRACE_ID_HEADER.as_str()] .to_str() .unwrap(); assert_eq!(generated_trace.len(), 32); assert_ne!(generated_trace, "canary-invalid-traceparent"); assert!(!writer.output().contains("canary-invalid-traceparent")); } #[tokio::test(flavor = "current_thread")] async fn structured_boundary_error_carries_the_same_safe_ids() { let _tracing_test_guard = TRACING_TEST_LOCK.lock().await; let dispatch = tracing::Dispatch::new(tracing_subscriber::registry()); let _dispatch_guard = tracing::dispatcher::set_default(&dispatch); let response = error_probe_app() .oneshot( Request::builder() .uri("/error") .header("x-request-id", "request-boundary-error") .header( "traceparent", "00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01", ) .body(Body::empty()) .unwrap(), ) .await .unwrap(); assert_eq!(response.status(), StatusCode::BAD_REQUEST); assert_eq!( response.headers()[REQUEST_ID_HEADER], "request-boundary-error" ); assert_eq!( response.headers()[TRACE_ID_HEADER], "0af7651916cd43dd8448eb211c80319c" ); let body = to_bytes(response.into_body(), 16 * 1024).await.unwrap(); let payload: serde_json::Value = serde_json::from_slice(&body).unwrap(); assert_eq!(payload["error"]["request_id"], "request-boundary-error"); assert_eq!( payload["error"]["trace_id"], "0af7651916cd43dd8448eb211c80319c" ); } #[tokio::test(flavor = "current_thread")] async fn boundary_status_matrix_keeps_ids_and_redacts_internal_causes() { let _tracing_test_guard = TRACING_TEST_LOCK.lock().await; let writer = SharedLogWriter::default(); let subscriber = crank_observability::build_subscriber( ObservabilityConfig::new( ServiceIdentity::try_new("admin-api", "test", "test").unwrap(), "info", RedactionLimits::default(), ), writer.clone(), ) .unwrap(); let app = Router::new() .route( "/bad-request", get(|| async { Err::<(), _>(admin_api::error::ApiError::validation("invalid")) }), ) .route( "/unauthorized", get(|| async { Err::<(), _>(admin_api::error::ApiError::unauthorized("unauthorized")) }), ) .route( "/forbidden", get(|| async { Err::<(), _>(admin_api::error::ApiError::forbidden("forbidden")) }), ) .route( "/internal", get(|| async { Err::<(), _>(admin_api::error::ApiError::internal( "postgres://canary-user:canary-password@private-host/database", )) }), ) .route( "/rate-limited", get(|| async { StatusCode::TOO_MANY_REQUESTS }), ) .layer(axum::middleware::from_fn(apply_request_context)); let responses = async { let mut responses = Vec::new(); for (path, status) in [ ("/bad-request", StatusCode::BAD_REQUEST), ("/unauthorized", StatusCode::UNAUTHORIZED), ("/forbidden", StatusCode::FORBIDDEN), ("/missing", StatusCode::NOT_FOUND), ("/rate-limited", StatusCode::TOO_MANY_REQUESTS), ("/internal", StatusCode::INTERNAL_SERVER_ERROR), ] { let response = app .clone() .oneshot( Request::builder() .uri(path) .header(REQUEST_ID_HEADER.as_str(), "matrix-request-id") .body(Body::empty()) .unwrap(), ) .await .unwrap(); assert_eq!(response.status(), status); assert_eq!(response.headers()[REQUEST_ID_HEADER], "matrix-request-id"); assert_eq!(response.headers()[TRACE_ID_HEADER].as_bytes().len(), 32); responses.push(to_bytes(response.into_body(), 16 * 1024).await.unwrap()); } responses }; let dispatch = tracing::Dispatch::new(subscriber); let _dispatch_guard = tracing::dispatcher::set_default(&dispatch); let responses = responses.await; let combined = responses .iter() .flat_map(|body| body.iter().copied()) .collect::>(); assert!( !combined .windows("canary-password".len()) .any(|value| value == b"canary-password") ); assert!(!writer.output().contains("canary-password")); let internal_event: serde_json::Value = writer .output() .lines() .map(|line| serde_json::from_str(line).unwrap()) .find(|event: &serde_json::Value| event["event"] == "admin.response.internal_error") .unwrap(); assert_eq!(internal_event["request_id"], "matrix-request-id"); assert_eq!(internal_event["trace_id"].as_str().unwrap().len(), 32); } #[tokio::test(flavor = "current_thread")] async fn covers_valid_invalid_and_absent_traceparent() { let _tracing_test_guard = TRACING_TEST_LOCK.lock().await; global::set_text_map_propagator(TraceContextPropagator::new()); let provider = SdkTracerProvider::builder().build(); let tracer = provider.tracer("admin-request-context-test"); let subscriber = tracing_subscriber::registry().with(tracing_opentelemetry::layer().with_tracer(tracer)); let dispatch = tracing::Dispatch::new(subscriber); let _dispatch_guard = tracing::dispatcher::set_default(&dispatch); let app = trace_probe_app(); let valid = observed_trace_id( app.clone() .oneshot( Request::builder() .uri("/trace") .header( "traceparent", "00-0af7651916cd43dd8448eb211c80319c-b7ad6b7169203331-01", ) .header(REQUEST_ID_HEADER.as_str(), "request-id-is-separate") .body(Body::empty()) .unwrap(), ) .await .unwrap(), ); let invalid = observed_trace_id( app.clone() .oneshot( Request::builder() .uri("/trace") .header("traceparent", "canary-invalid-traceparent") .body(Body::empty()) .unwrap(), ) .await .unwrap(), ); let absent = observed_trace_id( app.oneshot( Request::builder() .uri("/trace") .body(Body::empty()) .unwrap(), ) .await .unwrap(), ); assert_eq!(valid, "0af7651916cd43dd8448eb211c80319c"); assert_ne!(invalid, valid); assert_ne!(absent, valid); assert_ne!(invalid, absent); provider.shutdown().unwrap(); } #[tokio::test(flavor = "current_thread")] async fn replaces_multiple_request_id_headers_with_one_uuid_v7() { let _tracing_test_guard = TRACING_TEST_LOCK.lock().await; let dispatch = tracing::Dispatch::new(tracing_subscriber::registry()); let _dispatch_guard = tracing::dispatcher::set_default(&dispatch); let mut request = Request::builder() .uri("/probe") .body(Body::empty()) .unwrap(); request.headers_mut().append( REQUEST_ID_HEADER, HeaderValue::from_static("first-request-id"), ); request.headers_mut().append( REQUEST_ID_HEADER, 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 = probe_app().oneshot(request).await.unwrap(); let generated = response.headers()[REQUEST_ID_HEADER].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(Version::SortRand) ); let trace_id = response.headers()[TRACE_ID_HEADER].to_str().unwrap(); assert_ne!(trace_id, "0af7651916cd43dd8448eb211c80319c"); assert_ne!(trace_id, "1af7651916cd43dd8448eb211c80319c"); } fn probe_app() -> Router { Router::new() .route("/probe", get(|| async { "ok" })) .layer(axum::middleware::from_fn(apply_request_context)) } fn error_probe_app() -> Router { Router::new() .route( "/error", get(|| async { Err::<(), _>(admin_api::error::ApiError::validation("invalid")) }), ) .layer(axum::middleware::from_fn(apply_request_context)) } fn trace_probe_app() -> Router { Router::new() .route("/trace", get(observed_traceparent)) .layer(axum::middleware::from_fn(apply_request_context)) } async fn observed_traceparent() -> HeaderMap { let mut trace_headers = HeaderMap::new(); inject_current_trace_context(&mut trace_headers); let mut response_headers = HeaderMap::new(); if let Some(traceparent) = trace_headers.remove("traceparent") { response_headers.insert("x-observed-traceparent", traceparent); } response_headers } fn observed_trace_id(response: axum::response::Response) -> String { let traceparent = response.headers()["x-observed-traceparent"] .to_str() .unwrap(); traceparent[3..35].to_owned() } #[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(()) } }