Skip to main content

claude_codex/providers/codex/
websocket.rs

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
21// ---------------------------------------------------------------------------
22// Constants
23// ---------------------------------------------------------------------------
24
25pub 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
34// Terminal WebSocket event types that signal the request is done
35const 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// ---------------------------------------------------------------------------
45// Errors
46// ---------------------------------------------------------------------------
47
48#[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
85// ---------------------------------------------------------------------------
86// Pool
87// ---------------------------------------------------------------------------
88
89struct 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    // Evict oldest if at capacity
117    if guard.len() >= MAX_POOL_ENTRIES
118        && let Some(oldest_key) = guard.keys().next().cloned()
119    {
120        guard.remove(&oldest_key);
121    }
122    // Evict expired entries
123    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
128// ---------------------------------------------------------------------------
129// URL conversion
130// ---------------------------------------------------------------------------
131
132pub 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" => { /* already a ws scheme */ }
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
152// ---------------------------------------------------------------------------
153// Header rewriting
154// ---------------------------------------------------------------------------
155
156pub 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        // Skip hop-by-hop headers
161        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    // Rewrite openai-beta for WebSocket protocol
170    ws.insert("openai-beta", WEBSOCKET_PROTOCOL_HEADER.parse().unwrap());
171    // Ensure WebSocket key is present
172    if !ws.contains_key("sec-websocket-key") {
173        ws.insert("sec-websocket-key", generate_key().parse().unwrap());
174    }
175    ws
176}
177
178// ---------------------------------------------------------------------------
179// SSE framing
180// ---------------------------------------------------------------------------
181
182fn 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
193// ---------------------------------------------------------------------------
194// Terminal event detection
195// ---------------------------------------------------------------------------
196
197fn 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    // Case-insensitive message check
222    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// ---------------------------------------------------------------------------
249// Main request function
250// ---------------------------------------------------------------------------
251
252#[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    // Check pool for existing connection
298    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        // Use pooled connection
305        let mut ws_guard = entry.ws.lock().await;
306        // Check if connection is still alive by sending a ping
307        if ws_guard.send(Message::Ping(vec![])).await.is_err() {
308            invalidate_codex_websocket_pool_key(pool_key.unwrap());
309            // Fall through to new connection
310            connect_with_timeout(&ws_url, headers, connect_timeout_ms).await?
311        } else {
312            // Connection is alive, send the request through it
313            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            // Collect events
328            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            // Handle previous response missing
335            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            // Extract status from error events
346            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            // Write traffic metadata
353            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    // New connection path (not pooled or pool miss)
369    let entry = Arc::new(PoolEntry {
370        ws: Arc::new(AsyncMutex::new(ws_stream)),
371        created_at: now_ms(),
372    });
373
374    // Send the request
375    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        // Pool the connection if we have a key and it was successful
406        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        // Write traffic metadata
420        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
628// ---------------------------------------------------------------------------
629// Connection helper
630// ---------------------------------------------------------------------------
631
632async 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    // Build an http::Request with the given headers for the WebSocket upgrade
644    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    // Copy over the codex headers
655    for (key, value) in headers.iter() {
656        let key_str = key.as_str().to_lowercase();
657        // Skip headers already set for WebSocket upgrade
658        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
722// ---------------------------------------------------------------------------
723// Event collection
724// ---------------------------------------------------------------------------
725
726struct 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                // Parse JSON
800                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                        // Write invalid JSON as-is
813                        sse_body.extend_from_slice(&encode_sse(&text));
814                        continue;
815                    }
816                };
817
818                // Convert to SSE bytes
819                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                // Check for terminal events
830                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                // Reject binary frames
844                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                // Respond to ping automatically, continue
857                let _ = ws.send(Message::Pong(data)).await;
858                continue;
859            }
860            Some(Ok(Message::Pong(_))) => {
861                continue;
862            }
863            Some(Ok(Message::Frame(_))) => {
864                // Raw frame passthrough - continue
865                continue;
866            }
867            Some(Ok(Message::Close(_))) => {
868                // Connection closed - invalidate pool
869                if let Some(key) = pool_key {
870                    invalidate_codex_websocket_pool_key(key);
871                }
872                break;
873            }
874            Some(Err(e)) => {
875                // Stream error - invalidate pool
876                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                // Stream ended - invalidate pool
889                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// ---------------------------------------------------------------------------
1107// Tests
1108// ---------------------------------------------------------------------------
1109
1110#[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        // Verify pool operations work through the public API
1259        // We insert an entry directly into the pool, then invalidate it
1260        {
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 a connected TcpStream pair with connect_async which returns
1508        // WebSocketStream<MaybeTlsStream<TcpStream>>
1509        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                // Accept WebSocket handshake
1517                let _ = tokio_tungstenite::accept_async(socket).await;
1518                // Keep alive
1519                futures_util::future::pending::<()>().await;
1520            });
1521            // Use connect_async to get MaybeTlsStream
1522            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}