Skip to main content

codex_api/endpoint/
responses_websocket.rs

1use crate::auth::SharedAuthProvider;
2use crate::common::ResponseEvent;
3use crate::common::ResponseStream;
4use crate::common::ResponsesWsRequest;
5use crate::common::SafetyBufferingTreatment;
6use crate::common::WS_REQUEST_HEADER_TRACEPARENT_CLIENT_METADATA_KEY;
7use crate::error::ApiError;
8use crate::provider::Provider;
9use crate::rate_limits::parse_rate_limit_event;
10use crate::safety_buffering::treatment_from_headers;
11use crate::sse::ResponsesStreamEvent;
12use crate::sse::process_responses_event;
13use crate::telemetry::WebsocketTelemetry;
14use codex_client::TransportError;
15use codex_http_client::HttpClientFactory;
16use codex_websocket_client::WebSocketConnection;
17use codex_websocket_client::WebSocketConnector;
18use futures::SinkExt;
19use futures::StreamExt;
20use http::HeaderMap;
21use http::HeaderName;
22use http::HeaderValue;
23use http::StatusCode;
24use serde::Deserialize;
25use serde_json::Value;
26use serde_json::map::Map as JsonMap;
27use std::sync::Arc;
28use std::sync::OnceLock;
29use std::time::Duration;
30use tokio::sync::Mutex;
31use tokio::sync::mpsc;
32use tokio::sync::oneshot;
33use tokio::time::Instant;
34use tokio_tungstenite::tungstenite::Error as WsError;
35use tokio_tungstenite::tungstenite::Message;
36use tokio_tungstenite::tungstenite::client::IntoClientRequest;
37use tokio_tungstenite::tungstenite::protocol::CloseFrame;
38use tracing::Instrument;
39use tracing::Span;
40use tracing::debug;
41use tracing::error;
42use tracing::info;
43use tracing::instrument;
44use tungstenite::extensions::ExtensionsConfig;
45use tungstenite::extensions::compression::deflate::DeflateConfig;
46use tungstenite::protocol::WebSocketConfig;
47use url::Url;
48
49struct WsStream {
50    tx_command: mpsc::Sender<WsCommand>,
51    rx_message: mpsc::UnboundedReceiver<Result<Message, WsError>>,
52    pump_task: tokio::task::JoinHandle<()>,
53}
54
55enum WsCommand {
56    Send {
57        message: Message,
58        tx_result: oneshot::Sender<Result<(), WsError>>,
59    },
60}
61
62impl WsStream {
63    fn new(inner: WebSocketConnection) -> Self {
64        let (tx_command, mut rx_command) = mpsc::channel::<WsCommand>(32);
65        let (tx_message, rx_message) = mpsc::unbounded_channel::<Result<Message, WsError>>();
66
67        let pump_task = tokio::spawn(async move {
68            let mut inner = inner;
69            loop {
70                tokio::select! {
71                    command = rx_command.recv() => {
72                        let Some(command) = command else {
73                            break;
74                        };
75                        match command {
76                            WsCommand::Send { message, tx_result } => {
77                                let result = inner.send(message).await;
78                                let should_break = result.is_err();
79                                let _ = tx_result.send(result);
80                                if should_break {
81                                    break;
82                                }
83                            }
84                        }
85                    }
86                    message = inner.next() => {
87                        let Some(message) = message else {
88                            break;
89                        };
90                        match message {
91                            Ok(Message::Ping(payload)) => {
92                                if let Err(err) = inner.send(Message::Pong(payload)).await {
93                                    let _ = tx_message.send(Err(err));
94                                    break;
95                                }
96                            }
97                            Ok(Message::Pong(_)) => {}
98                            Ok(message @ (Message::Text(_)
99                            | Message::Binary(_)
100                            | Message::Close(_)
101                            | Message::Frame(_))) => {
102                                let is_close = matches!(message, Message::Close(_));
103                                if tx_message.send(Ok(message)).is_err() {
104                                    break;
105                                }
106                                if is_close {
107                                    break;
108                                }
109                            }
110                            Err(err) => {
111                                let _ = tx_message.send(Err(err));
112                                break;
113                            }
114                        }
115                    }
116                }
117            }
118        });
119
120        Self {
121            tx_command,
122            rx_message,
123            pump_task,
124        }
125    }
126
127    async fn request(
128        &self,
129        make_command: impl FnOnce(oneshot::Sender<Result<(), WsError>>) -> WsCommand,
130    ) -> Result<(), WsError> {
131        let (tx_result, rx_result) = oneshot::channel();
132        if self.tx_command.send(make_command(tx_result)).await.is_err() {
133            return Err(WsError::ConnectionClosed);
134        }
135        rx_result.await.unwrap_or(Err(WsError::ConnectionClosed))
136    }
137
138    async fn send(&self, message: Message) -> Result<(), WsError> {
139        self.request(|tx_result| WsCommand::Send { message, tx_result })
140            .await
141    }
142
143    async fn next(&mut self) -> Option<Result<Message, WsError>> {
144        self.rx_message.recv().await
145    }
146}
147
148impl Drop for WsStream {
149    fn drop(&mut self) {
150        self.pump_task.abort();
151    }
152}
153
154const X_CODEX_TURN_STATE_HEADER: &str = "x-codex-turn-state";
155const X_MODELS_ETAG_HEADER: &str = "x-models-etag";
156const X_REASONING_INCLUDED_HEADER: &str = "x-reasoning-included";
157const OPENAI_MODEL_HEADER: &str = "openai-model";
158const WEBSOCKET_CONNECTION_LIMIT_REACHED_CODE: &str = "websocket_connection_limit_reached";
159const WEBSOCKET_CONNECTION_LIMIT_REACHED_MESSAGE: &str = "Responses websocket connection limit reached (60 minutes). Create a new websocket connection to continue.";
160const RESPONSES_WEBSOCKET_TIMING_KIND: &str = "responsesapi.websocket_timing";
161const RESPONSES_WEBSOCKET_TIMING_EVENT_TARGET: &str = "codex_api::responses_websocket_timing";
162const SESSION_ID_CLIENT_METADATA_KEY: &str = "session_id";
163const THREAD_ID_CLIENT_METADATA_KEY: &str = "thread_id";
164const TURN_ID_CLIENT_METADATA_KEY: &str = "turn_id";
165const WS_STREAM_REQUEST_START_MS_CLIENT_METADATA_KEY: &str = "x-codex-ws-stream-request-start-ms";
166
167struct ResponsesWebsocketTimingLogContext {
168    model: String,
169    session_id: Option<String>,
170    thread_id: Option<String>,
171    turn_id: Option<String>,
172    traceparent: Option<String>,
173    previous_response_id: Option<String>,
174    request_start_ms: Option<String>,
175    warmup: bool,
176    connection_reused: bool,
177}
178
179pub struct ResponsesWebsocketConnection {
180    stream: Arc<Mutex<Option<WsStream>>>,
181    // TODO (pakrym): is this the right place for timeout?
182    idle_timeout: Duration,
183    server_reasoning_included: bool,
184    models_etag: Option<String>,
185    server_model: Option<String>,
186    telemetry: Option<Arc<dyn WebsocketTelemetry>>,
187}
188
189impl std::fmt::Debug for ResponsesWebsocketConnection {
190    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
191        f.debug_struct("ResponsesWebsocketConnection")
192            .field("stream", &"<ws-stream>")
193            .field("idle_timeout", &self.idle_timeout)
194            .field("server_reasoning_included", &self.server_reasoning_included)
195            .field("models_etag", &self.models_etag)
196            .field("server_model", &self.server_model)
197            .field("telemetry", &self.telemetry.as_ref().map(|_| "<telemetry>"))
198            .finish()
199    }
200}
201
202impl ResponsesWebsocketConnection {
203    fn new(
204        stream: WsStream,
205        idle_timeout: Duration,
206        server_reasoning_included: bool,
207        models_etag: Option<String>,
208        server_model: Option<String>,
209        telemetry: Option<Arc<dyn WebsocketTelemetry>>,
210    ) -> Self {
211        Self {
212            stream: Arc::new(Mutex::new(Some(stream))),
213            idle_timeout,
214            server_reasoning_included,
215            models_etag,
216            server_model,
217            telemetry,
218        }
219    }
220
221    pub async fn is_closed(&self) -> bool {
222        self.stream.lock().await.is_none()
223    }
224
225    #[instrument(
226        name = "responses_websocket.stream_request",
227        level = "info",
228        skip_all,
229        fields(transport = "responses_websocket", api.path = "responses")
230    )]
231    pub async fn stream_request(
232        &self,
233        request: ResponsesWsRequest<'_>,
234        connection_reused: bool,
235        turn_state: Option<Arc<OnceLock<String>>>,
236    ) -> Result<ResponseStream, ApiError> {
237        let (tx_event, rx_event) =
238            mpsc::channel::<std::result::Result<ResponseEvent, ApiError>>(1600);
239        let stream = Arc::clone(&self.stream);
240        let idle_timeout = self.idle_timeout;
241        let server_reasoning_included = self.server_reasoning_included;
242        let models_etag = self.models_etag.clone();
243        let server_model = self.server_model.clone();
244        let telemetry = self.telemetry.clone();
245        let ResponsesWsRequest::ResponseCreate(ws_request) = &request;
246        let client_metadata = ws_request.client_metadata.as_ref();
247        let timing_log_context = ResponsesWebsocketTimingLogContext {
248            model: ws_request.model.to_string(),
249            session_id: client_metadata
250                .and_then(|metadata| metadata.get(SESSION_ID_CLIENT_METADATA_KEY))
251                .cloned(),
252            thread_id: client_metadata
253                .and_then(|metadata| metadata.get(THREAD_ID_CLIENT_METADATA_KEY))
254                .cloned(),
255            turn_id: client_metadata
256                .and_then(|metadata| metadata.get(TURN_ID_CLIENT_METADATA_KEY))
257                .cloned(),
258            traceparent: client_metadata
259                .and_then(|metadata| {
260                    metadata.get(WS_REQUEST_HEADER_TRACEPARENT_CLIENT_METADATA_KEY)
261                })
262                .cloned(),
263            previous_response_id: ws_request.previous_response_id.clone(),
264            request_start_ms: client_metadata
265                .and_then(|metadata| metadata.get(WS_STREAM_REQUEST_START_MS_CLIENT_METADATA_KEY))
266                .cloned(),
267            warmup: ws_request.generate == Some(false),
268            connection_reused,
269        };
270        let request_text = serialize_websocket_request(&request)?;
271
272        let current_span = Span::current();
273        tokio::spawn(
274            #[expect(
275                clippy::await_holding_invalid_type,
276                reason = "the guard serializes exclusive use of the websocket stream for the lifetime of the response stream"
277            )]
278            async move {
279                if let Some(model) = server_model {
280                    let _ = tx_event.send(Ok(ResponseEvent::ServerModel(model))).await;
281                }
282                if let Some(etag) = models_etag {
283                    let _ = tx_event.send(Ok(ResponseEvent::ModelsEtag(etag))).await;
284                }
285                if server_reasoning_included {
286                    let _ = tx_event
287                        .send(Ok(ResponseEvent::ServerReasoningIncluded(true)))
288                        .await;
289                }
290                let mut guard = stream.lock().await;
291                let result = {
292                    let Some(ws_stream) = guard.as_mut() else {
293                        let _ = tx_event
294                            .send(Err(ApiError::Stream(
295                                "websocket connection is closed".to_string(),
296                            )))
297                            .await;
298                        return;
299                    };
300
301                    run_websocket_response_stream(
302                        ws_stream,
303                        tx_event.clone(),
304                        request_text,
305                        idle_timeout,
306                        telemetry,
307                        turn_state.as_deref(),
308                        &timing_log_context,
309                    )
310                    .await
311                };
312
313                if let Err(err) = result {
314                    // A terminal stream error should reach the caller immediately. Waiting for a
315                    // graceful close handshake here can stall indefinitely and mask the error.
316                    let failed_stream = guard.take();
317                    drop(guard);
318                    drop(failed_stream);
319                    let _ = tx_event.send(Err(err)).await;
320                }
321            }
322            .instrument(current_span),
323        );
324
325        Ok(ResponseStream {
326            rx_event,
327            upstream_request_id: None,
328        })
329    }
330}
331
332/// Client for connecting to the Responses WebSocket endpoint for one provider.
333pub struct ResponsesWebsocketClient {
334    provider: Provider,
335    auth: SharedAuthProvider,
336}
337
338/// Close frame information captured by a handshake probe.
339#[derive(Debug, Clone, PartialEq, Eq)]
340pub struct ResponsesWebsocketClose {
341    /// WebSocket close code returned by the server.
342    pub code: String,
343    /// Human-readable close reason returned by the server.
344    pub reason: String,
345}
346
347/// Result of a handshake-only Responses WebSocket probe.
348#[derive(Debug, Clone, PartialEq, Eq)]
349pub struct ResponsesWebsocketProbe {
350    /// Redacted by callers before displaying or serializing support reports.
351    pub url: String,
352    /// HTTP status returned by the successful WebSocket upgrade.
353    pub status: StatusCode,
354    /// Whether the server reported reasoning support in the upgrade response.
355    pub reasoning_included: bool,
356    /// Whether the server returned a model catalog ETag in the upgrade response.
357    pub models_etag_present: bool,
358    /// Whether the server returned a server-selected model in the upgrade response.
359    pub server_model_present: bool,
360    /// Close frame received immediately after upgrade, when one arrives quickly.
361    pub immediate_close: Option<ResponsesWebsocketClose>,
362}
363
364impl ResponsesWebsocketClient {
365    /// Creates a Responses WebSocket client for an already-resolved provider and auth source.
366    pub fn new(provider: Provider, auth: SharedAuthProvider) -> Self {
367        Self { provider, auth }
368    }
369
370    #[instrument(
371        name = "responses_websocket.connect",
372        level = "info",
373        skip_all,
374        fields(transport = "responses_websocket", api.path = "responses")
375    )]
376    pub async fn connect(
377        &self,
378        http_client_factory: &HttpClientFactory,
379        extra_headers: HeaderMap,
380        default_headers: HeaderMap,
381        turn_state: Option<Arc<OnceLock<String>>>,
382        telemetry: Option<Arc<dyn WebsocketTelemetry>>,
383    ) -> Result<ResponsesWebsocketConnection, ApiError> {
384        let ws_url = self
385            .provider
386            .websocket_url_for_path("responses")
387            .map_err(|err| ApiError::Stream(format!("failed to build websocket URL: {err}")))?;
388
389        let mut headers =
390            merge_request_headers(&self.provider.headers, extra_headers, default_headers);
391        self.auth.add_auth_headers(&mut headers);
392
393        let (stream, _status, server_reasoning_included, models_etag, server_model) =
394            connect_websocket(ws_url, headers, http_client_factory, turn_state.clone()).await?;
395        Ok(ResponsesWebsocketConnection::new(
396            stream,
397            self.provider.stream_idle_timeout,
398            server_reasoning_included,
399            models_etag,
400            server_model,
401            telemetry,
402        ))
403    }
404
405    /// Opens a WebSocket connection long enough to validate the upgrade response.
406    ///
407    /// The probe uses the same URL construction, headers, authentication, TLS,
408    /// and custom-CA path as a real Responses WebSocket connection, but it does
409    /// not send a request frame. After the HTTP 101 upgrade succeeds, it waits
410    /// briefly for an immediate server close frame so diagnostics can distinguish
411    /// a usable connection from a policy rejection that closes right away.
412    pub async fn probe_handshake(
413        &self,
414        http_client_factory: &HttpClientFactory,
415        extra_headers: HeaderMap,
416        default_headers: HeaderMap,
417        immediate_close_timeout: Duration,
418    ) -> Result<ResponsesWebsocketProbe, ApiError> {
419        let ws_url = self
420            .provider
421            .websocket_url_for_path("responses")
422            .map_err(|err| ApiError::Stream(format!("failed to build websocket URL: {err}")))?;
423
424        let mut headers =
425            merge_request_headers(&self.provider.headers, extra_headers, default_headers);
426        self.auth.add_auth_headers(&mut headers);
427
428        let (mut stream, status, reasoning_included, models_etag, server_model) =
429            connect_websocket(
430                ws_url.clone(),
431                headers,
432                http_client_factory,
433                /*turn_state*/ None,
434            )
435            .await?;
436        let immediate_close = tokio::time::timeout(immediate_close_timeout, stream.next())
437            .await
438            .ok()
439            .flatten()
440            .transpose()
441            .map_err(|err| {
442                ApiError::Stream(format!("failed to read websocket probe event: {err}"))
443            })?
444            .and_then(immediate_close_from_message);
445
446        Ok(ResponsesWebsocketProbe {
447            url: ws_url.to_string(),
448            status,
449            reasoning_included,
450            models_etag_present: models_etag.is_some(),
451            server_model_present: server_model.is_some(),
452            immediate_close,
453        })
454    }
455}
456
457fn immediate_close_from_message(message: Message) -> Option<ResponsesWebsocketClose> {
458    let Message::Close(frame) = message else {
459        return None;
460    };
461    frame.map(close_frame_to_probe)
462}
463
464fn close_frame_to_probe(frame: CloseFrame) -> ResponsesWebsocketClose {
465    ResponsesWebsocketClose {
466        code: frame.code.to_string(),
467        reason: frame.reason.to_string(),
468    }
469}
470
471fn merge_request_headers(
472    provider_headers: &HeaderMap,
473    extra_headers: HeaderMap,
474    default_headers: HeaderMap,
475) -> HeaderMap {
476    let mut headers = provider_headers.clone();
477    headers.extend(extra_headers);
478    for (name, value) in &default_headers {
479        if let http::header::Entry::Vacant(entry) = headers.entry(name) {
480            entry.insert(value.clone());
481        }
482    }
483    headers
484}
485
486async fn connect_websocket(
487    url: Url,
488    headers: HeaderMap,
489    http_client_factory: &HttpClientFactory,
490    turn_state: Option<Arc<OnceLock<String>>>,
491) -> Result<(WsStream, StatusCode, bool, Option<String>, Option<String>), ApiError> {
492    info!("connecting to websocket: {url}");
493
494    let mut request = url
495        .as_str()
496        .into_client_request()
497        .map_err(|err| ApiError::Stream(format!("failed to build websocket request: {err}")))?;
498    request.headers_mut().extend(headers);
499
500    let connector = WebSocketConnector::new(http_client_factory)
501        .map_err(|err| ApiError::Stream(format!("failed to configure websocket TLS: {err}")))?;
502    let response = connector.connect(request, websocket_config()).await;
503
504    let (stream, response) = match response {
505        Ok((stream, response)) => {
506            info!(
507                "successfully connected to websocket: {url}, headers: {:?}",
508                response.headers()
509            );
510            (stream, response)
511        }
512        Err(err) => {
513            error!("failed to connect to websocket: {err}, url: {url}");
514            return Err(map_ws_error(err, &url));
515        }
516    };
517
518    let reasoning_included = response.headers().contains_key(X_REASONING_INCLUDED_HEADER);
519    let models_etag = response
520        .headers()
521        .get(X_MODELS_ETAG_HEADER)
522        .and_then(|value| value.to_str().ok())
523        .map(ToString::to_string);
524    let server_model = response
525        .headers()
526        .get(OPENAI_MODEL_HEADER)
527        .and_then(|value| value.to_str().ok())
528        .map(ToString::to_string);
529    if let Some(turn_state) = turn_state
530        && let Some(header_value) = response
531            .headers()
532            .get(X_CODEX_TURN_STATE_HEADER)
533            .and_then(|value| value.to_str().ok())
534    {
535        let _ = turn_state.set(header_value.to_string());
536    }
537    Ok((
538        WsStream::new(stream),
539        response.status(),
540        reasoning_included,
541        models_etag,
542        server_model,
543    ))
544}
545
546fn websocket_config() -> WebSocketConfig {
547    let mut extensions = ExtensionsConfig::default();
548    extensions.permessage_deflate = Some(DeflateConfig::default());
549
550    let mut config = WebSocketConfig::default();
551    config.extensions = extensions;
552    config
553}
554
555fn map_ws_error(err: WsError, url: &Url) -> ApiError {
556    match err {
557        WsError::Http(response) => {
558            let status = response.status();
559            let headers = response.headers().clone();
560            let body = response
561                .body()
562                .as_ref()
563                .and_then(|bytes| String::from_utf8(bytes.clone()).ok());
564            ApiError::Transport(TransportError::Http {
565                status,
566                url: Some(url.to_string()),
567                headers: Some(headers),
568                body,
569            })
570        }
571        WsError::ConnectionClosed | WsError::AlreadyClosed => {
572            ApiError::Stream("websocket closed".to_string())
573        }
574        WsError::Io(err) => ApiError::Transport(TransportError::Network(err.to_string())),
575        other => ApiError::Transport(TransportError::Network(other.to_string())),
576    }
577}
578
579#[derive(Debug, Deserialize)]
580struct WrappedWebsocketError {
581    code: Option<String>,
582    message: Option<String>,
583}
584
585#[derive(Debug, Deserialize)]
586struct WrappedWebsocketErrorEvent {
587    #[serde(rename = "type")]
588    kind: String,
589    #[serde(alias = "status_code")]
590    status: Option<u16>,
591    #[serde(default)]
592    error: Option<WrappedWebsocketError>,
593    #[serde(default)]
594    headers: Option<JsonMap<String, Value>>,
595}
596
597fn parse_wrapped_websocket_error_event(payload: &str) -> Option<WrappedWebsocketErrorEvent> {
598    let event: WrappedWebsocketErrorEvent = serde_json::from_str(payload).ok()?;
599    if event.kind != "error" {
600        return None;
601    }
602    Some(event)
603}
604
605fn map_wrapped_websocket_error_event(
606    event: WrappedWebsocketErrorEvent,
607    original_payload: String,
608) -> Option<ApiError> {
609    let WrappedWebsocketErrorEvent {
610        status,
611        error,
612        headers,
613        ..
614    } = event;
615
616    if let Some(error) = error.as_ref()
617        && let Some(code) = error.code.as_deref()
618        && code == WEBSOCKET_CONNECTION_LIMIT_REACHED_CODE
619    {
620        return Some(ApiError::Retryable {
621            message: error
622                .message
623                .clone()
624                .unwrap_or_else(|| WEBSOCKET_CONNECTION_LIMIT_REACHED_MESSAGE.to_string()),
625            delay: None,
626        });
627    }
628
629    let status = StatusCode::from_u16(status?).ok()?;
630    if status.is_success() {
631        return None;
632    }
633
634    Some(ApiError::Transport(TransportError::Http {
635        status,
636        url: None,
637        headers: headers.as_ref().map(json_headers_to_http_headers),
638        body: Some(original_payload),
639    }))
640}
641
642fn json_headers_to_http_headers(headers: &JsonMap<String, Value>) -> HeaderMap {
643    let mut mapped = HeaderMap::new();
644    for (name, value) in headers {
645        let Ok(header_name) = HeaderName::from_bytes(name.as_bytes()) else {
646            continue;
647        };
648        let Some(header_value) = json_header_value(value) else {
649            continue;
650        };
651        mapped.insert(header_name, header_value);
652    }
653    mapped
654}
655
656fn json_header_value(value: &Value) -> Option<HeaderValue> {
657    let value = match value {
658        Value::String(value) => value.clone(),
659        Value::Number(value) => value.to_string(),
660        Value::Bool(value) => value.to_string(),
661        _ => return None,
662    };
663    HeaderValue::from_str(&value).ok()
664}
665
666async fn run_websocket_response_stream(
667    ws_stream: &mut WsStream,
668    tx_event: mpsc::Sender<std::result::Result<ResponseEvent, ApiError>>,
669    request_text: String,
670    idle_timeout: Duration,
671    telemetry: Option<Arc<dyn WebsocketTelemetry>>,
672    turn_state: Option<&OnceLock<String>>,
673    timing_log_context: &ResponsesWebsocketTimingLogContext,
674) -> Result<(), ApiError> {
675    let mut last_server_model: Option<String> = None;
676    let mut safety_buffering_treatment = SafetyBufferingTreatment::default();
677    send_websocket_request(
678        ws_stream,
679        request_text,
680        idle_timeout,
681        telemetry.as_ref(),
682        timing_log_context.connection_reused,
683    )
684    .await?;
685
686    loop {
687        let poll_start = Instant::now();
688        let response = tokio::time::timeout(idle_timeout, ws_stream.next())
689            .await
690            .map_err(|_| ApiError::Stream("idle timeout waiting for websocket".into()));
691        if let Some(t) = telemetry.as_ref() {
692            t.on_ws_event(&response, poll_start.elapsed());
693        }
694        let message = match response {
695            Ok(Some(Ok(msg))) => msg,
696            Ok(Some(Err(err))) => {
697                return Err(ApiError::Stream(err.to_string()));
698            }
699            Ok(None) => {
700                return Err(ApiError::Stream(
701                    "stream closed before response.completed".into(),
702                ));
703            }
704            Err(err) => {
705                return Err(err);
706            }
707        };
708
709        match message {
710            Message::Text(text) => {
711                if let Some(wrapped_error) = parse_wrapped_websocket_error_event(&text)
712                    && let Some(error) =
713                        map_wrapped_websocket_error_event(wrapped_error, text.to_string())
714                {
715                    return Err(error);
716                }
717
718                let event = match serde_json::from_str::<ResponsesStreamEvent>(&text) {
719                    Ok(event) => event,
720                    Err(err) => {
721                        debug!("failed to parse websocket event: {err}, data: {text}");
722                        continue;
723                    }
724                };
725                emit_responses_websocket_timing_event(
726                    event.kind(),
727                    text.as_str(),
728                    timing_log_context,
729                );
730                if let Some(response_turn_state) = event.turn_state()
731                    && let Some(turn_state) = turn_state
732                {
733                    let _ = turn_state.set(response_turn_state);
734                }
735                let model_verifications = event.model_verifications();
736                let turn_moderation_metadata = event.turn_moderation_metadata();
737                let safety_buffering =
738                    safety_buffering_for_event(&event, &mut safety_buffering_treatment);
739                if event.kind() == "codex.rate_limits" {
740                    if let Some(snapshot) = parse_rate_limit_event(&text) {
741                        let _ = tx_event.send(Ok(ResponseEvent::RateLimits(snapshot))).await;
742                    }
743                    continue;
744                }
745                if let Some(model) = event.response_model()
746                    && last_server_model.as_deref() != Some(model.as_str())
747                {
748                    let _ = tx_event
749                        .send(Ok(ResponseEvent::ServerModel(model.clone())))
750                        .await;
751                    last_server_model = Some(model);
752                }
753                if let Some(verifications) = model_verifications
754                    && tx_event
755                        .send(Ok(ResponseEvent::ModelVerifications(verifications)))
756                        .await
757                        .is_err()
758                {
759                    return Err(ApiError::Stream(
760                        "response event consumer dropped".to_string(),
761                    ));
762                }
763                if let Some(metadata) = turn_moderation_metadata
764                    && tx_event
765                        .send(Ok(ResponseEvent::TurnModerationMetadata(metadata)))
766                        .await
767                        .is_err()
768                {
769                    return Err(ApiError::Stream(
770                        "response event consumer dropped".to_string(),
771                    ));
772                }
773                if let Some(buffering) = safety_buffering
774                    && tx_event
775                        .send(Ok(ResponseEvent::SafetyBuffering(buffering)))
776                        .await
777                        .is_err()
778                {
779                    return Err(ApiError::Stream(
780                        "response event consumer dropped".to_string(),
781                    ));
782                }
783                match process_responses_event(event) {
784                    Ok(Some(event)) => {
785                        let is_completed = matches!(event, ResponseEvent::Completed { .. });
786                        let _ = tx_event.send(Ok(event)).await;
787                        if is_completed {
788                            break;
789                        }
790                    }
791                    Ok(None) => {}
792                    Err(error) => {
793                        return Err(error.into_api_error());
794                    }
795                }
796            }
797            Message::Binary(_) => {
798                return Err(ApiError::Stream("unexpected binary websocket event".into()));
799            }
800            Message::Close(_) => {
801                return Err(ApiError::Stream(
802                    "websocket closed by server before response.completed".into(),
803                ));
804            }
805            Message::Frame(_) => {}
806            Message::Ping(_) | Message::Pong(_) => {}
807        }
808    }
809
810    Ok(())
811}
812
813fn emit_responses_websocket_timing_event(
814    kind: &str,
815    payload: &str,
816    context: &ResponsesWebsocketTimingLogContext,
817) {
818    if kind != RESPONSES_WEBSOCKET_TIMING_KIND {
819        return;
820    }
821
822    // This full payload is excluded from always-on sinks. Opt in with
823    // `RUST_LOG='codex_api::responses_websocket_timing=trace'`.
824    tracing::event!(
825        name: RESPONSES_WEBSOCKET_TIMING_KIND,
826        target: RESPONSES_WEBSOCKET_TIMING_EVENT_TARGET,
827        tracing::Level::TRACE,
828        model = context.model.as_str(),
829        session_id = context.session_id.as_deref().unwrap_or_default(),
830        thread_id = context.thread_id.as_deref().unwrap_or_default(),
831        turn_id = context.turn_id.as_deref().unwrap_or_default(),
832        traceparent = context.traceparent.as_deref().unwrap_or_default(),
833        previous_response_id = context.previous_response_id.as_deref().unwrap_or_default(),
834        request_start_ms = context.request_start_ms.as_deref().unwrap_or_default(),
835        warmup = context.warmup,
836        connection_reused = context.connection_reused,
837        payload,
838        "responses websocket timing"
839    );
840}
841
842fn safety_buffering_for_event(
843    event: &ResponsesStreamEvent,
844    treatment: &mut SafetyBufferingTreatment,
845) -> Option<crate::common::SafetyBuffering> {
846    if let Some(headers) = event.headers.as_ref().and_then(Value::as_object)
847        && let Some(updated_treatment) =
848            treatment_from_headers(&json_headers_to_http_headers(headers))
849    {
850        *treatment = updated_treatment;
851    }
852    event.safety_buffering(treatment)
853}
854
855async fn send_websocket_request(
856    ws_stream: &WsStream,
857    request_text: String,
858    idle_timeout: Duration,
859    telemetry: Option<&Arc<dyn WebsocketTelemetry>>,
860    connection_reused: bool,
861) -> Result<(), ApiError> {
862    let request_start = Instant::now();
863    let result = tokio::time::timeout(
864        idle_timeout,
865        ws_stream.send(Message::Text(request_text.into())),
866    )
867    .await
868    .map_err(|_| ApiError::Stream("idle timeout sending websocket request".into()))
869    .and_then(|result| {
870        result.map_err(|err| ApiError::Stream(format!("failed to send websocket request: {err}")))
871    });
872
873    if let Some(t) = telemetry.as_ref() {
874        t.on_ws_request(
875            request_start.elapsed(),
876            result.as_ref().err(),
877            connection_reused,
878        );
879    }
880
881    result?;
882
883    Ok(())
884}
885
886fn serialize_websocket_request(request: &ResponsesWsRequest<'_>) -> Result<String, ApiError> {
887    serde_json::to_string(request)
888        .map_err(|err| ApiError::Stream(format!("failed to encode websocket request: {err}")))
889}
890
891#[cfg(test)]
892mod tests {
893    use super::*;
894    use crate::common::ResponseCreateWsRequest;
895    use crate::common::ResponsesApiRequest;
896    use codex_protocol::ResponseItemId;
897    use codex_protocol::models::ContentItem;
898    use codex_protocol::models::ResponseItem;
899    use pretty_assertions::assert_eq;
900    use serde_json::json;
901    use std::collections::HashMap;
902
903    #[test]
904    fn direct_serialization_preserves_websocket_request_payload() {
905        let api_request = ResponsesApiRequest {
906            model: "gpt-test".to_string(),
907            instructions: "Use the available tools.".to_string(),
908            input: vec![ResponseItem::Message {
909                id: Some(ResponseItemId::with_suffix("msg", "1")),
910                role: "user".to_string(),
911                content: vec![ContentItem::InputText {
912                    text: "hello".to_string(),
913                }],
914                phase: None,
915                internal_chat_message_metadata_passthrough: None,
916            }],
917            tools: Some(vec![json!({
918                "type": "function",
919                "name": "lookup",
920                "parameters": {"type": "object"}
921            })]),
922            tool_choice: "auto".to_string(),
923            parallel_tool_calls: true,
924            reasoning: None,
925            store: false,
926            stream: true,
927            stream_options: None,
928            include: vec!["reasoning.encrypted_content".to_string()],
929            service_tier: Some("priority".to_string()),
930            prompt_cache_key: Some("cache-key".to_string()),
931            text: None,
932            client_metadata: Some(HashMap::from([(
933                "traceparent".to_string(),
934                "00-0123456789abcdef0123456789abcdef-0123456789abcdef-01".to_string(),
935            )])),
936        };
937        let request = ResponsesWsRequest::ResponseCreate(ResponseCreateWsRequest {
938            previous_response_id: Some("resp-1".to_string()),
939            generate: Some(false),
940            ..ResponseCreateWsRequest::from(&api_request)
941        });
942
943        let mut expected_payload =
944            serde_json::to_value(&api_request).expect("serialize responses API request");
945        expected_payload["type"] = json!("response.create");
946        expected_payload["previous_response_id"] = json!("resp-1");
947        expected_payload["generate"] = json!(false);
948        let request_text =
949            serialize_websocket_request(&request).expect("serialize websocket request");
950        let wire_payload =
951            serde_json::from_str::<Value>(&request_text).expect("parse websocket request");
952
953        assert_eq!(wire_payload, expected_payload);
954    }
955
956    #[test]
957    fn websocket_config_enables_permessage_deflate() {
958        let config = websocket_config();
959        assert!(config.extensions.permessage_deflate.is_some());
960    }
961
962    #[test]
963    fn parse_wrapped_websocket_error_event_maps_to_transport_http() {
964        let payload = json!({
965            "type": "error",
966            "status": 429,
967            "error": {
968                "type": "usage_limit_reached",
969                "message": "The usage limit has been reached",
970                "plan_type": "pro",
971                "resets_at": 1738888888
972            },
973            "headers": {
974                "x-codex-primary-used-percent": "100.0",
975                "x-codex-primary-window-minutes": 15
976            }
977        })
978        .to_string();
979
980        let wrapped_error = parse_wrapped_websocket_error_event(&payload)
981            .expect("expected websocket error payload to be parsed");
982        let api_error = map_wrapped_websocket_error_event(wrapped_error, payload)
983            .expect("expected websocket error payload to map to ApiError");
984
985        let ApiError::Transport(TransportError::Http {
986            status,
987            headers,
988            body,
989            ..
990        }) = api_error
991        else {
992            panic!("expected ApiError::Transport(Http)");
993        };
994
995        assert_eq!(status, StatusCode::TOO_MANY_REQUESTS);
996        let headers = headers.expect("expected headers");
997        assert_eq!(
998            headers
999                .get("x-codex-primary-used-percent")
1000                .and_then(|value| value.to_str().ok()),
1001            Some("100.0")
1002        );
1003        assert_eq!(
1004            headers
1005                .get("x-codex-primary-window-minutes")
1006                .and_then(|value| value.to_str().ok()),
1007            Some("15")
1008        );
1009        let body = body.expect("expected body");
1010        assert!(body.contains("usage_limit_reached"));
1011        assert!(body.contains("The usage limit has been reached"));
1012    }
1013
1014    #[test]
1015    fn parse_wrapped_websocket_error_event_ignores_non_error_payloads() {
1016        let payload = json!({
1017            "type": "response.created",
1018            "response": {
1019                "id": "resp-1"
1020            }
1021        })
1022        .to_string();
1023
1024        let wrapped_error = parse_wrapped_websocket_error_event(&payload);
1025        assert!(wrapped_error.is_none());
1026    }
1027
1028    #[test]
1029    fn parse_wrapped_websocket_error_event_with_status_maps_invalid_request() {
1030        let payload = json!({
1031            "type": "error",
1032            "status": 400,
1033            "error": {
1034                "type": "invalid_request_error",
1035                "message": "Model does not support image inputs"
1036            }
1037        })
1038        .to_string();
1039
1040        let wrapped_error = parse_wrapped_websocket_error_event(&payload)
1041            .expect("expected websocket error payload to be parsed");
1042        let api_error = map_wrapped_websocket_error_event(wrapped_error, payload)
1043            .expect("expected websocket error payload to map to ApiError");
1044        let ApiError::Transport(TransportError::Http { status, body, .. }) = api_error else {
1045            panic!("expected ApiError::Transport(Http)");
1046        };
1047        assert_eq!(status, StatusCode::BAD_REQUEST);
1048        let body = body.expect("expected body");
1049        assert!(body.contains("invalid_request_error"));
1050        assert!(body.contains("Model does not support image inputs"));
1051    }
1052
1053    #[test]
1054    fn parse_wrapped_websocket_error_event_with_connection_limit_maps_retryable() {
1055        let payload = json!({
1056            "type": "error",
1057            "status": 400,
1058            "error": {
1059                "type": "invalid_request_error",
1060                "code": "websocket_connection_limit_reached",
1061                "message": "Responses websocket connection limit reached (60 minutes). Create a new websocket connection to continue."
1062            }
1063        })
1064        .to_string();
1065
1066        let wrapped_error = parse_wrapped_websocket_error_event(&payload)
1067            .expect("expected websocket error payload to be parsed");
1068        let api_error = map_wrapped_websocket_error_event(wrapped_error, payload)
1069            .expect("expected websocket error payload to map to ApiError");
1070        let ApiError::Retryable { message, delay } = api_error else {
1071            panic!("expected ApiError::Retryable");
1072        };
1073        assert_eq!(message, WEBSOCKET_CONNECTION_LIMIT_REACHED_MESSAGE);
1074        assert_eq!(delay, None);
1075    }
1076
1077    #[test]
1078    fn parse_wrapped_websocket_error_event_without_status_is_not_mapped() {
1079        let payload = json!({
1080            "type": "error",
1081            "error": {
1082                "type": "usage_limit_reached",
1083                "message": "The usage limit has been reached"
1084            },
1085            "headers": {
1086                "x-codex-primary-used-percent": "100.0",
1087                "x-codex-primary-window-minutes": 15
1088            }
1089        })
1090        .to_string();
1091
1092        let wrapped_error = parse_wrapped_websocket_error_event(&payload)
1093            .expect("expected websocket error payload to be parsed");
1094        let api_error = map_wrapped_websocket_error_event(wrapped_error, payload);
1095        assert!(api_error.is_none());
1096    }
1097
1098    #[test]
1099    fn merge_request_headers_matches_http_precedence() {
1100        let mut provider_headers = HeaderMap::new();
1101        provider_headers.insert(
1102            "originator",
1103            HeaderValue::from_static("provider-originator"),
1104        );
1105        provider_headers.insert("x-priority", HeaderValue::from_static("provider"));
1106
1107        let mut extra_headers = HeaderMap::new();
1108        extra_headers.insert("x-priority", HeaderValue::from_static("extra"));
1109
1110        let mut default_headers = HeaderMap::new();
1111        default_headers.insert("originator", HeaderValue::from_static("default-originator"));
1112        default_headers.insert("x-priority", HeaderValue::from_static("default"));
1113        default_headers.insert("x-default-only", HeaderValue::from_static("default-only"));
1114
1115        let merged = merge_request_headers(&provider_headers, extra_headers, default_headers);
1116
1117        assert_eq!(
1118            merged.get("originator"),
1119            Some(&HeaderValue::from_static("provider-originator"))
1120        );
1121        assert_eq!(
1122            merged.get("x-priority"),
1123            Some(&HeaderValue::from_static("extra"))
1124        );
1125        assert_eq!(
1126            merged.get("x-default-only"),
1127            Some(&HeaderValue::from_static("default-only"))
1128        );
1129    }
1130
1131    #[test]
1132    fn websocket_safety_buffering_uses_event_before_header_fallback() {
1133        let metadata: ResponsesStreamEvent = serde_json::from_value(json!({
1134            "type": "codex.response.metadata",
1135            "headers": {
1136                "x-codex-safety-buffering-enabled": "true",
1137                "x-codex-safety-buffering-faster-model": "gpt-fast-header"
1138            }
1139        }))
1140        .expect("deserialize treatment metadata");
1141        let event: ResponsesStreamEvent = serde_json::from_value(json!({
1142            "type": "response.output_text.delta",
1143            "safety_buffering": {
1144                "use_cases": ["cyber"],
1145                "reasons": ["user_risk"],
1146                "retry_model": "gpt-fast-wire"
1147            }
1148        }))
1149        .expect("deserialize safety buffering event");
1150        let mut treatment = SafetyBufferingTreatment::default();
1151
1152        assert!(safety_buffering_for_event(&metadata, &mut treatment).is_none());
1153        let buffering = safety_buffering_for_event(&event, &mut treatment)
1154            .expect("expected safety buffering payload");
1155
1156        assert_eq!(
1157            buffering,
1158            crate::common::SafetyBuffering {
1159                use_cases: vec!["cyber".to_string()],
1160                reasons: vec!["user_risk".to_string()],
1161                show_buffering_ui: true,
1162                faster_model: Some("gpt-fast-wire".to_string()),
1163            }
1164        );
1165    }
1166
1167    #[test]
1168    fn websocket_safety_buffering_event_controls_visibility_when_header_disables_it() {
1169        let metadata: ResponsesStreamEvent = serde_json::from_value(json!({
1170            "type": "codex.response.metadata",
1171            "headers": {
1172                "x-codex-safety-buffering-enabled": "false",
1173                "x-codex-safety-buffering-faster-model": "gpt-fast-header"
1174            }
1175        }))
1176        .expect("deserialize treatment metadata");
1177        let event: ResponsesStreamEvent = serde_json::from_value(json!({
1178            "type": "response.output_text.delta",
1179            "safety_buffering": {
1180                "use_cases": ["cyber"],
1181                "reasons": ["user_risk"]
1182            }
1183        }))
1184        .expect("deserialize safety buffering event");
1185        let mut treatment = SafetyBufferingTreatment::default();
1186
1187        assert!(safety_buffering_for_event(&metadata, &mut treatment).is_none());
1188        let buffering = safety_buffering_for_event(&event, &mut treatment)
1189            .expect("expected safety buffering payload");
1190
1191        assert_eq!(
1192            buffering,
1193            crate::common::SafetyBuffering {
1194                use_cases: vec!["cyber".to_string()],
1195                reasons: vec!["user_risk".to_string()],
1196                show_buffering_ui: true,
1197                faster_model: Some("gpt-fast-header".to_string()),
1198            }
1199        );
1200    }
1201}