Skip to main content

claude_codex/providers/codex/
websocket.rs

1use std::collections::HashMap;
2use std::pin::Pin;
3use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
4use std::sync::{Arc, Mutex};
5use std::task::{Context, Poll};
6use std::time::{Duration, Instant};
7
8use futures_util::{SinkExt, StreamExt};
9use http::HeaderMap;
10use hyper_util::client::proxy::matcher::Matcher as ProxyMatcher;
11use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, ReadBuf};
12use tokio::net::TcpStream;
13use tokio::sync::mpsc;
14use tokio::sync::{Mutex as AsyncMutex, OwnedMutexGuard};
15use tokio_tungstenite::{
16    WebSocketStream,
17    tungstenite::{
18        Message,
19        client::IntoClientRequest,
20        handshake::{client::generate_key, derive_accept_key},
21        protocol::Role,
22    },
23};
24
25use crate::logging::create_logger;
26use crate::provider::RequestContext;
27use crate::request_identity::ConversationIdentity;
28use crate::retry::sleep as retry_sleep;
29use crate::traffic::TrafficCapture;
30
31use super::client::{
32    ActualTransport, CodexError, CodexErrorOrigin, CodexResponse, OwnerAwareCodexResponse,
33};
34use super::continuation::ContinuationReservation;
35
36// ---------------------------------------------------------------------------
37// Constants
38// ---------------------------------------------------------------------------
39
40pub const WEBSOCKET_PROTOCOL_HEADER: &str = "responses_websockets=2026-02-06";
41pub const WEBSOCKET_CONNECT_TIMEOUT_MS: u64 = 15_000;
42pub const WEBSOCKET_IDLE_TIMEOUT_MS: u64 = 300_000;
43pub const WEBSOCKET_RESPONSE_START_TIMEOUT_DETAIL: &str = "websocket_response_start_timeout";
44pub const WEBSOCKET_MISSING_TERMINAL_DETAIL: &str = "websocket_missing_terminal";
45pub const WEBSOCKET_KEEPALIVE_FAILURE_DETAIL: &str = "websocket_keepalive_failure";
46pub const WEBSOCKET_CONTINUATION_SOCKET_MISSING_DETAIL: &str =
47    "websocket_continuation_socket_missing";
48pub(super) const WEBSOCKET_PROXY_TUNNEL_REJECTED_DETAIL: &str = "websocket_proxy_tunnel_rejected";
49
50const POOL_IDLE_TTL_MS: u64 = 30 * 60 * 1000;
51const MAX_POOL_ENTRIES: usize = 10_000;
52const POOL_CONNECT_CLEANUP_THRESHOLD: usize = 50;
53const POOL_CONNECT_CLEANUP_TARGET: usize = 40;
54const MAX_CONNECT_RESPONSE_HEADER_BYTES: usize = 8 * 1024;
55const WEBSOCKET_CONNECT_START_SPACING: Duration = Duration::from_secs(1);
56const WEBSOCKET_CONNECT_FORBIDDEN_COOLDOWN: Duration = Duration::from_secs(3);
57const WEBSOCKET_KEEPALIVE_INTERVAL: Duration = Duration::from_secs(30);
58const WEBSOCKET_KEEPALIVE_SEND_TIMEOUT: Duration = Duration::from_secs(10);
59
60// Terminal WebSocket event types that signal the request is done
61const TERMINAL_EVENTS: &[&str] = &[
62    "response.completed",
63    "response.incomplete",
64    "response.failed",
65    "error",
66];
67
68pub type CodexWebSocketEventReceiver =
69    tokio::sync::mpsc::Receiver<Result<serde_json::Value, CodexError>>;
70
71pub(crate) struct CodexWebSocketEventStream {
72    receiver: CodexWebSocketEventReceiver,
73    socket_id: Arc<AtomicU64>,
74    full_context_retry: Arc<AtomicBool>,
75    provider_retry_handoff: Arc<AtomicBool>,
76}
77
78#[derive(Clone)]
79pub(crate) struct CodexWebSocketSocketIdPublisher {
80    socket_id: Arc<AtomicU64>,
81    full_context_retry: Arc<AtomicBool>,
82    provider_retry_handoff: Arc<AtomicBool>,
83}
84
85impl CodexWebSocketEventStream {
86    pub(crate) fn pending(
87        receiver: CodexWebSocketEventReceiver,
88    ) -> (Self, CodexWebSocketSocketIdPublisher) {
89        let socket_id = Arc::new(AtomicU64::new(0));
90        let full_context_retry = Arc::new(AtomicBool::new(false));
91        let provider_retry_handoff = Arc::new(AtomicBool::new(false));
92        (
93            Self {
94                receiver,
95                socket_id: socket_id.clone(),
96                full_context_retry: full_context_retry.clone(),
97                provider_retry_handoff: provider_retry_handoff.clone(),
98            },
99            CodexWebSocketSocketIdPublisher {
100                socket_id,
101                full_context_retry,
102                provider_retry_handoff,
103            },
104        )
105    }
106
107    pub(crate) async fn recv(&mut self) -> Option<Result<serde_json::Value, CodexError>> {
108        self.receiver.recv().await
109    }
110
111    pub(crate) fn socket_id(&self) -> Option<u64> {
112        match self.socket_id.load(Ordering::Acquire) {
113            0 => None,
114            socket_id => Some(socket_id),
115        }
116    }
117
118    pub(crate) fn used_full_context_retry(&self) -> bool {
119        self.full_context_retry.load(Ordering::Acquire)
120    }
121
122    pub(crate) fn mark_provider_retry_handoff(&self) {
123        self.provider_retry_handoff.store(true, Ordering::Release);
124    }
125
126    pub(crate) fn into_receiver(self) -> CodexWebSocketEventReceiver {
127        self.receiver
128    }
129
130    pub(crate) fn replace_receiver(
131        &mut self,
132        receiver: CodexWebSocketEventReceiver,
133    ) -> CodexWebSocketEventReceiver {
134        std::mem::replace(&mut self.receiver, receiver)
135    }
136}
137
138impl CodexWebSocketSocketIdPublisher {
139    pub(super) fn publish(&self, socket_id: Option<u64>) {
140        self.socket_id
141            .store(socket_id.unwrap_or(0), Ordering::Release);
142    }
143
144    pub(super) fn mark_full_context_retry(&self) {
145        self.full_context_retry.store(true, Ordering::Release);
146    }
147
148    pub(super) fn is_provider_retry_handoff(&self) -> bool {
149        self.provider_retry_handoff.load(Ordering::Acquire)
150    }
151}
152
153trait WebSocketIo: AsyncRead + AsyncWrite + Unpin + Send {}
154
155impl<T> WebSocketIo for T where T: AsyncRead + AsyncWrite + Unpin + Send {}
156
157type BoxedWebSocketIo = Box<dyn WebSocketIo>;
158type CodexWebSocketStream = WebSocketStream<BoxedWebSocketIo>;
159
160pub(super) struct WebSocketProxyConfig {
161    matcher: ProxyMatcher,
162    tls_config: Arc<rustls::ClientConfig>,
163}
164
165#[derive(Clone)]
166struct WebSocketProxyRoute {
167    uri: http::Uri,
168    basic_auth: Option<http::HeaderValue>,
169}
170
171impl WebSocketProxyConfig {
172    pub(super) fn new(
173        http_proxy: Option<&str>,
174        https_proxy: Option<&str>,
175        all_proxy: Option<&str>,
176        no_proxy: Option<&str>,
177    ) -> Self {
178        let mut builder = ProxyMatcher::builder();
179        if let Some(proxy) = all_proxy {
180            builder = builder.all(proxy.to_string());
181        }
182        if let Some(proxy) = http_proxy {
183            builder = builder.http(proxy.to_string());
184        }
185        if let Some(proxy) = https_proxy {
186            builder = builder.https(proxy.to_string());
187        }
188        if let Some(no_proxy) = no_proxy {
189            builder = builder.no(no_proxy.to_string());
190        }
191        Self {
192            matcher: builder.build(),
193            tls_config: websocket_tls_config(),
194        }
195    }
196
197    #[cfg(test)]
198    pub(super) fn direct() -> Self {
199        Self::new(None, None, None, None)
200    }
201
202    pub(super) fn uses_proxy_for(&self, websocket_url: &str) -> bool {
203        let Ok(http_url) = to_http_upgrade_url(websocket_url) else {
204            return true;
205        };
206        let Ok(destination) = http_url.parse::<http::Uri>() else {
207            return true;
208        };
209        self.matcher.intercept(&destination).is_some()
210    }
211
212    fn http_connect_route(
213        &self,
214        websocket_url: &str,
215    ) -> Result<Option<WebSocketProxyRoute>, CodexError> {
216        let http_url = to_http_upgrade_url(websocket_url).map_err(|error| CodexError {
217            status: 0,
218            message: error.message,
219            detail: None,
220            retry_after: None,
221            origin: CodexErrorOrigin::WebSocketHandshake,
222        })?;
223        if !http_url.starts_with("https://") {
224            return Ok(None);
225        }
226        let destination = http_url.parse::<http::Uri>().map_err(|_| {
227            websocket_protocol_error("WebSocket destination URL could not be routed")
228        })?;
229        let Some(proxy) = self.matcher.intercept(&destination) else {
230            return Ok(None);
231        };
232        if !matches!(proxy.uri().scheme_str(), Some("http" | "https")) {
233            return Ok(None);
234        }
235        Ok(Some(WebSocketProxyRoute {
236            uri: proxy.uri().clone(),
237            basic_auth: proxy.basic_auth().cloned(),
238        }))
239    }
240}
241
242static WEBSOCKET_TLS_CONFIG: once_cell::sync::Lazy<Arc<rustls::ClientConfig>> =
243    once_cell::sync::Lazy::new(|| {
244        let mut roots = rustls::RootCertStore::empty();
245        let native = rustls_native_certs::load_native_certs();
246        let load_error_count = native.errors.len();
247        let (_, parse_error_count) = roots.add_parsable_certificates(native.certs);
248        if load_error_count > 0 || parse_error_count > 0 {
249            let mut fields = serde_json::Map::new();
250            fields.insert("loadErrorCount".into(), serde_json::json!(load_error_count));
251            fields.insert(
252                "parseErrorCount".into(),
253                serde_json::json!(parse_error_count),
254            );
255            create_logger("codex").warn("native_certificate_load_errors", Some(fields));
256        }
257        roots.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned());
258        Arc::new(
259            rustls::ClientConfig::builder()
260                .with_root_certificates(roots)
261                .with_no_client_auth(),
262        )
263    });
264
265pub(super) fn websocket_tls_config() -> Arc<rustls::ClientConfig> {
266    WEBSOCKET_TLS_CONFIG.clone()
267}
268
269// ---------------------------------------------------------------------------
270// Errors
271// ---------------------------------------------------------------------------
272
273#[derive(Debug, Clone)]
274pub struct CodexWebSocketError {
275    pub message: String,
276    pub status: Option<u16>,
277    pub code: Option<String>,
278    pub retry_after: Option<String>,
279    pub request_sent: bool,
280}
281
282impl CodexWebSocketError {
283    pub fn new(message: String) -> Self {
284        Self {
285            message,
286            status: None,
287            code: None,
288            retry_after: None,
289            request_sent: false,
290        }
291    }
292
293    pub fn with_status(mut self, status: u16) -> Self {
294        self.status = Some(status);
295        self
296    }
297
298    pub fn with_code(mut self, code: String) -> Self {
299        self.code = Some(code);
300        self
301    }
302}
303
304impl std::fmt::Display for CodexWebSocketError {
305    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
306        write!(f, "Codex WebSocket error: {}", self.message)
307    }
308}
309
310// ---------------------------------------------------------------------------
311// Pool
312// ---------------------------------------------------------------------------
313
314struct PoolEntry {
315    ws: Arc<AsyncMutex<CodexWebSocketStream>>,
316    socket_id: u64,
317    created_at: u64,
318    last_activity: AtomicU64,
319}
320
321impl PoolEntry {
322    fn new(ws: CodexWebSocketStream) -> Self {
323        Self {
324            ws: Arc::new(AsyncMutex::new(ws)),
325            socket_id: next_monotonic_nonzero(&NEXT_SOCKET_ID, "WebSocket ID"),
326            created_at: now_ms(),
327            last_activity: AtomicU64::new(next_pool_activity()),
328        }
329    }
330
331    fn touch(&self) {
332        self.last_activity
333            .fetch_max(next_pool_activity(), Ordering::Relaxed);
334    }
335}
336
337static NEXT_SOCKET_ID: AtomicU64 = AtomicU64::new(0);
338static POOLED_VALIDATION_SEQUENCE: AtomicU64 = AtomicU64::new(0);
339static POOL_ACTIVITY_SEQUENCE: AtomicU64 = AtomicU64::new(1);
340static WS_POOL: once_cell::sync::Lazy<Mutex<HashMap<ConversationIdentity, Arc<PoolEntry>>>> =
341    once_cell::sync::Lazy::new(|| Mutex::new(HashMap::new()));
342#[cfg(test)]
343static WS_POOL_TEST_LOCK: AsyncMutex<()> = AsyncMutex::const_new(());
344static WS_CONNECT_GATE: once_cell::sync::Lazy<WebSocketConnectGate> =
345    once_cell::sync::Lazy::new(|| WebSocketConnectGate::new(WEBSOCKET_CONNECT_START_SPACING));
346
347fn next_monotonic_nonzero(sequence: &AtomicU64, label: &str) -> u64 {
348    let previous = sequence
349        .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |value| {
350            value.checked_add(1)
351        })
352        .unwrap_or_else(|_| panic!("{label} sequence exhausted"));
353    previous + 1
354}
355
356fn next_pool_activity() -> u64 {
357    POOL_ACTIVITY_SEQUENCE.fetch_add(1, Ordering::Relaxed)
358}
359
360fn now_ms() -> u64 {
361    std::time::SystemTime::now()
362        .duration_since(std::time::UNIX_EPOCH)
363        .unwrap_or_default()
364        .as_millis() as u64
365}
366
367pub fn clear_codex_websocket_pool_for_tests() {
368    let mut guard = WS_POOL.lock().unwrap();
369    guard.clear();
370}
371
372#[cfg(test)]
373pub(crate) async fn lock_codex_websocket_pool_for_tests() -> tokio::sync::MutexGuard<'static, ()> {
374    WS_POOL_TEST_LOCK.lock().await
375}
376
377#[cfg(test)]
378pub(crate) fn pooled_socket_id_for_tests(owner: &ConversationIdentity) -> Option<u64> {
379    WS_POOL
380        .lock()
381        .unwrap()
382        .get(owner)
383        .map(|entry| entry.socket_id)
384}
385
386pub fn invalidate_codex_websocket_pool_owner(owner: &ConversationIdentity) {
387    let mut guard = WS_POOL.lock().unwrap();
388    guard.remove(owner);
389}
390
391#[deprecated(note = "use typed conversation ownership internally")]
392pub fn invalidate_codex_websocket_pool_key(session_id: &str) {
393    let mut guard = WS_POOL.lock().unwrap();
394    guard.retain(|owner, _| match owner {
395        ConversationIdentity::Main(owner_session_id)
396        | ConversationIdentity::Agent(owner_session_id, _) => owner_session_id != session_id,
397    });
398}
399
400pub(crate) fn invalidate_codex_websocket_pool_turn_for_owner(
401    owner: &ConversationIdentity,
402    turn_id: Option<u64>,
403) {
404    let reservation = ContinuationReservation::for_owner_turn(Some(owner), turn_id);
405    super::continuation::with_current_turn_for_owner(&reservation, || {
406        invalidate_codex_websocket_pool_owner(owner)
407    });
408}
409
410#[deprecated(note = "use typed conversation ownership internally")]
411pub fn invalidate_codex_websocket_pool_turn(session_id: &str, turn_id: Option<u64>) {
412    let owner = ConversationIdentity::Main(session_id.to_owned());
413    invalidate_codex_websocket_pool_turn_for_owner(&owner, turn_id);
414}
415
416fn invalidate_pool_entry(owner: &ConversationIdentity, entry: &Arc<PoolEntry>) {
417    let mut guard = WS_POOL.lock().unwrap();
418    if guard
419        .get(owner)
420        .is_some_and(|pooled| Arc::ptr_eq(pooled, entry))
421    {
422        guard.remove(owner);
423    }
424}
425
426fn invalidate_pool_owner(owner: Option<&ConversationIdentity>, entry: Option<&Arc<PoolEntry>>) {
427    let Some(owner) = owner else {
428        return;
429    };
430    match entry {
431        Some(entry) => invalidate_pool_entry(owner, entry),
432        None => invalidate_codex_websocket_pool_owner(owner),
433    }
434}
435
436fn reservation_pool_owner(
437    reservation: Option<&ContinuationReservation>,
438) -> Option<&ConversationIdentity> {
439    let reservation = reservation?;
440    if reservation.candidate().disabled_reason.as_deref() == Some("disabled") {
441        return None;
442    }
443    reservation.owner()
444}
445
446fn pool_take_for_turn(reservation: &ContinuationReservation) -> Option<Arc<PoolEntry>> {
447    let owner = reservation.owner()?;
448    super::continuation::if_current_turn_for_owner(reservation, || {
449        WS_POOL.lock().ok()?.remove(owner)
450    })
451    .flatten()
452}
453
454fn take_pool_entry_for_request(
455    reservation: Option<&ContinuationReservation>,
456) -> Result<Option<Arc<PoolEntry>>, CodexError> {
457    let candidate = reservation.map(ContinuationReservation::candidate);
458    let requires_origin = candidate
459        .and_then(|candidate| candidate.previous_response_id.as_deref())
460        .is_some();
461    let expected_socket_id = reservation.and_then(ContinuationReservation::origin_socket_id);
462    let pool_owner = reservation_pool_owner(reservation);
463    let pooled = reservation
464        .filter(|_| pool_owner.is_some())
465        .and_then(pool_take_for_turn);
466
467    if requires_origin
468        && (pool_owner.is_none()
469            || expected_socket_id.is_none()
470            || pooled
471                .as_ref()
472                .is_none_or(|entry| Some(entry.socket_id) != expected_socket_id))
473    {
474        if let (Some(owner), Some(entry)) = (pool_owner, pooled.as_ref()) {
475            pool_insert_if_vacant_or_same(owner.clone(), entry.clone());
476        }
477        return Err(continuation_socket_missing_error());
478    }
479
480    Ok(pooled)
481}
482
483fn pool_insert_for_turn(reservation: &ContinuationReservation, entry: Arc<PoolEntry>) -> bool {
484    let Some(owner) = reservation_pool_owner(Some(reservation)).cloned() else {
485        return false;
486    };
487    super::continuation::if_current_turn_for_owner(reservation, || {
488        pool_insert_if_vacant_or_same(owner, entry)
489    })
490    .unwrap_or(false)
491}
492
493fn pool_remove_entry(owner: &ConversationIdentity, entry: &Arc<PoolEntry>) {
494    invalidate_pool_entry(owner, entry);
495}
496
497pub(super) fn invalidate_codex_websocket_pool_socket(
498    reservation: &ContinuationReservation,
499    socket_id: Option<u64>,
500) {
501    let Some(owner) = reservation_pool_owner(Some(reservation)) else {
502        return;
503    };
504    let Some(socket_id) = socket_id else {
505        return;
506    };
507    let entry = WS_POOL.lock().ok().and_then(|pool| {
508        pool.get(owner)
509            .filter(|entry| entry.socket_id == socket_id)
510            .cloned()
511    });
512    if let Some(entry) = entry {
513        pool_remove_entry(owner, &entry);
514    }
515}
516
517fn pool_insert_if_vacant_or_same(owner: ConversationIdentity, entry: Arc<PoolEntry>) -> bool {
518    entry.touch();
519    let mut guard = WS_POOL.lock().unwrap();
520    if let Some(existing) = guard.get(&owner) {
521        return Arc::ptr_eq(existing, &entry);
522    }
523    if guard.len() >= MAX_POOL_ENTRIES
524        && let Some(oldest_owner) = guard.keys().next().cloned()
525    {
526        guard.remove(&oldest_owner);
527    }
528    let now = now_ms();
529    guard.retain(|_, pooled| now.saturating_sub(pooled.created_at) < POOL_IDLE_TTL_MS);
530    guard.insert(owner, entry);
531    true
532}
533
534#[cfg(test)]
535fn pool_insert(owner: ConversationIdentity, entry: Arc<PoolEntry>) {
536    entry.touch();
537    let mut guard = WS_POOL.lock().unwrap();
538    // Evict oldest if at capacity
539    if guard.len() >= MAX_POOL_ENTRIES
540        && let Some(oldest_owner) = guard.keys().next().cloned()
541    {
542        guard.remove(&oldest_owner);
543    }
544    // Evict expired entries
545    let now = now_ms();
546    guard.retain(|_, entry| now.saturating_sub(entry.created_at) < POOL_IDLE_TTL_MS);
547    guard.insert(owner, entry);
548}
549
550fn cleanup_pool_before_connect() {
551    let removed = {
552        let mut guard = WS_POOL.lock().unwrap();
553        if guard.len() <= POOL_CONNECT_CLEANUP_THRESHOLD {
554            return;
555        }
556
557        let remove_count = guard.len() - POOL_CONNECT_CLEANUP_TARGET;
558        let mut candidates: Vec<_> = guard
559            .iter()
560            .filter(|(_, entry)| Arc::strong_count(entry) == 1)
561            .map(|(owner, entry)| (owner.clone(), entry.last_activity.load(Ordering::Relaxed)))
562            .collect();
563        candidates.sort_unstable_by_key(|(_, activity)| *activity);
564        candidates
565            .into_iter()
566            .take(remove_count)
567            .filter_map(|(owner, _)| guard.remove(&owner))
568            .collect::<Vec<_>>()
569    };
570    drop(removed);
571}
572
573struct WebSocketConnectGate {
574    last_start: AsyncMutex<Option<tokio::time::Instant>>,
575    start_spacing: Duration,
576}
577
578impl WebSocketConnectGate {
579    fn new(start_spacing: Duration) -> Self {
580        Self {
581            last_start: AsyncMutex::new(None),
582            start_spacing,
583        }
584    }
585
586    async fn wait_to_start(&self, before_start: impl std::future::Future<Output = ()>) {
587        let mut last_start = self.last_start.lock().await;
588        if let Some(previous) = *last_start {
589            tokio::time::sleep_until(previous + self.start_spacing).await;
590        }
591        before_start.await;
592        *last_start = Some(tokio::time::Instant::now());
593    }
594}
595
596// ---------------------------------------------------------------------------
597// URL conversion
598// ---------------------------------------------------------------------------
599
600pub fn to_websocket_url(url: &str) -> Result<String, CodexWebSocketError> {
601    let mut parsed = url::Url::parse(url)
602        .map_err(|e| CodexWebSocketError::new(format!("Failed to parse URL: {e}")))?;
603    match parsed.scheme() {
604        "http" => parsed.set_scheme("ws").map_err(|_| {
605            CodexWebSocketError::new("Unsupported Codex WebSocket URL scheme".to_string())
606        })?,
607        "https" => parsed.set_scheme("wss").map_err(|_| {
608            CodexWebSocketError::new("Unsupported Codex WebSocket URL scheme".to_string())
609        })?,
610        "ws" | "wss" => { /* already a ws scheme */ }
611        other => {
612            return Err(CodexWebSocketError::new(format!(
613                "Unsupported Codex WebSocket URL scheme: {other}"
614            )));
615        }
616    }
617    Ok(parsed.to_string())
618}
619
620fn to_http_upgrade_url(url: &str) -> Result<String, CodexWebSocketError> {
621    let mut parsed = url::Url::parse(url)
622        .map_err(|e| CodexWebSocketError::new(format!("Failed to parse URL: {e}")))?;
623    match parsed.scheme() {
624        "ws" => parsed.set_scheme("http").map_err(|_| {
625            CodexWebSocketError::new("Unsupported Codex WebSocket URL scheme".to_string())
626        })?,
627        "wss" => parsed.set_scheme("https").map_err(|_| {
628            CodexWebSocketError::new("Unsupported Codex WebSocket URL scheme".to_string())
629        })?,
630        other => {
631            return Err(CodexWebSocketError::new(format!(
632                "Unsupported Codex WebSocket URL scheme: {other}"
633            )));
634        }
635    }
636    Ok(parsed.to_string())
637}
638
639// ---------------------------------------------------------------------------
640// Header rewriting
641// ---------------------------------------------------------------------------
642
643pub fn codex_websocket_headers(http_headers: &HeaderMap) -> HeaderMap {
644    let mut ws = HeaderMap::new();
645    for (key, value) in http_headers.iter() {
646        let key_str = key.as_str().to_lowercase();
647        // Skip hop-by-hop headers
648        if matches!(
649            key_str.as_str(),
650            "content-length"
651                | "content-type"
652                | "accept"
653                | "connection"
654                | "upgrade"
655                | "proxy-authorization"
656        ) {
657            continue;
658        }
659        ws.insert(key.clone(), value.clone());
660    }
661    // Rewrite openai-beta for WebSocket protocol
662    ws.insert("openai-beta", WEBSOCKET_PROTOCOL_HEADER.parse().unwrap());
663    // Ensure WebSocket key is present
664    if !ws.contains_key("sec-websocket-key") {
665        ws.insert("sec-websocket-key", generate_key().parse().unwrap());
666    }
667    ws
668}
669
670// ---------------------------------------------------------------------------
671// SSE framing
672// ---------------------------------------------------------------------------
673
674fn encode_sse(text: &str) -> Vec<u8> {
675    let mut out = String::new();
676    for line in text.lines() {
677        out.push_str("data: ");
678        out.push_str(line);
679        out.push('\n');
680    }
681    out.push('\n');
682    out.into_bytes()
683}
684
685// ---------------------------------------------------------------------------
686// Terminal event detection
687// ---------------------------------------------------------------------------
688
689pub(super) fn is_terminal_event(payload: &serde_json::Value) -> bool {
690    match payload.get("type").and_then(|v| v.as_str()) {
691        Some(t) => TERMINAL_EVENTS.contains(&t),
692        None => false,
693    }
694}
695
696fn is_response_event(payload: &serde_json::Value) -> bool {
697    match payload.get("type").and_then(|v| v.as_str()) {
698        Some("error") => true,
699        Some(t) => t.starts_with("response."),
700        None => false,
701    }
702}
703
704fn is_previous_response_missing(payload: &serde_json::Value) -> bool {
705    let error = super::events::event_error(payload);
706    if error
707        .and_then(|error| error.get("code"))
708        .and_then(|value| value.as_str())
709        == Some("previous_response_not_found")
710    {
711        return true;
712    }
713    // Case-insensitive message check
714    if let Some(msg) = error
715        .and_then(|error| error.get("message"))
716        .and_then(|value| value.as_str())
717    {
718        let lower = msg.to_lowercase();
719        if lower.contains("previous response") && lower.contains("not found") {
720            return true;
721        }
722    }
723    false
724}
725
726pub(super) fn event_error_status(payload: &serde_json::Value) -> Option<u16> {
727    super::events::classify_event_failure(payload).and_then(|failure| failure.explicit_status)
728}
729
730#[allow(dead_code)]
731fn extract_retry_after(payload: &serde_json::Value) -> Option<String> {
732    payload
733        .get("error")
734        .and_then(|e| e.get("retry_after"))
735        .and_then(|v| v.as_str())
736        .map(|s| s.to_string())
737}
738
739// ---------------------------------------------------------------------------
740// Main request function
741// ---------------------------------------------------------------------------
742
743#[allow(clippy::too_many_arguments)]
744pub(super) async fn codex_websocket_request(
745    websocket_client: &reqwest::Client,
746    proxy_config: &WebSocketProxyConfig,
747    url: &str,
748    headers: &HeaderMap,
749    body_value: &serde_json::Value,
750    _ctx: &RequestContext,
751    traffic: Option<&TrafficCapture>,
752    connect_timeout_ms: u64,
753    idle_timeout_ms: u64,
754    reservation: Option<&ContinuationReservation>,
755) -> Result<OwnerAwareCodexResponse, CodexError> {
756    let continuation = reservation.map(ContinuationReservation::candidate);
757    let pool_owner = reservation_pool_owner(reservation);
758    let ws_url = to_websocket_url(url).map_err(|e| CodexError {
759        status: 0,
760        message: e.message,
761        detail: None,
762        retry_after: None,
763        origin: CodexErrorOrigin::WebSocketHandshake,
764    })?;
765    let body_json = serde_json::to_string(body_value).unwrap_or_default();
766    if let Some(tc) = traffic {
767        tc.write_json("020-upstream-request", body_value);
768        tc.write_json(
769            "021-upstream-request-metadata",
770            &serde_json::json!({
771                "provider": "codex",
772                "transport": "websocket",
773                "url": ws_url,
774                "method": "GET",
775                "headers": headers_to_json(headers),
776                "size": summarize_json_request_size(body_value, &body_json),
777                "continuation": {
778                    "previousResponseId": continuation
779                        .and_then(|c| c.previous_response_id.as_deref()),
780                    "inputDeltaCount": continuation
781                        .and_then(|c| c.input_delta.as_ref())
782                        .map(|items| items.len()),
783                    "disabledReason": continuation
784                        .and_then(|c| c.disabled_reason.as_deref()),
785                },
786            }),
787        );
788    }
789    let started_at = Instant::now();
790
791    let requires_origin = continuation
792        .and_then(|candidate| candidate.previous_response_id.as_deref())
793        .is_some();
794    let pooled = take_pool_entry_for_request(reservation)?;
795    let mut used_pooled = pooled.is_some();
796    let mut entry = if let Some(entry) = pooled {
797        entry
798    } else {
799        Arc::new(PoolEntry::new(
800            connect_with_timeout(
801                websocket_client,
802                proxy_config,
803                &ws_url,
804                headers,
805                connect_timeout_ms,
806            )
807            .await?,
808        ))
809    };
810    let mut guard = entry.ws.clone().lock_owned().await;
811
812    if used_pooled
813        && validate_pooled_websocket(&mut guard, connect_timeout_ms)
814            .await
815            .is_err()
816    {
817        drop(guard);
818        if let Some(owner) = pool_owner {
819            pool_remove_entry(owner, &entry);
820        }
821        if requires_origin {
822            return Err(continuation_socket_missing_error());
823        }
824        entry = Arc::new(PoolEntry::new(
825            connect_with_timeout(
826                websocket_client,
827                proxy_config,
828                &ws_url,
829                headers,
830                connect_timeout_ms,
831            )
832            .await?,
833        ));
834        guard = entry.ws.clone().lock_owned().await;
835        used_pooled = false;
836    }
837
838    guard
839        .send(Message::Text(body_json))
840        .await
841        .map_err(|error| {
842            if let Some(owner) = pool_owner {
843                pool_remove_entry(owner, &entry);
844            }
845            CodexError {
846                status: 0,
847                message: format!("WebSocket send error: {error}"),
848                detail: None,
849                retry_after: None,
850                origin: CodexErrorOrigin::WebSocket,
851            }
852        })?;
853
854    let collected = collect_ws_events(
855        &mut guard,
856        idle_timeout_ms,
857        pool_owner,
858        Some(&entry),
859        traffic,
860    )
861    .await;
862    drop(guard);
863    let (sse_body, terminal_event) = match collected {
864        Ok(result) => result,
865        Err(error) => {
866            if let Some(owner) = pool_owner {
867                pool_remove_entry(owner, &entry);
868            }
869            return Err(error);
870        }
871    };
872    let Some(terminal_event) = terminal_event else {
873        if let Some(owner) = pool_owner {
874            pool_remove_entry(owner, &entry);
875        }
876        return Err(missing_terminal_error());
877    };
878
879    if is_previous_response_missing(&terminal_event.payload) {
880        if let Some(owner) = pool_owner {
881            pool_remove_entry(owner, &entry);
882        }
883        return Err(CodexError {
884            status: 0,
885            message: "Previous response not found".to_string(),
886            detail: Some("previous_response_not_found".to_string()),
887            retry_after: None,
888            origin: CodexErrorOrigin::WebSocket,
889        });
890    }
891
892    let completed = terminal_event.event_type == "response.completed";
893    let origin_reinserted = if completed {
894        reservation.is_some_and(|reservation| pool_insert_for_turn(reservation, entry.clone()))
895    } else {
896        if let Some(owner) = pool_owner {
897            pool_remove_entry(owner, &entry);
898        }
899        false
900    };
901    let status = if terminal_event.event_type == "error" {
902        event_error_status(&terminal_event.payload).unwrap_or(500)
903    } else {
904        200
905    };
906
907    if let Some(tc) = traffic {
908        write_websocket_metadata_capture(tc, &ws_url, reservation, used_pooled);
909        write_websocket_response_capture(tc, status, started_at.elapsed(), &sse_body);
910    }
911
912    Ok(OwnerAwareCodexResponse::new(
913        CodexResponse {
914            body: sse_body,
915            status,
916            headers: vec![],
917            transport: ActualTransport::WebSocket,
918        },
919        origin_reinserted.then_some(entry.socket_id),
920    ))
921}
922
923async fn validate_pooled_websocket<S>(
924    websocket: &mut WebSocketStream<S>,
925    timeout_ms: u64,
926) -> Result<(), String>
927where
928    S: AsyncRead + AsyncWrite + Unpin,
929{
930    let nonce = next_monotonic_nonzero(&POOLED_VALIDATION_SEQUENCE, "pooled validation")
931        .to_be_bytes()
932        .to_vec();
933    websocket
934        .send(Message::Ping(nonce.clone()))
935        .await
936        .map_err(|error| error.to_string())?;
937    tokio::time::timeout(Duration::from_millis(timeout_ms), async {
938        loop {
939            match websocket.next().await {
940                Some(Ok(Message::Pong(payload))) if payload == nonce => return Ok(()),
941                Some(Ok(Message::Ping(payload))) => websocket
942                    .send(Message::Pong(payload))
943                    .await
944                    .map_err(|error| error.to_string())?,
945                Some(Ok(Message::Pong(_))) => {
946                    return Err("unexpected Pong during pooled validation".to_string());
947                }
948                Some(Ok(_)) => {
949                    return Err("unexpected frame during pooled validation".to_string());
950                }
951                Some(Err(error)) => return Err(error.to_string()),
952                None => return Err("connection closed during pooled validation".to_string()),
953            }
954        }
955    })
956    .await
957    .map_err(|_| "validation timeout".to_string())?
958}
959
960pub(super) struct ReadyWebSocket {
961    ws_url: String,
962    guard: OwnedMutexGuard<CodexWebSocketStream>,
963    entry: Arc<PoolEntry>,
964    used_pooled: bool,
965    reservation: Option<ContinuationReservation>,
966    traffic: Option<Arc<TrafficCapture>>,
967    idle_timeout_ms: u64,
968}
969
970#[allow(clippy::too_many_arguments)]
971pub(super) async fn prepare_codex_websocket(
972    websocket_client: &reqwest::Client,
973    proxy_config: &WebSocketProxyConfig,
974    url: &str,
975    headers: &HeaderMap,
976    traffic: Option<Arc<TrafficCapture>>,
977    reservation: Option<&ContinuationReservation>,
978    connect_timeout_ms: u64,
979    idle_timeout_ms: u64,
980) -> Result<ReadyWebSocket, CodexError> {
981    let pool_owner = reservation_pool_owner(reservation);
982    let continuation = reservation.map(ContinuationReservation::candidate);
983    let ws_url = to_websocket_url(url).map_err(|error| CodexError {
984        status: 0,
985        message: error.message,
986        detail: None,
987        retry_after: None,
988        origin: CodexErrorOrigin::WebSocketHandshake,
989    })?;
990    let requires_origin = continuation
991        .and_then(|candidate| candidate.previous_response_id.as_deref())
992        .is_some();
993    let pooled = take_pool_entry_for_request(reservation)?;
994    let used_pooled = pooled.is_some();
995    let entry = if let Some(entry) = pooled {
996        entry
997    } else {
998        let stream = connect_with_timeout(
999            websocket_client,
1000            proxy_config,
1001            &ws_url,
1002            headers,
1003            connect_timeout_ms,
1004        )
1005        .await?;
1006        Arc::new(PoolEntry::new(stream))
1007    };
1008    let mut guard = entry.ws.clone().lock_owned().await;
1009    if used_pooled
1010        && let Err(detail) = validate_pooled_websocket(&mut guard, connect_timeout_ms).await
1011    {
1012        drop(guard);
1013        if let Some(owner) = pool_owner {
1014            pool_remove_entry(owner, &entry);
1015        }
1016        return Err(if requires_origin {
1017            continuation_socket_missing_error()
1018        } else {
1019            pooled_validation_error(detail)
1020        });
1021    }
1022    Ok(ReadyWebSocket {
1023        ws_url,
1024        guard,
1025        entry,
1026        used_pooled,
1027        reservation: reservation.cloned(),
1028        traffic,
1029        idle_timeout_ms,
1030    })
1031}
1032
1033fn pooled_validation_error(detail: String) -> CodexError {
1034    CodexError {
1035        status: 0,
1036        message: format!("WebSocket pooled connection validation failed: {detail}"),
1037        detail: None,
1038        retry_after: None,
1039        origin: CodexErrorOrigin::WebSocketHandshake,
1040    }
1041}
1042
1043pub(super) fn start_codex_websocket_events(
1044    ready: ReadyWebSocket,
1045    body_value: &serde_json::Value,
1046    body_json: String,
1047    headers: &HeaderMap,
1048    reservation: Option<&ContinuationReservation>,
1049) -> CodexWebSocketEventStream {
1050    if let Some(tc) = ready.traffic.as_deref() {
1051        write_websocket_metadata_capture(tc, &ready.ws_url, reservation, ready.used_pooled);
1052        tc.write_json("020-upstream-request", body_value);
1053        tc.write_json(
1054            "021-upstream-request-metadata",
1055            &serde_json::json!({
1056                "provider": "codex",
1057                "transport": "websocket",
1058                "headers": headers_to_json(headers),
1059                "size": summarize_json_request_size(body_value, &body_json),
1060            }),
1061        );
1062    }
1063    let (tx, rx) = mpsc::channel(64);
1064    let (receiver, socket_id_publisher) = CodexWebSocketEventStream::pending(rx);
1065    tokio::spawn(async move {
1066        let ReadyWebSocket {
1067            ws_url: _,
1068            mut guard,
1069            entry,
1070            used_pooled: _,
1071            reservation,
1072            traffic,
1073            idle_timeout_ms,
1074        } = ready;
1075        let pool_owner = reservation_pool_owner(reservation.as_ref());
1076        if let Err(error) = guard.send(Message::Text(body_json)).await {
1077            drop(guard);
1078            if let Some(owner) = pool_owner {
1079                pool_remove_entry(owner, &entry);
1080            }
1081            socket_id_publisher.publish(None);
1082            let _ = tx
1083                .send(Err(CodexError {
1084                    status: 0,
1085                    message: format!("WebSocket send error: {error}"),
1086                    detail: None,
1087                    retry_after: None,
1088                    origin: CodexErrorOrigin::WebSocket,
1089                }))
1090                .await;
1091            return;
1092        }
1093        let (reusable, terminal_item) =
1094            stream_ws_events(&mut guard, idle_timeout_ms, traffic, &tx).await;
1095        drop(guard);
1096
1097        let origin_reinserted = if reusable {
1098            reservation
1099                .as_ref()
1100                .is_some_and(|reservation| pool_insert_for_turn(reservation, entry.clone()))
1101        } else {
1102            if let Some(owner) = pool_owner {
1103                pool_remove_entry(owner, &entry);
1104            }
1105            false
1106        };
1107        socket_id_publisher.publish(origin_reinserted.then_some(entry.socket_id));
1108        if let Some(item) = terminal_item
1109            && tx.send(item).await.is_err()
1110            && origin_reinserted
1111            && let Some(owner) = pool_owner
1112        {
1113            pool_remove_entry(owner, &entry);
1114        }
1115    });
1116    receiver
1117}
1118
1119#[allow(clippy::too_many_arguments)]
1120pub(super) async fn codex_websocket_event_stream(
1121    websocket_client: &reqwest::Client,
1122    proxy_config: &WebSocketProxyConfig,
1123    url: &str,
1124    headers: &HeaderMap,
1125    body_value: &serde_json::Value,
1126    _ctx: &RequestContext,
1127    traffic: Option<Arc<TrafficCapture>>,
1128    connect_timeout_ms: u64,
1129    idle_timeout_ms: u64,
1130    reservation: Option<&ContinuationReservation>,
1131) -> Result<CodexWebSocketEventStream, CodexError> {
1132    let body_json = serde_json::to_string(body_value).map_err(|error| CodexError {
1133        status: 500,
1134        message: "Failed to serialize WebSocket request".to_string(),
1135        detail: Some(error.to_string()),
1136        retry_after: None,
1137        origin: CodexErrorOrigin::WebSocketHandshake,
1138    })?;
1139    let ready = prepare_codex_websocket(
1140        websocket_client,
1141        proxy_config,
1142        url,
1143        headers,
1144        traffic,
1145        reservation,
1146        connect_timeout_ms,
1147        idle_timeout_ms,
1148    )
1149    .await?;
1150    Ok(start_codex_websocket_events(
1151        ready,
1152        body_value,
1153        body_json,
1154        headers,
1155        reservation,
1156    ))
1157}
1158
1159fn continuation_socket_missing_error() -> CodexError {
1160    CodexError {
1161        status: 0,
1162        message: "Previous response socket is no longer available".to_string(),
1163        detail: Some(WEBSOCKET_CONTINUATION_SOCKET_MISSING_DETAIL.to_string()),
1164        retry_after: None,
1165        origin: CodexErrorOrigin::WebSocketHandshake,
1166    }
1167}
1168
1169fn missing_terminal_error() -> CodexError {
1170    CodexError {
1171        status: 0,
1172        message: "WebSocket connection closed before terminal Codex response event".to_string(),
1173        detail: Some(WEBSOCKET_MISSING_TERMINAL_DETAIL.to_string()),
1174        retry_after: None,
1175        origin: CodexErrorOrigin::WebSocket,
1176    }
1177}
1178
1179fn response_start_timeout_error(timeout_ms: u64) -> CodexError {
1180    CodexError {
1181        status: 0,
1182        message: format!("WebSocket response start timeout after {timeout_ms}ms"),
1183        detail: Some(WEBSOCKET_RESPONSE_START_TIMEOUT_DETAIL.to_string()),
1184        retry_after: None,
1185        origin: CodexErrorOrigin::WebSocket,
1186    }
1187}
1188
1189fn write_websocket_metadata_capture(
1190    traffic: &TrafficCapture,
1191    ws_url: &str,
1192    reservation: Option<&ContinuationReservation>,
1193    pooled: bool,
1194) {
1195    let pool_owner = reservation_pool_owner(reservation);
1196    let continuation = reservation.map(ContinuationReservation::candidate);
1197    traffic.write_json(
1198        "022-upstream-websocket-metadata",
1199        &serde_json::json!({
1200            "provider": "codex",
1201            "transport": "websocket",
1202            "url": ws_url,
1203            "poolingEnabled": pool_owner.is_some(),
1204            "pooled": pooled,
1205            "continuation": {
1206                "previousResponseId": continuation
1207                    .and_then(|c| c.previous_response_id.as_deref()),
1208                "inputDeltaCount": continuation
1209                    .and_then(|c| c.input_delta.as_ref())
1210                    .map(|items| items.len()),
1211                "disabledReason": continuation
1212                    .and_then(|c| c.disabled_reason.as_deref()),
1213            },
1214        }),
1215    );
1216}
1217
1218fn write_websocket_response_capture(
1219    traffic: &TrafficCapture,
1220    status: u16,
1221    elapsed: Duration,
1222    sse_body: &[u8],
1223) {
1224    traffic.write_json(
1225        "030-upstream-response-headers",
1226        &serde_json::json!({
1227            "status": status,
1228            "elapsedMs": elapsed.as_millis(),
1229            "headers": {
1230                "content-type": "text/event-stream",
1231            },
1232        }),
1233    );
1234    if status >= 400 {
1235        traffic.write_text(
1236            "031-upstream-error-body",
1237            &String::from_utf8_lossy(sse_body),
1238        );
1239    } else {
1240        traffic.write_bytes("032-upstream-response-body.sse", sse_body);
1241    }
1242}
1243
1244// ---------------------------------------------------------------------------
1245// Connection helper
1246// ---------------------------------------------------------------------------
1247
1248const MAX_HANDSHAKE_ERROR_DETAIL_BYTES: usize = 1024;
1249const GENERIC_HANDSHAKE_ERROR_DETAIL: &str = "WebSocket upgrade was rejected";
1250
1251fn handshake_error_detail(body: Option<&[u8]>) -> String {
1252    let Some(value) = body.and_then(|body| serde_json::from_slice::<serde_json::Value>(body).ok())
1253    else {
1254        return GENERIC_HANDSHAKE_ERROR_DETAIL.to_string();
1255    };
1256    let Some(message) = value
1257        .pointer("/error/message")
1258        .or_else(|| value.get("message"))
1259        .and_then(|value| value.as_str())
1260    else {
1261        return GENERIC_HANDSHAKE_ERROR_DETAIL.to_string();
1262    };
1263    let sanitized: String = message
1264        .chars()
1265        .filter(|ch| !ch.is_control() || matches!(ch, ' '))
1266        .collect();
1267    let mut end = sanitized.len().min(MAX_HANDSHAKE_ERROR_DETAIL_BYTES);
1268    while !sanitized.is_char_boundary(end) {
1269        end -= 1;
1270    }
1271    sanitized[..end].to_string()
1272}
1273
1274fn header_has_token(headers: &HeaderMap, name: &str, expected: &str) -> bool {
1275    headers.get_all(name).iter().any(|value| {
1276        value.to_str().ok().is_some_and(|value| {
1277            value
1278                .split(',')
1279                .any(|token| token.trim().eq_ignore_ascii_case(expected))
1280        })
1281    })
1282}
1283
1284fn requested_subprotocols(headers: &HeaderMap) -> Vec<String> {
1285    headers
1286        .get_all(http::header::SEC_WEBSOCKET_PROTOCOL)
1287        .iter()
1288        .filter_map(|value| value.to_str().ok())
1289        .flat_map(|value| value.split(','))
1290        .map(str::trim)
1291        .filter(|value| !value.is_empty())
1292        .map(str::to_string)
1293        .collect()
1294}
1295
1296fn websocket_protocol_error(message: &str) -> CodexError {
1297    CodexError {
1298        status: 0,
1299        message: message.to_string(),
1300        detail: None,
1301        retry_after: None,
1302        origin: CodexErrorOrigin::WebSocketHandshake,
1303    }
1304}
1305
1306fn validate_websocket_upgrade(
1307    version: http::Version,
1308    headers: &HeaderMap,
1309    websocket_key: &str,
1310    requested_subprotocols: &[String],
1311) -> Result<(), CodexError> {
1312    if version != http::Version::HTTP_11 {
1313        return Err(websocket_protocol_error(
1314            "WebSocket upgrade response did not use HTTP/1.1",
1315        ));
1316    }
1317    if !header_has_token(headers, http::header::UPGRADE.as_str(), "websocket") {
1318        return Err(websocket_protocol_error(
1319            "WebSocket upgrade response is missing Upgrade: websocket",
1320        ));
1321    }
1322    if !header_has_token(headers, http::header::CONNECTION.as_str(), "upgrade") {
1323        return Err(websocket_protocol_error(
1324            "WebSocket upgrade response is missing Connection: Upgrade",
1325        ));
1326    }
1327
1328    let expected_accept = derive_accept_key(websocket_key.as_bytes());
1329    let mut accept_values = headers.get_all(http::header::SEC_WEBSOCKET_ACCEPT).iter();
1330    let accept = accept_values.next().and_then(|value| value.to_str().ok());
1331    if accept_values.next().is_some() || accept != Some(expected_accept.as_str()) {
1332        return Err(websocket_protocol_error(
1333            "WebSocket upgrade response has an invalid Sec-WebSocket-Accept",
1334        ));
1335    }
1336    if headers.contains_key(http::header::SEC_WEBSOCKET_EXTENSIONS) {
1337        return Err(websocket_protocol_error(
1338            "WebSocket upgrade response selected an unsolicited extension",
1339        ));
1340    }
1341
1342    let mut response_protocols = headers.get_all(http::header::SEC_WEBSOCKET_PROTOCOL).iter();
1343    let response_protocol = response_protocols
1344        .next()
1345        .map(|value| value.to_str().map(str::trim));
1346    if response_protocols.next().is_some() {
1347        return Err(websocket_protocol_error(
1348            "WebSocket upgrade response contains multiple subprotocols",
1349        ));
1350    }
1351    match response_protocol {
1352        None if requested_subprotocols.is_empty() => {}
1353        None => {
1354            return Err(websocket_protocol_error(
1355                "WebSocket upgrade response omitted the requested subprotocol",
1356            ));
1357        }
1358        Some(Err(_)) => {
1359            return Err(websocket_protocol_error(
1360                "WebSocket upgrade response contains an invalid subprotocol",
1361            ));
1362        }
1363        Some(Ok(_)) if requested_subprotocols.is_empty() => {
1364            return Err(websocket_protocol_error(
1365                "WebSocket upgrade response selected an unsolicited subprotocol",
1366            ));
1367        }
1368        Some(Ok(protocol))
1369            if !requested_subprotocols
1370                .iter()
1371                .any(|requested| requested == protocol) =>
1372        {
1373            return Err(websocket_protocol_error(
1374                "WebSocket upgrade response selected an unsupported subprotocol",
1375            ));
1376        }
1377        Some(Ok(_)) => {}
1378    }
1379
1380    Ok(())
1381}
1382
1383async fn bounded_handshake_error_body(mut response: reqwest::Response) -> Vec<u8> {
1384    let mut body = Vec::new();
1385    while body.len() < MAX_HANDSHAKE_ERROR_DETAIL_BYTES {
1386        let chunk = match response.chunk().await {
1387            Ok(Some(chunk)) => chunk,
1388            Ok(None) | Err(_) => break,
1389        };
1390        let remaining = MAX_HANDSHAKE_ERROR_DETAIL_BYTES - body.len();
1391        body.extend_from_slice(&chunk[..chunk.len().min(remaining)]);
1392    }
1393    body
1394}
1395
1396fn error_chain_contains(error: &(dyn std::error::Error + 'static), expected: &str) -> bool {
1397    let expected = expected.to_ascii_lowercase();
1398    let mut current = Some(error);
1399    while let Some(error) = current {
1400        if error.to_string().to_ascii_lowercase().contains(&expected) {
1401            return true;
1402        }
1403        current = error.source();
1404    }
1405    false
1406}
1407
1408fn reqwest_handshake_error(error: reqwest::Error) -> CodexError {
1409    let proxy_auth_required =
1410        error.is_connect() && error_chain_contains(&error, "proxy authorization required");
1411    let proxy_tunnel_rejected =
1412        error.is_connect() && error_chain_contains(&error, "tunnel error: unsuccessful");
1413    let status = if proxy_auth_required {
1414        http::StatusCode::PROXY_AUTHENTICATION_REQUIRED.as_u16()
1415    } else {
1416        error.status().map(|status| status.as_u16()).unwrap_or(0)
1417    };
1418    let message = if proxy_auth_required {
1419        "WebSocket proxy authentication failed"
1420    } else if proxy_tunnel_rejected {
1421        "WebSocket proxy tunnel was rejected"
1422    } else if error.is_timeout() {
1423        "WebSocket upgrade request timed out"
1424    } else if error.is_connect() {
1425        "WebSocket connection failed"
1426    } else {
1427        "WebSocket upgrade request failed"
1428    };
1429    CodexError {
1430        status,
1431        message: message.to_string(),
1432        detail: proxy_tunnel_rejected.then(|| WEBSOCKET_PROXY_TUNNEL_REJECTED_DETAIL.to_string()),
1433        retry_after: None,
1434        origin: CodexErrorOrigin::WebSocketHandshake,
1435    }
1436}
1437
1438struct PrefixedIo {
1439    prefix: Vec<u8>,
1440    position: usize,
1441    inner: BoxedWebSocketIo,
1442}
1443
1444impl AsyncRead for PrefixedIo {
1445    fn poll_read(
1446        self: Pin<&mut Self>,
1447        cx: &mut Context<'_>,
1448        buffer: &mut ReadBuf<'_>,
1449    ) -> Poll<std::io::Result<()>> {
1450        let this = self.get_mut();
1451        if this.position < this.prefix.len() {
1452            let available = &this.prefix[this.position..];
1453            let len = available.len().min(buffer.remaining());
1454            buffer.put_slice(&available[..len]);
1455            this.position += len;
1456            return Poll::Ready(Ok(()));
1457        }
1458        Pin::new(&mut this.inner).poll_read(cx, buffer)
1459    }
1460}
1461
1462impl AsyncWrite for PrefixedIo {
1463    fn poll_write(
1464        self: Pin<&mut Self>,
1465        cx: &mut Context<'_>,
1466        buffer: &[u8],
1467    ) -> Poll<std::io::Result<usize>> {
1468        Pin::new(&mut self.get_mut().inner).poll_write(cx, buffer)
1469    }
1470
1471    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
1472        Pin::new(&mut self.get_mut().inner).poll_flush(cx)
1473    }
1474
1475    fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
1476        Pin::new(&mut self.get_mut().inner).poll_shutdown(cx)
1477    }
1478}
1479
1480fn skip_websocket_request_header(name: &http::HeaderName) -> bool {
1481    matches!(
1482        name.as_str(),
1483        "connection"
1484            | "upgrade"
1485            | "sec-websocket-key"
1486            | "sec-websocket-version"
1487            | "host"
1488            | "content-length"
1489            | "proxy-authorization"
1490    )
1491}
1492
1493fn tunneled_websocket_request(
1494    url: &str,
1495    headers: &HeaderMap,
1496    websocket_key: &str,
1497) -> Result<http::Request<()>, CodexError> {
1498    let mut request = url
1499        .into_client_request()
1500        .map_err(|_| websocket_protocol_error("WebSocket request URL was invalid"))?;
1501    *request.version_mut() = http::Version::HTTP_11;
1502    request.headers_mut().insert(
1503        http::header::SEC_WEBSOCKET_KEY,
1504        http::HeaderValue::from_str(websocket_key)
1505            .map_err(|_| websocket_protocol_error("WebSocket key was invalid"))?,
1506    );
1507    for (name, value) in headers {
1508        if !skip_websocket_request_header(name) {
1509            request.headers_mut().append(name.clone(), value.clone());
1510        }
1511    }
1512    Ok(request)
1513}
1514
1515fn tunnel_error(status: u16, retry_after: Option<String>) -> CodexError {
1516    let proxy_auth_required = status == http::StatusCode::PROXY_AUTHENTICATION_REQUIRED.as_u16();
1517    CodexError {
1518        status,
1519        message: if proxy_auth_required {
1520            "WebSocket proxy authentication failed".to_string()
1521        } else {
1522            "WebSocket proxy tunnel was rejected".to_string()
1523        },
1524        detail: if proxy_auth_required {
1525            Some(GENERIC_HANDSHAKE_ERROR_DETAIL.to_string())
1526        } else {
1527            Some(WEBSOCKET_PROXY_TUNNEL_REJECTED_DETAIL.to_string())
1528        },
1529        retry_after,
1530        origin: CodexErrorOrigin::WebSocketHandshake,
1531    }
1532}
1533
1534fn invalid_tunnel_response(message: &str) -> CodexError {
1535    CodexError {
1536        status: 0,
1537        message: message.to_string(),
1538        detail: Some(WEBSOCKET_PROXY_TUNNEL_REJECTED_DETAIL.to_string()),
1539        retry_after: None,
1540        origin: CodexErrorOrigin::WebSocketHandshake,
1541    }
1542}
1543
1544fn connect_response_header_end(response: &[u8]) -> Option<usize> {
1545    response
1546        .windows(4)
1547        .position(|window| window == b"\r\n\r\n")
1548        .map(|position| position + 4)
1549}
1550
1551fn parse_connect_response_head(response: &[u8]) -> Result<(u16, Option<String>), CodexError> {
1552    let status_line_end = response
1553        .windows(2)
1554        .position(|window| window == b"\r\n")
1555        .ok_or_else(|| invalid_tunnel_response("WebSocket proxy returned an invalid response"))?;
1556    let status_line = std::str::from_utf8(&response[..status_line_end])
1557        .map_err(|_| invalid_tunnel_response("WebSocket proxy returned an invalid response"))?;
1558    let mut parts = status_line.split_ascii_whitespace();
1559    let version = parts.next().unwrap_or_default();
1560    let status = parts.next().unwrap_or_default();
1561    if !matches!(version, "HTTP/1.0" | "HTTP/1.1")
1562        || status.len() != 3
1563        || !status.bytes().all(|byte| byte.is_ascii_digit())
1564    {
1565        return Err(invalid_tunnel_response(
1566            "WebSocket proxy returned an invalid response",
1567        ));
1568    }
1569    let status = status
1570        .parse::<u16>()
1571        .map_err(|_| invalid_tunnel_response("WebSocket proxy returned an invalid response"))?;
1572    let retry_after = response[status_line_end + 2..]
1573        .split(|byte| *byte == b'\n')
1574        .find_map(|line| {
1575            let line = line.strip_suffix(b"\r").unwrap_or(line);
1576            let separator = line.iter().position(|byte| *byte == b':')?;
1577            let (name, value) = line.split_at(separator);
1578            if !name.eq_ignore_ascii_case(b"retry-after") {
1579                return None;
1580            }
1581            std::str::from_utf8(&value[1..])
1582                .ok()
1583                .map(|value| value.trim().to_string())
1584        });
1585    Ok((status, retry_after))
1586}
1587
1588async fn establish_connect_tunnel(
1589    mut stream: BoxedWebSocketIo,
1590    authority: &str,
1591    basic_auth: Option<&http::HeaderValue>,
1592) -> Result<BoxedWebSocketIo, CodexError> {
1593    let mut request = format!("CONNECT {authority} HTTP/1.1\r\nHost: {authority}\r\n").into_bytes();
1594    if let Some(auth) = basic_auth {
1595        request.extend_from_slice(b"Proxy-Authorization: ");
1596        request.extend_from_slice(auth.as_bytes());
1597        request.extend_from_slice(b"\r\n");
1598    }
1599    request.extend_from_slice(b"\r\n");
1600    stream.write_all(&request).await.map_err(|_| CodexError {
1601        status: 0,
1602        message: "WebSocket proxy tunnel request failed".to_string(),
1603        detail: None,
1604        retry_after: None,
1605        origin: CodexErrorOrigin::WebSocketHandshake,
1606    })?;
1607
1608    let mut response = Vec::new();
1609    loop {
1610        if let Some(header_end) = connect_response_header_end(&response) {
1611            let (status, retry_after) = parse_connect_response_head(&response[..header_end])?;
1612            if (100..200).contains(&status) {
1613                response.drain(..header_end);
1614                continue;
1615            }
1616            if (200..300).contains(&status) {
1617                let prefix = response.split_off(header_end);
1618                return Ok(Box::new(PrefixedIo {
1619                    prefix,
1620                    position: 0,
1621                    inner: stream,
1622                }));
1623            }
1624            return Err(tunnel_error(status, retry_after));
1625        }
1626        if response.len() == MAX_CONNECT_RESPONSE_HEADER_BYTES {
1627            return Err(invalid_tunnel_response(
1628                "WebSocket proxy response headers were too large",
1629            ));
1630        }
1631        let remaining = MAX_CONNECT_RESPONSE_HEADER_BYTES - response.len();
1632        let mut buffer = [0_u8; 1024];
1633        let capacity = remaining.min(buffer.len());
1634        let read = stream
1635            .read(&mut buffer[..capacity])
1636            .await
1637            .map_err(|_| invalid_tunnel_response("WebSocket proxy response could not be read"))?;
1638        if read == 0 {
1639            return Err(invalid_tunnel_response(
1640                "WebSocket proxy closed the tunnel response early",
1641            ));
1642        }
1643        response.extend_from_slice(&buffer[..read]);
1644    }
1645}
1646
1647async fn tls_connect(
1648    stream: BoxedWebSocketIo,
1649    host: &str,
1650    tls_config: Arc<rustls::ClientConfig>,
1651    peer: &str,
1652) -> Result<BoxedWebSocketIo, CodexError> {
1653    let server_name = rustls::pki_types::ServerName::try_from(host.to_string())
1654        .map_err(|_| websocket_protocol_error("WebSocket TLS host name was invalid"))?;
1655    let stream = tokio_rustls::TlsConnector::from(tls_config)
1656        .connect(server_name, stream)
1657        .await
1658        .map_err(|_| CodexError {
1659            status: 0,
1660            message: format!("WebSocket TLS connection to {peer} failed"),
1661            detail: None,
1662            retry_after: None,
1663            origin: CodexErrorOrigin::WebSocketHandshake,
1664        })?;
1665    Ok(Box::new(stream))
1666}
1667
1668async fn connect_to_http_proxy(
1669    route: &WebSocketProxyRoute,
1670    tls_config: Arc<rustls::ClientConfig>,
1671) -> Result<BoxedWebSocketIo, CodexError> {
1672    let scheme = route.uri.scheme_str().unwrap_or_default();
1673    let host = route
1674        .uri
1675        .host()
1676        .ok_or_else(|| websocket_protocol_error("WebSocket proxy URL did not contain a host"))?;
1677    let host = host.trim_start_matches('[').trim_end_matches(']');
1678    let port = route
1679        .uri
1680        .port_u16()
1681        .unwrap_or(if scheme == "https" { 443 } else { 80 });
1682    let stream = TcpStream::connect((host, port))
1683        .await
1684        .map_err(|_| CodexError {
1685            status: 0,
1686            message: "WebSocket proxy connection failed".to_string(),
1687            detail: None,
1688            retry_after: None,
1689            origin: CodexErrorOrigin::WebSocketHandshake,
1690        })?;
1691    let stream: BoxedWebSocketIo = Box::new(stream);
1692    if scheme == "https" {
1693        tls_connect(stream, host, tls_config, "proxy").await
1694    } else {
1695        Ok(stream)
1696    }
1697}
1698
1699fn websocket_destination(url: &str) -> Result<(String, String), CodexError> {
1700    let destination = url::Url::parse(url)
1701        .map_err(|_| websocket_protocol_error("WebSocket destination URL was invalid"))?;
1702    let host = destination
1703        .host_str()
1704        .ok_or_else(|| websocket_protocol_error("WebSocket destination URL had no host"))?;
1705    let port = destination
1706        .port_or_known_default()
1707        .ok_or_else(|| websocket_protocol_error("WebSocket destination URL had no usable port"))?;
1708    let authority = match destination.host() {
1709        Some(url::Host::Ipv6(address)) => format!("[{address}]:{port}"),
1710        Some(_) => format!("{host}:{port}"),
1711        None => {
1712            return Err(websocket_protocol_error(
1713                "WebSocket destination URL had no host",
1714            ));
1715        }
1716    };
1717    Ok((host.to_string(), authority))
1718}
1719
1720fn tungstenite_handshake_error(error: tokio_tungstenite::tungstenite::Error) -> CodexError {
1721    if let tokio_tungstenite::tungstenite::Error::Http(response) = error {
1722        let status = response.status().as_u16();
1723        let retry_after = response
1724            .headers()
1725            .get(http::header::RETRY_AFTER)
1726            .and_then(|value| value.to_str().ok())
1727            .map(str::to_string);
1728        return CodexError {
1729            status,
1730            message: format!("WebSocket upgrade rejected with status {status}"),
1731            detail: Some(GENERIC_HANDSHAKE_ERROR_DETAIL.to_string()),
1732            retry_after,
1733            origin: CodexErrorOrigin::WebSocketHandshake,
1734        };
1735    }
1736    CodexError {
1737        status: 0,
1738        message: "WebSocket upgrade request failed".to_string(),
1739        detail: None,
1740        retry_after: None,
1741        origin: CodexErrorOrigin::WebSocketHandshake,
1742    }
1743}
1744
1745enum ConnectAttemptError {
1746    Origin(CodexError),
1747    ProxyTunnel(CodexError),
1748}
1749
1750impl ConnectAttemptError {
1751    fn is_origin_forbidden(&self) -> bool {
1752        matches!(
1753            self,
1754            Self::Origin(error) if error.status == http::StatusCode::FORBIDDEN.as_u16()
1755        )
1756    }
1757
1758    fn into_error(self) -> CodexError {
1759        match self {
1760            Self::Origin(error) | Self::ProxyTunnel(error) => error,
1761        }
1762    }
1763}
1764
1765async fn connect_via_http_proxy_tunnel(
1766    proxy_config: &WebSocketProxyConfig,
1767    route: WebSocketProxyRoute,
1768    url: &str,
1769    headers: &HeaderMap,
1770) -> Result<CodexWebSocketStream, ConnectAttemptError> {
1771    let (host, authority) = websocket_destination(url).map_err(ConnectAttemptError::Origin)?;
1772    let stream = connect_to_http_proxy(&route, proxy_config.tls_config.clone())
1773        .await
1774        .map_err(ConnectAttemptError::ProxyTunnel)?;
1775    let stream = establish_connect_tunnel(stream, &authority, route.basic_auth.as_ref())
1776        .await
1777        .map_err(ConnectAttemptError::ProxyTunnel)?;
1778    let stream = tls_connect(
1779        stream,
1780        &host,
1781        proxy_config.tls_config.clone(),
1782        "destination",
1783    )
1784    .await
1785    .map_err(ConnectAttemptError::Origin)?;
1786    let websocket_key = generate_key();
1787    let subprotocols = requested_subprotocols(headers);
1788    let request = tunneled_websocket_request(url, headers, &websocket_key)
1789        .map_err(ConnectAttemptError::Origin)?;
1790    let (websocket, response) = tokio_tungstenite::client_async(request, stream)
1791        .await
1792        .map_err(tungstenite_handshake_error)
1793        .map_err(ConnectAttemptError::Origin)?;
1794    validate_websocket_upgrade(
1795        response.version(),
1796        response.headers(),
1797        &websocket_key,
1798        &subprotocols,
1799    )
1800    .map_err(ConnectAttemptError::Origin)?;
1801    Ok(websocket)
1802}
1803
1804async fn connect_via_http_upgrade(
1805    websocket_client: &reqwest::Client,
1806    url: &str,
1807    headers: &HeaderMap,
1808) -> Result<CodexWebSocketStream, CodexError> {
1809    let http_url = to_http_upgrade_url(url).map_err(|error| CodexError {
1810        status: 0,
1811        message: error.message,
1812        detail: None,
1813        retry_after: None,
1814        origin: CodexErrorOrigin::WebSocketHandshake,
1815    })?;
1816    let websocket_key = generate_key();
1817    let subprotocols = requested_subprotocols(headers);
1818    let mut request = websocket_client
1819        .get(http_url)
1820        .version(http::Version::HTTP_11)
1821        .header(http::header::CONNECTION, "Upgrade")
1822        .header(http::header::UPGRADE, "websocket")
1823        .header(http::header::SEC_WEBSOCKET_VERSION, "13")
1824        .header(http::header::SEC_WEBSOCKET_KEY, &websocket_key);
1825
1826    for (key, value) in headers {
1827        if skip_websocket_request_header(key) {
1828            continue;
1829        }
1830        request = request.header(key.clone(), value.clone());
1831    }
1832
1833    let response = request.send().await.map_err(reqwest_handshake_error)?;
1834    if response.status() != http::StatusCode::SWITCHING_PROTOCOLS {
1835        let status = response.status().as_u16();
1836        let retry_after = response
1837            .headers()
1838            .get(http::header::RETRY_AFTER)
1839            .and_then(|value| value.to_str().ok())
1840            .map(str::to_string);
1841        let detail = if status == http::StatusCode::PROXY_AUTHENTICATION_REQUIRED.as_u16() {
1842            GENERIC_HANDSHAKE_ERROR_DETAIL.to_string()
1843        } else {
1844            let body = bounded_handshake_error_body(response).await;
1845            handshake_error_detail(Some(&body))
1846        };
1847        return Err(CodexError {
1848            status,
1849            message: format!("WebSocket upgrade rejected with status {status}"),
1850            detail: Some(detail),
1851            retry_after,
1852            origin: CodexErrorOrigin::WebSocketHandshake,
1853        });
1854    }
1855
1856    validate_websocket_upgrade(
1857        response.version(),
1858        response.headers(),
1859        &websocket_key,
1860        &subprotocols,
1861    )?;
1862    let upgraded = response
1863        .upgrade()
1864        .await
1865        .map_err(|_| websocket_protocol_error("WebSocket upgrade stream was not available"))?;
1866    let upgraded: BoxedWebSocketIo = Box::new(upgraded);
1867    Ok(WebSocketStream::from_raw_socket(upgraded, Role::Client, None).await)
1868}
1869
1870async fn connect_once(
1871    websocket_client: &reqwest::Client,
1872    proxy_config: &WebSocketProxyConfig,
1873    url: &str,
1874    headers: &HeaderMap,
1875) -> Result<CodexWebSocketStream, ConnectAttemptError> {
1876    if let Some(route) = proxy_config
1877        .http_connect_route(url)
1878        .map_err(ConnectAttemptError::Origin)?
1879    {
1880        connect_via_http_proxy_tunnel(proxy_config, route, url, headers).await
1881    } else {
1882        connect_via_http_upgrade(websocket_client, url, headers)
1883            .await
1884            .map_err(ConnectAttemptError::Origin)
1885    }
1886}
1887
1888fn connect_timeout_error(connect_timeout: Duration) -> CodexError {
1889    websocket_protocol_error(&format!(
1890        "WebSocket connect timeout after {}ms",
1891        connect_timeout.as_millis()
1892    ))
1893}
1894
1895async fn connect_with_policy<T, P, PFut, F, Fut>(
1896    gate: &WebSocketConnectGate,
1897    connect_timeout: Duration,
1898    forbidden_cooldown: Duration,
1899    mut before_start: P,
1900    mut connect: F,
1901) -> Result<T, CodexError>
1902where
1903    P: FnMut() -> PFut,
1904    PFut: std::future::Future<Output = ()>,
1905    F: FnMut() -> Fut,
1906    Fut: std::future::Future<Output = Result<T, ConnectAttemptError>>,
1907{
1908    for attempt in 0..=1 {
1909        gate.wait_to_start(before_start()).await;
1910        let result = tokio::time::timeout(connect_timeout, connect())
1911            .await
1912            .unwrap_or_else(|_| {
1913                Err(ConnectAttemptError::Origin(connect_timeout_error(
1914                    connect_timeout,
1915                )))
1916            });
1917        if result
1918            .as_ref()
1919            .is_err_and(ConnectAttemptError::is_origin_forbidden)
1920            && attempt == 0
1921        {
1922            retry_sleep(u64::try_from(forbidden_cooldown.as_millis()).unwrap_or(u64::MAX)).await;
1923            continue;
1924        }
1925        return result.map_err(ConnectAttemptError::into_error);
1926    }
1927    unreachable!()
1928}
1929
1930async fn connect_with_timeout(
1931    websocket_client: &reqwest::Client,
1932    proxy_config: &WebSocketProxyConfig,
1933    url: &str,
1934    headers: &HeaderMap,
1935    connect_timeout_ms: u64,
1936) -> Result<CodexWebSocketStream, CodexError> {
1937    connect_with_policy(
1938        &WS_CONNECT_GATE,
1939        Duration::from_millis(connect_timeout_ms),
1940        WEBSOCKET_CONNECT_FORBIDDEN_COOLDOWN,
1941        || async { cleanup_pool_before_connect() },
1942        || connect_once(websocket_client, proxy_config, url, headers),
1943    )
1944    .await
1945}
1946
1947// ---------------------------------------------------------------------------
1948// Event collection
1949// ---------------------------------------------------------------------------
1950
1951enum WebSocketRead {
1952    Frame(Option<Result<Message, tokio_tungstenite::tungstenite::Error>>),
1953    Timeout,
1954    KeepaliveError(String),
1955}
1956
1957async fn read_ws_frame_with_keepalive<S>(
1958    ws: &mut WebSocketStream<S>,
1959    read_timeout: Duration,
1960    keepalive_interval: Duration,
1961) -> WebSocketRead
1962where
1963    S: AsyncRead + AsyncWrite + Unpin,
1964{
1965    let timeout = tokio::time::sleep(read_timeout);
1966    tokio::pin!(timeout);
1967    let mut keepalive = tokio::time::interval(keepalive_interval);
1968    keepalive.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
1969    keepalive.tick().await;
1970
1971    loop {
1972        tokio::select! {
1973            biased;
1974            frame = ws.next() => return WebSocketRead::Frame(frame),
1975            _ = &mut timeout => return WebSocketRead::Timeout,
1976            _ = keepalive.tick() => {
1977                let send = ws.send(Message::Ping(Vec::new()));
1978                let send_timeout = tokio::time::sleep(WEBSOCKET_KEEPALIVE_SEND_TIMEOUT);
1979                tokio::pin!(send_timeout);
1980                tokio::select! {
1981                    biased;
1982                    result = send => {
1983                        if let Err(error) = result {
1984                            return WebSocketRead::KeepaliveError(error.to_string());
1985                        }
1986                    }
1987                    _ = &mut timeout => return WebSocketRead::Timeout,
1988                    _ = &mut send_timeout => {
1989                        return WebSocketRead::KeepaliveError(format!(
1990                            "send timed out after {}ms",
1991                            WEBSOCKET_KEEPALIVE_SEND_TIMEOUT.as_millis(),
1992                        ));
1993                    }
1994                }
1995            }
1996        }
1997    }
1998}
1999
2000struct WsEvent {
2001    event_type: String,
2002    payload: serde_json::Value,
2003}
2004
2005async fn collect_ws_events<S>(
2006    ws: &mut WebSocketStream<S>,
2007    idle_timeout_ms: u64,
2008    pool_owner: Option<&ConversationIdentity>,
2009    pool_entry: Option<&Arc<PoolEntry>>,
2010    traffic: Option<&TrafficCapture>,
2011) -> Result<(Vec<u8>, Option<WsEvent>), CodexError>
2012where
2013    S: AsyncRead + AsyncWrite + Unpin,
2014{
2015    collect_ws_events_with_keepalive_interval(
2016        ws,
2017        idle_timeout_ms,
2018        pool_owner,
2019        pool_entry,
2020        traffic,
2021        WEBSOCKET_KEEPALIVE_INTERVAL,
2022    )
2023    .await
2024}
2025
2026async fn collect_ws_events_with_keepalive_interval<S>(
2027    ws: &mut WebSocketStream<S>,
2028    idle_timeout_ms: u64,
2029    pool_owner: Option<&ConversationIdentity>,
2030    pool_entry: Option<&Arc<PoolEntry>>,
2031    traffic: Option<&TrafficCapture>,
2032    keepalive_interval: Duration,
2033) -> Result<(Vec<u8>, Option<WsEvent>), CodexError>
2034where
2035    S: AsyncRead + AsyncWrite + Unpin,
2036{
2037    let mut sse_body: Vec<u8> = Vec::new();
2038    let mut terminal_event: Option<WsEvent> = None;
2039    let response_event_budget = Duration::from_millis(idle_timeout_ms);
2040    let response_wait_started = Instant::now();
2041    let mut last_response_event_at = response_wait_started;
2042    let mut response_started = false;
2043
2044    loop {
2045        let response_deadline_started = if response_started {
2046            last_response_event_at
2047        } else {
2048            response_wait_started
2049        };
2050        let read_timeout = if response_started {
2051            match response_event_budget.checked_sub(response_deadline_started.elapsed()) {
2052                Some(remaining) if !remaining.is_zero() => remaining,
2053                _ => {
2054                    invalidate_pool_owner(pool_owner, pool_entry);
2055                    return Err(CodexError {
2056                        status: 0,
2057                        message: format!("WebSocket idle timeout after {idle_timeout_ms}ms"),
2058                        detail: None,
2059                        retry_after: None,
2060                        origin: CodexErrorOrigin::WebSocket,
2061                    });
2062                }
2063            }
2064        } else {
2065            match response_event_budget.checked_sub(response_deadline_started.elapsed()) {
2066                Some(remaining) if !remaining.is_zero() => remaining,
2067                _ => {
2068                    invalidate_pool_owner(pool_owner, pool_entry);
2069                    return Err(response_start_timeout_error(idle_timeout_ms));
2070                }
2071            }
2072        };
2073
2074        let frame = match read_ws_frame_with_keepalive(ws, read_timeout, keepalive_interval).await {
2075            WebSocketRead::Frame(frame) => frame,
2076            WebSocketRead::Timeout => {
2077                invalidate_pool_owner(pool_owner, pool_entry);
2078                return Err(if response_started {
2079                    CodexError {
2080                        status: 0,
2081                        message: format!("WebSocket idle timeout after {idle_timeout_ms}ms"),
2082                        detail: None,
2083                        retry_after: None,
2084                        origin: CodexErrorOrigin::WebSocket,
2085                    }
2086                } else {
2087                    response_start_timeout_error(idle_timeout_ms)
2088                });
2089            }
2090            WebSocketRead::KeepaliveError(error) => {
2091                invalidate_pool_owner(pool_owner, pool_entry);
2092                return Err(CodexError {
2093                    status: 0,
2094                    message: format!("WebSocket keepalive error: {error}"),
2095                    detail: Some(WEBSOCKET_KEEPALIVE_FAILURE_DETAIL.to_string()),
2096                    retry_after: None,
2097                    origin: CodexErrorOrigin::WebSocket,
2098                });
2099            }
2100        };
2101
2102        match frame {
2103            Some(Ok(Message::Text(text))) => {
2104                // Parse JSON
2105                let parsed: serde_json::Value = match serde_json::from_str(&text) {
2106                    Ok(v) => v,
2107                    Err(_) => {
2108                        if let Some(tc) = traffic {
2109                            tc.write_json_event(
2110                                "040-upstream-event",
2111                                &serde_json::json!({
2112                                    "unparseable": true,
2113                                    "data": text,
2114                                }),
2115                            );
2116                        }
2117                        // Write invalid JSON as-is
2118                        sse_body.extend_from_slice(&encode_sse(&text));
2119                        continue;
2120                    }
2121                };
2122
2123                // Convert to SSE bytes
2124                sse_body.extend_from_slice(&encode_sse(&text));
2125                if let Some(tc) = traffic {
2126                    tc.write_json_event("040-upstream-event", &parsed);
2127                }
2128
2129                if is_response_event(&parsed) {
2130                    response_started = true;
2131                    last_response_event_at = Instant::now();
2132                }
2133
2134                // Check for terminal events
2135                if is_terminal_event(&parsed) {
2136                    terminal_event = Some(WsEvent {
2137                        event_type: parsed
2138                            .get("type")
2139                            .and_then(|v| v.as_str())
2140                            .unwrap_or("unknown")
2141                            .to_string(),
2142                        payload: parsed,
2143                    });
2144                    break;
2145                }
2146            }
2147            Some(Ok(Message::Binary(_))) => {
2148                // Reject binary frames
2149                invalidate_pool_owner(pool_owner, pool_entry);
2150                return Err(CodexError {
2151                    status: 0,
2152                    message: "WebSocket binary frames not supported".to_string(),
2153                    detail: None,
2154                    retry_after: None,
2155                    origin: CodexErrorOrigin::WebSocket,
2156                });
2157            }
2158            Some(Ok(Message::Ping(data))) => {
2159                // Respond to ping automatically, continue
2160                let _ = ws.send(Message::Pong(data)).await;
2161                continue;
2162            }
2163            Some(Ok(Message::Pong(_))) => {
2164                continue;
2165            }
2166            Some(Ok(Message::Frame(_))) => {
2167                // Raw frame passthrough - continue
2168                continue;
2169            }
2170            Some(Ok(Message::Close(_))) => {
2171                // Connection closed - invalidate pool
2172                invalidate_pool_owner(pool_owner, pool_entry);
2173                break;
2174            }
2175            Some(Err(e)) => {
2176                // Stream error - invalidate pool
2177                invalidate_pool_owner(pool_owner, pool_entry);
2178                return Err(CodexError {
2179                    status: 0,
2180                    message: format!("WebSocket stream error: {e}"),
2181                    detail: None,
2182                    retry_after: None,
2183                    origin: CodexErrorOrigin::WebSocket,
2184                });
2185            }
2186            None => {
2187                // Stream ended - invalidate pool
2188                invalidate_pool_owner(pool_owner, pool_entry);
2189                break;
2190            }
2191        }
2192    }
2193
2194    Ok((sse_body, terminal_event))
2195}
2196
2197async fn stream_ws_events<S>(
2198    ws: &mut WebSocketStream<S>,
2199    idle_timeout_ms: u64,
2200    traffic: Option<Arc<TrafficCapture>>,
2201    tx: &mpsc::Sender<Result<serde_json::Value, CodexError>>,
2202) -> (bool, Option<Result<serde_json::Value, CodexError>>)
2203where
2204    S: AsyncRead + AsyncWrite + Unpin,
2205{
2206    let started_at = Instant::now();
2207    let mut sse_body: Vec<u8> = Vec::new();
2208    let response_event_budget = Duration::from_millis(idle_timeout_ms);
2209    let response_wait_started = Instant::now();
2210    let mut last_response_event_at = response_wait_started;
2211    let mut response_started = false;
2212    let mut status = 200u16;
2213    let mut reusable = false;
2214    let mut terminal_item = None;
2215
2216    loop {
2217        let response_deadline_started = if response_started {
2218            last_response_event_at
2219        } else {
2220            response_wait_started
2221        };
2222        let read_timeout =
2223            match response_event_budget.checked_sub(response_deadline_started.elapsed()) {
2224                Some(remaining) if !remaining.is_zero() => remaining,
2225                _ => {
2226                    terminal_item = Some(Err(if response_started {
2227                        CodexError {
2228                            status: 0,
2229                            message: format!("WebSocket idle timeout after {idle_timeout_ms}ms"),
2230                            detail: None,
2231                            retry_after: None,
2232                            origin: CodexErrorOrigin::WebSocket,
2233                        }
2234                    } else {
2235                        response_start_timeout_error(idle_timeout_ms)
2236                    }));
2237                    break;
2238                }
2239            };
2240
2241        let frame = tokio::select! {
2242            biased;
2243            _ = tx.closed() => break,
2244            frame = read_ws_frame_with_keepalive(
2245                ws,
2246                read_timeout,
2247                WEBSOCKET_KEEPALIVE_INTERVAL,
2248            ) => match frame {
2249                WebSocketRead::Frame(frame) => frame,
2250                WebSocketRead::Timeout => {
2251                    terminal_item = Some(Err(if response_started {
2252                        CodexError {
2253                            status: 0,
2254                            message: format!("WebSocket idle timeout after {idle_timeout_ms}ms"),
2255                            detail: None,
2256                            retry_after: None,
2257                            origin: CodexErrorOrigin::WebSocket,
2258                        }
2259                    } else {
2260                        response_start_timeout_error(idle_timeout_ms)
2261                    }));
2262                    break;
2263                }
2264                WebSocketRead::KeepaliveError(error) => {
2265                    terminal_item = Some(Err(CodexError {
2266                        status: 0,
2267                        message: format!("WebSocket keepalive error: {error}"),
2268                        detail: None,
2269                        retry_after: None,
2270                        origin: CodexErrorOrigin::WebSocket,
2271                    }));
2272                    break;
2273                }
2274            },
2275        };
2276
2277        match frame {
2278            Some(Ok(Message::Text(text))) => {
2279                let parsed: serde_json::Value = match serde_json::from_str(&text) {
2280                    Ok(value) => value,
2281                    Err(_) => {
2282                        if let Some(tc) = traffic.as_deref() {
2283                            tc.write_json_event(
2284                                "040-upstream-event",
2285                                &serde_json::json!({
2286                                    "unparseable": true,
2287                                    "data": text,
2288                                }),
2289                            );
2290                        }
2291                        sse_body.extend_from_slice(&encode_sse(&text));
2292                        continue;
2293                    }
2294                };
2295
2296                sse_body.extend_from_slice(&encode_sse(&text));
2297                if let Some(tc) = traffic.as_deref() {
2298                    tc.write_json_event("040-upstream-event", &parsed);
2299                }
2300
2301                if is_response_event(&parsed) {
2302                    response_started = true;
2303                    last_response_event_at = Instant::now();
2304                }
2305                if parsed.get("type").and_then(|value| value.as_str()) == Some("error") {
2306                    status = event_error_status(&parsed).unwrap_or(500);
2307                }
2308
2309                if is_terminal_event(&parsed) {
2310                    if is_previous_response_missing(&parsed) {
2311                        terminal_item = Some(Err(CodexError {
2312                            status: 0,
2313                            message: "Previous response not found".to_string(),
2314                            detail: Some("previous_response_not_found".to_string()),
2315                            retry_after: None,
2316                            origin: CodexErrorOrigin::WebSocket,
2317                        }));
2318                    } else {
2319                        reusable = parsed.get("type").and_then(|value| value.as_str())
2320                            == Some("response.completed");
2321                        terminal_item = Some(Ok(parsed));
2322                    }
2323                    break;
2324                }
2325
2326                if tx.send(Ok(parsed)).await.is_err() {
2327                    break;
2328                }
2329            }
2330            Some(Ok(Message::Binary(_))) => {
2331                terminal_item = Some(Err(CodexError {
2332                    status: 0,
2333                    message: "WebSocket binary frames not supported".to_string(),
2334                    detail: None,
2335                    retry_after: None,
2336                    origin: CodexErrorOrigin::WebSocket,
2337                }));
2338                break;
2339            }
2340            Some(Ok(Message::Ping(data))) => {
2341                let _ = ws.send(Message::Pong(data)).await;
2342            }
2343            Some(Ok(Message::Pong(_))) | Some(Ok(Message::Frame(_))) => {}
2344            Some(Ok(Message::Close(_))) | None => {
2345                terminal_item = Some(Err(missing_terminal_error()));
2346                break;
2347            }
2348            Some(Err(error)) => {
2349                terminal_item = Some(Err(CodexError {
2350                    status: 0,
2351                    message: format!("WebSocket stream error: {error}"),
2352                    detail: None,
2353                    retry_after: None,
2354                    origin: CodexErrorOrigin::WebSocket,
2355                }));
2356                break;
2357            }
2358        }
2359    }
2360
2361    if let Some(tc) = traffic.as_deref() {
2362        write_websocket_response_capture(tc, status, started_at.elapsed(), &sse_body);
2363    }
2364    (reusable, terminal_item)
2365}
2366
2367fn headers_to_json(headers: &HeaderMap) -> serde_json::Value {
2368    let mut out = serde_json::Map::new();
2369    for (key, value) in headers.iter() {
2370        out.insert(
2371            key.to_string(),
2372            serde_json::Value::String(value.to_str().unwrap_or("").to_string()),
2373        );
2374    }
2375    serde_json::Value::Object(out)
2376}
2377
2378fn summarize_json_request_size(body: &serde_json::Value, body_json: &str) -> serde_json::Value {
2379    serde_json::json!({
2380        "bytes": body_json.len(),
2381        "inputCount": body
2382            .get("input")
2383            .and_then(|v| v.as_array())
2384            .map(|items| items.len()),
2385        "toolCount": body
2386            .get("tools")
2387            .and_then(|v| v.as_array())
2388            .map(|items| items.len()),
2389    })
2390}
2391
2392// ---------------------------------------------------------------------------
2393// Tests
2394// ---------------------------------------------------------------------------
2395
2396#[cfg(test)]
2397mod tests {
2398    use std::sync::atomic::{AtomicBool, AtomicUsize};
2399
2400    use super::*;
2401
2402    #[test]
2403    fn provider_retry_handoff_is_attempt_local() {
2404        let (_tx_a, rx_a) = mpsc::channel(1);
2405        let (stream_a, publisher_a) = CodexWebSocketEventStream::pending(rx_a);
2406        let (_tx_b, rx_b) = mpsc::channel(1);
2407        let (_stream_b, publisher_b) = CodexWebSocketEventStream::pending(rx_b);
2408
2409        assert!(!publisher_a.is_provider_retry_handoff());
2410        assert!(!publisher_b.is_provider_retry_handoff());
2411        stream_a.mark_provider_retry_handoff();
2412        assert!(publisher_a.is_provider_retry_handoff());
2413        assert!(!publisher_b.is_provider_retry_handoff());
2414    }
2415
2416    fn main_owner(session_id: &str) -> ConversationIdentity {
2417        ConversationIdentity::Main(session_id.to_string())
2418    }
2419
2420    fn agent_owner(session_id: &str, agent_id: &str) -> ConversationIdentity {
2421        ConversationIdentity::Agent(session_id.to_string(), agent_id.to_string())
2422    }
2423
2424    fn test_continuation(
2425        owner: Option<ConversationIdentity>,
2426        turn_id: Option<u64>,
2427        previous_response_id: Option<&str>,
2428        origin_socket_id: Option<u64>,
2429    ) -> ContinuationReservation {
2430        ContinuationReservation::new(
2431            super::super::continuation::ContinuationCandidate {
2432                turn_id,
2433                previous_response_id: previous_response_id.map(str::to_string),
2434                input_delta: Some(vec![]),
2435                input_delta_count: 0,
2436                disabled_reason: None,
2437            },
2438            owner,
2439            origin_socket_id,
2440        )
2441    }
2442
2443    fn continuation_request() -> super::super::translate::request::ResponsesRequest {
2444        super::super::translate::request::ResponsesRequest {
2445            model: "gpt-5.6-sol".to_string(),
2446            instructions: None,
2447            input: vec![],
2448            tools: None,
2449            tool_choice: None,
2450            store: false,
2451            stream: true,
2452            parallel_tool_calls: true,
2453            include: None,
2454            client_metadata: None,
2455            service_tier: None,
2456            prompt_cache_key: None,
2457            text: super::super::translate::request::ResponsesText {
2458                verbosity: None,
2459                format: None,
2460            },
2461            reasoning: None,
2462        }
2463    }
2464
2465    fn test_websocket_client() -> reqwest::Client {
2466        reqwest::Client::builder()
2467            .http1_only()
2468            .redirect(reqwest::redirect::Policy::none())
2469            .no_proxy()
2470            .build()
2471            .unwrap()
2472    }
2473
2474    fn handshake_error(status: u16) -> CodexError {
2475        CodexError {
2476            status,
2477            message: format!("handshake failed with status {status}"),
2478            detail: None,
2479            retry_after: None,
2480            origin: CodexErrorOrigin::WebSocketHandshake,
2481        }
2482    }
2483
2484    fn origin_forbidden_error() -> ConnectAttemptError {
2485        ConnectAttemptError::Origin(handshake_error(http::StatusCode::FORBIDDEN.as_u16()))
2486    }
2487
2488    struct DropProbeIo {
2489        probe: Option<(Arc<AtomicBool>, Arc<AtomicBool>)>,
2490    }
2491
2492    impl Drop for DropProbeIo {
2493        fn drop(&mut self) {
2494            if let Some((dropped, pool_was_unlocked)) = &self.probe {
2495                pool_was_unlocked.store(WS_POOL.try_lock().is_ok(), Ordering::SeqCst);
2496                dropped.store(true, Ordering::SeqCst);
2497            }
2498        }
2499    }
2500
2501    impl AsyncRead for DropProbeIo {
2502        fn poll_read(
2503            self: Pin<&mut Self>,
2504            _cx: &mut Context<'_>,
2505            _buffer: &mut ReadBuf<'_>,
2506        ) -> Poll<std::io::Result<()>> {
2507            Poll::Pending
2508        }
2509    }
2510
2511    impl AsyncWrite for DropProbeIo {
2512        fn poll_write(
2513            self: Pin<&mut Self>,
2514            _cx: &mut Context<'_>,
2515            buffer: &[u8],
2516        ) -> Poll<std::io::Result<usize>> {
2517            Poll::Ready(Ok(buffer.len()))
2518        }
2519
2520        fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
2521            Poll::Ready(Ok(()))
2522        }
2523
2524        fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
2525            Poll::Ready(Ok(()))
2526        }
2527    }
2528
2529    struct FailingWriteIo;
2530
2531    impl AsyncRead for FailingWriteIo {
2532        fn poll_read(
2533            self: Pin<&mut Self>,
2534            _cx: &mut Context<'_>,
2535            _buffer: &mut ReadBuf<'_>,
2536        ) -> Poll<std::io::Result<()>> {
2537            Poll::Pending
2538        }
2539    }
2540
2541    impl AsyncWrite for FailingWriteIo {
2542        fn poll_write(
2543            self: Pin<&mut Self>,
2544            _cx: &mut Context<'_>,
2545            _buffer: &[u8],
2546        ) -> Poll<std::io::Result<usize>> {
2547            Poll::Ready(Err(std::io::Error::new(
2548                std::io::ErrorKind::BrokenPipe,
2549                "test write failed",
2550            )))
2551        }
2552
2553        fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
2554            Poll::Ready(Ok(()))
2555        }
2556
2557        fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
2558            Poll::Ready(Ok(()))
2559        }
2560    }
2561
2562    async fn raw_test_stream(
2563        probe: Option<(Arc<AtomicBool>, Arc<AtomicBool>)>,
2564    ) -> CodexWebSocketStream {
2565        let io: BoxedWebSocketIo = Box::new(DropProbeIo { probe });
2566        WebSocketStream::from_raw_socket(io, Role::Client, None).await
2567    }
2568
2569    fn shared_pool_entry(ws: &Arc<AsyncMutex<CodexWebSocketStream>>) -> Arc<PoolEntry> {
2570        Arc::new(PoolEntry {
2571            ws: ws.clone(),
2572            socket_id: next_monotonic_nonzero(&NEXT_SOCKET_ID, "WebSocket ID"),
2573            created_at: now_ms(),
2574            last_activity: AtomicU64::new(next_pool_activity()),
2575        })
2576    }
2577
2578    #[tokio::test]
2579    async fn connect_cleanup_uses_threshold_target_lru_and_skips_leases() {
2580        let _pool_test_guard = lock_codex_websocket_pool_for_tests().await;
2581        clear_codex_websocket_pool_for_tests();
2582        let ws = Arc::new(AsyncMutex::new(raw_test_stream(None).await));
2583        {
2584            let mut guard = WS_POOL.lock().unwrap();
2585            for index in 0..POOL_CONNECT_CLEANUP_THRESHOLD {
2586                guard.insert(
2587                    main_owner(&format!("entry-{index:02}")),
2588                    shared_pool_entry(&ws),
2589                );
2590            }
2591        }
2592
2593        cleanup_pool_before_connect();
2594        assert_eq!(
2595            WS_POOL.lock().unwrap().len(),
2596            POOL_CONNECT_CLEANUP_THRESHOLD
2597        );
2598
2599        WS_POOL
2600            .lock()
2601            .unwrap()
2602            .insert(main_owner("entry-50"), shared_pool_entry(&ws));
2603        WS_POOL
2604            .lock()
2605            .unwrap()
2606            .get(&main_owner("entry-00"))
2607            .unwrap()
2608            .touch();
2609        let leased = WS_POOL
2610            .lock()
2611            .unwrap()
2612            .get(&main_owner("entry-01"))
2613            .unwrap()
2614            .clone();
2615
2616        cleanup_pool_before_connect();
2617
2618        let guard = WS_POOL.lock().unwrap();
2619        assert_eq!(guard.len(), POOL_CONNECT_CLEANUP_TARGET);
2620        assert!(guard.contains_key(&main_owner("entry-00")));
2621        assert!(guard.contains_key(&main_owner("entry-01")));
2622        for index in 2..=12 {
2623            assert!(!guard.contains_key(&main_owner(&format!("entry-{index:02}"))));
2624        }
2625        assert!(guard.contains_key(&main_owner("entry-13")));
2626        assert!(guard.contains_key(&main_owner("entry-50")));
2627        drop(guard);
2628        drop(leased);
2629        clear_codex_websocket_pool_for_tests();
2630    }
2631
2632    #[tokio::test]
2633    async fn connect_cleanup_drops_final_socket_owners_after_unlocking_pool() {
2634        let _pool_test_guard = lock_codex_websocket_pool_for_tests().await;
2635        clear_codex_websocket_pool_for_tests();
2636        let dropped = Arc::new(AtomicBool::new(false));
2637        let pool_was_unlocked = Arc::new(AtomicBool::new(false));
2638        let shared_ws = Arc::new(AsyncMutex::new(raw_test_stream(None).await));
2639        let probe_stream =
2640            raw_test_stream(Some((dropped.clone(), pool_was_unlocked.clone()))).await;
2641        {
2642            let mut guard = WS_POOL.lock().unwrap();
2643            guard.insert(
2644                main_owner("entry-00"),
2645                Arc::new(PoolEntry::new(probe_stream)),
2646            );
2647            for index in 1..=POOL_CONNECT_CLEANUP_THRESHOLD {
2648                guard.insert(
2649                    main_owner(&format!("entry-{index:02}")),
2650                    shared_pool_entry(&shared_ws),
2651                );
2652            }
2653        }
2654
2655        cleanup_pool_before_connect();
2656
2657        assert!(dropped.load(Ordering::SeqCst));
2658        assert!(pool_was_unlocked.load(Ordering::SeqCst));
2659        clear_codex_websocket_pool_for_tests();
2660    }
2661
2662    #[tokio::test(start_paused = true)]
2663    async fn connect_gate_spaces_starts_without_serializing_handshakes() {
2664        let gate = Arc::new(WebSocketConnectGate::new(Duration::from_secs(1)));
2665        let preflight_started = Arc::new(tokio::sync::Notify::new());
2666        let release_first = Arc::new(tokio::sync::Notify::new());
2667        let (started_tx, mut started_rx) = mpsc::unbounded_channel();
2668
2669        let first_gate = gate.clone();
2670        let first_preflight = preflight_started.clone();
2671        let first_release = release_first.clone();
2672        let first_tx = started_tx.clone();
2673        let first = tokio::spawn(async move {
2674            first_gate
2675                .wait_to_start(async move {
2676                    first_preflight.notify_one();
2677                    tokio::time::sleep(Duration::from_secs(2)).await;
2678                })
2679                .await;
2680            first_tx.send(tokio::time::Instant::now()).unwrap();
2681            first_release.notified().await;
2682        });
2683        preflight_started.notified().await;
2684
2685        let second_gate = gate.clone();
2686        let second = tokio::spawn(async move {
2687            second_gate.wait_to_start(async {}).await;
2688            started_tx.send(tokio::time::Instant::now()).unwrap();
2689        });
2690        let first_started = started_rx.recv().await.unwrap();
2691        let second_started = started_rx.recv().await.unwrap();
2692
2693        assert!(second_started.duration_since(first_started) >= Duration::from_secs(1));
2694        assert!(!first.is_finished());
2695        second.await.unwrap();
2696        release_first.notify_one();
2697        first.await.unwrap();
2698    }
2699
2700    #[tokio::test(start_paused = true)]
2701    async fn connect_gate_spaces_after_non_403_failure_without_retrying() {
2702        let gate = WebSocketConnectGate::new(Duration::from_secs(1));
2703        let attempts = AtomicUsize::new(0);
2704        let starts = Mutex::new(Vec::new());
2705
2706        let error = connect_with_policy(
2707            &gate,
2708            Duration::from_secs(10),
2709            Duration::from_secs(3),
2710            || async {},
2711            || {
2712                attempts.fetch_add(1, Ordering::SeqCst);
2713                starts.lock().unwrap().push(tokio::time::Instant::now());
2714                async {
2715                    Err::<(), ConnectAttemptError>(ConnectAttemptError::Origin(handshake_error(
2716                        http::StatusCode::BAD_GATEWAY.as_u16(),
2717                    )))
2718                }
2719            },
2720        )
2721        .await
2722        .unwrap_err();
2723        assert_eq!(error.status, http::StatusCode::BAD_GATEWAY.as_u16());
2724        assert_eq!(attempts.load(Ordering::SeqCst), 1);
2725
2726        connect_with_policy(
2727            &gate,
2728            Duration::from_secs(10),
2729            Duration::from_secs(3),
2730            || async {},
2731            || {
2732                starts.lock().unwrap().push(tokio::time::Instant::now());
2733                async { Ok::<(), ConnectAttemptError>(()) }
2734            },
2735        )
2736        .await
2737        .unwrap();
2738
2739        let starts = starts.lock().unwrap();
2740        assert!(starts[1].duration_since(starts[0]) >= Duration::from_secs(1));
2741    }
2742
2743    #[tokio::test(start_paused = true)]
2744    async fn origin_403_waits_for_cooldown_retries_once_and_returns_second_error() {
2745        let gate = WebSocketConnectGate::new(Duration::from_secs(1));
2746        let attempts = AtomicUsize::new(0);
2747        let starts = Mutex::new(Vec::new());
2748
2749        let error = connect_with_policy(
2750            &gate,
2751            Duration::from_secs(10),
2752            Duration::from_secs(3),
2753            || async {},
2754            || {
2755                let attempt = attempts.fetch_add(1, Ordering::SeqCst);
2756                starts.lock().unwrap().push(tokio::time::Instant::now());
2757                async move {
2758                    if attempt == 0 {
2759                        Err::<(), ConnectAttemptError>(origin_forbidden_error())
2760                    } else {
2761                        Err::<(), ConnectAttemptError>(ConnectAttemptError::Origin(CodexError {
2762                            status: http::StatusCode::FORBIDDEN.as_u16(),
2763                            message: "second forbidden".to_string(),
2764                            detail: Some("second-detail".to_string()),
2765                            retry_after: Some("7".to_string()),
2766                            origin: CodexErrorOrigin::WebSocketHandshake,
2767                        }))
2768                    }
2769                }
2770            },
2771        )
2772        .await
2773        .unwrap_err();
2774
2775        assert_eq!(error.status, http::StatusCode::FORBIDDEN.as_u16());
2776        assert_eq!(error.message, "second forbidden");
2777        assert_eq!(error.detail.as_deref(), Some("second-detail"));
2778        assert_eq!(error.retry_after.as_deref(), Some("7"));
2779        assert_eq!(error.origin, CodexErrorOrigin::WebSocketHandshake);
2780        assert_eq!(attempts.load(Ordering::SeqCst), 2);
2781        let starts = starts.lock().unwrap();
2782        assert!(starts[1].duration_since(starts[0]) >= Duration::from_secs(3));
2783    }
2784
2785    #[tokio::test(start_paused = true)]
2786    async fn proxy_connect_403_is_not_retried() {
2787        let gate = WebSocketConnectGate::new(Duration::ZERO);
2788        let attempts = AtomicUsize::new(0);
2789
2790        let error = connect_with_policy(
2791            &gate,
2792            Duration::from_secs(10),
2793            Duration::from_secs(3),
2794            || async {},
2795            || {
2796                attempts.fetch_add(1, Ordering::SeqCst);
2797                async {
2798                    Err::<(), ConnectAttemptError>(ConnectAttemptError::ProxyTunnel(
2799                        handshake_error(http::StatusCode::FORBIDDEN.as_u16()),
2800                    ))
2801                }
2802            },
2803        )
2804        .await
2805        .unwrap_err();
2806
2807        assert_eq!(error.status, http::StatusCode::FORBIDDEN.as_u16());
2808        assert_eq!(attempts.load(Ordering::SeqCst), 1);
2809    }
2810
2811    #[tokio::test]
2812    async fn explicit_proxy_connect_403_has_private_tunnel_provenance() {
2813        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
2814        let proxy_addr = listener.local_addr().unwrap();
2815        let proxy = tokio::spawn(async move {
2816            let (mut socket, _) = listener.accept().await.unwrap();
2817            let mut request = Vec::new();
2818            let mut buffer = [0_u8; 1024];
2819            while connect_response_header_end(&request).is_none() {
2820                let read = socket.read(&mut buffer).await.unwrap();
2821                assert!(read > 0);
2822                request.extend_from_slice(&buffer[..read]);
2823            }
2824            socket
2825                .write_all(b"HTTP/1.1 403 Forbidden\r\nContent-Length: 0\r\n\r\n")
2826                .await
2827                .unwrap();
2828        });
2829        let proxy_url = format!("http://{proxy_addr}");
2830        let proxy_config = WebSocketProxyConfig::new(None, Some(&proxy_url), None, None);
2831        let client = test_websocket_client();
2832
2833        let error = match connect_once(
2834            &client,
2835            &proxy_config,
2836            "wss://codex.invalid/backend-api/codex/responses",
2837            &HeaderMap::new(),
2838        )
2839        .await
2840        {
2841            Ok(_) => panic!("proxy CONNECT rejection should fail"),
2842            Err(error) => error,
2843        };
2844
2845        match error {
2846            ConnectAttemptError::ProxyTunnel(error) => {
2847                assert_eq!(error.status, http::StatusCode::FORBIDDEN.as_u16());
2848                assert_eq!(
2849                    error.detail.as_deref(),
2850                    Some(WEBSOCKET_PROXY_TUNNEL_REJECTED_DETAIL)
2851                );
2852                assert_eq!(error.origin, CodexErrorOrigin::WebSocketHandshake);
2853            }
2854            ConnectAttemptError::Origin(_) => panic!("proxy rejection was classified as origin"),
2855        }
2856        proxy.await.unwrap();
2857    }
2858
2859    #[tokio::test]
2860    async fn origin_403_with_spoofed_proxy_detail_still_retries_once() {
2861        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
2862        let addr = listener.local_addr().unwrap();
2863        let server = tokio::spawn(async move {
2864            for _ in 0..2 {
2865                let (mut socket, _) = listener.accept().await.unwrap();
2866                let mut request = Vec::new();
2867                let mut buffer = [0_u8; 1024];
2868                while connect_response_header_end(&request).is_none() {
2869                    let read = socket.read(&mut buffer).await.unwrap();
2870                    assert!(read > 0);
2871                    request.extend_from_slice(&buffer[..read]);
2872                }
2873                let body = format!(
2874                    "{{\"error\":{{\"message\":\"{WEBSOCKET_PROXY_TUNNEL_REJECTED_DETAIL}\"}}}}"
2875                );
2876                let response = format!(
2877                    "HTTP/1.1 403 Forbidden\r\nContent-Type: application/json\r\nContent-Length: {}\r\n\r\n{body}",
2878                    body.len()
2879                );
2880                socket.write_all(response.as_bytes()).await.unwrap();
2881            }
2882        });
2883        let client = test_websocket_client();
2884        let proxy_config = WebSocketProxyConfig::direct();
2885        let url = format!("ws://{addr}/backend-api/codex/responses");
2886        let headers = HeaderMap::new();
2887        let gate = WebSocketConnectGate::new(Duration::ZERO);
2888
2889        let error = match connect_with_policy(
2890            &gate,
2891            Duration::from_secs(1),
2892            Duration::ZERO,
2893            || async {},
2894            || connect_once(&client, &proxy_config, &url, &headers),
2895        )
2896        .await
2897        {
2898            Ok(_) => panic!("spoofed origin rejection should fail after one retry"),
2899            Err(error) => error,
2900        };
2901
2902        assert_eq!(error.status, http::StatusCode::FORBIDDEN.as_u16());
2903        assert_eq!(
2904            error.detail.as_deref(),
2905            Some(WEBSOCKET_PROXY_TUNNEL_REJECTED_DETAIL)
2906        );
2907        server.await.unwrap();
2908    }
2909
2910    #[tokio::test(start_paused = true)]
2911    async fn timeout_applies_to_each_network_attempt_not_waits_or_cooldown() {
2912        let gate = WebSocketConnectGate::new(Duration::from_secs(1));
2913        let attempts = AtomicUsize::new(0);
2914        let started_at = tokio::time::Instant::now();
2915
2916        let error = connect_with_policy(
2917            &gate,
2918            Duration::from_secs(2),
2919            Duration::from_secs(3),
2920            || async {},
2921            || {
2922                let attempt = attempts.fetch_add(1, Ordering::SeqCst);
2923                async move {
2924                    if attempt == 0 {
2925                        tokio::time::sleep(Duration::from_secs(1)).await;
2926                        Err::<(), ConnectAttemptError>(origin_forbidden_error())
2927                    } else {
2928                        tokio::time::sleep(Duration::from_secs(3)).await;
2929                        Ok(())
2930                    }
2931                }
2932            },
2933        )
2934        .await
2935        .unwrap_err();
2936
2937        assert_eq!(attempts.load(Ordering::SeqCst), 2);
2938        assert_eq!(error.message, "WebSocket connect timeout after 2000ms");
2939        assert_eq!(started_at.elapsed(), Duration::from_secs(6));
2940    }
2941
2942    #[test]
2943    fn websocket_tls_configuration_is_shared() {
2944        let first = websocket_tls_config();
2945        let second = websocket_tls_config();
2946        assert!(Arc::ptr_eq(&first, &second));
2947    }
2948
2949    #[test]
2950    fn event_error_status_requires_error_event_and_checks_numeric_fallbacks() {
2951        assert_eq!(
2952            event_error_status(&serde_json::json!({
2953                "type": "response.failed",
2954                "status": "failed",
2955                "status_code": 401
2956            })),
2957            Some(401)
2958        );
2959        assert_eq!(
2960            event_error_status(&serde_json::json!({
2961                "type": "response.completed",
2962                "status_code": 401
2963            })),
2964            None
2965        );
2966        assert_eq!(
2967            event_error_status(&serde_json::json!({
2968                "type": "error",
2969                "error": {"status": 401}
2970            })),
2971            Some(401)
2972        );
2973    }
2974
2975    #[test]
2976    fn websocket_url_conversion() {
2977        assert_eq!(
2978            to_websocket_url("https://example.test/codex").unwrap(),
2979            "wss://example.test/codex"
2980        );
2981        assert_eq!(
2982            to_websocket_url("http://example.test/codex").unwrap(),
2983            "ws://example.test/codex"
2984        );
2985        assert_eq!(
2986            to_websocket_url("wss://example.test/codex").unwrap(),
2987            "wss://example.test/codex"
2988        );
2989        assert!(to_websocket_url("ftp://example.test/codex").is_err());
2990    }
2991
2992    #[test]
2993    fn websocket_upgrade_url_preserves_authority_path_and_query() {
2994        assert_eq!(
2995            to_http_upgrade_url("wss://chatgpt.com/backend-api/codex/responses?mode=live").unwrap(),
2996            "https://chatgpt.com/backend-api/codex/responses?mode=live"
2997        );
2998        assert_eq!(
2999            to_http_upgrade_url("ws://127.0.0.1:4141/backend-api/codex/responses").unwrap(),
3000            "http://127.0.0.1:4141/backend-api/codex/responses"
3001        );
3002        assert_eq!(
3003            to_http_upgrade_url("ws://[::1]:4141/path").unwrap(),
3004            "http://[::1]:4141/path"
3005        );
3006    }
3007
3008    #[test]
3009    fn validates_tokenized_websocket_upgrade_headers() {
3010        let key = "dGhlIHNhbXBsZSBub25jZQ==";
3011        let mut headers = HeaderMap::new();
3012        headers.insert(http::header::UPGRADE, "h2c, WebSocket".parse().unwrap());
3013        headers.insert(
3014            http::header::CONNECTION,
3015            "keep-alive, Upgrade".parse().unwrap(),
3016        );
3017        headers.insert(
3018            http::header::SEC_WEBSOCKET_ACCEPT,
3019            derive_accept_key(key.as_bytes()).parse().unwrap(),
3020        );
3021        headers.insert(
3022            http::header::SEC_WEBSOCKET_PROTOCOL,
3023            "responses".parse().unwrap(),
3024        );
3025
3026        validate_websocket_upgrade(
3027            http::Version::HTTP_11,
3028            &headers,
3029            key,
3030            &["responses".to_string()],
3031        )
3032        .unwrap();
3033    }
3034
3035    #[test]
3036    fn rejects_invalid_websocket_upgrade_headers() {
3037        let key = "dGhlIHNhbXBsZSBub25jZQ==";
3038        let valid = || {
3039            let mut headers = HeaderMap::new();
3040            headers.insert(http::header::UPGRADE, "websocket".parse().unwrap());
3041            headers.insert(http::header::CONNECTION, "Upgrade".parse().unwrap());
3042            headers.insert(
3043                http::header::SEC_WEBSOCKET_ACCEPT,
3044                derive_accept_key(key.as_bytes()).parse().unwrap(),
3045            );
3046            headers
3047        };
3048
3049        let mut missing_upgrade = valid();
3050        missing_upgrade.remove(http::header::UPGRADE);
3051        assert!(
3052            validate_websocket_upgrade(http::Version::HTTP_11, &missing_upgrade, key, &[]).is_err()
3053        );
3054
3055        let mut missing_connection = valid();
3056        missing_connection.remove(http::header::CONNECTION);
3057        assert!(
3058            validate_websocket_upgrade(http::Version::HTTP_11, &missing_connection, key, &[])
3059                .is_err()
3060        );
3061
3062        let mut wrong_accept = valid();
3063        wrong_accept.insert(http::header::SEC_WEBSOCKET_ACCEPT, "wrong".parse().unwrap());
3064        assert!(
3065            validate_websocket_upgrade(http::Version::HTTP_11, &wrong_accept, key, &[]).is_err()
3066        );
3067
3068        let unsolicited_extension = {
3069            let mut headers = valid();
3070            headers.insert(
3071                http::header::SEC_WEBSOCKET_EXTENSIONS,
3072                "permessage-deflate".parse().unwrap(),
3073            );
3074            headers
3075        };
3076        assert!(
3077            validate_websocket_upgrade(http::Version::HTTP_11, &unsolicited_extension, key, &[],)
3078                .is_err()
3079        );
3080
3081        assert!(validate_websocket_upgrade(http::Version::HTTP_10, &valid(), key, &[]).is_err());
3082
3083        let mut unsolicited_protocol = valid();
3084        unsolicited_protocol.insert(
3085            http::header::SEC_WEBSOCKET_PROTOCOL,
3086            "unexpected".parse().unwrap(),
3087        );
3088        assert!(
3089            validate_websocket_upgrade(http::Version::HTTP_11, &unsolicited_protocol, key, &[])
3090                .is_err()
3091        );
3092    }
3093
3094    #[test]
3095    fn websocket_headers_rewrite_beta() {
3096        let mut headers = http::HeaderMap::new();
3097        headers.insert("openai-beta", "responses=experimental".parse().unwrap());
3098        headers.insert("content-length", "10".parse().unwrap());
3099        headers.insert("authorization", "Bearer tok".parse().unwrap());
3100        headers.insert(
3101            http::header::PROXY_AUTHORIZATION,
3102            "Basic dXNlcjpwYXNz".parse().unwrap(),
3103        );
3104        let ws = codex_websocket_headers(&headers);
3105        assert_eq!(ws.get("openai-beta").unwrap(), WEBSOCKET_PROTOCOL_HEADER);
3106        assert!(!ws.contains_key("content-length"));
3107        assert!(!ws.contains_key(http::header::PROXY_AUTHORIZATION));
3108        assert_eq!(ws.get("authorization").unwrap(), "Bearer tok");
3109    }
3110
3111    #[test]
3112    fn websocket_headers_strips_accept() {
3113        let mut headers = http::HeaderMap::new();
3114        headers.insert(http::header::ACCEPT, "text/event-stream".parse().unwrap());
3115        let ws = codex_websocket_headers(&headers);
3116        assert!(!ws.contains_key(http::header::ACCEPT.as_str()));
3117    }
3118
3119    #[test]
3120    fn websocket_headers_adds_sec_key() {
3121        let headers = http::HeaderMap::new();
3122        let ws = codex_websocket_headers(&headers);
3123        assert!(ws.contains_key("sec-websocket-key"));
3124    }
3125
3126    #[test]
3127    fn encode_sse_single_line() {
3128        let result = encode_sse(r#"{"type":"test","data":"hello"}"#);
3129        let expected = b"data: {\"type\":\"test\",\"data\":\"hello\"}\n\n";
3130        assert_eq!(result, expected);
3131    }
3132
3133    #[test]
3134    fn encode_sse_multi_line() {
3135        let result = encode_sse("line1\nline2");
3136        assert_eq!(
3137            String::from_utf8(result).unwrap(),
3138            "data: line1\ndata: line2\n\n"
3139        );
3140    }
3141
3142    #[test]
3143    fn is_terminal_event_detection() {
3144        let completed = serde_json::json!({"type": "response.completed"});
3145        assert!(is_terminal_event(&completed));
3146
3147        let delta = serde_json::json!({"type": "response.output_text.delta"});
3148        assert!(!is_terminal_event(&delta));
3149
3150        let error = serde_json::json!({"type": "error", "error": {"message": "fail"}});
3151        assert!(is_terminal_event(&error));
3152    }
3153
3154    #[test]
3155    fn is_response_event_detection() {
3156        let rate_limits = serde_json::json!({"type": "codex.rate_limits"});
3157        assert!(!is_response_event(&rate_limits));
3158
3159        let output = serde_json::json!({"type": "response.output_text.delta"});
3160        assert!(is_response_event(&output));
3161
3162        let error = serde_json::json!({"type": "error", "error": {"message": "fail"}});
3163        assert!(is_response_event(&error));
3164    }
3165
3166    #[test]
3167    fn is_previous_response_missing_detection() {
3168        let by_code = serde_json::json!({
3169            "type": "error",
3170            "error": {"code": "previous_response_not_found", "message": "not found"}
3171        });
3172        assert!(is_previous_response_missing(&by_code));
3173
3174        let by_msg = serde_json::json!({
3175            "type": "error",
3176            "error": {"message": "The previous response was not found"}
3177        });
3178        assert!(is_previous_response_missing(&by_msg));
3179
3180        let nested = serde_json::json!({
3181            "type": "response.failed",
3182            "response": {
3183                "error": {
3184                    "code": "previous_response_not_found",
3185                    "message": "Previous response not found"
3186                }
3187            }
3188        });
3189        assert!(is_previous_response_missing(&nested));
3190
3191        let unrelated = serde_json::json!({"type": "error", "error": {"message": "rate limited"}});
3192        assert!(!is_previous_response_missing(&unrelated));
3193    }
3194
3195    #[test]
3196    fn websocket_metadata_does_not_serialize_typed_owner() {
3197        let temp = tempfile::tempdir().unwrap();
3198        let traffic = crate::traffic::test_capture(temp.path().join("traffic"));
3199        let owner = agent_owner("session-secret", "agent-secret");
3200        let reservation = test_continuation(Some(owner), None, None, None);
3201
3202        write_websocket_metadata_capture(
3203            &traffic,
3204            "wss://example.invalid/responses",
3205            Some(&reservation),
3206            false,
3207        );
3208
3209        let artifact = std::fs::read_dir(traffic.root())
3210            .unwrap()
3211            .next()
3212            .unwrap()
3213            .unwrap()
3214            .path();
3215        let captured = std::fs::read_to_string(artifact).unwrap();
3216        assert!(captured.contains("poolingEnabled"));
3217        assert!(!captured.contains("poolKey"));
3218        assert!(!captured.contains("session-secret"));
3219        assert!(!captured.contains("agent-secret"));
3220    }
3221
3222    #[tokio::test]
3223    async fn pool_checkout_is_exclusive_and_removal_is_identity_safe() {
3224        let _pool_test_guard = lock_codex_websocket_pool_for_tests().await;
3225        clear_codex_websocket_pool_for_tests();
3226        let first = Arc::new(PoolEntry::new(create_dummy_stream_async().await));
3227        let owner = main_owner("exclusive");
3228        pool_insert(owner.clone(), first.clone());
3229        assert!(Arc::ptr_eq(
3230            &WS_POOL.lock().unwrap().remove(&owner).unwrap(),
3231            &first
3232        ));
3233        assert!(WS_POOL.lock().unwrap().remove(&owner).is_none());
3234
3235        let replacement = Arc::new(PoolEntry::new(create_dummy_stream_async().await));
3236        pool_insert(owner.clone(), replacement.clone());
3237        pool_remove_entry(&owner, &first);
3238        let reservation = test_continuation(Some(owner.clone()), None, None, None);
3239        invalidate_codex_websocket_pool_socket(&reservation, Some(first.socket_id));
3240        assert!(Arc::ptr_eq(
3241            WS_POOL.lock().unwrap().get(&owner).unwrap(),
3242            &replacement
3243        ));
3244        clear_codex_websocket_pool_for_tests();
3245    }
3246
3247    #[tokio::test]
3248    async fn pooled_validation_uses_unique_nonce_and_rejects_stale_pong() {
3249        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
3250        let addr = listener.local_addr().unwrap();
3251        let (queue_stale_tx, queue_stale_rx) = tokio::sync::oneshot::channel();
3252        let (stale_queued_tx, stale_queued_rx) = tokio::sync::oneshot::channel();
3253        let (release_peer_tx, release_peer_rx) = tokio::sync::oneshot::channel();
3254        let peer = tokio::spawn(async move {
3255            let (socket, _) = listener.accept().await.unwrap();
3256            let mut websocket = tokio_tungstenite::accept_async(socket).await.unwrap();
3257            let prior_nonce = match websocket.next().await {
3258                Some(Ok(Message::Ping(payload))) => payload,
3259                other => panic!("unexpected first validation frame: {other:?}"),
3260            };
3261            assert_eq!(prior_nonce.len(), std::mem::size_of::<u64>());
3262            assert_ne!(prior_nonce, 0_u64.to_be_bytes());
3263            websocket
3264                .send(Message::Pong(prior_nonce.clone()))
3265                .await
3266                .unwrap();
3267
3268            queue_stale_rx.await.unwrap();
3269            websocket
3270                .send(Message::Pong(prior_nonce.clone()))
3271                .await
3272                .unwrap();
3273            stale_queued_tx.send(()).unwrap();
3274
3275            let current_nonce = match websocket.next().await {
3276                Some(Ok(Message::Ping(payload))) => payload,
3277                other => panic!("unexpected second validation frame: {other:?}"),
3278            };
3279            assert_ne!(current_nonce, prior_nonce);
3280            assert_ne!(current_nonce, 0_u64.to_be_bytes());
3281            release_peer_rx.await.unwrap();
3282        });
3283        let (mut websocket, _) = tokio_tungstenite::connect_async(format!("ws://{addr}/"))
3284            .await
3285            .unwrap();
3286
3287        validate_pooled_websocket(&mut websocket, 1_000)
3288            .await
3289            .unwrap();
3290        queue_stale_tx.send(()).unwrap();
3291        stale_queued_rx.await.unwrap();
3292        let error = validate_pooled_websocket(&mut websocket, 100)
3293            .await
3294            .unwrap_err();
3295        assert_eq!(error, "unexpected Pong during pooled validation");
3296
3297        release_peer_tx.send(()).unwrap();
3298        peer.await.unwrap();
3299    }
3300
3301    #[tokio::test]
3302    async fn pool_entries_receive_monotonic_nonzero_socket_ids() {
3303        let first = PoolEntry::new(raw_test_stream(None).await);
3304        let second = PoolEntry::new(raw_test_stream(None).await);
3305
3306        assert_ne!(first.socket_id, 0);
3307        assert!(second.socket_id > first.socket_id);
3308    }
3309
3310    #[tokio::test]
3311    async fn continuation_rejects_and_preserves_same_owner_replacement() {
3312        let _registry_guard =
3313            super::super::continuation::lock_continuation_registry_for_async_tests().await;
3314        let _pool_test_guard = lock_codex_websocket_pool_for_tests().await;
3315        let owner = agent_owner("replacement-session", "replacement-agent");
3316        super::super::continuation::clear_continuation_for_owner(Some(&owner));
3317        invalidate_codex_websocket_pool_owner(&owner);
3318        let request = continuation_request();
3319        let reserved = super::super::continuation::continuation_candidate_for_owner(
3320            Some(&owner),
3321            &request,
3322            true,
3323        );
3324        let replacement = Arc::new(PoolEntry::new(raw_test_stream(None).await));
3325        pool_insert(owner.clone(), replacement.clone());
3326        let continuation = test_continuation(
3327            Some(owner.clone()),
3328            reserved.turn_id(),
3329            Some("resp_origin"),
3330            Some(replacement.socket_id.checked_add(1).unwrap()),
3331        );
3332
3333        let error = match prepare_codex_websocket(
3334            &test_websocket_client(),
3335            &WebSocketProxyConfig::direct(),
3336            "ws://127.0.0.1:9/responses",
3337            &HeaderMap::new(),
3338            None,
3339            Some(&continuation),
3340            50,
3341            50,
3342        )
3343        .await
3344        {
3345            Ok(_) => panic!("replacement socket must not satisfy continuation provenance"),
3346            Err(error) => error,
3347        };
3348
3349        assert_eq!(
3350            error.detail.as_deref(),
3351            Some(WEBSOCKET_CONTINUATION_SOCKET_MISSING_DETAIL)
3352        );
3353        assert!(!error.message.contains("replacement-session"));
3354        assert!(!error.message.contains("replacement-agent"));
3355        assert!(Arc::ptr_eq(
3356            WS_POOL.lock().unwrap().get(&owner).unwrap(),
3357            &replacement
3358        ));
3359
3360        invalidate_codex_websocket_pool_owner(&owner);
3361        super::super::continuation::abort_continuation_for_owner(&reserved);
3362    }
3363
3364    #[tokio::test]
3365    async fn dead_exact_origin_removes_only_that_arc_and_preserves_replacement() {
3366        let _registry_guard =
3367            super::super::continuation::lock_continuation_registry_for_async_tests().await;
3368        let _pool_test_guard = lock_codex_websocket_pool_for_tests().await;
3369        let owner = main_owner("dead-origin-session");
3370        super::super::continuation::clear_continuation_for_owner(Some(&owner));
3371        invalidate_codex_websocket_pool_owner(&owner);
3372        let request = continuation_request();
3373        let reserved = super::super::continuation::continuation_candidate_for_owner(
3374            Some(&owner),
3375            &request,
3376            true,
3377        );
3378        let exact = Arc::new(PoolEntry::new(raw_test_stream(None).await));
3379        let replacement = Arc::new(PoolEntry::new(raw_test_stream(None).await));
3380        pool_insert(owner.clone(), exact.clone());
3381        let continuation = test_continuation(
3382            Some(owner.clone()),
3383            reserved.turn_id(),
3384            Some("resp_exact"),
3385            Some(exact.socket_id),
3386        );
3387
3388        let replacement_owner = owner.clone();
3389        let replacement_for_task = replacement.clone();
3390        let mut insert_replacement = tokio::spawn(async move {
3391            tokio::time::timeout(Duration::from_secs(1), async move {
3392                loop {
3393                    if !WS_POOL.lock().unwrap().contains_key(&replacement_owner) {
3394                        pool_insert(replacement_owner, replacement_for_task);
3395                        return;
3396                    }
3397                    tokio::task::yield_now().await;
3398                }
3399            })
3400            .await
3401        });
3402        let error = match prepare_codex_websocket(
3403            &test_websocket_client(),
3404            &WebSocketProxyConfig::direct(),
3405            "ws://127.0.0.1:9/responses",
3406            &HeaderMap::new(),
3407            None,
3408            Some(&continuation),
3409            50,
3410            50,
3411        )
3412        .await
3413        {
3414            Ok(_) => panic!("dead continuation origin must be rejected before request send"),
3415            Err(error) => error,
3416        };
3417        match tokio::time::timeout(Duration::from_secs(2), &mut insert_replacement).await {
3418            Ok(Ok(Ok(()))) => {}
3419            Ok(Ok(Err(_))) => panic!(
3420                "replacement polling timed out; currently pooled socket ID: {:?}",
3421                pooled_socket_id_for_tests(&owner)
3422            ),
3423            Ok(Err(error)) => panic!(
3424                "replacement polling task failed ({error}); currently pooled socket ID: {:?}",
3425                pooled_socket_id_for_tests(&owner)
3426            ),
3427            Err(_) => {
3428                insert_replacement.abort();
3429                let abort_result = insert_replacement.await;
3430                panic!(
3431                    "replacement polling join timed out ({abort_result:?}); currently pooled socket ID: {:?}",
3432                    pooled_socket_id_for_tests(&owner)
3433                );
3434            }
3435        }
3436
3437        assert_eq!(
3438            error.detail.as_deref(),
3439            Some(WEBSOCKET_CONTINUATION_SOCKET_MISSING_DETAIL)
3440        );
3441        assert!(Arc::ptr_eq(
3442            WS_POOL.lock().unwrap().get(&owner).unwrap(),
3443            &replacement
3444        ));
3445        assert!(!Arc::ptr_eq(
3446            WS_POOL.lock().unwrap().get(&owner).unwrap(),
3447            &exact
3448        ));
3449
3450        invalidate_codex_websocket_pool_owner(&owner);
3451        super::super::continuation::abort_continuation_for_owner(&reserved);
3452    }
3453
3454    #[tokio::test]
3455    async fn completed_terminal_is_published_after_origin_returns_to_pool() {
3456        let _registry_guard =
3457            super::super::continuation::lock_continuation_registry_for_async_tests().await;
3458        let _pool_test_guard = lock_codex_websocket_pool_for_tests().await;
3459        let owner = agent_owner("terminal-order-session", "terminal-order-agent");
3460        super::super::continuation::clear_continuation_for_owner(Some(&owner));
3461        invalidate_codex_websocket_pool_owner(&owner);
3462        let request = continuation_request();
3463        let continuation = super::super::continuation::continuation_candidate_for_owner(
3464            Some(&owner),
3465            &request,
3466            true,
3467        );
3468        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
3469        let addr = listener.local_addr().unwrap();
3470        let (release_response_tx, release_response_rx) = tokio::sync::oneshot::channel();
3471        let server = tokio::spawn(async move {
3472            let (socket, _) = listener.accept().await.unwrap();
3473            let mut websocket = tokio_tungstenite::accept_async(socket).await.unwrap();
3474            while let Some(Ok(message)) = websocket.next().await {
3475                match message {
3476                    Message::Ping(payload) => {
3477                        websocket.send(Message::Pong(payload)).await.unwrap();
3478                    }
3479                    Message::Text(_) => {
3480                        release_response_rx.await.unwrap();
3481                        websocket
3482                            .send(Message::Text(
3483                                serde_json::json!({
3484                                    "type": "response.completed",
3485                                    "response": {"id": "resp_terminal_order"}
3486                                })
3487                                .to_string(),
3488                            ))
3489                            .await
3490                            .unwrap();
3491                        return;
3492                    }
3493                    _ => {}
3494                }
3495            }
3496        });
3497        let context = RequestContext {
3498            req_id: "terminal-order-request".to_string(),
3499            session_id: Some("header-session-is-not-pool-owner".to_string()),
3500            session_seq: None,
3501            provider: "codex".to_string(),
3502            traffic: None,
3503            monitor: None,
3504            passthrough: None,
3505        };
3506        let mut events = codex_websocket_event_stream(
3507            &test_websocket_client(),
3508            &WebSocketProxyConfig::direct(),
3509            &format!("http://{addr}/responses"),
3510            &HeaderMap::new(),
3511            &serde_json::json!({"type":"response.create","input":[]}),
3512            &context,
3513            None,
3514            1_000,
3515            1_000,
3516            Some(&continuation),
3517        )
3518        .await
3519        .unwrap();
3520        assert_eq!(events.socket_id(), None);
3521        release_response_tx.send(()).unwrap();
3522        let terminal = events.recv().await.unwrap().unwrap();
3523
3524        assert_eq!(
3525            terminal.get("type").and_then(serde_json::Value::as_str),
3526            Some("response.completed")
3527        );
3528        let pooled = WS_POOL
3529            .lock()
3530            .unwrap()
3531            .get(&owner)
3532            .cloned()
3533            .expect("origin must be reusable before terminal publication");
3534        assert_eq!(events.socket_id(), Some(pooled.socket_id));
3535        server.await.unwrap();
3536
3537        invalidate_codex_websocket_pool_owner(&owner);
3538        super::super::continuation::abort_continuation_for_owner(&continuation);
3539    }
3540
3541    #[tokio::test]
3542    async fn dropped_receiver_removes_completed_origin_after_terminal_publication_fails() {
3543        let _registry_guard =
3544            super::super::continuation::lock_continuation_registry_for_async_tests().await;
3545        let _pool_test_guard = lock_codex_websocket_pool_for_tests().await;
3546        let owner = main_owner("dropped-terminal-receiver-session");
3547        super::super::continuation::clear_continuation_for_owner(Some(&owner));
3548        invalidate_codex_websocket_pool_owner(&owner);
3549        let request = continuation_request();
3550        let continuation = super::super::continuation::continuation_candidate_for_owner(
3551            Some(&owner),
3552            &request,
3553            true,
3554        );
3555        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
3556        let addr = listener.local_addr().unwrap();
3557        let (request_seen_tx, request_seen_rx) = tokio::sync::oneshot::channel();
3558        let (release_response_tx, release_response_rx) = tokio::sync::oneshot::channel();
3559        let server = tokio::spawn(async move {
3560            let (socket, _) = listener.accept().await.unwrap();
3561            let mut websocket = tokio_tungstenite::accept_async(socket).await.unwrap();
3562            while let Some(Ok(message)) = websocket.next().await {
3563                match message {
3564                    Message::Ping(payload) => {
3565                        websocket.send(Message::Pong(payload)).await.unwrap();
3566                    }
3567                    Message::Text(_) => {
3568                        request_seen_tx.send(()).unwrap();
3569                        release_response_rx.await.unwrap();
3570                        websocket
3571                            .send(Message::Text(
3572                                serde_json::json!({
3573                                    "type": "response.completed",
3574                                    "response": {"id": "resp_dropped_receiver"}
3575                                })
3576                                .to_string(),
3577                            ))
3578                            .await
3579                            .unwrap();
3580                        return;
3581                    }
3582                    _ => {}
3583                }
3584            }
3585        });
3586        let headers = HeaderMap::new();
3587        let body = serde_json::json!({"type":"response.create","input":[]});
3588        let ready = prepare_codex_websocket(
3589            &test_websocket_client(),
3590            &WebSocketProxyConfig::direct(),
3591            &format!("http://{addr}/responses"),
3592            &headers,
3593            None,
3594            Some(&continuation),
3595            1_000,
3596            1_000,
3597        )
3598        .await
3599        .unwrap();
3600        let exact = ready.entry.clone();
3601        let events = start_codex_websocket_events(
3602            ready,
3603            &body,
3604            serde_json::to_string(&body).unwrap(),
3605            &headers,
3606            Some(&continuation),
3607        );
3608
3609        request_seen_rx.await.unwrap();
3610        drop(events);
3611        release_response_tx.send(()).unwrap();
3612        server.await.unwrap();
3613        tokio::time::timeout(Duration::from_secs(1), async {
3614            while Arc::strong_count(&exact) != 1 {
3615                tokio::task::yield_now().await;
3616            }
3617        })
3618        .await
3619        .expect("failed terminal publication must release the exact completed pool entry");
3620
3621        assert!(
3622            WS_POOL
3623                .lock()
3624                .unwrap()
3625                .get(&owner)
3626                .is_none_or(|pooled| !Arc::ptr_eq(pooled, &exact)),
3627            "the exact completed socket must not remain pooled"
3628        );
3629        super::super::continuation::abort_continuation_for_owner(&continuation);
3630    }
3631
3632    #[tokio::test]
3633    async fn pool_invalidation() {
3634        let _pool_test_guard = lock_codex_websocket_pool_for_tests().await;
3635        clear_codex_websocket_pool_for_tests();
3636        let first_stream = create_dummy_stream_async().await;
3637        let second_stream = create_dummy_stream_async().await;
3638        let first_owner = agent_owner("test-session", "first-agent");
3639        let sibling_owner = agent_owner("test-session", "sibling-agent");
3640        {
3641            let mut guard = WS_POOL.lock().unwrap();
3642            guard.insert(first_owner.clone(), Arc::new(PoolEntry::new(first_stream)));
3643            guard.insert(
3644                sibling_owner.clone(),
3645                Arc::new(PoolEntry::new(second_stream)),
3646            );
3647        }
3648        assert!(WS_POOL.lock().unwrap().contains_key(&first_owner));
3649        assert!(WS_POOL.lock().unwrap().contains_key(&sibling_owner));
3650
3651        invalidate_codex_websocket_pool_owner(&first_owner);
3652        assert!(!WS_POOL.lock().unwrap().contains_key(&first_owner));
3653        assert!(WS_POOL.lock().unwrap().contains_key(&sibling_owner));
3654        clear_codex_websocket_pool_for_tests();
3655    }
3656
3657    #[tokio::test]
3658    #[allow(deprecated)]
3659    async fn session_key_invalidation_removes_main_and_agents_only_for_exact_session() {
3660        let _pool_test_guard = lock_codex_websocket_pool_for_tests().await;
3661        clear_codex_websocket_pool_for_tests();
3662        let session_id = "session-key";
3663        let main = main_owner(session_id);
3664        let first_agent = agent_owner(session_id, "first-agent");
3665        let second_agent = agent_owner(session_id, "second-agent");
3666        let other_session = main_owner("session-key-other");
3667        for owner in [
3668            main.clone(),
3669            first_agent.clone(),
3670            second_agent.clone(),
3671            other_session.clone(),
3672        ] {
3673            pool_insert(
3674                owner,
3675                Arc::new(PoolEntry::new(create_dummy_stream_async().await)),
3676            );
3677        }
3678
3679        invalidate_codex_websocket_pool_key(session_id);
3680
3681        let pool = WS_POOL.lock().unwrap();
3682        assert!(!pool.contains_key(&main));
3683        assert!(!pool.contains_key(&first_agent));
3684        assert!(!pool.contains_key(&second_agent));
3685        assert!(pool.contains_key(&other_session));
3686        drop(pool);
3687        clear_codex_websocket_pool_for_tests();
3688    }
3689
3690    #[tokio::test]
3691    async fn missing_owner_or_turn_cannot_mutate_pool_state() {
3692        let _registry_guard =
3693            super::super::continuation::lock_continuation_registry_for_async_tests().await;
3694        let _pool_test_guard = lock_codex_websocket_pool_for_tests().await;
3695        let owner = main_owner("missing-owner-turn-pool");
3696        super::super::continuation::clear_continuation_for_owner(Some(&owner));
3697        clear_codex_websocket_pool_for_tests();
3698        let request = continuation_request();
3699        let current = super::super::continuation::continuation_candidate_for_owner(
3700            Some(&owner),
3701            &request,
3702            true,
3703        );
3704        let pooled = Arc::new(PoolEntry::new(create_dummy_stream_async().await));
3705        pool_insert(owner.clone(), pooled.clone());
3706        let missing_owner = test_continuation(None, current.turn_id(), None, None);
3707        let missing_turn = test_continuation(Some(owner.clone()), None, None, None);
3708
3709        assert!(pool_take_for_turn(&missing_owner).is_none());
3710        assert!(pool_take_for_turn(&missing_turn).is_none());
3711        assert!(!pool_insert_for_turn(
3712            &missing_owner,
3713            Arc::new(PoolEntry::new(create_dummy_stream_async().await)),
3714        ));
3715        assert!(!pool_insert_for_turn(
3716            &missing_turn,
3717            Arc::new(PoolEntry::new(create_dummy_stream_async().await)),
3718        ));
3719        invalidate_codex_websocket_pool_turn_for_owner(&owner, None);
3720        invalidate_codex_websocket_pool_socket(&missing_owner, Some(pooled.socket_id));
3721
3722        assert!(Arc::ptr_eq(
3723            WS_POOL.lock().unwrap().get(&owner).unwrap(),
3724            &pooled
3725        ));
3726        clear_codex_websocket_pool_for_tests();
3727        super::super::continuation::abort_continuation_for_owner(&current);
3728    }
3729
3730    #[tokio::test]
3731    async fn websocket_connect_401_is_pre_request_handshake_error() {
3732        let _pool_test_guard = lock_codex_websocket_pool_for_tests().await;
3733        use tokio::io::{AsyncReadExt, AsyncWriteExt};
3734        use tokio::net::TcpListener;
3735
3736        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3737        let addr = listener.local_addr().unwrap();
3738        tokio::spawn(async move {
3739            let (mut socket, _) = listener.accept().await.unwrap();
3740            let mut buf = [0_u8; 2048];
3741            let _ = socket.read(&mut buf).await;
3742            socket
3743                .write_all(b"HTTP/1.1 401 Unauthorized\r\nContent-Length: 13\r\n\r\npolicy denied")
3744                .await
3745                .unwrap();
3746        });
3747
3748        let client = test_websocket_client();
3749        let err = match connect_with_timeout(
3750            &client,
3751            &WebSocketProxyConfig::direct(),
3752            &format!("ws://{addr}/backend-api/codex/responses"),
3753            &HeaderMap::new(),
3754            1_000,
3755        )
3756        .await
3757        {
3758            Ok(_) => panic!("expected unauthorized websocket handshake to fail"),
3759            Err(err) => err,
3760        };
3761
3762        assert_eq!(err.status, 401);
3763        assert_eq!(err.detail.as_deref(), Some(GENERIC_HANDSHAKE_ERROR_DETAIL));
3764        assert_eq!(err.origin, CodexErrorOrigin::WebSocketHandshake);
3765    }
3766
3767    #[tokio::test]
3768    async fn websocket_connect_502_preserves_retry_metadata() {
3769        let _pool_test_guard = lock_codex_websocket_pool_for_tests().await;
3770        use tokio::io::{AsyncReadExt, AsyncWriteExt};
3771        use tokio::net::TcpListener;
3772
3773        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3774        let addr = listener.local_addr().unwrap();
3775        tokio::spawn(async move {
3776            let (mut socket, _) = listener.accept().await.unwrap();
3777            let mut buf = [0_u8; 2048];
3778            let _ = socket.read(&mut buf).await;
3779            socket
3780                .write_all(
3781                    b"HTTP/1.1 502 Bad Gateway\r\nRetry-After: 3\r\nContent-Length: 0\r\n\r\n",
3782                )
3783                .await
3784                .unwrap();
3785        });
3786
3787        let client = test_websocket_client();
3788        let err = match connect_with_timeout(
3789            &client,
3790            &WebSocketProxyConfig::direct(),
3791            &format!("ws://{addr}/backend-api/codex/responses"),
3792            &HeaderMap::new(),
3793            1_000,
3794        )
3795        .await
3796        {
3797            Ok(_) => panic!("expected websocket handshake to fail"),
3798            Err(err) => err,
3799        };
3800
3801        assert_eq!(err.status, 502);
3802        assert_eq!(err.detail.as_deref(), Some(GENERIC_HANDSHAKE_ERROR_DETAIL));
3803        assert_eq!(err.retry_after.as_deref(), Some("3"));
3804        assert_eq!(err.origin, CodexErrorOrigin::WebSocketHandshake);
3805    }
3806
3807    #[tokio::test]
3808    async fn websocket_connects_through_explicit_http_proxy() {
3809        let _pool_test_guard = lock_codex_websocket_pool_for_tests().await;
3810        use tokio::io::{AsyncReadExt, AsyncWriteExt};
3811        use tokio::net::TcpListener;
3812
3813        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3814        let proxy_addr = listener.local_addr().unwrap();
3815        let (captured_tx, captured_rx) = tokio::sync::oneshot::channel();
3816        let proxy = tokio::spawn(async move {
3817            let (mut stream, _) = listener.accept().await.unwrap();
3818            let mut request = Vec::new();
3819            let mut buffer = [0_u8; 1024];
3820            while !request.windows(4).any(|window| window == b"\r\n\r\n") {
3821                let read = stream.read(&mut buffer).await.unwrap();
3822                assert!(read > 0);
3823                request.extend_from_slice(&buffer[..read]);
3824            }
3825            let request_text = String::from_utf8(request).unwrap();
3826            let key = request_text
3827                .lines()
3828                .find_map(|line| {
3829                    let (name, value) = line.split_once(':')?;
3830                    name.eq_ignore_ascii_case("sec-websocket-key")
3831                        .then(|| value.trim().to_string())
3832                })
3833                .unwrap();
3834            let _ = captured_tx.send(request_text);
3835            let response = format!(
3836                "HTTP/1.1 101 Switching Protocols\r\nConnection: Upgrade\r\nUpgrade: websocket\r\nSec-WebSocket-Accept: {}\r\n\r\n",
3837                derive_accept_key(key.as_bytes())
3838            );
3839            stream.write_all(response.as_bytes()).await.unwrap();
3840
3841            let mut websocket = WebSocketStream::from_raw_socket(stream, Role::Server, None).await;
3842            assert_eq!(
3843                websocket.next().await.unwrap().unwrap(),
3844                Message::Text("hello".to_string())
3845            );
3846            websocket
3847                .send(Message::Text("proxy-ok".to_string()))
3848                .await
3849                .unwrap();
3850        });
3851
3852        let client = reqwest::Client::builder()
3853            .http1_only()
3854            .redirect(reqwest::redirect::Policy::none())
3855            .proxy(
3856                reqwest::Proxy::http(format!("http://proxy-user:proxy-pass@{proxy_addr}")).unwrap(),
3857            )
3858            .build()
3859            .unwrap();
3860        let mut websocket = connect_with_timeout(
3861            &client,
3862            &WebSocketProxyConfig::direct(),
3863            "ws://codex.invalid/backend-api/codex/responses",
3864            &HeaderMap::new(),
3865            2_000,
3866        )
3867        .await
3868        .unwrap();
3869        websocket
3870            .send(Message::Text("hello".to_string()))
3871            .await
3872            .unwrap();
3873        assert_eq!(
3874            websocket.next().await.unwrap().unwrap(),
3875            Message::Text("proxy-ok".to_string())
3876        );
3877
3878        let captured = captured_rx.await.unwrap();
3879        assert!(
3880            captured.starts_with("GET http://codex.invalid/backend-api/codex/responses HTTP/1.1")
3881        );
3882        assert!(
3883            captured
3884                .to_ascii_lowercase()
3885                .contains("proxy-authorization: basic ")
3886        );
3887        proxy.await.unwrap();
3888    }
3889
3890    #[tokio::test]
3891    async fn connect_tunnel_accepts_fragmented_non_200_success() {
3892        let (client, mut proxy) = tokio::io::duplex(4096);
3893        let proxy_task = tokio::spawn(async move {
3894            let mut request = Vec::new();
3895            let mut buffer = [0_u8; 256];
3896            while !request.windows(4).any(|window| window == b"\r\n\r\n") {
3897                let read = proxy.read(&mut buffer).await.unwrap();
3898                assert!(read > 0);
3899                request.extend_from_slice(&buffer[..read]);
3900            }
3901            assert!(request.starts_with(b"CONNECT codex.invalid:4443 HTTP/1.1\r\n"));
3902            proxy.write_all(b"HTTP/1.").await.unwrap();
3903            tokio::task::yield_now().await;
3904            proxy
3905                .write_all(b"1 204 No Content\r\nX-Proxy: ok\r\n\r\nprefixed")
3906                .await
3907                .unwrap();
3908        });
3909
3910        let stream: BoxedWebSocketIo = Box::new(client);
3911        let mut stream = establish_connect_tunnel(stream, "codex.invalid:4443", None)
3912            .await
3913            .unwrap();
3914        let mut prefix = [0_u8; 8];
3915        stream.read_exact(&mut prefix).await.unwrap();
3916        assert_eq!(&prefix, b"prefixed");
3917        proxy_task.await.unwrap();
3918    }
3919
3920    #[tokio::test]
3921    async fn connect_tunnel_classifies_fragmented_proxy_authentication() {
3922        let (client, mut proxy) = tokio::io::duplex(4096);
3923        tokio::spawn(async move {
3924            let mut request = [0_u8; 512];
3925            let _ = proxy.read(&mut request).await.unwrap();
3926            proxy.write_all(b"HTTP/1.1 4").await.unwrap();
3927            tokio::task::yield_now().await;
3928            proxy
3929                .write_all(b"07 Proxy Authentication Required\r\n\r\n")
3930                .await
3931                .unwrap();
3932        });
3933
3934        let stream: BoxedWebSocketIo = Box::new(client);
3935        let error = match establish_connect_tunnel(stream, "codex.invalid:4443", None).await {
3936            Ok(_) => panic!("proxy authentication should be rejected"),
3937            Err(error) => error,
3938        };
3939        assert_eq!(
3940            error.status,
3941            http::StatusCode::PROXY_AUTHENTICATION_REQUIRED.as_u16()
3942        );
3943        assert_eq!(error.message, "WebSocket proxy authentication failed");
3944    }
3945
3946    #[test]
3947    fn websocket_proxy_routing_honors_no_proxy_and_leaves_socks_to_reqwest() {
3948        let http_proxy = "http://proxy.example:8080";
3949        let config = WebSocketProxyConfig::new(None, Some(http_proxy), None, None);
3950        assert!(config.uses_proxy_for("wss://codex.invalid/responses"));
3951        assert!(
3952            config
3953                .http_connect_route("wss://codex.invalid/responses")
3954                .unwrap()
3955                .is_some()
3956        );
3957
3958        let bypass = WebSocketProxyConfig::new(None, Some(http_proxy), None, Some("codex.invalid"));
3959        assert!(!bypass.uses_proxy_for("wss://codex.invalid/responses"));
3960        assert!(
3961            bypass
3962                .http_connect_route("wss://codex.invalid/responses")
3963                .unwrap()
3964                .is_none()
3965        );
3966
3967        let socks =
3968            WebSocketProxyConfig::new(None, Some("socks5h://proxy.example:1080"), None, None);
3969        assert!(socks.uses_proxy_for("wss://codex.invalid/responses"));
3970        assert!(
3971            socks
3972                .http_connect_route("wss://codex.invalid/responses")
3973                .unwrap()
3974                .is_none()
3975        );
3976    }
3977
3978    #[tokio::test]
3979    async fn websocket_wss_uses_http_connect_without_leaking_proxy_credentials() {
3980        let _pool_test_guard = lock_codex_websocket_pool_for_tests().await;
3981        use tokio::io::{AsyncReadExt, AsyncWriteExt};
3982        use tokio::net::TcpListener;
3983
3984        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3985        let proxy_addr = listener.local_addr().unwrap();
3986        let (captured_tx, captured_rx) = tokio::sync::oneshot::channel();
3987        let proxy = tokio::spawn(async move {
3988            let (mut stream, _) = listener.accept().await.unwrap();
3989            let mut request = Vec::new();
3990            let mut buffer = [0_u8; 1024];
3991            while !request.windows(4).any(|window| window == b"\r\n\r\n") {
3992                let read = stream.read(&mut buffer).await.unwrap();
3993                assert!(read > 0);
3994                request.extend_from_slice(&buffer[..read]);
3995            }
3996            let _ = captured_tx.send(String::from_utf8(request).unwrap());
3997            stream
3998                .write_all(
3999                    b"HTTP/1.1 407 Proxy Authentication Required\r\nContent-Length: 0\r\n\r\n",
4000                )
4001                .await
4002                .unwrap();
4003        });
4004
4005        let proxy_url = format!("http://secret-user:secret-pass@{proxy_addr}");
4006        let client = reqwest::Client::builder()
4007            .http1_only()
4008            .redirect(reqwest::redirect::Policy::none())
4009            .proxy(reqwest::Proxy::https(&proxy_url).unwrap())
4010            .build()
4011            .unwrap();
4012        let proxy_config = WebSocketProxyConfig::new(None, Some(&proxy_url), None, None);
4013        let error = match connect_with_timeout(
4014            &client,
4015            &proxy_config,
4016            "wss://codex.invalid:4443/backend-api/codex/responses",
4017            &HeaderMap::new(),
4018            2_000,
4019        )
4020        .await
4021        {
4022            Ok(_) => panic!("proxy rejection should fail the WebSocket connection"),
4023            Err(error) => error,
4024        };
4025
4026        let captured = captured_rx.await.unwrap();
4027        assert!(captured.starts_with("CONNECT codex.invalid:4443 HTTP/1.1"));
4028        assert!(
4029            captured
4030                .to_ascii_lowercase()
4031                .contains("proxy-authorization: basic ")
4032        );
4033        assert!(!error.message.contains("secret-user"));
4034        assert!(!error.message.contains("secret-pass"));
4035        assert!(
4036            !error
4037                .detail
4038                .as_deref()
4039                .unwrap_or_default()
4040                .contains("secret")
4041        );
4042        proxy.await.unwrap();
4043    }
4044
4045    #[tokio::test]
4046    async fn binary_frame_invalidates_pool_owner() {
4047        let _pool_test_guard = lock_codex_websocket_pool_for_tests().await;
4048        clear_codex_websocket_pool_for_tests();
4049        let pooled_stream = create_dummy_stream_async().await;
4050        let owner = agent_owner("binary-session", "binary-agent");
4051        {
4052            let mut guard = WS_POOL.lock().unwrap();
4053            guard.insert(owner.clone(), Arc::new(PoolEntry::new(pooled_stream)));
4054        }
4055
4056        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
4057        let addr = listener.local_addr().unwrap();
4058        tokio::spawn(async move {
4059            let (stream, _) = listener.accept().await.unwrap();
4060            let mut ws = tokio_tungstenite::accept_async(stream).await.unwrap();
4061            ws.send(Message::Binary(vec![1, 2, 3])).await.unwrap();
4062        });
4063
4064        let (mut ws, _) = tokio_tungstenite::connect_async(format!("ws://{addr}/"))
4065            .await
4066            .unwrap();
4067        let err = match collect_ws_events(&mut ws, 1_000, Some(&owner), None, None).await {
4068            Ok(_) => panic!("expected binary frame to fail"),
4069            Err(err) => err,
4070        };
4071
4072        assert!(err.message.contains("binary frames"));
4073        assert!(!WS_POOL.lock().unwrap().contains_key(&owner));
4074    }
4075
4076    #[tokio::test]
4077    async fn response_start_timeout_ignores_rate_limits_and_pings() {
4078        let _pool_test_guard = lock_codex_websocket_pool_for_tests().await;
4079        clear_codex_websocket_pool_for_tests();
4080        let pooled_stream = create_dummy_stream_async().await;
4081        let owner = main_owner("start-timeout-session");
4082        {
4083            let mut guard = WS_POOL.lock().unwrap();
4084            guard.insert(owner.clone(), Arc::new(PoolEntry::new(pooled_stream)));
4085        }
4086
4087        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
4088        let addr = listener.local_addr().unwrap();
4089        tokio::spawn(async move {
4090            let (stream, _) = listener.accept().await.unwrap();
4091            let mut ws = tokio_tungstenite::accept_async(stream).await.unwrap();
4092            ws.send(Message::Text(
4093                r#"{"type":"codex.rate_limits","rate_limits":{"allowed":true}}"#.into(),
4094            ))
4095            .await
4096            .unwrap();
4097            loop {
4098                if ws.send(Message::Ping(Vec::new())).await.is_err() {
4099                    break;
4100                }
4101                tokio::time::sleep(Duration::from_millis(10)).await;
4102            }
4103        });
4104
4105        let (mut ws, _) = tokio_tungstenite::connect_async(format!("ws://{addr}/"))
4106            .await
4107            .unwrap();
4108        let err = match collect_ws_events(&mut ws, 50, Some(&owner), None, None).await {
4109            Ok(_) => panic!("expected response start timeout"),
4110            Err(err) => err,
4111        };
4112
4113        assert_eq!(
4114            err.detail.as_deref(),
4115            Some(WEBSOCKET_RESPONSE_START_TIMEOUT_DETAIL)
4116        );
4117        assert!(!WS_POOL.lock().unwrap().contains_key(&owner));
4118    }
4119
4120    #[tokio::test]
4121    async fn response_idle_timeout_ignores_pings_after_response_event() {
4122        let _pool_test_guard = lock_codex_websocket_pool_for_tests().await;
4123        clear_codex_websocket_pool_for_tests();
4124        let pooled_stream = create_dummy_stream_async().await;
4125        let owner = main_owner("response-idle-session");
4126        {
4127            let mut guard = WS_POOL.lock().unwrap();
4128            guard.insert(owner.clone(), Arc::new(PoolEntry::new(pooled_stream)));
4129        }
4130
4131        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
4132        let addr = listener.local_addr().unwrap();
4133        tokio::spawn(async move {
4134            let (stream, _) = listener.accept().await.unwrap();
4135            let mut ws = tokio_tungstenite::accept_async(stream).await.unwrap();
4136            ws.send(Message::Text(
4137                r#"{"type":"response.output_item.added","output_index":0,"item":{"type":"message"}}"#
4138                    .into(),
4139            ))
4140            .await
4141            .unwrap();
4142            loop {
4143                if ws.send(Message::Ping(Vec::new())).await.is_err() {
4144                    break;
4145                }
4146                tokio::time::sleep(Duration::from_millis(10)).await;
4147            }
4148        });
4149
4150        let (mut ws, _) = tokio_tungstenite::connect_async(format!("ws://{addr}/"))
4151            .await
4152            .unwrap();
4153        let err = match collect_ws_events(&mut ws, 50, Some(&owner), None, None).await {
4154            Ok(_) => panic!("expected response idle timeout"),
4155            Err(err) => err,
4156        };
4157
4158        assert!(err.message.contains("idle timeout"));
4159        assert_eq!(err.detail, None);
4160        assert!(!WS_POOL.lock().unwrap().contains_key(&owner));
4161    }
4162
4163    #[tokio::test(start_paused = true)]
4164    async fn silent_response_wait_sends_keepalive_ping() {
4165        let (client_io, server_io) = tokio::io::duplex(1024);
4166        let mut client = WebSocketStream::from_raw_socket(client_io, Role::Client, None).await;
4167        let mut server = WebSocketStream::from_raw_socket(server_io, Role::Server, None).await;
4168        let server_task = tokio::spawn(async move {
4169            let frame = server.next().await.unwrap().unwrap();
4170            assert!(matches!(frame, Message::Ping(_)));
4171            server
4172                .send(Message::Text(r#"{"type":"response.created"}"#.into()))
4173                .await
4174                .unwrap();
4175        });
4176
4177        let frame = read_ws_frame_with_keepalive(
4178            &mut client,
4179            Duration::from_secs(60),
4180            WEBSOCKET_KEEPALIVE_INTERVAL,
4181        )
4182        .await;
4183
4184        assert!(matches!(
4185            frame,
4186            WebSocketRead::Frame(Some(Ok(Message::Text(text))))
4187                if text.contains("response.created")
4188        ));
4189        server_task.await.unwrap();
4190    }
4191
4192    #[tokio::test(start_paused = true)]
4193    async fn failed_keepalive_write_returns_stable_error() {
4194        let mut client = WebSocketStream::from_raw_socket(FailingWriteIo, Role::Client, None).await;
4195
4196        let error = match collect_ws_events_with_keepalive_interval(
4197            &mut client,
4198            60_000,
4199            None,
4200            None,
4201            None,
4202            WEBSOCKET_KEEPALIVE_INTERVAL,
4203        )
4204        .await
4205        {
4206            Ok(_) => panic!("expected keepalive failure"),
4207            Err(error) => error,
4208        };
4209
4210        assert_eq!(
4211            error.detail.as_deref(),
4212            Some(WEBSOCKET_KEEPALIVE_FAILURE_DETAIL)
4213        );
4214        assert!(error.message.contains("keepalive error"));
4215    }
4216
4217    #[tokio::test]
4218    async fn keepalive_pongs_do_not_extend_response_start_timeout() {
4219        let (client_io, server_io) = tokio::io::duplex(1024);
4220        let mut client = WebSocketStream::from_raw_socket(client_io, Role::Client, None).await;
4221        let mut server = WebSocketStream::from_raw_socket(server_io, Role::Server, None).await;
4222        let ping_count = Arc::new(AtomicUsize::new(0));
4223        let server_ping_count = ping_count.clone();
4224        let server_task = tokio::spawn(async move {
4225            while let Some(Ok(frame)) = server.next().await {
4226                if let Message::Ping(payload) = frame {
4227                    server_ping_count.fetch_add(1, Ordering::SeqCst);
4228                    if server.send(Message::Pong(payload)).await.is_err() {
4229                        break;
4230                    }
4231                }
4232            }
4233        });
4234
4235        let error = match collect_ws_events_with_keepalive_interval(
4236            &mut client,
4237            50,
4238            None,
4239            None,
4240            None,
4241            Duration::from_millis(10),
4242        )
4243        .await
4244        {
4245            Ok(_) => panic!("expected response start timeout"),
4246            Err(error) => error,
4247        };
4248
4249        assert_eq!(
4250            error.detail.as_deref(),
4251            Some(WEBSOCKET_RESPONSE_START_TIMEOUT_DETAIL)
4252        );
4253        assert!(ping_count.load(Ordering::SeqCst) >= 2);
4254        server_task.abort();
4255    }
4256
4257    async fn create_dummy_stream_async() -> CodexWebSocketStream {
4258        use tokio::net::TcpListener;
4259
4260        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4261        let addr = listener.local_addr().unwrap();
4262        tokio::spawn(async move {
4263            let (socket, _) = listener.accept().await.unwrap();
4264            let _ = tokio_tungstenite::accept_async(socket).await;
4265            futures_util::future::pending::<()>().await;
4266        });
4267        let url = format!("ws://{addr}/");
4268        let client = test_websocket_client();
4269        connect_with_timeout(
4270            &client,
4271            &WebSocketProxyConfig::direct(),
4272            &url,
4273            &HeaderMap::new(),
4274            1_000,
4275        )
4276        .await
4277        .unwrap()
4278    }
4279
4280    #[test]
4281    fn handshake_details_are_structured_sanitized_and_bounded() {
4282        let message = format!("safe\n{}", "x".repeat(MAX_HANDSHAKE_ERROR_DETAIL_BYTES * 2));
4283        let body = serde_json::to_vec(&serde_json::json!({
4284            "error": { "message": message }
4285        }))
4286        .unwrap();
4287        let detail = handshake_error_detail(Some(&body));
4288        assert!(!detail.contains('\n'));
4289        assert!(detail.len() <= MAX_HANDSHAKE_ERROR_DETAIL_BYTES);
4290        assert!(detail.starts_with("safe"));
4291    }
4292
4293    #[test]
4294    fn handshake_details_reject_unstructured_and_binary_bodies() {
4295        for body in [
4296            b"<html>denied</html>".to_vec(),
4297            vec![0xff, 0xfe],
4298            b"{".to_vec(),
4299        ] {
4300            assert_eq!(
4301                handshake_error_detail(Some(&body)),
4302                GENERIC_HANDSHAKE_ERROR_DETAIL
4303            );
4304        }
4305    }
4306}