use std::sync::Arc; use std::time::Duration; use futures_util::SinkExt; use futures_util::StreamExt; use http::HeaderMap; use http::HeaderValue; use serde::Deserialize; use serde_json::Value; use serde_json::json; use serde_json::map::Map as JsonMap; use tokio::net::TcpStream; use tokio::sync::Mutex; use tokio::sync::mpsc; use tokio::sync::oneshot; use tokio::time::Instant; use tokio_tungstenite::MaybeTlsStream; use tokio_tungstenite::WebSocketStream; use tokio_tungstenite::connect_async; use tokio_tungstenite::tungstenite::Error as WsError; use tokio_tungstenite::tungstenite::Message; use tokio_tungstenite::tungstenite::client::IntoClientRequest; use url::Url; use crate::ProviderError; use crate::ProviderEvent; use crate::ProviderEventStream; use crate::ProviderId; use crate::ReasoningFormat; use crate::ReasoningProvenance; use crate::ResponseHeaders; use crate::error::retry_after_from_header_value; use crate::error::retry_after_from_headers; use super::SharedTurnState; use super::sse::StreamState; use super::sse::parse_json_event; const X_CODEX_TURN_STATE_HEADER: &str = "x-codex-turn-state"; const WEBSOCKET_CONNECTION_LIMIT_REACHED_CODE: &str = "websocket_connection_limit_reached"; const WEBSOCKET_CONNECTION_LIMIT_REACHED_MESSAGE: &str = "Responses websocket connection limit reached (60 minutes). Create a new websocket connection to continue."; pub trait ResponsesWebsocketTelemetry: Send + Sync { fn on_ws_request( &self, duration: Duration, error: Option<&ProviderError>, connection_reused: bool, ); fn on_ws_event( &self, result: &Result>, ProviderError>, duration: Duration, ); } struct WsStream { tx_command: mpsc::Sender, rx_message: mpsc::UnboundedReceiver>, pump_task: tokio::task::JoinHandle<()>, } enum WsCommand { Send { message: Message, tx_result: oneshot::Sender>, }, } impl WsStream { fn new(inner: WebSocketStream>) -> Self { let (tx_command, mut rx_command) = mpsc::channel::(32); let (tx_message, rx_message) = mpsc::unbounded_channel::>(); let pump_task = tokio::spawn(async move { let mut inner = inner; loop { tokio::select! { command = rx_command.recv() => { let Some(command) = command else { break; }; match command { WsCommand::Send { message, tx_result } => { let result = inner.send(message).await; let should_break = result.is_err(); let _ = tx_result.send(result); if should_break { break; } } } } message = inner.next() => { let Some(message) = message else { break; }; match message { Ok(Message::Ping(payload)) => { if let Err(err) = inner.send(Message::Pong(payload)).await { let _ = tx_message.send(Err(err)); break; } } Ok(Message::Pong(_)) => {} Ok(message @ (Message::Text(_) | Message::Binary(_) | Message::Close(_) | Message::Frame(_))) => { let is_close = matches!(message, Message::Close(_)); if tx_message.send(Ok(message)).is_err() { break; } if is_close { break; } } Err(err) => { let _ = tx_message.send(Err(err)); break; } } } } } }); Self { tx_command, rx_message, pump_task, } } async fn send(&self, message: Message) -> Result<(), WsError> { let (tx_result, rx_result) = oneshot::channel(); if self .tx_command .send(WsCommand::Send { message, tx_result }) .await .is_err() { return Err(WsError::ConnectionClosed); } rx_result.await.unwrap_or(Err(WsError::ConnectionClosed)) } async fn next(&mut self) -> Option> { self.rx_message.recv().await } } impl Drop for WsStream { fn drop(&mut self) { self.pump_task.abort(); } } #[derive(Clone)] pub struct ResponsesWebsocketConnection { stream: Arc>>, idle_timeout: Duration, response_headers: ResponseHeaders, telemetry: Option>, } impl std::fmt::Debug for ResponsesWebsocketConnection { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("ResponsesWebsocketConnection") .field("stream", &"") .field("idle_timeout", &self.idle_timeout) .field("response_headers", &self.response_headers) .field("telemetry", &self.telemetry.as_ref().map(|_| "")) .finish() } } impl ResponsesWebsocketConnection { pub async fn connect( url: Url, headers: HeaderMap, turn_state: Option, idle_timeout: Duration, telemetry: Option>, ) -> Result { let mut request = url .as_str() .into_client_request() .map_err(|error| ProviderError::InvalidRequest(error.to_string()))?; request.headers_mut().extend(headers); let (stream, response) = connect_async(request) .await .map_err(|error| map_ws_error(error, &url))?; let header_value = response .headers() .get(X_CODEX_TURN_STATE_HEADER) .and_then(|value| value.to_str().ok()); if let (Some(turn_state), Some(header_value)) = (turn_state, header_value) { *turn_state .lock() .unwrap_or_else(std::sync::PoisonError::into_inner) = Some(header_value.to_string()); } let response_headers = ResponseHeaders { values: response .headers() .iter() .filter_map(|(name, value)| { value .to_str() .ok() .map(|value| (name.as_str().to_string(), value.to_string())) }) .collect(), }; Ok(Self { stream: Arc::new(Mutex::new(Some(WsStream::new(stream)))), idle_timeout, response_headers, telemetry, }) } pub async fn is_closed(&self) -> bool { self.stream.lock().await.is_none() } pub async fn stream_request( &self, request_body: Value, connection_reused: bool, ) -> Result { let requested_model = request_body .get("model") .and_then(Value::as_str) .unwrap_or_default() .to_string(); self.stream_request_with_provenance( request_body, connection_reused, ReasoningProvenance { provider: ProviderId::new("openai"), model: requested_model, format: ReasoningFormat::OpenAiEncrypted, }, ) .await } pub(crate) async fn stream_request_with_provenance( &self, request_body: Value, connection_reused: bool, provenance: ReasoningProvenance, ) -> Result { let (tx_event, rx_event) = mpsc::unbounded_channel::>(); let stream = Arc::clone(&self.stream); let idle_timeout = self.idle_timeout; let response_headers = self.response_headers.clone(); let telemetry = self.telemetry.clone(); let request_text = serde_json::to_string(&request_body).map_err(ProviderError::Serialize)?; tokio::spawn(async move { if tx_event .send(Ok(ProviderEvent::ResponseHeaders(response_headers))) .is_err() { return; } let mut guard = stream.lock().await; let result = { let Some(ws_stream) = guard.as_mut() else { let _ = tx_event.send(Err(ProviderError::MalformedStream( "websocket connection is closed".to_string(), ))); return; }; run_websocket_response_stream( ws_stream, tx_event.clone(), request_text, idle_timeout, telemetry, connection_reused, StreamState::new(provenance.provider, provenance.model), ) .await }; if let Err(err) = result { let failed_stream = guard.take(); drop(guard); drop(failed_stream); let _ = tx_event.send(Err(err)); } }); Ok(rx_event) } } pub fn merge_request_headers( provider_headers: &HeaderMap, extra_headers: HeaderMap, default_headers: HeaderMap, ) -> HeaderMap { let mut headers = provider_headers.clone(); headers.extend(extra_headers); for (name, value) in &default_headers { if let http::header::Entry::Vacant(entry) = headers.entry(name) { entry.insert(value.clone()); } } headers } pub(crate) fn response_create_frame(mut response: Value) -> Value { let Value::Object(response_object) = &mut response else { return json!({ "type": "response.create", "response": response, }); }; response_object.insert( "type".to_string(), Value::String("response.create".to_string()), ); response_object .entry("instructions") .or_insert_with(|| Value::String(String::new())); response } async fn run_websocket_response_stream( ws_stream: &mut WsStream, tx_event: mpsc::UnboundedSender>, request_text: String, idle_timeout: Duration, telemetry: Option>, connection_reused: bool, mut state: StreamState, ) -> Result<(), ProviderError> { let request_start = Instant::now(); let send_result = ws_stream.send(Message::Text(request_text.into())).await; let send_error = send_result .as_ref() .err() .map(|error| ProviderError::MalformedStream(error.to_string())); if let Some(t) = telemetry.as_ref() { t.on_ws_request( request_start.elapsed(), send_error.as_ref(), connection_reused, ); } send_result.map_err(|error| ProviderError::MalformedStream(error.to_string()))?; loop { let poll_start = Instant::now(); let message_result = tokio::time::timeout(idle_timeout, ws_stream.next()) .await .map_err(|_| { ProviderError::MalformedStream("idle timeout waiting for websocket".into()) }); if let Some(t) = telemetry.as_ref() { t.on_ws_event(&message_result, poll_start.elapsed()); } let message = match message_result { Ok(Some(Ok(message))) => message, Ok(Some(Err(error))) => { return Err(ProviderError::MalformedStream(error.to_string())); } Ok(None) => { return Err(ProviderError::MalformedStream( "stream closed before response.completed".to_string(), )); } Err(error) => return Err(error), }; match message { Message::Text(text) => { if let Some(mapped) = parse_wrapped_websocket_error_event(&text) { let receiver_closed = mapped .headers .map(|headers| { tx_event .send(Ok(ProviderEvent::ResponseHeaders(headers))) .is_err() }) .unwrap_or(false); if receiver_closed { return Ok(()); } return Err(mapped.error); } let mut saw_message_stopped = false; for event in parse_json_event(&text, &mut state)? { if matches!(event, ProviderEvent::MessageStopped) { saw_message_stopped = true; } if tx_event.send(Ok(event)).is_err() { return Ok(()); } } if saw_message_stopped { break; } } Message::Binary(_) => { return Err(ProviderError::MalformedStream( "unexpected binary websocket event".to_string(), )); } Message::Close(_) => { return Err(ProviderError::MalformedStream( "websocket closed by server before response.completed".to_string(), )); } Message::Frame(_) => {} Message::Ping(_) | Message::Pong(_) => {} } } Ok(()) } /// Delay hint for connection-level websocket failures. Short enough that a /// tunnel blip or service restart is retried promptly, long enough to avoid /// hammering a target that is still coming back up. const WS_CONNECTION_RETRY_DELAY: Duration = Duration::from_millis(750); fn map_ws_error(error: WsError, url: &Url) -> ProviderError { match error { WsError::Http(response) => { let status = response.status(); let retry_after = retry_after_from_headers(response.headers()); let body = response .body() .as_ref() .and_then(|bytes| String::from_utf8(bytes.clone()).ok()) .unwrap_or_else(|| format!("websocket connection failed for {url}")); ProviderError::Http { status, body, retry_after, } } // Connection lost mid-handshake: the peer went away before the upgrade // completed. That is a transient liveness failure (restart, tunnel // teardown), not a request the caller must fix — surface it as retryable. WsError::ConnectionClosed | WsError::AlreadyClosed => ProviderError::Retryable { message: format!("websocket connection closed while connecting to {url}"), delay: Some(WS_CONNECTION_RETRY_DELAY), }, // Transport-level failure establishing the connection (e.g. connection // refused when an SSH tunnel is down, DNS blips). These clear on their // own once the path is back, so the whole turn is worth retrying rather // than falling straight through to a terminal failure. WsError::Io(error) => ProviderError::Retryable { message: format!("websocket connection failed for {url}: {error}"), delay: Some(WS_CONNECTION_RETRY_DELAY), }, other => ProviderError::InvalidResponse(other.to_string()), } } #[derive(Debug)] struct MappedWebsocketError { error: ProviderError, headers: Option, } #[derive(Debug, Deserialize)] struct WrappedWebsocketError { code: Option, message: Option, } #[derive(Debug, Deserialize)] struct WrappedWebsocketErrorEvent { #[serde(rename = "type")] kind: String, #[serde(alias = "status_code")] status: Option, #[serde(default)] error: Option, #[serde(default)] headers: Option>, } fn parse_wrapped_websocket_error_event(payload: &str) -> Option { let event: WrappedWebsocketErrorEvent = serde_json::from_str(payload).ok()?; if event.kind != "error" { return None; } if let Some(error) = event .error .as_ref() .filter(|error| error.code.as_deref() == Some(WEBSOCKET_CONNECTION_LIMIT_REACHED_CODE)) { return Some(MappedWebsocketError { error: ProviderError::Retryable { message: error .message .clone() .unwrap_or_else(|| WEBSOCKET_CONNECTION_LIMIT_REACHED_MESSAGE.to_string()), delay: None, }, headers: event.headers.map(response_headers_from_json), }); } let status = reqwest::StatusCode::from_u16(event.status?).ok()?; let body = payload.to_string(); let headers = event.headers.map(response_headers_from_json); // The upgrade's own headers are long gone by the time a rate limit arrives // as a wrapped frame, so the only `Retry-After` on this path is the one the // frame echoed. Read it here or it is lost. let retry_after = headers.as_ref().and_then(retry_after_from_response_headers); Some(MappedWebsocketError { error: ProviderError::Http { status, body, retry_after, }, headers, }) } /// Reads `Retry-After` out of headers a frame carried as JSON. fn retry_after_from_response_headers(headers: &ResponseHeaders) -> Option { headers .values .iter() .find(|(name, _)| name.eq_ignore_ascii_case("retry-after")) .and_then(|(_, value)| retry_after_from_header_value(value)) } fn response_headers_from_json(headers: JsonMap) -> ResponseHeaders { ResponseHeaders { values: headers .into_iter() .filter_map(|(name, value)| { json_header_value(value) .and_then(|value| value.to_str().ok().map(|value| (name, value.to_string()))) }) .collect(), } } fn json_header_value(value: Value) -> Option { let value = match value { Value::String(value) => value, Value::Number(value) => value.to_string(), Value::Bool(value) => value.to_string(), _ => return None, }; HeaderValue::from_str(&value).ok() } #[cfg(test)] mod tests { use super::*; use serde_json::json; #[test] fn merge_request_headers_matches_http_precedence() { let mut provider_headers = HeaderMap::new(); provider_headers.insert( "originator", HeaderValue::from_static("provider-originator"), ); provider_headers.insert("x-priority", HeaderValue::from_static("provider")); let mut extra_headers = HeaderMap::new(); extra_headers.insert("x-priority", HeaderValue::from_static("extra")); let mut default_headers = HeaderMap::new(); default_headers.insert("originator", HeaderValue::from_static("default-originator")); default_headers.insert("x-priority", HeaderValue::from_static("default")); default_headers.insert("x-default-only", HeaderValue::from_static("default-only")); let merged = merge_request_headers(&provider_headers, extra_headers, default_headers); assert_eq!( merged.get("originator"), Some(&HeaderValue::from_static("provider-originator")) ); assert_eq!( merged.get("x-priority"), Some(&HeaderValue::from_static("extra")) ); assert_eq!( merged.get("x-default-only"), Some(&HeaderValue::from_static("default-only")) ); } #[test] fn response_create_frame_adds_type_to_responses_request_payload() { let frame = response_create_frame(json!({ "model": "gpt-5", "input": [] })); assert_eq!(frame["type"], "response.create"); assert_eq!(frame["model"], "gpt-5"); assert_eq!(frame["input"], json!([])); assert_eq!(frame["instructions"], ""); assert!(frame.get("response").is_none()); } #[test] fn wrapped_websocket_error_preserves_rate_limit_headers() { let payload = json!({ "type": "error", "status": 429, "error": { "type": "usage_limit_reached", "message": "The usage limit has been reached" }, "headers": { "x-codex-primary-used-percent": "100.0" } }) .to_string(); let mapped = parse_wrapped_websocket_error_event(&payload) .expect("error payload should map to a provider error"); let ProviderError::Http { status, .. } = mapped.error else { panic!("expected ProviderError::Http"); }; assert_eq!(status, reqwest::StatusCode::TOO_MANY_REQUESTS); assert_eq!( mapped .headers .expect("rate limit headers should be present") .values, vec![( "x-codex-primary-used-percent".to_string(), "100.0".to_string() )] ); } #[test] fn ws_io_error_maps_to_retryable_with_delay() { let url = Url::parse("wss://example.test/responses").expect("url parses"); let io = std::io::Error::new(std::io::ErrorKind::ConnectionRefused, "Connection refused"); let mapped = map_ws_error(WsError::Io(io), &url); let ProviderError::Retryable { message, delay } = mapped else { panic!("expected ProviderError::Retryable for a connection-level io error"); }; assert_eq!(delay, Some(WS_CONNECTION_RETRY_DELAY)); // The io error string is preserved so operators can see the real cause. assert!(message.contains("Connection refused"), "message: {message}"); } #[test] fn ws_connection_closed_maps_to_retryable() { let url = Url::parse("wss://example.test/responses").expect("url parses"); for error in [WsError::ConnectionClosed, WsError::AlreadyClosed] { let mapped = map_ws_error(error, &url); assert!( matches!(mapped, ProviderError::Retryable { delay: Some(_), .. }), "connection-closed handshake failure should be retryable", ); } } #[test] fn websocket_connection_limit_maps_to_retryable() { let payload = json!({ "type": "error", "status": 400, "error": { "type": "invalid_request_error", "code": "websocket_connection_limit_reached", "message": WEBSOCKET_CONNECTION_LIMIT_REACHED_MESSAGE } }) .to_string(); let mapped = parse_wrapped_websocket_error_event(&payload) .expect("connection limit payload should map"); let ProviderError::Retryable { message, delay } = mapped.error else { panic!("expected ProviderError::Retryable"); }; assert_eq!(message, WEBSOCKET_CONNECTION_LIMIT_REACHED_MESSAGE); assert_eq!(delay, None); } }