diff --git a/codex-rs/codex-api/src/common.rs b/codex-rs/codex-api/src/common.rs index 31b4dcdb44..12ce54626d 100644 --- a/codex-rs/codex-api/src/common.rs +++ b/codex-rs/codex-api/src/common.rs @@ -5,6 +5,7 @@ use codex_protocol::models::ResponseItem; use codex_protocol::openai_models::ReasoningEffort as ReasoningEffortConfig; use codex_protocol::protocol::RateLimitSnapshot; use codex_protocol::protocol::TokenUsage; +use codex_protocol::protocol::W3cTraceContext; use futures::Stream; use serde::Deserialize; use serde::Serialize; @@ -179,6 +180,7 @@ impl From<&ResponsesApiRequest> for ResponseCreateWsRequest { text: request.text.clone(), generate: None, client_metadata: None, + trace: None, } } } @@ -207,6 +209,8 @@ pub struct ResponseCreateWsRequest { pub generate: Option, #[serde(skip_serializing_if = "Option::is_none")] pub client_metadata: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub trace: Option, } #[derive(Debug, Serialize)] diff --git a/codex-rs/codex-api/src/endpoint/responses_websocket.rs b/codex-rs/codex-api/src/endpoint/responses_websocket.rs index 7c90882a96..c01058abdf 100644 --- a/codex-rs/codex-api/src/endpoint/responses_websocket.rs +++ b/codex-rs/codex-api/src/endpoint/responses_websocket.rs @@ -10,7 +10,6 @@ use crate::sse::responses::ResponsesStreamEvent; use crate::sse::responses::process_responses_event; use crate::telemetry::WebsocketTelemetry; use codex_client::TransportError; -use codex_client::inject_current_span_trace_headers; use codex_utils_rustls_provider::ensure_rustls_crypto_provider; use futures::SinkExt; use futures::StreamExt; @@ -308,7 +307,6 @@ impl ResponsesWebsocketClient { let mut headers = merge_request_headers(&self.provider.headers, extra_headers, default_headers); add_auth_headers_to_header_map(&self.auth, &mut headers); - inject_current_span_trace_headers(&mut headers); let (stream, server_reasoning_included, models_etag, server_model) = connect_websocket(ws_url, headers, turn_state.clone()).await?; diff --git a/codex-rs/codex-client/src/default_client.rs b/codex-rs/codex-client/src/default_client.rs index c84e5647ed..56b3ce4b16 100644 --- a/codex-rs/codex-client/src/default_client.rs +++ b/codex-rs/codex-client/src/default_client.rs @@ -154,15 +154,14 @@ impl<'a> Injector for HeaderMapInjector<'a> { } } -pub fn inject_current_span_trace_headers(headers: &mut HeaderMap) { - global::get_text_map_propagator(|prop| { - prop.inject_context(&Span::current().context(), &mut HeaderMapInjector(headers)); - }); -} - fn trace_headers() -> HeaderMap { let mut headers = HeaderMap::new(); - inject_current_span_trace_headers(&mut headers); + global::get_text_map_propagator(|prop| { + prop.inject_context( + &Span::current().context(), + &mut HeaderMapInjector(&mut headers), + ); + }); headers } diff --git a/codex-rs/codex-client/src/lib.rs b/codex-rs/codex-client/src/lib.rs index 513018e118..089d777c3a 100644 --- a/codex-rs/codex-client/src/lib.rs +++ b/codex-rs/codex-client/src/lib.rs @@ -8,7 +8,6 @@ mod transport; pub use crate::default_client::CodexHttpClient; pub use crate::default_client::CodexRequestBuilder; -pub use crate::default_client::inject_current_span_trace_headers; pub use crate::error::StreamError; pub use crate::error::TransportError; pub use crate::request::Request; diff --git a/codex-rs/core/src/client.rs b/codex-rs/core/src/client.rs index fe8514f89a..a150f92ffe 100644 --- a/codex-rs/core/src/client.rs +++ b/codex-rs/core/src/client.rs @@ -58,6 +58,7 @@ use codex_api::create_text_param_for_request; use codex_api::error::ApiError; use codex_api::requests::responses::Compression; use codex_otel::SessionTelemetry; +use codex_otel::current_span_w3c_trace_context; use codex_protocol::ThreadId; use codex_protocol::config_types::ReasoningSummary as ReasoningSummaryConfig; @@ -886,6 +887,7 @@ impl ModelClientSession { )?; let mut ws_payload = ResponseCreateWsRequest { client_metadata: build_ws_client_metadata(turn_metadata_header), + trace: current_span_w3c_trace_context(), ..ResponseCreateWsRequest::from(&request) }; if warmup { diff --git a/codex-rs/core/tests/suite/client_websockets.rs b/codex-rs/core/tests/suite/client_websockets.rs index b659ea2de2..1895bbc1f2 100755 --- a/codex-rs/core/tests/suite/client_websockets.rs +++ b/codex-rs/core/tests/suite/client_websockets.rs @@ -26,6 +26,7 @@ use codex_protocol::openai_models::ReasoningEffort as ReasoningEffortConfig; use codex_protocol::protocol::EventMsg; use codex_protocol::protocol::Op; use codex_protocol::protocol::SessionSource; +use codex_protocol::protocol::W3cTraceContext; use codex_protocol::user_input::UserInput; use core_test_support::load_default_config_for_test; use core_test_support::responses::WebSocketConnectionConfig; @@ -60,6 +61,31 @@ const OPENAI_BETA_HEADER: &str = "OpenAI-Beta"; const WS_V2_BETA_HEADER_VALUE: &str = "responses_websockets=2026-02-06"; const X_CLIENT_REQUEST_ID_HEADER: &str = "x-client-request-id"; +fn trace_id(traceparent: &str) -> &str { + traceparent + .split('-') + .nth(1) + .expect("traceparent missing trace id") +} + +fn assert_request_trace_matches(body: &serde_json::Value, expected_trace: &W3cTraceContext) { + let trace = body["trace"].as_object().expect("missing trace payload"); + let actual_traceparent = trace + .get("traceparent") + .and_then(serde_json::Value::as_str) + .expect("missing traceparent"); + let expected_traceparent = expected_trace + .traceparent + .as_deref() + .expect("missing expected traceparent"); + + assert_eq!(trace_id(actual_traceparent), trace_id(expected_traceparent)); + assert_eq!( + trace.get("tracestate").and_then(serde_json::Value::as_str), + expected_trace.tracestate.as_deref() + ); +} + struct WebsocketTestHarness { _codex_home: TempDir, client: ModelClient, @@ -109,7 +135,80 @@ async fn responses_websocket_streams_request() { #[tokio::test(flavor = "multi_thread", worker_threads = 2)] #[serial] -async fn responses_websocket_propagates_current_span_trace_headers() { +async fn responses_websocket_reuses_connection_with_per_turn_trace_payloads() { + skip_if_no_network!(); + + global::set_text_map_propagator(TraceContextPropagator::new()); + + let provider = SdkTracerProvider::builder().build(); + let tracer = provider.tracer("client-websocket-test"); + let subscriber = + tracing_subscriber::registry().with(tracing_opentelemetry::layer().with_tracer(tracer)); + let _guard = subscriber.set_default(); + + let server = start_websocket_server(vec![vec![ + vec![ev_response_created("resp-1"), ev_completed("resp-1")], + vec![ev_response_created("resp-2"), ev_completed("resp-2")], + ]]) + .await; + + let harness = websocket_harness(&server).await; + let prompt_one = prompt_with_input(vec![message_item("hello")]); + let prompt_two = prompt_with_input(vec![message_item("again")]); + + let first_trace = { + let mut client_session = harness.client.new_session(); + async { + let expected_trace = + current_span_w3c_trace_context().expect("current span should have trace context"); + stream_until_complete(&mut client_session, &harness, &prompt_one).await; + expected_trace + } + .instrument(tracing::info_span!("client.websocket.turn_one")) + .await + }; + + let second_trace = { + let mut client_session = harness.client.new_session(); + async { + let expected_trace = + current_span_w3c_trace_context().expect("current span should have trace context"); + stream_until_complete(&mut client_session, &harness, &prompt_two).await; + expected_trace + } + .instrument(tracing::info_span!("client.websocket.turn_two")) + .await + }; + + assert_eq!(server.handshakes().len(), 1); + let connection = server.single_connection(); + assert_eq!(connection.len(), 2); + + let first_request = connection + .first() + .expect("missing first request") + .body_json(); + let second_request = connection + .get(1) + .expect("missing second request") + .body_json(); + assert_request_trace_matches(&first_request, &first_trace); + assert_request_trace_matches(&second_request, &second_trace); + + let first_traceparent = first_request["trace"]["traceparent"] + .as_str() + .expect("missing first traceparent"); + let second_traceparent = second_request["trace"]["traceparent"] + .as_str() + .expect("missing second traceparent"); + assert_ne!(trace_id(first_traceparent), trace_id(second_traceparent)); + + server.shutdown().await; +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[serial] +async fn responses_websocket_preconnect_does_not_replace_turn_trace_payload() { skip_if_no_network!(); global::set_text_map_propagator(TraceContextPropagator::new()); @@ -128,6 +227,10 @@ async fn responses_websocket_propagates_current_span_trace_headers() { let harness = websocket_harness(&server).await; let mut client_session = harness.client.new_session(); + client_session + .preconnect_websocket(&harness.session_telemetry, &harness.model_info) + .await + .expect("websocket preconnect failed"); let prompt = prompt_with_input(vec![message_item("hello")]); let expected_trace = async { @@ -139,22 +242,11 @@ async fn responses_websocket_propagates_current_span_trace_headers() { .instrument(tracing::info_span!("client.websocket.request")) .await; - let handshake = server.single_handshake(); - let handshake_traceparent = handshake - .header("traceparent") - .expect("missing traceparent header"); - let expected_traceparent = expected_trace - .traceparent - .clone() - .expect("missing expected traceparent"); - assert_eq!( - handshake_traceparent.split('-').nth(1), - expected_traceparent.split('-').nth(1) - ); - assert_eq!( - handshake.header("tracestate"), - expected_trace.tracestate.clone() - ); + assert_eq!(server.handshakes().len(), 1); + let connection = server.single_connection(); + assert_eq!(connection.len(), 1); + let request = connection.first().expect("missing request").body_json(); + assert_request_trace_matches(&request, &expected_trace); server.shutdown().await; }