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 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 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
332pub struct ResponsesWebsocketClient {
334 provider: Provider,
335 auth: SharedAuthProvider,
336}
337
338#[derive(Debug, Clone, PartialEq, Eq)]
340pub struct ResponsesWebsocketClose {
341 pub code: String,
343 pub reason: String,
345}
346
347#[derive(Debug, Clone, PartialEq, Eq)]
349pub struct ResponsesWebsocketProbe {
350 pub url: String,
352 pub status: StatusCode,
354 pub reasoning_included: bool,
356 pub models_etag_present: bool,
358 pub server_model_present: bool,
360 pub immediate_close: Option<ResponsesWebsocketClose>,
362}
363
364impl ResponsesWebsocketClient {
365 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 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 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 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}