This commit is contained in:
Owen Lin
2026-03-12 15:55:46 -07:00
parent b1e3b1d08d
commit 4fa0c979bd
6 changed files with 121 additions and 27 deletions

View File

@@ -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<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub client_metadata: Option<HashMap<String, String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub trace: Option<W3cTraceContext>,
}
#[derive(Debug, Serialize)]

View File

@@ -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<A: AuthProvider> ResponsesWebsocketClient<A> {
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?;

View File

@@ -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
}

View File

@@ -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;

View File

@@ -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 {

View File

@@ -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;
}