From e9e5bea81b680931b32a987ec00e836f05b37caa Mon Sep 17 00:00:00 2001 From: nicholasclark-openai Date: Wed, 25 Mar 2026 12:37:11 -0700 Subject: [PATCH] Propagate RMCP trace context Add traceparent/tracestate propagation for RMCP streamable HTTP and stdio calls, and cover it with transport-level tests. Co-authored-by: Codex --- codex-rs/Cargo.lock | 5 + codex-rs/core/tests/suite/rmcp_client.rs | 4 + codex-rs/rmcp-client/Cargo.toml | 5 + .../rmcp-client/src/bin/test_stdio_server.rs | 12 +- .../src/bin/test_streamable_http_server.rs | 23 ++ codex-rs/rmcp-client/src/rmcp_client.rs | 235 ++++++++++++++++-- codex-rs/rmcp-client/tests/resources.rs | 85 +++++++ 7 files changed, 354 insertions(+), 15 deletions(-) diff --git a/codex-rs/Cargo.lock b/codex-rs/Cargo.lock index fb1710987e..e12b3c7cba 100644 --- a/codex-rs/Cargo.lock +++ b/codex-rs/Cargo.lock @@ -2514,6 +2514,7 @@ dependencies = [ "axum", "codex-client", "codex-keyring-store", + "codex-otel", "codex-protocol", "codex-utils-cargo-bin", "codex-utils-home-dir", @@ -2521,6 +2522,8 @@ dependencies = [ "futures", "keyring", "oauth2", + "opentelemetry", + "opentelemetry_sdk", "pretty_assertions", "reqwest", "rmcp", @@ -2535,6 +2538,8 @@ dependencies = [ "tiny_http", "tokio", "tracing", + "tracing-opentelemetry", + "tracing-subscriber", "urlencoding", "webbrowser", "which 8.0.0", diff --git a/codex-rs/core/tests/suite/rmcp_client.rs b/codex-rs/core/tests/suite/rmcp_client.rs index 6cbf9521ba..12f9569077 100644 --- a/codex-rs/core/tests/suite/rmcp_client.rs +++ b/codex-rs/core/tests/suite/rmcp_client.rs @@ -36,6 +36,7 @@ use core_test_support::responses::mount_sse_once; use core_test_support::skip_if_no_network; use core_test_support::stdio_server_bin; use core_test_support::test_codex::test_codex; +use core_test_support::tracing::install_test_tracing; use core_test_support::wait_for_event; use core_test_support::wait_for_event_with_timeout; use reqwest::Client; @@ -687,6 +688,8 @@ async fn stdio_server_propagates_whitelisted_env_vars() -> anyhow::Result<()> { async fn streamable_http_tool_call_round_trip() -> anyhow::Result<()> { skip_if_no_network!(Ok(())); + let _trace_test_context = install_test_tracing("rmcp-integration-tests"); + let server = responses::start_mock_server().await; let call_id = "call-456"; @@ -733,6 +736,7 @@ async fn streamable_http_tool_call_round_trip() -> anyhow::Result<()> { .kill_on_drop(true) .env("MCP_STREAMABLE_HTTP_BIND_ADDR", &bind_addr) .env("MCP_TEST_VALUE", expected_env_value) + .env("MCP_EXPECT_TRACEPARENT", "1") .spawn()?; wait_for_streamable_http_server(&mut http_server_child, &bind_addr, Duration::from_secs(5)) diff --git a/codex-rs/rmcp-client/Cargo.toml b/codex-rs/rmcp-client/Cargo.toml index 4b20e9d6eb..3316549299 100644 --- a/codex-rs/rmcp-client/Cargo.toml +++ b/codex-rs/rmcp-client/Cargo.toml @@ -15,6 +15,7 @@ axum = { workspace = true, default-features = false, features = [ ] } codex-client = { workspace = true } codex-keyring-store = { workspace = true } +codex-otel = { workspace = true } codex-protocol = { workspace = true } codex-utils-pty = { workspace = true } codex-utils-home-dir = { workspace = true } @@ -60,9 +61,13 @@ which = { workspace = true } [dev-dependencies] codex-utils-cargo-bin = { workspace = true } +opentelemetry = { workspace = true } +opentelemetry_sdk = { workspace = true } pretty_assertions = { workspace = true } serial_test = { workspace = true } tempfile = { workspace = true } +tracing-opentelemetry = { workspace = true } +tracing-subscriber = { workspace = true } [target.'cfg(target_os = "linux")'.dependencies] keyring = { workspace = true, features = ["linux-native-async-persistent"] } diff --git a/codex-rs/rmcp-client/src/bin/test_stdio_server.rs b/codex-rs/rmcp-client/src/bin/test_stdio_server.rs index d9202babae..ba356ce1f7 100644 --- a/codex-rs/rmcp-client/src/bin/test_stdio_server.rs +++ b/codex-rs/rmcp-client/src/bin/test_stdio_server.rs @@ -315,7 +315,7 @@ impl ServerHandler for TestToolServer { async fn call_tool( &self, request: CallToolRequestParams, - _context: rmcp::service::RequestContext, + context: rmcp::service::RequestContext, ) -> Result { match request.name.as_ref() { "echo" | "echo-tool" => { @@ -333,9 +333,19 @@ impl ServerHandler for TestToolServer { }; let env_snapshot: HashMap = std::env::vars().collect(); + let traceparent = context + .meta + .get("x-codex-traceparent") + .and_then(serde_json::Value::as_str); + let tracestate = context + .meta + .get("x-codex-tracestate") + .and_then(serde_json::Value::as_str); let structured_content = json!({ "echo": format!("ECHOING: {}", args.message), "env": env_snapshot.get("MCP_TEST_VALUE"), + "traceparent": traceparent, + "tracestate": tracestate, }); Ok(CallToolResult { diff --git a/codex-rs/rmcp-client/src/bin/test_streamable_http_server.rs b/codex-rs/rmcp-client/src/bin/test_streamable_http_server.rs index 50ea06c7fa..af9ba7f978 100644 --- a/codex-rs/rmcp-client/src/bin/test_streamable_http_server.rs +++ b/codex-rs/rmcp-client/src/bin/test_streamable_http_server.rs @@ -58,6 +58,7 @@ struct TestToolServer { const MEMO_URI: &str = "memo://codex/example-note"; const MEMO_CONTENT: &str = "This is a sample MCP resource served by the rmcp test server."; const MCP_SESSION_ID_HEADER: &str = "mcp-session-id"; +const TRACEPARENT_HEADER: &str = "traceparent"; const SESSION_POST_FAILURE_CONTROL_PATH: &str = "/test/control/session-post-failure"; impl TestToolServer { @@ -347,6 +348,11 @@ async fn main() -> Result<(), Box> { } else { router }; + let router = if std::env::var("MCP_EXPECT_TRACEPARENT").is_ok() { + router.layer(middleware::from_fn(require_traceparent_on_session_post)) + } else { + router + }; axum::serve(listener, router).await?; task::yield_now().await; @@ -389,6 +395,23 @@ async fn arm_session_post_failure( Ok(StatusCode::NO_CONTENT) } +async fn require_traceparent_on_session_post(request: Request, next: Next) -> Response { + if request.uri().path() != "/mcp" + || request.method() != Method::POST + || !request.headers().contains_key(MCP_SESSION_ID_HEADER) + { + return next.run(request).await; + } + + if request.headers().contains_key(TRACEPARENT_HEADER) { + next.run(request).await + } else { + let mut response = Response::new(Body::from("missing traceparent header")); + *response.status_mut() = StatusCode::BAD_REQUEST; + response + } +} + async fn fail_session_post_when_armed( State(state): State, request: Request, diff --git a/codex-rs/rmcp-client/src/rmcp_client.rs b/codex-rs/rmcp-client/src/rmcp_client.rs index aa460c21b9..3ade1177ac 100644 --- a/codex-rs/rmcp-client/src/rmcp_client.rs +++ b/codex-rs/rmcp-client/src/rmcp_client.rs @@ -10,6 +10,9 @@ use std::time::Duration; use anyhow::Result; use anyhow::anyhow; use codex_client::build_reqwest_client_with_custom_ca; +use codex_otel::current_span_w3c_trace_context; +use codex_otel::span_w3c_trace_context; +use codex_protocol::protocol::W3cTraceContext; use futures::FutureExt; use futures::StreamExt; use futures::future::BoxFuture; @@ -64,6 +67,8 @@ use tokio::io::BufReader; use tokio::process::Command; use tokio::sync::Mutex; use tokio::time; +use tracing::Instrument; +use tracing::field::Empty; use tracing::info; use tracing::warn; @@ -82,6 +87,14 @@ const JSON_MIME_TYPE: &str = "application/json"; const HEADER_LAST_EVENT_ID: &str = "Last-Event-Id"; const HEADER_SESSION_ID: &str = "Mcp-Session-Id"; const NON_JSON_RESPONSE_BODY_PREVIEW_BYTES: usize = 8_192; +const TRACEPARENT_HEADER: &str = "traceparent"; +const TRACESTATE_HEADER: &str = "tracestate"; +const TRACEPARENT_META_KEY: &str = "x-codex-traceparent"; +const TRACESTATE_META_KEY: &str = "x-codex-tracestate"; + +tokio::task_local! { + static SERVICE_OPERATION_TRACE_CONTEXT: Option; +} #[derive(Clone)] struct StreamableHttpResponseClient { @@ -98,6 +111,24 @@ impl StreamableHttpResponseClient { ) -> StreamableHttpError { StreamableHttpError::Client(StreamableHttpResponseClientError::from(error)) } + + fn apply_trace_context( + request: reqwest::RequestBuilder, + trace: Option<&W3cTraceContext>, + ) -> reqwest::RequestBuilder { + let Some(trace) = trace else { + return request; + }; + + let mut request = request; + if let Some(traceparent) = trace.traceparent.as_deref() { + request = request.header(TRACEPARENT_HEADER, traceparent); + } + if let Some(tracestate) = trace.tracestate.as_deref() { + request = request.header(TRACESTATE_HEADER, tracestate); + } + request + } } fn build_http_client(default_headers: &HeaderMap) -> Result { @@ -123,6 +154,7 @@ impl StreamableHttpClient for StreamableHttpResponseClient { session_id: Option>, auth_token: Option, ) -> std::result::Result> { + let trace = current_service_operation_trace_context(); let mut request = self .inner .post(uri.as_ref()) @@ -133,6 +165,7 @@ impl StreamableHttpClient for StreamableHttpResponseClient { if let Some(session_id_value) = session_id.as_ref() { request = request.header(HEADER_SESSION_ID, session_id_value.as_ref()); } + request = Self::apply_trace_context(request, trace.as_ref()); let response = request .json(&message) @@ -224,10 +257,12 @@ impl StreamableHttpClient for StreamableHttpResponseClient { session: Arc, auth_token: Option, ) -> std::result::Result<(), StreamableHttpError> { + let trace = current_service_operation_trace_context(); let mut request_builder = self.inner.delete(uri.as_ref()); if let Some(auth_header) = auth_token { request_builder = request_builder.bearer_auth(auth_header); } + request_builder = Self::apply_trace_context(request_builder, trace.as_ref()); let response = request_builder .header(HEADER_SESSION_ID, session.as_ref()) .send() @@ -254,6 +289,7 @@ impl StreamableHttpClient for StreamableHttpResponseClient { BoxStream<'static, std::result::Result>, StreamableHttpError, > { + let trace = current_service_operation_trace_context(); let mut request_builder = self .inner .get(uri.as_ref()) @@ -265,6 +301,7 @@ impl StreamableHttpClient for StreamableHttpResponseClient { if let Some(auth_header) = auth_token { request_builder = request_builder.bearer_auth(auth_header); } + request_builder = Self::apply_trace_context(request_builder, trace.as_ref()); let response = request_builder .send() @@ -733,6 +770,8 @@ impl RmcpClient { let rmcp_params = rmcp_params.clone(); let meta = meta.clone(); async move { + let trace = current_service_operation_trace_context(); + let meta = merge_trace_context_into_meta(meta, trace.as_ref()); let result = service .peer() .send_request_with_option( @@ -1052,41 +1091,104 @@ impl RmcpClient { Fut: std::future::Future>, { let service = self.service().await?; - match Self::run_service_operation_once(Arc::clone(&service), label, timeout, &operation) - .await + let operation_span = self.service_operation_span(label); + let operation_trace = span_w3c_trace_context(&operation_span); + match Self::run_service_operation_once( + Arc::clone(&service), + label, + timeout, + operation_span, + operation_trace, + &operation, + ) + .await { Ok(result) => Ok(result), Err(error) if Self::is_session_expired_404(&error) => { self.reinitialize_after_session_expiry(&service).await?; let recovered_service = self.service().await?; - Self::run_service_operation_once(recovered_service, label, timeout, &operation) - .await - .map_err(Into::into) + let operation_span = self.service_operation_span(label); + let operation_trace = span_w3c_trace_context(&operation_span); + Self::run_service_operation_once( + recovered_service, + label, + timeout, + operation_span, + operation_trace, + &operation, + ) + .await + .map_err(Into::into) } Err(error) => Err(error.into()), } } + fn service_operation_span(&self, label: &str) -> tracing::Span { + let span = tracing::info_span!( + "mcp.client.operation", + otel.kind = "client", + rpc.system = "jsonrpc", + rpc.method = label, + mcp.transport = Empty, + mcp.server.name = Empty, + server.address = Empty, + server.port = Empty, + ); + + match &self.transport_recipe { + TransportRecipe::Stdio { .. } => { + span.record("mcp.transport", "stdio"); + } + TransportRecipe::StreamableHttp { + server_name, url, .. + } => { + span.record("mcp.transport", "streamable_http"); + span.record("mcp.server.name", server_name.as_str()); + if let Ok(parsed_url) = reqwest::Url::parse(url) { + if let Some(host) = parsed_url.host_str() { + span.record("server.address", host); + } + if let Some(port) = parsed_url.port_or_known_default() { + span.record("server.port", port as i64); + } + } + } + } + + span + } + async fn run_service_operation_once( service: Arc>, label: &str, timeout: Option, + operation_span: tracing::Span, + operation_trace: Option, operation: &F, ) -> std::result::Result where F: Fn(Arc>) -> Fut, Fut: std::future::Future>, { - match timeout { - Some(duration) => time::timeout(duration, operation(service)) + SERVICE_OPERATION_TRACE_CONTEXT + .scope(operation_trace, async move { + async move { + match timeout { + Some(duration) => time::timeout(duration, operation(service)) + .await + .map_err(|_| ClientOperationError::Timeout { + label: label.to_string(), + duration, + })? + .map_err(ClientOperationError::from), + None => operation(service).await.map_err(ClientOperationError::from), + } + } + .instrument(operation_span) .await - .map_err(|_| ClientOperationError::Timeout { - label: label.to_string(), - duration, - })? - .map_err(ClientOperationError::from), - None => operation(service).await.map_err(ClientOperationError::from), - } + }) + .await } fn is_session_expired_404(error: &ClientOperationError) -> bool { @@ -1161,6 +1263,39 @@ impl RmcpClient { } } +fn merge_trace_context_into_meta( + meta: Option, + trace: Option<&W3cTraceContext>, +) -> Option { + let mut meta = meta.unwrap_or_default(); + let Some(trace) = trace else { + return (!meta.is_empty()).then_some(meta); + }; + + if let Some(traceparent) = trace.traceparent.as_ref() { + meta.insert( + TRACEPARENT_META_KEY.to_string(), + serde_json::Value::String(traceparent.clone()), + ); + } + if let Some(tracestate) = trace.tracestate.as_ref() { + meta.insert( + TRACESTATE_META_KEY.to_string(), + serde_json::Value::String(tracestate.clone()), + ); + } + + (!meta.is_empty()).then_some(meta) +} + +fn current_service_operation_trace_context() -> Option { + SERVICE_OPERATION_TRACE_CONTEXT + .try_with(Clone::clone) + .ok() + .flatten() + .or_else(current_span_w3c_trace_context) +} + async fn create_oauth_transport_and_runtime( server_name: &str, url: &str, @@ -1207,3 +1342,75 @@ async fn create_oauth_transport_and_runtime( Ok((transport, runtime)) } + +#[cfg(test)] +mod trace_tests { + use super::StreamableHttpResponseClient; + use super::TRACEPARENT_HEADER; + use super::TRACEPARENT_META_KEY; + use super::TRACESTATE_HEADER; + use super::TRACESTATE_META_KEY; + use super::merge_trace_context_into_meta; + use codex_protocol::protocol::W3cTraceContext; + use pretty_assertions::assert_eq; + use reqwest::Client; + use serde_json::Value; + + #[test] + fn merge_trace_context_into_meta_preserves_existing_fields() { + let trace = W3cTraceContext { + traceparent: Some("00-aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa-bbbbbbbbbbbbbbbb-01".into()), + tracestate: Some("vendor=value".into()), + }; + let meta = rmcp::model::Meta(serde_json::Map::from_iter([( + "existing".to_string(), + Value::String("value".into()), + )])); + + let merged = merge_trace_context_into_meta(Some(meta), Some(&trace)).expect("meta"); + + assert_eq!( + merged, + rmcp::model::Meta(serde_json::Map::from_iter([ + ("existing".to_string(), Value::String("value".into())), + ( + TRACEPARENT_META_KEY.to_string(), + Value::String("00-aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa-bbbbbbbbbbbbbbbb-01".into()) + ), + ( + TRACESTATE_META_KEY.to_string(), + Value::String("vendor=value".into()) + ), + ])) + ); + } + + #[test] + fn apply_trace_context_injects_http_headers() { + let trace = W3cTraceContext { + traceparent: Some("00-aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa-bbbbbbbbbbbbbbbb-01".into()), + tracestate: Some("vendor=value".into()), + }; + let request = StreamableHttpResponseClient::apply_trace_context( + Client::new().post("http://example.com"), + Some(&trace), + ) + .build() + .expect("request"); + + assert_eq!( + request + .headers() + .get(TRACEPARENT_HEADER) + .and_then(|value| value.to_str().ok()), + Some("00-aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa-bbbbbbbbbbbbbbbb-01") + ); + assert_eq!( + request + .headers() + .get(TRACESTATE_HEADER) + .and_then(|value| value.to_str().ok()), + Some("vendor=value") + ); + } +} diff --git a/codex-rs/rmcp-client/tests/resources.rs b/codex-rs/rmcp-client/tests/resources.rs index ba1a8e4310..861909e1c5 100644 --- a/codex-rs/rmcp-client/tests/resources.rs +++ b/codex-rs/rmcp-client/tests/resources.rs @@ -7,6 +7,10 @@ use codex_rmcp_client::ElicitationResponse; use codex_rmcp_client::RmcpClient; use codex_utils_cargo_bin::CargoBinError; use futures::FutureExt as _; +use opentelemetry::global; +use opentelemetry::trace::TracerProvider as _; +use opentelemetry_sdk::propagation::TraceContextPropagator; +use opentelemetry_sdk::trace::SdkTracerProvider; use rmcp::model::AnnotateAble; use rmcp::model::ClientCapabilities; use rmcp::model::ElicitationCapability; @@ -18,6 +22,10 @@ use rmcp::model::ProtocolVersion; use rmcp::model::ReadResourceRequestParams; use rmcp::model::ResourceContents; use serde_json::json; +use tracing::Instrument; +use tracing::dispatcher::DefaultGuard; +use tracing_subscriber::layer::SubscriberExt; +use tracing_subscriber::util::SubscriberInitExt; const RESOURCE_URI: &str = "memo://codex/example-note"; @@ -53,6 +61,25 @@ fn init_params() -> InitializeRequestParams { } } +struct TestTracingContext { + _provider: SdkTracerProvider, + _guard: DefaultGuard, +} + +fn install_test_tracing(tracer_name: &str) -> TestTracingContext { + global::set_text_map_propagator(TraceContextPropagator::new()); + + let provider = SdkTracerProvider::builder().build(); + let tracer = provider.tracer(tracer_name.to_string()); + let subscriber = + tracing_subscriber::registry().with(tracing_opentelemetry::layer().with_tracer(tracer)); + + TestTracingContext { + _provider: provider, + _guard: subscriber.set_default(), + } +} + #[tokio::test(flavor = "multi_thread", worker_threads = 1)] async fn rmcp_client_can_list_and_read_resources() -> anyhow::Result<()> { let client = RmcpClient::new_stdio_client( @@ -149,3 +176,61 @@ async fn rmcp_client_can_list_and_read_resources() -> anyhow::Result<()> { Ok(()) } + +#[tokio::test(flavor = "current_thread")] +async fn stdio_tool_call_propagates_trace_metadata() -> anyhow::Result<()> { + let _trace = install_test_tracing("rmcp-stdio-trace-test"); + let client = RmcpClient::new_stdio_client( + stdio_server_bin()?.into(), + Vec::::new(), + None, + &[], + None, + ) + .await?; + + client + .initialize( + init_params(), + Some(Duration::from_secs(5)), + Box::new(|_, _| { + async { + Ok(ElicitationResponse { + action: ElicitationAction::Accept, + content: Some(json!({})), + meta: None, + }) + } + .boxed() + }), + ) + .await?; + + let result = async { + client + .call_tool( + "echo".to_string(), + Some(json!({ "message": "ping" })), + None, + Some(Duration::from_secs(5)), + ) + .await + } + .instrument(tracing::info_span!("rmcp.client.trace_test")) + .await?; + + assert_eq!(result.is_error, Some(false)); + let structured = result.structured_content.expect("structured content"); + assert_eq!(structured["echo"], json!("ECHOING: ping")); + assert_eq!(structured["env"], serde_json::Value::Null); + assert!( + structured["tracestate"].is_null() + || structured["tracestate"].as_str().is_some_and(str::is_empty) + ); + let traceparent = structured["traceparent"] + .as_str() + .expect("traceparent should be propagated via request metadata"); + assert!(traceparent.starts_with("00-")); + + Ok(()) +}