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
36pub 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
60const 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#[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
310struct 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 if guard.len() >= MAX_POOL_ENTRIES
540 && let Some(oldest_owner) = guard.keys().next().cloned()
541 {
542 guard.remove(&oldest_owner);
543 }
544 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
596pub 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" => { }
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
639pub 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 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 ws.insert("openai-beta", WEBSOCKET_PROTOCOL_HEADER.parse().unwrap());
663 if !ws.contains_key("sec-websocket-key") {
665 ws.insert("sec-websocket-key", generate_key().parse().unwrap());
666 }
667 ws
668}
669
670fn 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
685pub(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 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#[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
1244const 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
1947enum 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 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 sse_body.extend_from_slice(&encode_sse(&text));
2119 continue;
2120 }
2121 };
2122
2123 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 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 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 let _ = ws.send(Message::Pong(data)).await;
2161 continue;
2162 }
2163 Some(Ok(Message::Pong(_))) => {
2164 continue;
2165 }
2166 Some(Ok(Message::Frame(_))) => {
2167 continue;
2169 }
2170 Some(Ok(Message::Close(_))) => {
2171 invalidate_pool_owner(pool_owner, pool_entry);
2173 break;
2174 }
2175 Some(Err(e)) => {
2176 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 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#[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(¤t);
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}