1use std::collections::HashMap;
2use std::sync::{Arc, Mutex};
3use std::time::{Duration, Instant};
4
5use futures_util::{SinkExt, StreamExt};
6use http::HeaderMap;
7use tokio::net::TcpStream;
8use tokio::sync::Mutex as AsyncMutex;
9use tokio::sync::mpsc;
10use tokio_tungstenite::{
11 MaybeTlsStream, WebSocketStream, connect_async,
12 tungstenite::{self, Message, handshake::client::generate_key},
13};
14
15use crate::provider::RequestContext;
16use crate::traffic::TrafficCapture;
17
18use super::client::{CodexError, CodexErrorOrigin, CodexResponse};
19use super::continuation::ContinuationCandidate;
20
21pub const WEBSOCKET_PROTOCOL_HEADER: &str = "responses_websockets=2026-02-06";
26pub const WEBSOCKET_CONNECT_TIMEOUT_MS: u64 = 15_000;
27pub const WEBSOCKET_IDLE_TIMEOUT_MS: u64 = 300_000;
28pub const WEBSOCKET_RESPONSE_START_TIMEOUT_DETAIL: &str = "websocket_response_start_timeout";
29pub const WEBSOCKET_MISSING_TERMINAL_DETAIL: &str = "websocket_missing_terminal";
30
31const POOL_IDLE_TTL_MS: u64 = 30 * 60 * 1000;
32const MAX_POOL_ENTRIES: usize = 10_000;
33
34const TERMINAL_EVENTS: &[&str] = &[
36 "response.completed",
37 "response.incomplete",
38 "response.failed",
39 "error",
40];
41
42pub type CodexWebSocketEventReceiver = mpsc::Receiver<Result<serde_json::Value, CodexError>>;
43
44#[derive(Debug, Clone)]
49pub struct CodexWebSocketError {
50 pub message: String,
51 pub status: Option<u16>,
52 pub code: Option<String>,
53 pub retry_after: Option<String>,
54 pub request_sent: bool,
55}
56
57impl CodexWebSocketError {
58 pub fn new(message: String) -> Self {
59 Self {
60 message,
61 status: None,
62 code: None,
63 retry_after: None,
64 request_sent: false,
65 }
66 }
67
68 pub fn with_status(mut self, status: u16) -> Self {
69 self.status = Some(status);
70 self
71 }
72
73 pub fn with_code(mut self, code: String) -> Self {
74 self.code = Some(code);
75 self
76 }
77}
78
79impl std::fmt::Display for CodexWebSocketError {
80 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
81 write!(f, "Codex WebSocket error: {}", self.message)
82 }
83}
84
85struct PoolEntry {
90 ws: Arc<AsyncMutex<WebSocketStream<MaybeTlsStream<TcpStream>>>>,
91 created_at: u64,
92}
93
94static WS_POOL: once_cell::sync::Lazy<Mutex<HashMap<String, Arc<PoolEntry>>>> =
95 once_cell::sync::Lazy::new(|| Mutex::new(HashMap::new()));
96
97fn now_ms() -> u64 {
98 std::time::SystemTime::now()
99 .duration_since(std::time::UNIX_EPOCH)
100 .unwrap_or_default()
101 .as_millis() as u64
102}
103
104pub fn clear_codex_websocket_pool_for_tests() {
105 let mut guard = WS_POOL.lock().unwrap();
106 guard.clear();
107}
108
109pub fn invalidate_codex_websocket_pool_key(session_id: &str) {
110 let mut guard = WS_POOL.lock().unwrap();
111 guard.remove(session_id);
112}
113
114fn pool_insert(key: String, entry: Arc<PoolEntry>) {
115 let mut guard = WS_POOL.lock().unwrap();
116 if guard.len() >= MAX_POOL_ENTRIES
118 && let Some(oldest_key) = guard.keys().next().cloned()
119 {
120 guard.remove(&oldest_key);
121 }
122 let now = now_ms();
124 guard.retain(|_, e| now.saturating_sub(e.created_at) < POOL_IDLE_TTL_MS);
125 guard.insert(key, entry);
126}
127
128pub fn to_websocket_url(url: &str) -> Result<String, CodexWebSocketError> {
133 let mut parsed = url::Url::parse(url)
134 .map_err(|e| CodexWebSocketError::new(format!("Failed to parse URL: {e}")))?;
135 match parsed.scheme() {
136 "http" => parsed.set_scheme("ws").map_err(|_| {
137 CodexWebSocketError::new("Unsupported Codex WebSocket URL scheme".to_string())
138 })?,
139 "https" => parsed.set_scheme("wss").map_err(|_| {
140 CodexWebSocketError::new("Unsupported Codex WebSocket URL scheme".to_string())
141 })?,
142 "ws" | "wss" => { }
143 other => {
144 return Err(CodexWebSocketError::new(format!(
145 "Unsupported Codex WebSocket URL scheme: {other}"
146 )));
147 }
148 }
149 Ok(parsed.to_string())
150}
151
152pub fn codex_websocket_headers(http_headers: &HeaderMap) -> HeaderMap {
157 let mut ws = HeaderMap::new();
158 for (key, value) in http_headers.iter() {
159 let key_str = key.as_str().to_lowercase();
160 if matches!(
162 key_str.as_str(),
163 "content-length" | "content-type" | "accept" | "connection" | "upgrade"
164 ) {
165 continue;
166 }
167 ws.insert(key.clone(), value.clone());
168 }
169 ws.insert("openai-beta", WEBSOCKET_PROTOCOL_HEADER.parse().unwrap());
171 if !ws.contains_key("sec-websocket-key") {
173 ws.insert("sec-websocket-key", generate_key().parse().unwrap());
174 }
175 ws
176}
177
178fn encode_sse(text: &str) -> Vec<u8> {
183 let mut out = String::new();
184 for line in text.lines() {
185 out.push_str("data: ");
186 out.push_str(line);
187 out.push('\n');
188 }
189 out.push('\n');
190 out.into_bytes()
191}
192
193fn is_terminal_event(payload: &serde_json::Value) -> bool {
198 match payload.get("type").and_then(|v| v.as_str()) {
199 Some(t) => TERMINAL_EVENTS.contains(&t),
200 None => false,
201 }
202}
203
204fn is_response_event(payload: &serde_json::Value) -> bool {
205 match payload.get("type").and_then(|v| v.as_str()) {
206 Some("error") => true,
207 Some(t) => t.starts_with("response."),
208 None => false,
209 }
210}
211
212fn is_previous_response_missing(payload: &serde_json::Value) -> bool {
213 if let Some(code) = payload
214 .get("error")
215 .and_then(|e| e.get("code"))
216 .and_then(|v| v.as_str())
217 && code == "previous_response_not_found"
218 {
219 return true;
220 }
221 if let Some(msg) = payload
223 .get("error")
224 .and_then(|e| e.get("message"))
225 .and_then(|v| v.as_str())
226 {
227 let lower = msg.to_lowercase();
228 if lower.contains("previous response") && lower.contains("not found") {
229 return true;
230 }
231 }
232 false
233}
234
235pub(super) fn event_error_status(payload: &serde_json::Value) -> Option<u16> {
236 super::events::classify_event_failure(payload).and_then(|failure| failure.explicit_status)
237}
238
239#[allow(dead_code)]
240fn extract_retry_after(payload: &serde_json::Value) -> Option<String> {
241 payload
242 .get("error")
243 .and_then(|e| e.get("retry_after"))
244 .and_then(|v| v.as_str())
245 .map(|s| s.to_string())
246}
247
248#[allow(clippy::too_many_arguments)]
253pub async fn codex_websocket_request(
254 url: &str,
255 headers: &HeaderMap,
256 body_value: &serde_json::Value,
257 _ctx: &RequestContext,
258 traffic: Option<&TrafficCapture>,
259 pool_key: Option<&str>,
260 connect_timeout_ms: u64,
261 idle_timeout_ms: u64,
262 continuation: Option<&ContinuationCandidate>,
263) -> Result<CodexResponse, CodexError> {
264 let ws_url = to_websocket_url(url).map_err(|e| CodexError {
265 status: 0,
266 message: e.message,
267 detail: None,
268 retry_after: None,
269 origin: CodexErrorOrigin::WebSocketHandshake,
270 })?;
271 let body_json = serde_json::to_string(body_value).unwrap_or_default();
272 if let Some(tc) = traffic {
273 tc.write_json("020-upstream-request", body_value);
274 tc.write_json(
275 "021-upstream-request-metadata",
276 &serde_json::json!({
277 "provider": "codex",
278 "transport": "websocket",
279 "url": ws_url,
280 "method": "GET",
281 "headers": headers_to_json(headers),
282 "size": summarize_json_request_size(body_value, &body_json),
283 "continuation": {
284 "previousResponseId": continuation
285 .and_then(|c| c.previous_response_id.as_deref()),
286 "inputDeltaCount": continuation
287 .and_then(|c| c.input_delta.as_ref())
288 .map(|items| items.len()),
289 "disabledReason": continuation
290 .and_then(|c| c.disabled_reason.as_deref()),
291 },
292 }),
293 );
294 }
295 let started_at = Instant::now();
296
297 let pooled = pool_key.and_then(|key| {
299 let guard = WS_POOL.lock().ok()?;
300 guard.get(key).cloned()
301 });
302
303 let (ws_stream, _response) = if let Some(entry) = pooled {
304 let mut ws_guard = entry.ws.lock().await;
306 if ws_guard.send(Message::Ping(vec![])).await.is_err() {
308 invalidate_codex_websocket_pool_key(pool_key.unwrap());
309 connect_with_timeout(&ws_url, headers, connect_timeout_ms).await?
311 } else {
312 let ws_msg = Message::Text(body_json.clone());
314 ws_guard.send(ws_msg).await.map_err(|e| {
315 if let Some(key) = pool_key {
316 invalidate_codex_websocket_pool_key(key);
317 }
318 CodexError {
319 status: 0,
320 message: format!("WebSocket send error: {e}"),
321 detail: None,
322 retry_after: None,
323 origin: CodexErrorOrigin::WebSocket,
324 }
325 })?;
326
327 let (sse_body, terminal_event) =
329 collect_ws_events(&mut ws_guard, idle_timeout_ms, pool_key, traffic).await?;
330 let Some(terminal_event) = terminal_event else {
331 return Err(missing_terminal_error());
332 };
333
334 if is_previous_response_missing(&terminal_event.payload) {
336 return Err(CodexError {
337 status: 0,
338 message: "Previous response not found".to_string(),
339 detail: Some("previous_response_not_found".to_string()),
340 retry_after: None,
341 origin: CodexErrorOrigin::WebSocket,
342 });
343 }
344
345 let status = if terminal_event.event_type == "error" {
347 event_error_status(&terminal_event.payload).unwrap_or(500)
348 } else {
349 200
350 };
351
352 if let Some(tc) = traffic {
354 write_websocket_metadata_capture(tc, &ws_url, pool_key, continuation, true);
355 write_websocket_response_capture(tc, status, started_at.elapsed(), &sse_body);
356 }
357
358 return Ok(CodexResponse {
359 body: sse_body,
360 status,
361 headers: vec![],
362 });
363 }
364 } else {
365 connect_with_timeout(&ws_url, headers, connect_timeout_ms).await?
366 };
367
368 let entry = Arc::new(PoolEntry {
370 ws: Arc::new(AsyncMutex::new(ws_stream)),
371 created_at: now_ms(),
372 });
373
374 let msg = Message::Text(body_json);
376 {
377 let mut ws_guard = entry.ws.lock().await;
378 ws_guard.send(msg).await.map_err(|e| CodexError {
379 status: 0,
380 message: format!("WebSocket send error: {e}"),
381 detail: None,
382 retry_after: None,
383 origin: CodexErrorOrigin::WebSocket,
384 })?;
385
386 let (sse_body, terminal_event) =
387 collect_ws_events(&mut ws_guard, idle_timeout_ms, pool_key, traffic).await?;
388 let Some(terminal_event) = terminal_event else {
389 return Err(missing_terminal_error());
390 };
391
392 if is_previous_response_missing(&terminal_event.payload) {
393 if let Some(key) = pool_key {
394 invalidate_codex_websocket_pool_key(key);
395 }
396 return Err(CodexError {
397 status: 0,
398 message: "Previous response not found".to_string(),
399 detail: Some("previous_response_not_found".to_string()),
400 retry_after: None,
401 origin: CodexErrorOrigin::WebSocket,
402 });
403 }
404
405 if let Some(key) = pool_key {
407 let should_pool = terminal_event.event_type == "response.completed";
408 if should_pool {
409 pool_insert(key.to_string(), entry.clone());
410 }
411 }
412
413 let status = if terminal_event.event_type == "error" {
414 event_error_status(&terminal_event.payload).unwrap_or(500)
415 } else {
416 200
417 };
418
419 if let Some(tc) = traffic {
421 write_websocket_metadata_capture(tc, &ws_url, pool_key, continuation, false);
422 write_websocket_response_capture(tc, status, started_at.elapsed(), &sse_body);
423 }
424
425 Ok(CodexResponse {
426 body: sse_body,
427 status,
428 headers: vec![],
429 })
430 }
431}
432
433#[allow(clippy::too_many_arguments)]
434pub async fn codex_websocket_event_stream(
435 url: &str,
436 headers: &HeaderMap,
437 body_value: &serde_json::Value,
438 _ctx: &RequestContext,
439 traffic: Option<Arc<TrafficCapture>>,
440 pool_key: Option<&str>,
441 connect_timeout_ms: u64,
442 idle_timeout_ms: u64,
443 continuation: Option<&ContinuationCandidate>,
444) -> Result<CodexWebSocketEventReceiver, CodexError> {
445 let ws_url = to_websocket_url(url).map_err(|e| CodexError {
446 status: 0,
447 message: e.message,
448 detail: None,
449 retry_after: None,
450 origin: CodexErrorOrigin::WebSocket,
451 })?;
452 let body_json = serde_json::to_string(body_value).unwrap_or_default();
453 if let Some(tc) = traffic.as_deref() {
454 tc.write_json("020-upstream-request", body_value);
455 tc.write_json(
456 "021-upstream-request-metadata",
457 &serde_json::json!({
458 "provider": "codex",
459 "transport": "websocket",
460 "url": ws_url,
461 "method": "GET",
462 "headers": headers_to_json(headers),
463 "size": summarize_json_request_size(body_value, &body_json),
464 "continuation": {
465 "previousResponseId": continuation
466 .and_then(|c| c.previous_response_id.as_deref()),
467 "inputDeltaCount": continuation
468 .and_then(|c| c.input_delta.as_ref())
469 .map(|items| items.len()),
470 "disabledReason": continuation
471 .and_then(|c| c.disabled_reason.as_deref()),
472 },
473 }),
474 );
475 }
476
477 let pooled = pool_key.and_then(|key| {
478 let guard = WS_POOL.lock().ok()?;
479 guard.get(key).cloned()
480 });
481 let used_pooled = pooled.is_some();
482 let entry = if let Some(entry) = pooled {
483 entry
484 } else {
485 let (ws_stream, _) = connect_with_timeout(&ws_url, headers, connect_timeout_ms).await?;
486 Arc::new(PoolEntry {
487 ws: Arc::new(AsyncMutex::new(ws_stream)),
488 created_at: now_ms(),
489 })
490 };
491
492 if let Some(tc) = traffic.as_deref() {
493 write_websocket_metadata_capture(tc, &ws_url, pool_key, continuation, used_pooled);
494 }
495
496 let (tx, rx) = mpsc::channel(64);
497 let pool_key = pool_key.map(str::to_string);
498 let ws = entry.ws.clone();
499 tokio::spawn(async move {
500 let mut ws_guard = ws.lock_owned().await;
501 if used_pooled && ws_guard.send(Message::Ping(vec![])).await.is_err() {
502 if let Some(key) = pool_key.as_deref() {
503 invalidate_codex_websocket_pool_key(key);
504 }
505 let _ = tx
506 .send(Err(CodexError {
507 status: 0,
508 message: "WebSocket send error: failed to ping pooled connection".to_string(),
509 detail: None,
510 retry_after: None,
511 origin: CodexErrorOrigin::WebSocket,
512 }))
513 .await;
514 return;
515 }
516 if let Err(e) = ws_guard.send(Message::Text(body_json)).await {
517 if let Some(key) = pool_key.as_deref() {
518 invalidate_codex_websocket_pool_key(key);
519 }
520 let _ = tx
521 .send(Err(CodexError {
522 status: 0,
523 message: format!("WebSocket send error: {e}"),
524 detail: None,
525 retry_after: None,
526 origin: CodexErrorOrigin::WebSocket,
527 }))
528 .await;
529 return;
530 }
531
532 let reusable = stream_ws_events(
533 &mut ws_guard,
534 idle_timeout_ms,
535 pool_key.as_deref(),
536 traffic,
537 tx,
538 )
539 .await;
540
541 if let Some(key) = pool_key.as_deref() {
542 if reusable {
543 if !used_pooled {
544 pool_insert(key.to_string(), entry.clone());
545 }
546 } else {
547 invalidate_codex_websocket_pool_key(key);
548 }
549 }
550 });
551 Ok(rx)
552}
553
554fn missing_terminal_error() -> CodexError {
555 CodexError {
556 status: 0,
557 message: "WebSocket connection closed before terminal Codex response event".to_string(),
558 detail: Some(WEBSOCKET_MISSING_TERMINAL_DETAIL.to_string()),
559 retry_after: None,
560 origin: CodexErrorOrigin::WebSocket,
561 }
562}
563
564fn response_start_timeout_error(timeout_ms: u64) -> CodexError {
565 CodexError {
566 status: 0,
567 message: format!("WebSocket response start timeout after {timeout_ms}ms"),
568 detail: Some(WEBSOCKET_RESPONSE_START_TIMEOUT_DETAIL.to_string()),
569 retry_after: None,
570 origin: CodexErrorOrigin::WebSocket,
571 }
572}
573
574fn write_websocket_metadata_capture(
575 traffic: &TrafficCapture,
576 ws_url: &str,
577 pool_key: Option<&str>,
578 continuation: Option<&ContinuationCandidate>,
579 pooled: bool,
580) {
581 traffic.write_json(
582 "022-upstream-websocket-metadata",
583 &serde_json::json!({
584 "provider": "codex",
585 "transport": "websocket",
586 "url": ws_url,
587 "poolKey": pool_key,
588 "pooled": pooled,
589 "continuation": {
590 "previousResponseId": continuation
591 .and_then(|c| c.previous_response_id.as_deref()),
592 "inputDeltaCount": continuation
593 .and_then(|c| c.input_delta.as_ref())
594 .map(|items| items.len()),
595 "disabledReason": continuation
596 .and_then(|c| c.disabled_reason.as_deref()),
597 },
598 }),
599 );
600}
601
602fn write_websocket_response_capture(
603 traffic: &TrafficCapture,
604 status: u16,
605 elapsed: Duration,
606 sse_body: &[u8],
607) {
608 traffic.write_json(
609 "030-upstream-response-headers",
610 &serde_json::json!({
611 "status": status,
612 "elapsedMs": elapsed.as_millis(),
613 "headers": {
614 "content-type": "text/event-stream",
615 },
616 }),
617 );
618 if status >= 400 {
619 traffic.write_text(
620 "031-upstream-error-body",
621 &String::from_utf8_lossy(sse_body),
622 );
623 } else {
624 traffic.write_bytes("032-upstream-response-body.sse", sse_body);
625 }
626}
627
628async fn connect_with_timeout(
633 url: &str,
634 headers: &HeaderMap,
635 connect_timeout_ms: u64,
636) -> Result<
637 (
638 WebSocketStream<MaybeTlsStream<TcpStream>>,
639 tungstenite::handshake::client::Response,
640 ),
641 CodexError,
642> {
643 let host = websocket_host_header(url);
645 let mut req_builder = http::Request::builder()
646 .uri(url)
647 .method("GET")
648 .header("Host", host)
649 .header("Connection", "Upgrade")
650 .header("Upgrade", "websocket")
651 .header("Sec-WebSocket-Version", "13")
652 .header("Sec-WebSocket-Key", generate_key());
653
654 for (key, value) in headers.iter() {
656 let key_str = key.as_str().to_lowercase();
657 if matches!(
659 key_str.as_str(),
660 "connection" | "upgrade" | "sec-websocket-key" | "sec-websocket-version" | "host"
661 ) {
662 continue;
663 }
664 req_builder = req_builder.header(key.as_str(), value.as_bytes());
665 }
666
667 let request = req_builder.body(()).map_err(|e| CodexError {
668 status: 0,
669 message: format!("Failed to build WebSocket request: {e}"),
670 detail: None,
671 retry_after: None,
672 origin: CodexErrorOrigin::WebSocket,
673 })?;
674
675 let connect_fut = connect_async(request);
676 tokio::time::timeout(Duration::from_millis(connect_timeout_ms), connect_fut)
677 .await
678 .map_err(|_| CodexError {
679 status: 0,
680 message: format!("WebSocket connect timeout after {connect_timeout_ms}ms"),
681 detail: None,
682 retry_after: None,
683 origin: CodexErrorOrigin::WebSocketHandshake,
684 })?
685 .map_err(|e| {
686 let (status, retry_after, detail) = match &e {
687 tungstenite::Error::Http(response) => {
688 let detail = response
689 .body()
690 .as_ref()
691 .and_then(|body| String::from_utf8(body.clone()).ok())
692 .filter(|body| !body.trim().is_empty());
693 (
694 Some(response.status().as_u16()),
695 response
696 .headers()
697 .get(http::header::RETRY_AFTER)
698 .and_then(|value| value.to_str().ok())
699 .map(str::to_string),
700 detail,
701 )
702 }
703 _ => (None, None, None),
704 };
705 CodexError {
706 status: status.unwrap_or(0),
707 message: format!("WebSocket connect error: {e}"),
708 detail,
709 retry_after,
710 origin: CodexErrorOrigin::WebSocketHandshake,
711 }
712 })
713}
714
715fn websocket_host_header(url: &str) -> String {
716 let Ok(parsed) = url::Url::parse(url) else {
717 return String::new();
718 };
719 parsed[url::Position::BeforeHost..url::Position::AfterPort].to_string()
720}
721
722struct WsEvent {
727 event_type: String,
728 payload: serde_json::Value,
729}
730
731async fn collect_ws_events(
732 ws: &mut WebSocketStream<MaybeTlsStream<TcpStream>>,
733 idle_timeout_ms: u64,
734 pool_key: Option<&str>,
735 traffic: Option<&TrafficCapture>,
736) -> Result<(Vec<u8>, Option<WsEvent>), CodexError> {
737 let mut sse_body: Vec<u8> = Vec::new();
738 let mut terminal_event: Option<WsEvent> = None;
739 let response_event_budget = Duration::from_millis(idle_timeout_ms);
740 let response_wait_started = Instant::now();
741 let mut last_response_event_at = response_wait_started;
742 let mut response_started = false;
743
744 loop {
745 let response_deadline_started = if response_started {
746 last_response_event_at
747 } else {
748 response_wait_started
749 };
750 let read_timeout = if response_started {
751 match response_event_budget.checked_sub(response_deadline_started.elapsed()) {
752 Some(remaining) if !remaining.is_zero() => remaining,
753 _ => {
754 if let Some(key) = pool_key {
755 invalidate_codex_websocket_pool_key(key);
756 }
757 return Err(CodexError {
758 status: 0,
759 message: format!("WebSocket idle timeout after {idle_timeout_ms}ms"),
760 detail: None,
761 retry_after: None,
762 origin: CodexErrorOrigin::WebSocket,
763 });
764 }
765 }
766 } else {
767 match response_event_budget.checked_sub(response_deadline_started.elapsed()) {
768 Some(remaining) if !remaining.is_zero() => remaining,
769 _ => {
770 if let Some(key) = pool_key {
771 invalidate_codex_websocket_pool_key(key);
772 }
773 return Err(response_start_timeout_error(idle_timeout_ms));
774 }
775 }
776 };
777
778 let timeout = tokio::time::timeout(read_timeout, ws.next());
779
780 let frame = timeout.await.map_err(|_| {
781 if let Some(key) = pool_key {
782 invalidate_codex_websocket_pool_key(key);
783 }
784 if response_started {
785 CodexError {
786 status: 0,
787 message: format!("WebSocket idle timeout after {idle_timeout_ms}ms"),
788 detail: None,
789 retry_after: None,
790 origin: CodexErrorOrigin::WebSocket,
791 }
792 } else {
793 response_start_timeout_error(idle_timeout_ms)
794 }
795 })?;
796
797 match frame {
798 Some(Ok(Message::Text(text))) => {
799 let parsed: serde_json::Value = match serde_json::from_str(&text) {
801 Ok(v) => v,
802 Err(_) => {
803 if let Some(tc) = traffic {
804 tc.write_json_event(
805 "040-upstream-event",
806 &serde_json::json!({
807 "unparseable": true,
808 "data": text,
809 }),
810 );
811 }
812 sse_body.extend_from_slice(&encode_sse(&text));
814 continue;
815 }
816 };
817
818 sse_body.extend_from_slice(&encode_sse(&text));
820 if let Some(tc) = traffic {
821 tc.write_json_event("040-upstream-event", &parsed);
822 }
823
824 if is_response_event(&parsed) {
825 response_started = true;
826 last_response_event_at = Instant::now();
827 }
828
829 if is_terminal_event(&parsed) {
831 terminal_event = Some(WsEvent {
832 event_type: parsed
833 .get("type")
834 .and_then(|v| v.as_str())
835 .unwrap_or("unknown")
836 .to_string(),
837 payload: parsed,
838 });
839 break;
840 }
841 }
842 Some(Ok(Message::Binary(_))) => {
843 if let Some(key) = pool_key {
845 invalidate_codex_websocket_pool_key(key);
846 }
847 return Err(CodexError {
848 status: 0,
849 message: "WebSocket binary frames not supported".to_string(),
850 detail: None,
851 retry_after: None,
852 origin: CodexErrorOrigin::WebSocket,
853 });
854 }
855 Some(Ok(Message::Ping(data))) => {
856 let _ = ws.send(Message::Pong(data)).await;
858 continue;
859 }
860 Some(Ok(Message::Pong(_))) => {
861 continue;
862 }
863 Some(Ok(Message::Frame(_))) => {
864 continue;
866 }
867 Some(Ok(Message::Close(_))) => {
868 if let Some(key) = pool_key {
870 invalidate_codex_websocket_pool_key(key);
871 }
872 break;
873 }
874 Some(Err(e)) => {
875 if let Some(key) = pool_key {
877 invalidate_codex_websocket_pool_key(key);
878 }
879 return Err(CodexError {
880 status: 0,
881 message: format!("WebSocket stream error: {e}"),
882 detail: None,
883 retry_after: None,
884 origin: CodexErrorOrigin::WebSocket,
885 });
886 }
887 None => {
888 if let Some(key) = pool_key {
890 invalidate_codex_websocket_pool_key(key);
891 }
892 break;
893 }
894 }
895 }
896
897 Ok((sse_body, terminal_event))
898}
899
900async fn stream_ws_events(
901 ws: &mut WebSocketStream<MaybeTlsStream<TcpStream>>,
902 idle_timeout_ms: u64,
903 pool_key: Option<&str>,
904 traffic: Option<Arc<TrafficCapture>>,
905 tx: mpsc::Sender<Result<serde_json::Value, CodexError>>,
906) -> bool {
907 let started_at = Instant::now();
908 let mut sse_body: Vec<u8> = Vec::new();
909 let response_event_budget = Duration::from_millis(idle_timeout_ms);
910 let response_wait_started = Instant::now();
911 let mut last_response_event_at = response_wait_started;
912 let mut response_started = false;
913 let mut status = 200u16;
914 let mut reusable = false;
915
916 loop {
917 let response_deadline_started = if response_started {
918 last_response_event_at
919 } else {
920 response_wait_started
921 };
922 let read_timeout =
923 match response_event_budget.checked_sub(response_deadline_started.elapsed()) {
924 Some(remaining) if !remaining.is_zero() => remaining,
925 _ => {
926 if let Some(key) = pool_key {
927 invalidate_codex_websocket_pool_key(key);
928 }
929 let err = if response_started {
930 CodexError {
931 status: 0,
932 message: format!("WebSocket idle timeout after {idle_timeout_ms}ms"),
933 detail: None,
934 retry_after: None,
935 origin: CodexErrorOrigin::WebSocket,
936 }
937 } else {
938 response_start_timeout_error(idle_timeout_ms)
939 };
940 let _ = tx.send(Err(err)).await;
941 break;
942 }
943 };
944
945 let frame = match tokio::time::timeout(read_timeout, ws.next()).await {
946 Ok(frame) => frame,
947 Err(_) => {
948 if let Some(key) = pool_key {
949 invalidate_codex_websocket_pool_key(key);
950 }
951 let err = if response_started {
952 CodexError {
953 status: 0,
954 message: format!("WebSocket idle timeout after {idle_timeout_ms}ms"),
955 detail: None,
956 retry_after: None,
957 origin: CodexErrorOrigin::WebSocket,
958 }
959 } else {
960 response_start_timeout_error(idle_timeout_ms)
961 };
962 let _ = tx.send(Err(err)).await;
963 break;
964 }
965 };
966
967 match frame {
968 Some(Ok(Message::Text(text))) => {
969 let parsed: serde_json::Value = match serde_json::from_str(&text) {
970 Ok(v) => v,
971 Err(_) => {
972 if let Some(tc) = traffic.as_deref() {
973 tc.write_json_event(
974 "040-upstream-event",
975 &serde_json::json!({
976 "unparseable": true,
977 "data": text,
978 }),
979 );
980 }
981 sse_body.extend_from_slice(&encode_sse(&text));
982 continue;
983 }
984 };
985
986 sse_body.extend_from_slice(&encode_sse(&text));
987 if let Some(tc) = traffic.as_deref() {
988 tc.write_json_event("040-upstream-event", &parsed);
989 }
990
991 if is_response_event(&parsed) {
992 response_started = true;
993 last_response_event_at = Instant::now();
994 }
995
996 if parsed.get("type").and_then(|v| v.as_str()) == Some("error") {
997 status = event_error_status(&parsed).unwrap_or(500);
998 }
999 let terminal = is_terminal_event(&parsed);
1000 if terminal && is_previous_response_missing(&parsed) {
1001 if let Some(key) = pool_key {
1002 invalidate_codex_websocket_pool_key(key);
1003 }
1004 let _ = tx
1005 .send(Err(CodexError {
1006 status: 0,
1007 message: "Previous response not found".to_string(),
1008 detail: Some("previous_response_not_found".to_string()),
1009 retry_after: None,
1010 origin: CodexErrorOrigin::WebSocket,
1011 }))
1012 .await;
1013 break;
1014 }
1015 let event_type = parsed
1016 .get("type")
1017 .and_then(|v| v.as_str())
1018 .unwrap_or("unknown")
1019 .to_string();
1020 if tx.send(Ok(parsed)).await.is_err() {
1021 if let Some(key) = pool_key {
1022 invalidate_codex_websocket_pool_key(key);
1023 }
1024 break;
1025 }
1026 if terminal {
1027 reusable = event_type == "response.completed";
1028 break;
1029 }
1030 }
1031 Some(Ok(Message::Binary(_))) => {
1032 if let Some(key) = pool_key {
1033 invalidate_codex_websocket_pool_key(key);
1034 }
1035 let _ = tx
1036 .send(Err(CodexError {
1037 status: 0,
1038 message: "WebSocket binary frames not supported".to_string(),
1039 detail: None,
1040 retry_after: None,
1041 origin: CodexErrorOrigin::WebSocket,
1042 }))
1043 .await;
1044 break;
1045 }
1046 Some(Ok(Message::Ping(data))) => {
1047 let _ = ws.send(Message::Pong(data)).await;
1048 }
1049 Some(Ok(Message::Pong(_))) | Some(Ok(Message::Frame(_))) => {}
1050 Some(Ok(Message::Close(_))) | None => {
1051 if let Some(key) = pool_key {
1052 invalidate_codex_websocket_pool_key(key);
1053 }
1054 let _ = tx.send(Err(missing_terminal_error())).await;
1055 break;
1056 }
1057 Some(Err(e)) => {
1058 if let Some(key) = pool_key {
1059 invalidate_codex_websocket_pool_key(key);
1060 }
1061 let _ = tx
1062 .send(Err(CodexError {
1063 status: 0,
1064 message: format!("WebSocket stream error: {e}"),
1065 detail: None,
1066 retry_after: None,
1067 origin: CodexErrorOrigin::WebSocket,
1068 }))
1069 .await;
1070 break;
1071 }
1072 }
1073 }
1074
1075 if let Some(tc) = traffic.as_deref() {
1076 write_websocket_response_capture(tc, status, started_at.elapsed(), &sse_body);
1077 }
1078 reusable
1079}
1080
1081fn headers_to_json(headers: &HeaderMap) -> serde_json::Value {
1082 let mut out = serde_json::Map::new();
1083 for (key, value) in headers.iter() {
1084 out.insert(
1085 key.to_string(),
1086 serde_json::Value::String(value.to_str().unwrap_or("").to_string()),
1087 );
1088 }
1089 serde_json::Value::Object(out)
1090}
1091
1092fn summarize_json_request_size(body: &serde_json::Value, body_json: &str) -> serde_json::Value {
1093 serde_json::json!({
1094 "bytes": body_json.len(),
1095 "inputCount": body
1096 .get("input")
1097 .and_then(|v| v.as_array())
1098 .map(|items| items.len()),
1099 "toolCount": body
1100 .get("tools")
1101 .and_then(|v| v.as_array())
1102 .map(|items| items.len()),
1103 })
1104}
1105
1106#[cfg(test)]
1111mod tests {
1112 use super::*;
1113
1114 #[test]
1115 fn event_error_status_requires_error_event_and_checks_numeric_fallbacks() {
1116 assert_eq!(
1117 event_error_status(&serde_json::json!({
1118 "type": "response.failed",
1119 "status": "failed",
1120 "status_code": 401
1121 })),
1122 Some(401)
1123 );
1124 assert_eq!(
1125 event_error_status(&serde_json::json!({
1126 "type": "response.completed",
1127 "status_code": 401
1128 })),
1129 None
1130 );
1131 assert_eq!(
1132 event_error_status(&serde_json::json!({
1133 "type": "error",
1134 "error": {"status": 401}
1135 })),
1136 Some(401)
1137 );
1138 }
1139
1140 #[test]
1141 fn websocket_url_conversion() {
1142 assert_eq!(
1143 to_websocket_url("https://example.test/codex").unwrap(),
1144 "wss://example.test/codex"
1145 );
1146 assert_eq!(
1147 to_websocket_url("http://example.test/codex").unwrap(),
1148 "ws://example.test/codex"
1149 );
1150 assert_eq!(
1151 to_websocket_url("wss://example.test/codex").unwrap(),
1152 "wss://example.test/codex"
1153 );
1154 assert!(to_websocket_url("ftp://example.test/codex").is_err());
1155 }
1156
1157 #[test]
1158 fn websocket_host_header_preserves_explicit_port() {
1159 assert_eq!(
1160 websocket_host_header("wss://chatgpt.com/backend-api/codex/responses"),
1161 "chatgpt.com"
1162 );
1163 assert_eq!(
1164 websocket_host_header("ws://127.0.0.1:4141/backend-api/codex/responses"),
1165 "127.0.0.1:4141"
1166 );
1167 assert_eq!(websocket_host_header("ws://[::1]:4141/path"), "[::1]:4141");
1168 }
1169
1170 #[test]
1171 fn websocket_headers_rewrite_beta() {
1172 let mut headers = http::HeaderMap::new();
1173 headers.insert("openai-beta", "responses=experimental".parse().unwrap());
1174 headers.insert("content-length", "10".parse().unwrap());
1175 headers.insert("authorization", "Bearer tok".parse().unwrap());
1176 let ws = codex_websocket_headers(&headers);
1177 assert_eq!(ws.get("openai-beta").unwrap(), WEBSOCKET_PROTOCOL_HEADER);
1178 assert!(!ws.contains_key("content-length"));
1179 assert_eq!(ws.get("authorization").unwrap(), "Bearer tok");
1180 }
1181
1182 #[test]
1183 fn websocket_headers_strips_accept() {
1184 let mut headers = http::HeaderMap::new();
1185 headers.insert(http::header::ACCEPT, "text/event-stream".parse().unwrap());
1186 let ws = codex_websocket_headers(&headers);
1187 assert!(!ws.contains_key(http::header::ACCEPT.as_str()));
1188 }
1189
1190 #[test]
1191 fn websocket_headers_adds_sec_key() {
1192 let headers = http::HeaderMap::new();
1193 let ws = codex_websocket_headers(&headers);
1194 assert!(ws.contains_key("sec-websocket-key"));
1195 }
1196
1197 #[test]
1198 fn encode_sse_single_line() {
1199 let result = encode_sse(r#"{"type":"test","data":"hello"}"#);
1200 let expected = b"data: {\"type\":\"test\",\"data\":\"hello\"}\n\n";
1201 assert_eq!(result, expected);
1202 }
1203
1204 #[test]
1205 fn encode_sse_multi_line() {
1206 let result = encode_sse("line1\nline2");
1207 assert_eq!(
1208 String::from_utf8(result).unwrap(),
1209 "data: line1\ndata: line2\n\n"
1210 );
1211 }
1212
1213 #[test]
1214 fn is_terminal_event_detection() {
1215 let completed = serde_json::json!({"type": "response.completed"});
1216 assert!(is_terminal_event(&completed));
1217
1218 let delta = serde_json::json!({"type": "response.output_text.delta"});
1219 assert!(!is_terminal_event(&delta));
1220
1221 let error = serde_json::json!({"type": "error", "error": {"message": "fail"}});
1222 assert!(is_terminal_event(&error));
1223 }
1224
1225 #[test]
1226 fn is_response_event_detection() {
1227 let rate_limits = serde_json::json!({"type": "codex.rate_limits"});
1228 assert!(!is_response_event(&rate_limits));
1229
1230 let output = serde_json::json!({"type": "response.output_text.delta"});
1231 assert!(is_response_event(&output));
1232
1233 let error = serde_json::json!({"type": "error", "error": {"message": "fail"}});
1234 assert!(is_response_event(&error));
1235 }
1236
1237 #[test]
1238 fn is_previous_response_missing_detection() {
1239 let by_code = serde_json::json!({
1240 "type": "error",
1241 "error": {"code": "previous_response_not_found", "message": "not found"}
1242 });
1243 assert!(is_previous_response_missing(&by_code));
1244
1245 let by_msg = serde_json::json!({
1246 "type": "error",
1247 "error": {"message": "The previous response was not found"}
1248 });
1249 assert!(is_previous_response_missing(&by_msg));
1250
1251 let unrelated = serde_json::json!({"type": "error", "error": {"message": "rate limited"}});
1252 assert!(!is_previous_response_missing(&unrelated));
1253 }
1254
1255 #[test]
1256 fn pool_invalidation() {
1257 clear_codex_websocket_pool_for_tests();
1258 {
1261 let mut guard = WS_POOL.lock().unwrap();
1262 guard.insert(
1263 "test-session".to_string(),
1264 Arc::new(PoolEntry {
1265 ws: Arc::new(AsyncMutex::new(create_dummy_stream())),
1266 created_at: now_ms(),
1267 }),
1268 );
1269 }
1270 assert!(WS_POOL.lock().unwrap().contains_key("test-session"));
1271
1272 invalidate_codex_websocket_pool_key("test-session");
1273 assert!(!WS_POOL.lock().unwrap().contains_key("test-session"));
1274 }
1275
1276 #[tokio::test]
1277 async fn websocket_connect_401_is_pre_request_handshake_error() {
1278 use tokio::io::{AsyncReadExt, AsyncWriteExt};
1279 use tokio::net::TcpListener;
1280
1281 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1282 let addr = listener.local_addr().unwrap();
1283 tokio::spawn(async move {
1284 let (mut socket, _) = listener.accept().await.unwrap();
1285 let mut buf = [0_u8; 2048];
1286 let _ = socket.read(&mut buf).await;
1287 socket
1288 .write_all(b"HTTP/1.1 401 Unauthorized\r\nContent-Length: 13\r\n\r\npolicy denied")
1289 .await
1290 .unwrap();
1291 });
1292
1293 let err = match connect_with_timeout(
1294 &format!("ws://{addr}/backend-api/codex/responses"),
1295 &HeaderMap::new(),
1296 1_000,
1297 )
1298 .await
1299 {
1300 Ok(_) => panic!("expected unauthorized websocket handshake to fail"),
1301 Err(err) => err,
1302 };
1303
1304 assert_eq!(err.status, 401);
1305 assert_eq!(err.detail.as_deref(), Some("policy denied"));
1306 assert_eq!(err.origin, CodexErrorOrigin::WebSocketHandshake);
1307 }
1308
1309 #[tokio::test]
1310 async fn websocket_connect_502_preserves_retry_metadata() {
1311 use tokio::io::{AsyncReadExt, AsyncWriteExt};
1312 use tokio::net::TcpListener;
1313
1314 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1315 let addr = listener.local_addr().unwrap();
1316 tokio::spawn(async move {
1317 let (mut socket, _) = listener.accept().await.unwrap();
1318 let mut buf = [0_u8; 2048];
1319 let _ = socket.read(&mut buf).await;
1320 socket
1321 .write_all(
1322 b"HTTP/1.1 502 Bad Gateway\r\nRetry-After: 3\r\nContent-Length: 0\r\n\r\n",
1323 )
1324 .await
1325 .unwrap();
1326 });
1327
1328 let err = match connect_with_timeout(
1329 &format!("ws://{addr}/backend-api/codex/responses"),
1330 &HeaderMap::new(),
1331 1_000,
1332 )
1333 .await
1334 {
1335 Ok(_) => panic!("expected websocket handshake to fail"),
1336 Err(err) => err,
1337 };
1338
1339 assert_eq!(err.status, 502);
1340 assert_eq!(err.detail, None);
1341 assert_eq!(err.retry_after.as_deref(), Some("3"));
1342 assert_eq!(err.origin, CodexErrorOrigin::WebSocketHandshake);
1343 }
1344
1345 #[tokio::test]
1346 async fn binary_frame_invalidates_pool_key() {
1347 clear_codex_websocket_pool_for_tests();
1348 let pooled_stream = create_dummy_stream_async().await;
1349 {
1350 let mut guard = WS_POOL.lock().unwrap();
1351 guard.insert(
1352 "binary-session".to_string(),
1353 Arc::new(PoolEntry {
1354 ws: Arc::new(AsyncMutex::new(pooled_stream)),
1355 created_at: now_ms(),
1356 }),
1357 );
1358 }
1359
1360 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1361 let addr = listener.local_addr().unwrap();
1362 tokio::spawn(async move {
1363 let (stream, _) = listener.accept().await.unwrap();
1364 let mut ws = tokio_tungstenite::accept_async(stream).await.unwrap();
1365 ws.send(Message::Binary(vec![1, 2, 3])).await.unwrap();
1366 });
1367
1368 let (mut ws, _) = tokio_tungstenite::connect_async(format!("ws://{addr}/"))
1369 .await
1370 .unwrap();
1371 let err = match collect_ws_events(&mut ws, 1_000, Some("binary-session"), None).await {
1372 Ok(_) => panic!("expected binary frame to fail"),
1373 Err(err) => err,
1374 };
1375
1376 assert!(err.message.contains("binary frames"));
1377 assert!(!WS_POOL.lock().unwrap().contains_key("binary-session"));
1378 }
1379
1380 #[tokio::test]
1381 async fn response_start_timeout_ignores_rate_limits_and_pings() {
1382 clear_codex_websocket_pool_for_tests();
1383 let pooled_stream = create_dummy_stream_async().await;
1384 {
1385 let mut guard = WS_POOL.lock().unwrap();
1386 guard.insert(
1387 "start-timeout-session".to_string(),
1388 Arc::new(PoolEntry {
1389 ws: Arc::new(AsyncMutex::new(pooled_stream)),
1390 created_at: now_ms(),
1391 }),
1392 );
1393 }
1394
1395 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1396 let addr = listener.local_addr().unwrap();
1397 tokio::spawn(async move {
1398 let (stream, _) = listener.accept().await.unwrap();
1399 let mut ws = tokio_tungstenite::accept_async(stream).await.unwrap();
1400 ws.send(Message::Text(
1401 r#"{"type":"codex.rate_limits","rate_limits":{"allowed":true}}"#.into(),
1402 ))
1403 .await
1404 .unwrap();
1405 loop {
1406 if ws.send(Message::Ping(Vec::new())).await.is_err() {
1407 break;
1408 }
1409 tokio::time::sleep(Duration::from_millis(10)).await;
1410 }
1411 });
1412
1413 let (mut ws, _) = tokio_tungstenite::connect_async(format!("ws://{addr}/"))
1414 .await
1415 .unwrap();
1416 let err = match collect_ws_events(&mut ws, 50, Some("start-timeout-session"), None).await {
1417 Ok(_) => panic!("expected response start timeout"),
1418 Err(err) => err,
1419 };
1420
1421 assert_eq!(
1422 err.detail.as_deref(),
1423 Some(WEBSOCKET_RESPONSE_START_TIMEOUT_DETAIL)
1424 );
1425 assert!(
1426 !WS_POOL
1427 .lock()
1428 .unwrap()
1429 .contains_key("start-timeout-session")
1430 );
1431 }
1432
1433 #[tokio::test]
1434 async fn response_idle_timeout_ignores_pings_after_response_event() {
1435 clear_codex_websocket_pool_for_tests();
1436 let pooled_stream = create_dummy_stream_async().await;
1437 {
1438 let mut guard = WS_POOL.lock().unwrap();
1439 guard.insert(
1440 "response-idle-session".to_string(),
1441 Arc::new(PoolEntry {
1442 ws: Arc::new(AsyncMutex::new(pooled_stream)),
1443 created_at: now_ms(),
1444 }),
1445 );
1446 }
1447
1448 let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
1449 let addr = listener.local_addr().unwrap();
1450 tokio::spawn(async move {
1451 let (stream, _) = listener.accept().await.unwrap();
1452 let mut ws = tokio_tungstenite::accept_async(stream).await.unwrap();
1453 ws.send(Message::Text(
1454 r#"{"type":"response.output_item.added","output_index":0,"item":{"type":"message"}}"#
1455 .into(),
1456 ))
1457 .await
1458 .unwrap();
1459 loop {
1460 if ws.send(Message::Ping(Vec::new())).await.is_err() {
1461 break;
1462 }
1463 tokio::time::sleep(Duration::from_millis(10)).await;
1464 }
1465 });
1466
1467 let (mut ws, _) = tokio_tungstenite::connect_async(format!("ws://{addr}/"))
1468 .await
1469 .unwrap();
1470 let err = match collect_ws_events(&mut ws, 50, Some("response-idle-session"), None).await {
1471 Ok(_) => panic!("expected response idle timeout"),
1472 Err(err) => err,
1473 };
1474
1475 assert!(err.message.contains("idle timeout"));
1476 assert_eq!(err.detail, None);
1477 assert!(
1478 !WS_POOL
1479 .lock()
1480 .unwrap()
1481 .contains_key("response-idle-session")
1482 );
1483 }
1484
1485 async fn create_dummy_stream_async() -> WebSocketStream<MaybeTlsStream<TcpStream>> {
1486 use tokio::net::TcpListener;
1487
1488 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1489 let addr = listener.local_addr().unwrap();
1490 tokio::spawn(async move {
1491 let (socket, _) = listener.accept().await.unwrap();
1492 let _ = tokio_tungstenite::accept_async(socket).await;
1493 futures_util::future::pending::<()>().await;
1494 });
1495 let url = format!("ws://{addr}/");
1496 let (ws, _) = tokio::time::timeout(
1497 Duration::from_millis(1000),
1498 tokio_tungstenite::connect_async(&url),
1499 )
1500 .await
1501 .unwrap()
1502 .unwrap();
1503 ws
1504 }
1505
1506 fn create_dummy_stream() -> WebSocketStream<MaybeTlsStream<TcpStream>> {
1507 use tokio::net::TcpListener;
1510 let rt = tokio::runtime::Runtime::new().unwrap();
1511 rt.block_on(async {
1512 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1513 let addr = listener.local_addr().unwrap();
1514 let _conn = tokio::spawn(async move {
1515 let (socket, _) = listener.accept().await.unwrap();
1516 let _ = tokio_tungstenite::accept_async(socket).await;
1518 futures_util::future::pending::<()>().await;
1520 });
1521 let url = format!("ws://{}/", addr);
1523 let (ws, _) = tokio::time::timeout(
1524 Duration::from_millis(1000),
1525 tokio_tungstenite::connect_async(&url),
1526 )
1527 .await
1528 .unwrap()
1529 .unwrap();
1530 ws
1531 })
1532 }
1533}