diff --git a/CHANGELOG.md b/CHANGELOG.md index f9626311..16bb50c8 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,6 +5,10 @@ description: Release notes for claude-code-proxy. ## Unreleased +- Codex WebSocket connections honor `HTTP_PROXY`, `HTTPS_PROXY`, `ALL_PROXY`, + and `NO_PROXY`, restoring standard non-TUN HTTP proxy support for both plain + WebSocket Upgrade requests and WSS CONNECT tunnels. + ## v0.1.26 (2026-07-28) - Standard OpenAI clients can use Codex through the optional diff --git a/Cargo.lock b/Cargo.lock index 5ba4c988..dd7ec54a 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -333,6 +333,7 @@ dependencies = [ "hostname", "http", "http-body-util", + "hyper-util", "jiff", "once_cell", "predicates", @@ -341,6 +342,8 @@ dependencies = [ "ratatui", "regex-lite", "reqwest", + "rustls", + "rustls-native-certs", "serde", "serde_json", "sha2", @@ -348,12 +351,14 @@ dependencies = [ "thiserror 2.0.18", "time", "tokio", + "tokio-rustls", "tokio-tungstenite", "tower", "tracing", "tracing-subscriber", "url", "uuid", + "webpki-roots", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index bc92838d..dc565a29 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -26,7 +26,7 @@ tracing = "0.1" tracing-subscriber = { version = "0.3", features = ["env-filter", "fmt", "json"] } uuid = { version = "1", features = ["v4", "serde"] } time = { version = "0.3", features = ["formatting", "parsing"] } -reqwest = { version = "0.12", default-features = false, features = ["json", "stream", "rustls-tls", "blocking", "http2"] } +reqwest = { version = "0.12", default-features = false, features = ["json", "stream", "rustls-tls", "blocking", "http2", "socks"] } base64 = "0.22" rand = "0.8" sha2 = "0.10" @@ -34,6 +34,11 @@ hostname = "0.4" once_cell = "1.20" regex-lite = "0.1" tokio-tungstenite = { version = "0.24", default-features = false, features = ["connect", "rustls-tls-native-roots"] } +tokio-rustls = { version = "0.26", default-features = false, features = ["ring", "tls12"] } +rustls = { version = "0.23", default-features = false, features = ["ring", "std"] } +rustls-native-certs = "0.8" +webpki-roots = "1" +hyper-util = { version = "0.1", features = ["client-proxy"] } url = "2" prost = "0.13" flate2 = "1.0" diff --git a/docs/src/content/docs/providers/codex.md b/docs/src/content/docs/providers/codex.md index a2734b52..2acd6137 100644 --- a/docs/src/content/docs/providers/codex.md +++ b/docs/src/content/docs/providers/codex.md @@ -45,6 +45,8 @@ Claude Code summary compaction requests are capped at low effort by default beca WebSocket is the default transport. Set `CCP_CODEX_TRANSPORT=http` for HTTP SSE, or `auto` to use WebSocket with HTTP fallback only when setup fails before a request is sent. +WebSocket setup honors `HTTP_PROXY` for `ws://`, `HTTPS_PROXY` for the default `wss://` endpoint, `ALL_PROXY` as a fallback, and `NO_PROXY` exclusions. A normal HTTP proxy can therefore carry the default WebSocket connection with CONNECT; TUN mode is not required. Set proxy variables before starting the process and restart after changing them. For example, setting `HTTPS_PROXY` to `http://127.0.0.1:7890` sends HTTPS/WSS destinations through the HTTP proxy at port 7890; it does not require an `https://` proxy URL. + `CCP_CODEX_PREVIOUS_RESPONSE_ID=1` enables append-only WebSocket continuation. It reuses a session connection and sends `previous_response_id` only when the translated request shape and transcript extension are safe. State is in memory, keyed by Claude Code session ID. ## Server compaction diff --git a/docs/src/content/docs/reference/configuration.md b/docs/src/content/docs/reference/configuration.md index d571fa6d..f3194812 100644 --- a/docs/src/content/docs/reference/configuration.md +++ b/docs/src/content/docs/reference/configuration.md @@ -69,6 +69,28 @@ All keys are optional. An unreadable file, malformed JSON, or incompatible field Codex auto-review classifier requests use `gpt-5.6-luna` by default. Requests routed through other providers retain their requested model. `CCP_AUTO_REVIEW_MODEL` or `autoReviewModel` selects an explicit registered model for all detected classifier requests without changing the session's provider affinity. Normal messages, streaming requests, tool-using requests, and token counting retain their requested model. +## Outbound proxies + +Outbound HTTP requests and Codex WebSocket setup inherit standard proxy environment variables when each provider client is created: + +| Environment | Applies to | Purpose | +| --- | --- | --- | +| `HTTP_PROXY` / `http_proxy` | `http://` and `ws://` destinations | Routes plain HTTP and WebSocket Upgrade requests through the configured proxy. | +| `HTTPS_PROXY` / `https_proxy` | `https://` and `wss://` destinations | Routes TLS destinations through the configured proxy, normally with HTTP CONNECT. | +| `ALL_PROXY` / `all_proxy` | Any destination without a scheme-specific proxy | Provides the fallback proxy. | +| `NO_PROXY` / `no_proxy` | Matching destination hosts and IPs | Bypasses proxy routing. Supports comma-separated domains, subdomains, IP addresses, CIDR ranges, and `*`. | + +On case-sensitive platforms, uppercase names are checked before lowercase names when both spellings exist. Windows treats environment names as case-insensitive, so each uppercase/lowercase pair identifies one variable. Set these variables before starting claude-code-proxy; changing them requires a restart because clients and pooled WebSocket connections retain their startup route. For CGI safety, proxy environment variables are ignored when `REQUEST_METHOD` is present. The variable name describes the **destination** scheme, so this default-WSS configuration is valid even though the local proxy URL uses `http://`: + +| Variable | Value | +| --- | --- | +| `HTTP_PROXY` | `http://127.0.0.1:7890` | +| `HTTPS_PROXY` | `http://127.0.0.1:7890` | + +After setting both variables through the operating system, service manager, or shell, start `claude-code-proxy serve` in the same environment. + +Proxy URLs may use `http`, `https`, `socks4`, `socks4a`, `socks5`, or `socks5h`. HTTP proxy URLs can contain percent-encoded Basic credentials, for example `http://user:password@127.0.0.1:7890`; SOCKS5 and SOCKS5H URLs can contain username/password credentials for the SOCKS handshake. SOCKS4 and SOCKS4A are supported without URL credentials. Prefer a secret-management mechanism when available because environment variables may be visible to other local processes. WSS certificate verification uses both bundled public roots and the platform native root store, including locally installed enterprise proxy CAs. Malformed or unsupported proxy URLs fail provider startup rather than being ignored. Proxy failures do not silently retry with a direct connection; only `NO_PROXY` selects direct routing. OS proxy settings, PAC files, and WPAD are not read automatically. + ## Codex | Environment | Config key | Default | Purpose | diff --git a/docs/src/content/docs/using/troubleshooting.md b/docs/src/content/docs/using/troubleshooting.md index 362a74b8..e127052e 100644 --- a/docs/src/content/docs/using/troubleshooting.md +++ b/docs/src/content/docs/using/troubleshooting.md @@ -51,6 +51,19 @@ Set `CLAUDE_CODE_DISABLE_NONSTREAMING_FALLBACK=1` for Claude Code. Retrying a pa ## Codex WebSocket fails +For a non-TUN local HTTP proxy, set both destination-scheme variables before starting the proxy: + +| Variable | Value | +| --- | --- | +| `HTTP_PROXY` | `http://127.0.0.1:7890` | +| `HTTPS_PROXY` | `http://127.0.0.1:7890` | + +Set them through the operating system, service manager, or shell, then start `claude-code-proxy serve` in the same environment. + +The default `wss://chatgpt.com` connection uses `HTTPS_PROXY`; a working proxy should show `CONNECT chatgpt.com:443`. The `http://` value is normal: it describes how to reach the proxy, while `HTTPS_PROXY` describes which destinations use it. Restart claude-code-proxy after changing these variables because the client and pooled WebSocket route are created at startup. + +Check `NO_PROXY` when the proxy sees no request. Proxy connection, authentication, or CONNECT failure is returned as an error and never retried directly. Environment variables are supported; OS proxy settings and PAC/WPAD discovery are not automatic. + Use HTTP SSE to isolate transport behavior: ```sh diff --git a/src/logging.rs b/src/logging.rs index 01f35bd0..0053dbca 100644 --- a/src/logging.rs +++ b/src/logging.rs @@ -11,8 +11,9 @@ pub const MAX_LOG_BYTES: u64 = 20 * 1024 * 1024; static STDERR_SUPPRESSION_DEPTH: AtomicUsize = AtomicUsize::new(0); -pub const REDACT_KEYS: [&str; 14] = [ +pub const REDACT_KEYS: [&str; 15] = [ "authorization", + "proxy-authorization", "access", "access_token", "refresh", @@ -259,4 +260,15 @@ mod tests { drop(outer); assert!(should_mirror_to_stderr("warn")); } + + #[test] + fn redacts_proxy_authorization_case_insensitively() { + let redacted = redact_value(serde_json::json!({ + "Proxy-Authorization": "Basic dXNlcjpwYXNz", + "safe": "kept" + })); + + assert_eq!(redacted["safe"], "kept"); + assert_eq!(redacted["Proxy-Authorization"], "[redacted len=18]"); + } } diff --git a/src/providers/codex/client.rs b/src/providers/codex/client.rs index d3e7fb84..908ce957 100644 --- a/src/providers/codex/client.rs +++ b/src/providers/codex/client.rs @@ -249,17 +249,174 @@ const MAX_BUFFERED_TRANSPORT_RETRIES: u32 = 3; const MAX_BUFFERED_TRANSPORT_ATTEMPTS: u32 = MAX_BUFFERED_TRANSPORT_RETRIES + 1; const HTTP_RESPONSE_BODY_IDLE_TIMEOUT_MS: u64 = 300_000; -fn native_http_client() -> reqwest::Client { +#[derive(Clone)] +struct ProxyEnvironment { + http_proxy: Option, + https_proxy: Option, + all_proxy: Option, + no_proxy: Option, + no_proxy_value: Option, +} + +impl ProxyEnvironment { + fn from_env() -> Self { + if std::env::var_os("REQUEST_METHOD").is_some() { + return Self { + http_proxy: None, + https_proxy: None, + all_proxy: None, + no_proxy: None, + no_proxy_value: None, + }; + } + + let no_proxy_value = std::env::var("NO_PROXY") + .or_else(|_| std::env::var("no_proxy")) + .ok(); + Self { + http_proxy: proxy_env_value("HTTP_PROXY", "http_proxy") + .unwrap_or_else(|name| panic!("invalid {name} proxy URL")), + https_proxy: proxy_env_value("HTTPS_PROXY", "https_proxy") + .unwrap_or_else(|name| panic!("invalid {name} proxy URL")), + all_proxy: proxy_env_value("ALL_PROXY", "all_proxy") + .unwrap_or_else(|name| panic!("invalid {name} proxy URL")), + no_proxy: no_proxy_value + .as_deref() + .and_then(reqwest::NoProxy::from_string), + no_proxy_value, + } + } + + fn websocket_proxy_config(&self) -> super::websocket::WebSocketProxyConfig { + super::websocket::WebSocketProxyConfig::new( + self.http_proxy.as_deref(), + self.https_proxy.as_deref(), + self.all_proxy.as_deref(), + self.no_proxy_value.as_deref(), + ) + } + + fn apply(&self, mut builder: reqwest::ClientBuilder) -> reqwest::ClientBuilder { + builder = builder.no_proxy(); + if let Some(proxy) = self.http_proxy.as_deref() { + builder = builder.proxy( + reqwest::Proxy::http(proxy) + .expect("validated HTTP_PROXY URL") + .no_proxy(self.no_proxy.clone()), + ); + } + if let Some(proxy) = self.https_proxy.as_deref() { + builder = builder.proxy( + reqwest::Proxy::https(proxy) + .expect("validated HTTPS_PROXY URL") + .no_proxy(self.no_proxy.clone()), + ); + } + if let Some(proxy) = self.all_proxy.as_deref() { + builder = builder.proxy( + reqwest::Proxy::all(proxy) + .expect("validated ALL_PROXY URL") + .no_proxy(self.no_proxy.clone()), + ); + } + builder + } +} + +fn native_http_client(proxy_environment: &ProxyEnvironment) -> reqwest::Client { + proxy_environment + .apply( + reqwest::Client::builder() + .connect_timeout(Duration::from_secs(15)) + .redirect(reqwest::redirect::Policy::none()), + ) + .build() + .expect("failed to create native Responses HTTP client") +} + +fn proxy_env_value( + uppercase: &'static str, + lowercase: &'static str, +) -> Result, &'static str> { + let Some(raw) = std::env::var_os(uppercase).or_else(|| std::env::var_os(lowercase)) else { + return Ok(None); + }; + let raw = raw.into_string().map_err(|_| uppercase)?; + let raw = raw.trim(); + if raw.is_empty() { + return Ok(None); + } + normalize_proxy_url(raw).map(Some).ok_or(uppercase) +} + +fn normalize_proxy_url(raw: &str) -> Option { + let candidate = if raw.contains("://") { + raw.to_string() + } else { + format!("http://{raw}") + }; + let parsed = url::Url::parse(&candidate).ok()?; + if !matches!( + parsed.scheme(), + "http" | "https" | "socks4" | "socks4a" | "socks5" | "socks5h" + ) || parsed.host_str().is_none() + || matches!(parsed.scheme(), "socks4" | "socks4a") + && (!parsed.username().is_empty() || parsed.password().is_some()) + { + return None; + } + Some(parsed.to_string()) +} + +fn websocket_http_client(proxy_environment: &ProxyEnvironment) -> reqwest::Client { + let tls_config = super::websocket::websocket_tls_config(); + proxy_environment + .apply( + reqwest::Client::builder() + .http1_only() + .redirect(reqwest::redirect::Policy::none()) + .use_preconfigured_tls((*tls_config).clone()), + ) + .build() + .expect("failed to create Codex WebSocket HTTP client") +} + +fn custom_client_auto_http_fallback_enabled( + base_url: &str, + proxy_config: &super::websocket::WebSocketProxyConfig, +) -> bool { + let Ok(websocket_url) = super::websocket::to_websocket_url(base_url) else { + return false; + }; + !proxy_config.uses_proxy_for(&websocket_url) +} + +#[cfg(test)] +fn test_native_http_client() -> reqwest::Client { reqwest::Client::builder() .connect_timeout(Duration::from_secs(15)) .redirect(reqwest::redirect::Policy::none()) + .no_proxy() .build() - .expect("failed to create native Responses HTTP client") + .expect("failed to create test native Responses HTTP client") +} + +#[cfg(test)] +fn test_websocket_http_client() -> reqwest::Client { + reqwest::Client::builder() + .http1_only() + .redirect(reqwest::redirect::Policy::none()) + .no_proxy() + .build() + .expect("failed to create test WebSocket HTTP client") } pub struct CodexHttpClient { client: reqwest::Client, native_client: reqwest::Client, + websocket_client: reqwest::Client, + websocket_proxy_config: super::websocket::WebSocketProxyConfig, + auto_http_fallback_enabled: bool, auth_manager: CodexAuthManager, base_url: String, header_timeout_ms: u64, @@ -277,12 +434,16 @@ impl Default for CodexHttpClient { impl CodexHttpClient { pub fn new() -> Self { let timeout_ms = 60_000; + let proxy_environment = ProxyEnvironment::from_env(); Self { - client: reqwest::Client::builder() - .connect_timeout(Duration::from_secs(15)) + client: proxy_environment + .apply(reqwest::Client::builder().connect_timeout(Duration::from_secs(15))) .build() .expect("failed to create HTTP client"), - native_client: native_http_client(), + native_client: native_http_client(&proxy_environment), + websocket_client: websocket_http_client(&proxy_environment), + websocket_proxy_config: proxy_environment.websocket_proxy_config(), + auto_http_fallback_enabled: true, auth_manager: CodexAuthManager::new(file_store()), base_url: config::codex_base_url(CODEX_API_ENDPOINT), header_timeout_ms: timeout_ms, @@ -296,8 +457,15 @@ impl CodexHttpClient { auth_manager: CodexAuthManager, base_url: String, ) -> Self { + let proxy_environment = ProxyEnvironment::from_env(); + let websocket_proxy_config = proxy_environment.websocket_proxy_config(); + let auto_http_fallback_enabled = + custom_client_auto_http_fallback_enabled(&base_url, &websocket_proxy_config); Self { - native_client: native_http_client(), + native_client: native_http_client(&proxy_environment), + websocket_client: websocket_http_client(&proxy_environment), + websocket_proxy_config, + auto_http_fallback_enabled, client, auth_manager, base_url, @@ -316,7 +484,10 @@ impl CodexHttpClient { header_timeout_retries: u32, ) -> Self { Self { - native_client: native_http_client(), + native_client: test_native_http_client(), + websocket_client: test_websocket_http_client(), + websocket_proxy_config: super::websocket::WebSocketProxyConfig::direct(), + auto_http_fallback_enabled: true, client, auth_manager: CodexAuthManager::new(file_store()), base_url, @@ -486,6 +657,8 @@ impl CodexHttpClient { let ws_body = build_websocket_request(body, active_continuation.as_ref()); super::websocket::codex_websocket_request( + &self.websocket_client, + &self.websocket_proxy_config, &self.base_url, &ws_headers, &ws_body, @@ -506,6 +679,8 @@ impl CodexHttpClient { // Try WebSocket first let ws_result = super::websocket::codex_websocket_request( + &self.websocket_client, + &self.websocket_proxy_config, &self.base_url, &ws_headers, &ws_body, @@ -520,7 +695,9 @@ impl CodexHttpClient { match ws_result { Ok(response) => Ok(response), - Err(err) if should_fallback_to_http(&err) => { + Err(err) + if self.auto_http_fallback_enabled && should_fallback_to_http(&err) => + { // Fall back to HTTP only if WebSocket failed before sending let body_json = serde_json::to_string(body).map_err(|e| CodexError { @@ -835,6 +1012,8 @@ impl CodexHttpClient { }; let ws_body = build_websocket_request(&body, continuation.as_ref()); let start = super::websocket::codex_websocket_event_stream( + &self.websocket_client, + &self.websocket_proxy_config, &self.base_url, &ws_headers, &ws_body, @@ -1326,6 +1505,9 @@ fn log_buffered_retry_exhausted( fn is_retryable_transport_error(err: &CodexError) -> bool { if err.origin == CodexErrorOrigin::WebSocketHandshake { + if err.detail.as_deref() == Some(super::websocket::WEBSOCKET_PROXY_TUNNEL_REJECTED_DETAIL) { + return false; + } return err.status == 0 || should_retry_codex_status(err.status); } if err.detail.as_deref() == Some("websocket_pre_request") { @@ -1377,6 +1559,8 @@ fn should_refresh_after_unauthorized( fn should_fallback_to_http(err: &CodexError) -> bool { err.origin == CodexErrorOrigin::WebSocketHandshake + && err.status != http::StatusCode::PROXY_AUTHENTICATION_REQUIRED.as_u16() + && err.detail.as_deref() != Some(super::websocket::WEBSOCKET_PROXY_TUNNEL_REJECTED_DETAIL) } fn should_retry_without_continuation( @@ -1448,6 +1632,56 @@ mod tests { use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::TcpListener; + #[test] + fn normalizes_supported_proxy_urls() { + assert_eq!( + normalize_proxy_url("127.0.0.1:8080").as_deref(), + Some("http://127.0.0.1:8080/") + ); + assert_eq!( + normalize_proxy_url("https://user:pass@proxy.example:8443").as_deref(), + Some("https://user:pass@proxy.example:8443/") + ); + for scheme in ["socks4", "socks4a"] { + let proxy = format!("{scheme}://proxy.example:1080"); + assert_eq!(normalize_proxy_url(&proxy), Some(proxy)); + } + for scheme in ["socks5", "socks5h"] { + let proxy = format!("{scheme}://user:pass@proxy.example:1080"); + assert_eq!(normalize_proxy_url(&proxy), Some(proxy)); + } + } + + #[test] + fn rejects_malformed_or_unsupported_proxy_urls() { + assert!(normalize_proxy_url("http://").is_none()); + assert!(normalize_proxy_url("ftp://proxy.example:21").is_none()); + assert!(normalize_proxy_url("socks4://user@proxy.example:1080").is_none()); + assert!(normalize_proxy_url("socks4a://user:pass@proxy.example:1080").is_none()); + } + + #[test] + fn custom_client_auto_fallback_tracks_effective_proxy_route() { + let proxy = "http://proxy.example:8080"; + let proxied = + super::super::websocket::WebSocketProxyConfig::new(None, Some(proxy), None, None); + assert!(!custom_client_auto_http_fallback_enabled( + "https://codex.invalid/responses", + &proxied + )); + + let bypassed = super::super::websocket::WebSocketProxyConfig::new( + None, + Some(proxy), + None, + Some("codex.invalid"), + ); + assert!(custom_client_auto_http_fallback_enabled( + "https://codex.invalid/responses", + &bypassed + )); + } + fn http_test_auth() -> StoredAuth { StoredAuth { access: "test".into(), @@ -1470,7 +1704,7 @@ mod tests { fn http_test_client(base_url: String, body_idle_timeout_ms: u64) -> CodexHttpClient { CodexHttpClient::new_for_test( - reqwest::Client::new(), + reqwest::Client::builder().no_proxy().build().unwrap(), base_url, 100, body_idle_timeout_ms, @@ -1687,7 +1921,11 @@ mod tests { let mut request = [0_u8; 16 * 1024]; let read = websocket.read(&mut request).await.unwrap(); assert!(read > 0); - assert!(String::from_utf8_lossy(&request[..read]).contains("Upgrade: websocket")); + assert!( + String::from_utf8_lossy(&request[..read]) + .to_ascii_lowercase() + .contains("upgrade: websocket") + ); websocket .write_all( b"HTTP/1.1 401 Unauthorized\r\ncontent-length: 13\r\nconnection: close\r\n\r\npolicy denied", @@ -1920,6 +2158,22 @@ mod tests { assert!(is_retryable_transport_error(&err)); } + #[test] + fn proxy_tunnel_rejection_is_not_retried_or_used_for_http_fallback() { + let err = CodexError { + status: 0, + message: "WebSocket proxy tunnel was rejected".to_string(), + detail: Some( + super::super::websocket::WEBSOCKET_PROXY_TUNNEL_REJECTED_DETAIL.to_string(), + ), + retry_after: None, + origin: CodexErrorOrigin::WebSocketHandshake, + }; + + assert!(!is_retryable_transport_error(&err)); + assert!(!should_fallback_to_http(&err)); + } + #[test] fn websocket_pre_request_statusless_error_is_retryable() { let err = CodexError { diff --git a/src/providers/codex/mod.rs b/src/providers/codex/mod.rs index 0717421e..800ccd9b 100644 --- a/src/providers/codex/mod.rs +++ b/src/providers/codex/mod.rs @@ -985,6 +985,9 @@ fn is_codex_success_terminal_event(payload: &serde_json::Value) -> bool { fn retryable_live_start_codex_error(err: &client::CodexError) -> bool { if err.origin == client::CodexErrorOrigin::WebSocketHandshake { + if err.detail.as_deref() == Some(websocket::WEBSOCKET_PROXY_TUNNEL_REJECTED_DETAIL) { + return false; + } return err.status == 0 || matches!(err.status, 429 | 500 | 502 | 503 | 504 | 529); } matches!(err.status, 429 | 500 | 502 | 503 | 504 | 529) @@ -1563,6 +1566,19 @@ mod tests { assert!(retryable_live_start_codex_error(&err)); } + #[test] + fn live_start_proxy_tunnel_rejection_is_not_retryable() { + let err = client::CodexError { + status: 0, + message: "WebSocket proxy tunnel was rejected".to_string(), + detail: Some(websocket::WEBSOCKET_PROXY_TUNNEL_REJECTED_DETAIL.to_string()), + retry_after: None, + origin: client::CodexErrorOrigin::WebSocketHandshake, + }; + + assert!(!retryable_live_start_codex_error(&err)); + } + #[test] fn live_start_payload_retry_detection_covers_rate_limit_and_overload() { assert!(retryable_live_start_payload( diff --git a/src/providers/codex/websocket.rs b/src/providers/codex/websocket.rs index 27a9fc85..155f7ba7 100644 --- a/src/providers/codex/websocket.rs +++ b/src/providers/codex/websocket.rs @@ -1,17 +1,27 @@ use std::collections::HashMap; +use std::pin::Pin; use std::sync::{Arc, Mutex}; +use std::task::{Context, Poll}; use std::time::{Duration, Instant}; use futures_util::{SinkExt, StreamExt}; use http::HeaderMap; +use hyper_util::client::proxy::matcher::Matcher as ProxyMatcher; +use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt, ReadBuf}; use tokio::net::TcpStream; use tokio::sync::mpsc; use tokio::sync::{Mutex as AsyncMutex, OwnedMutexGuard}; use tokio_tungstenite::{ - MaybeTlsStream, WebSocketStream, connect_async, - tungstenite::{self, Message, handshake::client::generate_key}, + WebSocketStream, + tungstenite::{ + Message, + client::IntoClientRequest, + handshake::{client::generate_key, derive_accept_key}, + protocol::Role, + }, }; +use crate::logging::create_logger; use crate::provider::RequestContext; use crate::traffic::TrafficCapture; @@ -27,9 +37,11 @@ pub const WEBSOCKET_CONNECT_TIMEOUT_MS: u64 = 15_000; pub const WEBSOCKET_IDLE_TIMEOUT_MS: u64 = 300_000; pub const WEBSOCKET_RESPONSE_START_TIMEOUT_DETAIL: &str = "websocket_response_start_timeout"; pub const WEBSOCKET_MISSING_TERMINAL_DETAIL: &str = "websocket_missing_terminal"; +pub(super) const WEBSOCKET_PROXY_TUNNEL_REJECTED_DETAIL: &str = "websocket_proxy_tunnel_rejected"; const POOL_IDLE_TTL_MS: u64 = 30 * 60 * 1000; const MAX_POOL_ENTRIES: usize = 10_000; +const MAX_CONNECT_RESPONSE_HEADER_BYTES: usize = 8 * 1024; // Terminal WebSocket event types that signal the request is done const TERMINAL_EVENTS: &[&str] = &[ @@ -41,6 +53,122 @@ const TERMINAL_EVENTS: &[&str] = &[ pub type CodexWebSocketEventReceiver = mpsc::Receiver>; +trait WebSocketIo: AsyncRead + AsyncWrite + Unpin + Send {} + +impl WebSocketIo for T where T: AsyncRead + AsyncWrite + Unpin + Send {} + +type BoxedWebSocketIo = Box; +type CodexWebSocketStream = WebSocketStream; + +pub(super) struct WebSocketProxyConfig { + matcher: ProxyMatcher, + tls_config: Arc, +} + +#[derive(Clone)] +struct WebSocketProxyRoute { + uri: http::Uri, + basic_auth: Option, +} + +impl WebSocketProxyConfig { + pub(super) fn new( + http_proxy: Option<&str>, + https_proxy: Option<&str>, + all_proxy: Option<&str>, + no_proxy: Option<&str>, + ) -> Self { + let mut builder = ProxyMatcher::builder(); + if let Some(proxy) = all_proxy { + builder = builder.all(proxy.to_string()); + } + if let Some(proxy) = http_proxy { + builder = builder.http(proxy.to_string()); + } + if let Some(proxy) = https_proxy { + builder = builder.https(proxy.to_string()); + } + if let Some(no_proxy) = no_proxy { + builder = builder.no(no_proxy.to_string()); + } + Self { + matcher: builder.build(), + tls_config: websocket_tls_config(), + } + } + + #[cfg(test)] + pub(super) fn direct() -> Self { + Self::new(None, None, None, None) + } + + pub(super) fn uses_proxy_for(&self, websocket_url: &str) -> bool { + let Ok(http_url) = to_http_upgrade_url(websocket_url) else { + return true; + }; + let Ok(destination) = http_url.parse::() else { + return true; + }; + self.matcher.intercept(&destination).is_some() + } + + fn http_connect_route( + &self, + websocket_url: &str, + ) -> Result, CodexError> { + let http_url = to_http_upgrade_url(websocket_url).map_err(|error| CodexError { + status: 0, + message: error.message, + detail: None, + retry_after: None, + origin: CodexErrorOrigin::WebSocketHandshake, + })?; + if !http_url.starts_with("https://") { + return Ok(None); + } + let destination = http_url.parse::().map_err(|_| { + websocket_protocol_error("WebSocket destination URL could not be routed") + })?; + let Some(proxy) = self.matcher.intercept(&destination) else { + return Ok(None); + }; + if !matches!(proxy.uri().scheme_str(), Some("http" | "https")) { + return Ok(None); + } + Ok(Some(WebSocketProxyRoute { + uri: proxy.uri().clone(), + basic_auth: proxy.basic_auth().cloned(), + })) + } +} + +static WEBSOCKET_TLS_CONFIG: once_cell::sync::Lazy> = + once_cell::sync::Lazy::new(|| { + let mut roots = rustls::RootCertStore::empty(); + let native = rustls_native_certs::load_native_certs(); + let load_error_count = native.errors.len(); + let (_, parse_error_count) = roots.add_parsable_certificates(native.certs); + if load_error_count > 0 || parse_error_count > 0 { + let mut fields = serde_json::Map::new(); + fields.insert("loadErrorCount".into(), serde_json::json!(load_error_count)); + fields.insert( + "parseErrorCount".into(), + serde_json::json!(parse_error_count), + ); + create_logger("codex").warn("native_certificate_load_errors", Some(fields)); + } + roots.extend(webpki_roots::TLS_SERVER_ROOTS.iter().cloned()); + Arc::new( + rustls::ClientConfig::builder() + .with_root_certificates(roots) + .with_no_client_auth(), + ) + }); + +pub(super) fn websocket_tls_config() -> Arc { + WEBSOCKET_TLS_CONFIG.clone() +} + // --------------------------------------------------------------------------- // Errors // --------------------------------------------------------------------------- @@ -87,7 +215,7 @@ impl std::fmt::Display for CodexWebSocketError { // --------------------------------------------------------------------------- struct PoolEntry { - ws: Arc>>>, + ws: Arc>, created_at: u64, } @@ -196,6 +324,25 @@ pub fn to_websocket_url(url: &str) -> Result { Ok(parsed.to_string()) } +fn to_http_upgrade_url(url: &str) -> Result { + let mut parsed = url::Url::parse(url) + .map_err(|e| CodexWebSocketError::new(format!("Failed to parse URL: {e}")))?; + match parsed.scheme() { + "ws" => parsed.set_scheme("http").map_err(|_| { + CodexWebSocketError::new("Unsupported Codex WebSocket URL scheme".to_string()) + })?, + "wss" => parsed.set_scheme("https").map_err(|_| { + CodexWebSocketError::new("Unsupported Codex WebSocket URL scheme".to_string()) + })?, + other => { + return Err(CodexWebSocketError::new(format!( + "Unsupported Codex WebSocket URL scheme: {other}" + ))); + } + } + Ok(parsed.to_string()) +} + // --------------------------------------------------------------------------- // Header rewriting // --------------------------------------------------------------------------- @@ -207,7 +354,12 @@ pub fn codex_websocket_headers(http_headers: &HeaderMap) -> HeaderMap { // Skip hop-by-hop headers if matches!( key_str.as_str(), - "content-length" | "content-type" | "accept" | "connection" | "upgrade" + "content-length" + | "content-type" + | "accept" + | "connection" + | "upgrade" + | "proxy-authorization" ) { continue; } @@ -297,7 +449,9 @@ fn extract_retry_after(payload: &serde_json::Value) -> Option { // --------------------------------------------------------------------------- #[allow(clippy::too_many_arguments)] -pub async fn codex_websocket_request( +pub(super) async fn codex_websocket_request( + websocket_client: &reqwest::Client, + proxy_config: &WebSocketProxyConfig, url: &str, headers: &HeaderMap, body_value: &serde_json::Value, @@ -346,14 +500,21 @@ pub async fn codex_websocket_request( pool_get_for_turn(key, continuation.and_then(|candidate| candidate.turn_id)) }); - let (ws_stream, _response) = if let Some(entry) = pooled { + let ws_stream = if let Some(entry) = pooled { // Use pooled connection let mut ws_guard = entry.ws.lock().await; // Check if connection is still alive by sending a ping if ws_guard.send(Message::Ping(vec![])).await.is_err() { invalidate_pool_entry(pool_key.unwrap(), &entry); // Fall through to new connection - connect_with_timeout(&ws_url, headers, connect_timeout_ms).await? + connect_with_timeout( + websocket_client, + proxy_config, + &ws_url, + headers, + connect_timeout_ms, + ) + .await? } else { // Connection is alive, send the request through it let ws_msg = Message::Text(body_json.clone()); @@ -415,7 +576,14 @@ pub async fn codex_websocket_request( }); } } else { - connect_with_timeout(&ws_url, headers, connect_timeout_ms).await? + connect_with_timeout( + websocket_client, + proxy_config, + &ws_url, + headers, + connect_timeout_ms, + ) + .await? }; // New connection path (not pooled or pool miss) @@ -496,7 +664,7 @@ pub async fn codex_websocket_request( pub(super) struct ReadyWebSocket { ws_url: String, - guard: OwnedMutexGuard>>, + guard: OwnedMutexGuard, entry: Arc, used_pooled: bool, pool_key: Option, @@ -505,7 +673,10 @@ pub(super) struct ReadyWebSocket { idle_timeout_ms: u64, } +#[allow(clippy::too_many_arguments)] pub(super) async fn prepare_codex_websocket( + websocket_client: &reqwest::Client, + proxy_config: &WebSocketProxyConfig, url: &str, headers: &HeaderMap, traffic: Option>, @@ -526,7 +697,14 @@ pub(super) async fn prepare_codex_websocket( let entry = if let Some(entry) = pooled { entry } else { - let (stream, _) = connect_with_timeout(&ws_url, headers, connect_timeout_ms).await?; + let stream = connect_with_timeout( + websocket_client, + proxy_config, + &ws_url, + headers, + connect_timeout_ms, + ) + .await?; Arc::new(PoolEntry { ws: Arc::new(AsyncMutex::new(stream)), created_at: now_ms(), @@ -655,7 +833,9 @@ pub(super) fn start_codex_websocket_events( } #[allow(clippy::too_many_arguments)] -pub async fn codex_websocket_event_stream( +pub(super) async fn codex_websocket_event_stream( + websocket_client: &reqwest::Client, + proxy_config: &WebSocketProxyConfig, url: &str, headers: &HeaderMap, body_value: &serde_json::Value, @@ -674,6 +854,8 @@ pub async fn codex_websocket_event_stream( origin: CodexErrorOrigin::WebSocketHandshake, })?; let ready = prepare_codex_websocket( + websocket_client, + proxy_config, url, headers, traffic, @@ -773,7 +955,7 @@ fn write_websocket_response_capture( const MAX_HANDSHAKE_ERROR_DETAIL_BYTES: usize = 1024; const GENERIC_HANDSHAKE_ERROR_DETAIL: &str = "WebSocket upgrade was rejected"; -fn handshake_error_detail(body: Option<&Vec>) -> String { +fn handshake_error_detail(body: Option<&[u8]>) -> String { let Some(value) = body.and_then(|body| serde_json::from_slice::(body).ok()) else { return GENERIC_HANDSHAKE_ERROR_DETAIL.to_string(); @@ -796,90 +978,597 @@ fn handshake_error_detail(body: Option<&Vec>) -> String { sanitized[..end].to_string() } -async fn connect_with_timeout( +fn header_has_token(headers: &HeaderMap, name: &str, expected: &str) -> bool { + headers.get_all(name).iter().any(|value| { + value.to_str().ok().is_some_and(|value| { + value + .split(',') + .any(|token| token.trim().eq_ignore_ascii_case(expected)) + }) + }) +} + +fn requested_subprotocols(headers: &HeaderMap) -> Vec { + headers + .get_all(http::header::SEC_WEBSOCKET_PROTOCOL) + .iter() + .filter_map(|value| value.to_str().ok()) + .flat_map(|value| value.split(',')) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string) + .collect() +} + +fn websocket_protocol_error(message: &str) -> CodexError { + CodexError { + status: 0, + message: message.to_string(), + detail: None, + retry_after: None, + origin: CodexErrorOrigin::WebSocketHandshake, + } +} + +fn validate_websocket_upgrade( + version: http::Version, + headers: &HeaderMap, + websocket_key: &str, + requested_subprotocols: &[String], +) -> Result<(), CodexError> { + if version != http::Version::HTTP_11 { + return Err(websocket_protocol_error( + "WebSocket upgrade response did not use HTTP/1.1", + )); + } + if !header_has_token(headers, http::header::UPGRADE.as_str(), "websocket") { + return Err(websocket_protocol_error( + "WebSocket upgrade response is missing Upgrade: websocket", + )); + } + if !header_has_token(headers, http::header::CONNECTION.as_str(), "upgrade") { + return Err(websocket_protocol_error( + "WebSocket upgrade response is missing Connection: Upgrade", + )); + } + + let expected_accept = derive_accept_key(websocket_key.as_bytes()); + let mut accept_values = headers.get_all(http::header::SEC_WEBSOCKET_ACCEPT).iter(); + let accept = accept_values.next().and_then(|value| value.to_str().ok()); + if accept_values.next().is_some() || accept != Some(expected_accept.as_str()) { + return Err(websocket_protocol_error( + "WebSocket upgrade response has an invalid Sec-WebSocket-Accept", + )); + } + if headers.contains_key(http::header::SEC_WEBSOCKET_EXTENSIONS) { + return Err(websocket_protocol_error( + "WebSocket upgrade response selected an unsolicited extension", + )); + } + + let mut response_protocols = headers.get_all(http::header::SEC_WEBSOCKET_PROTOCOL).iter(); + let response_protocol = response_protocols + .next() + .map(|value| value.to_str().map(str::trim)); + if response_protocols.next().is_some() { + return Err(websocket_protocol_error( + "WebSocket upgrade response contains multiple subprotocols", + )); + } + match response_protocol { + None if requested_subprotocols.is_empty() => {} + None => { + return Err(websocket_protocol_error( + "WebSocket upgrade response omitted the requested subprotocol", + )); + } + Some(Err(_)) => { + return Err(websocket_protocol_error( + "WebSocket upgrade response contains an invalid subprotocol", + )); + } + Some(Ok(_)) if requested_subprotocols.is_empty() => { + return Err(websocket_protocol_error( + "WebSocket upgrade response selected an unsolicited subprotocol", + )); + } + Some(Ok(protocol)) + if !requested_subprotocols + .iter() + .any(|requested| requested == protocol) => + { + return Err(websocket_protocol_error( + "WebSocket upgrade response selected an unsupported subprotocol", + )); + } + Some(Ok(_)) => {} + } + + Ok(()) +} + +async fn bounded_handshake_error_body(mut response: reqwest::Response) -> Vec { + let mut body = Vec::new(); + while body.len() < MAX_HANDSHAKE_ERROR_DETAIL_BYTES { + let chunk = match response.chunk().await { + Ok(Some(chunk)) => chunk, + Ok(None) | Err(_) => break, + }; + let remaining = MAX_HANDSHAKE_ERROR_DETAIL_BYTES - body.len(); + body.extend_from_slice(&chunk[..chunk.len().min(remaining)]); + } + body +} + +fn error_chain_contains(error: &(dyn std::error::Error + 'static), expected: &str) -> bool { + let expected = expected.to_ascii_lowercase(); + let mut current = Some(error); + while let Some(error) = current { + if error.to_string().to_ascii_lowercase().contains(&expected) { + return true; + } + current = error.source(); + } + false +} + +fn reqwest_handshake_error(error: reqwest::Error) -> CodexError { + let proxy_auth_required = + error.is_connect() && error_chain_contains(&error, "proxy authorization required"); + let proxy_tunnel_rejected = + error.is_connect() && error_chain_contains(&error, "tunnel error: unsuccessful"); + let status = if proxy_auth_required { + http::StatusCode::PROXY_AUTHENTICATION_REQUIRED.as_u16() + } else { + error.status().map(|status| status.as_u16()).unwrap_or(0) + }; + let message = if proxy_auth_required { + "WebSocket proxy authentication failed" + } else if proxy_tunnel_rejected { + "WebSocket proxy tunnel was rejected" + } else if error.is_timeout() { + "WebSocket upgrade request timed out" + } else if error.is_connect() { + "WebSocket connection failed" + } else { + "WebSocket upgrade request failed" + }; + CodexError { + status, + message: message.to_string(), + detail: proxy_tunnel_rejected.then(|| WEBSOCKET_PROXY_TUNNEL_REJECTED_DETAIL.to_string()), + retry_after: None, + origin: CodexErrorOrigin::WebSocketHandshake, + } +} + +struct PrefixedIo { + prefix: Vec, + position: usize, + inner: BoxedWebSocketIo, +} + +impl AsyncRead for PrefixedIo { + fn poll_read( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + buffer: &mut ReadBuf<'_>, + ) -> Poll> { + let this = self.get_mut(); + if this.position < this.prefix.len() { + let available = &this.prefix[this.position..]; + let len = available.len().min(buffer.remaining()); + buffer.put_slice(&available[..len]); + this.position += len; + return Poll::Ready(Ok(())); + } + Pin::new(&mut this.inner).poll_read(cx, buffer) + } +} + +impl AsyncWrite for PrefixedIo { + fn poll_write( + self: Pin<&mut Self>, + cx: &mut Context<'_>, + buffer: &[u8], + ) -> Poll> { + Pin::new(&mut self.get_mut().inner).poll_write(cx, buffer) + } + + fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.get_mut().inner).poll_flush(cx) + } + + fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.get_mut().inner).poll_shutdown(cx) + } +} + +fn skip_websocket_request_header(name: &http::HeaderName) -> bool { + matches!( + name.as_str(), + "connection" + | "upgrade" + | "sec-websocket-key" + | "sec-websocket-version" + | "host" + | "content-length" + | "proxy-authorization" + ) +} + +fn tunneled_websocket_request( url: &str, headers: &HeaderMap, - connect_timeout_ms: u64, -) -> Result< - ( - WebSocketStream>, - tungstenite::handshake::client::Response, - ), - CodexError, -> { - // Build an http::Request with the given headers for the WebSocket upgrade - let host = websocket_host_header(url); - let mut req_builder = http::Request::builder() - .uri(url) - .method("GET") - .header("Host", host) - .header("Connection", "Upgrade") - .header("Upgrade", "websocket") - .header("Sec-WebSocket-Version", "13") - .header("Sec-WebSocket-Key", generate_key()); - - // Copy over the codex headers - for (key, value) in headers.iter() { - let key_str = key.as_str().to_lowercase(); - // Skip headers already set for WebSocket upgrade - if matches!( - key_str.as_str(), - "connection" | "upgrade" | "sec-websocket-key" | "sec-websocket-version" | "host" - ) { - continue; + websocket_key: &str, +) -> Result, CodexError> { + let mut request = url + .into_client_request() + .map_err(|_| websocket_protocol_error("WebSocket request URL was invalid"))?; + *request.version_mut() = http::Version::HTTP_11; + request.headers_mut().insert( + http::header::SEC_WEBSOCKET_KEY, + http::HeaderValue::from_str(websocket_key) + .map_err(|_| websocket_protocol_error("WebSocket key was invalid"))?, + ); + for (name, value) in headers { + if !skip_websocket_request_header(name) { + request.headers_mut().append(name.clone(), value.clone()); } - req_builder = req_builder.header(key.as_str(), value.as_bytes()); } + Ok(request) +} - let request = req_builder.body(()).map_err(|e| CodexError { +fn tunnel_error(status: u16, retry_after: Option) -> CodexError { + let proxy_auth_required = status == http::StatusCode::PROXY_AUTHENTICATION_REQUIRED.as_u16(); + CodexError { + status, + message: if proxy_auth_required { + "WebSocket proxy authentication failed".to_string() + } else { + "WebSocket proxy tunnel was rejected".to_string() + }, + detail: if proxy_auth_required { + Some(GENERIC_HANDSHAKE_ERROR_DETAIL.to_string()) + } else { + Some(WEBSOCKET_PROXY_TUNNEL_REJECTED_DETAIL.to_string()) + }, + retry_after, + origin: CodexErrorOrigin::WebSocketHandshake, + } +} + +fn invalid_tunnel_response(message: &str) -> CodexError { + CodexError { status: 0, - message: format!("Failed to build WebSocket request: {e}"), + message: message.to_string(), + detail: Some(WEBSOCKET_PROXY_TUNNEL_REJECTED_DETAIL.to_string()), + retry_after: None, + origin: CodexErrorOrigin::WebSocketHandshake, + } +} + +fn connect_response_header_end(response: &[u8]) -> Option { + response + .windows(4) + .position(|window| window == b"\r\n\r\n") + .map(|position| position + 4) +} + +fn parse_connect_response_head(response: &[u8]) -> Result<(u16, Option), CodexError> { + let status_line_end = response + .windows(2) + .position(|window| window == b"\r\n") + .ok_or_else(|| invalid_tunnel_response("WebSocket proxy returned an invalid response"))?; + let status_line = std::str::from_utf8(&response[..status_line_end]) + .map_err(|_| invalid_tunnel_response("WebSocket proxy returned an invalid response"))?; + let mut parts = status_line.split_ascii_whitespace(); + let version = parts.next().unwrap_or_default(); + let status = parts.next().unwrap_or_default(); + if !matches!(version, "HTTP/1.0" | "HTTP/1.1") + || status.len() != 3 + || !status.bytes().all(|byte| byte.is_ascii_digit()) + { + return Err(invalid_tunnel_response( + "WebSocket proxy returned an invalid response", + )); + } + let status = status + .parse::() + .map_err(|_| invalid_tunnel_response("WebSocket proxy returned an invalid response"))?; + let retry_after = response[status_line_end + 2..] + .split(|byte| *byte == b'\n') + .find_map(|line| { + let line = line.strip_suffix(b"\r").unwrap_or(line); + let separator = line.iter().position(|byte| *byte == b':')?; + let (name, value) = line.split_at(separator); + if !name.eq_ignore_ascii_case(b"retry-after") { + return None; + } + std::str::from_utf8(&value[1..]) + .ok() + .map(|value| value.trim().to_string()) + }); + Ok((status, retry_after)) +} + +async fn establish_connect_tunnel( + mut stream: BoxedWebSocketIo, + authority: &str, + basic_auth: Option<&http::HeaderValue>, +) -> Result { + let mut request = format!("CONNECT {authority} HTTP/1.1\r\nHost: {authority}\r\n").into_bytes(); + if let Some(auth) = basic_auth { + request.extend_from_slice(b"Proxy-Authorization: "); + request.extend_from_slice(auth.as_bytes()); + request.extend_from_slice(b"\r\n"); + } + request.extend_from_slice(b"\r\n"); + stream.write_all(&request).await.map_err(|_| CodexError { + status: 0, + message: "WebSocket proxy tunnel request failed".to_string(), detail: None, retry_after: None, - origin: CodexErrorOrigin::WebSocket, + origin: CodexErrorOrigin::WebSocketHandshake, })?; - let connect_fut = connect_async(request); - tokio::time::timeout(Duration::from_millis(connect_timeout_ms), connect_fut) + let mut response = Vec::new(); + loop { + if let Some(header_end) = connect_response_header_end(&response) { + let (status, retry_after) = parse_connect_response_head(&response[..header_end])?; + if (100..200).contains(&status) { + response.drain(..header_end); + continue; + } + if (200..300).contains(&status) { + let prefix = response.split_off(header_end); + return Ok(Box::new(PrefixedIo { + prefix, + position: 0, + inner: stream, + })); + } + return Err(tunnel_error(status, retry_after)); + } + if response.len() == MAX_CONNECT_RESPONSE_HEADER_BYTES { + return Err(invalid_tunnel_response( + "WebSocket proxy response headers were too large", + )); + } + let remaining = MAX_CONNECT_RESPONSE_HEADER_BYTES - response.len(); + let mut buffer = [0_u8; 1024]; + let capacity = remaining.min(buffer.len()); + let read = stream + .read(&mut buffer[..capacity]) + .await + .map_err(|_| invalid_tunnel_response("WebSocket proxy response could not be read"))?; + if read == 0 { + return Err(invalid_tunnel_response( + "WebSocket proxy closed the tunnel response early", + )); + } + response.extend_from_slice(&buffer[..read]); + } +} + +async fn tls_connect( + stream: BoxedWebSocketIo, + host: &str, + tls_config: Arc, + peer: &str, +) -> Result { + let server_name = rustls::pki_types::ServerName::try_from(host.to_string()) + .map_err(|_| websocket_protocol_error("WebSocket TLS host name was invalid"))?; + let stream = tokio_rustls::TlsConnector::from(tls_config) + .connect(server_name, stream) .await .map_err(|_| CodexError { status: 0, - message: format!("WebSocket connect timeout after {connect_timeout_ms}ms"), + message: format!("WebSocket TLS connection to {peer} failed"), detail: None, retry_after: None, origin: CodexErrorOrigin::WebSocketHandshake, - })? - .map_err(|e| { - let (status, retry_after, detail) = match &e { - tungstenite::Error::Http(response) => { - let detail = Some(handshake_error_detail(response.body().as_ref())); - ( - Some(response.status().as_u16()), - response - .headers() - .get(http::header::RETRY_AFTER) - .and_then(|value| value.to_str().ok()) - .map(str::to_string), - detail, - ) - } - _ => (None, None, None), - }; - CodexError { - status: status.unwrap_or(0), - message: format!("WebSocket connect error: {e}"), - detail, - retry_after, - origin: CodexErrorOrigin::WebSocketHandshake, - } - }) + })?; + Ok(Box::new(stream)) } -fn websocket_host_header(url: &str) -> String { - let Ok(parsed) = url::Url::parse(url) else { - return String::new(); +async fn connect_to_http_proxy( + route: &WebSocketProxyRoute, + tls_config: Arc, +) -> Result { + let scheme = route.uri.scheme_str().unwrap_or_default(); + let host = route + .uri + .host() + .ok_or_else(|| websocket_protocol_error("WebSocket proxy URL did not contain a host"))?; + let host = host.trim_start_matches('[').trim_end_matches(']'); + let port = route + .uri + .port_u16() + .unwrap_or(if scheme == "https" { 443 } else { 80 }); + let stream = TcpStream::connect((host, port)) + .await + .map_err(|_| CodexError { + status: 0, + message: "WebSocket proxy connection failed".to_string(), + detail: None, + retry_after: None, + origin: CodexErrorOrigin::WebSocketHandshake, + })?; + let stream: BoxedWebSocketIo = Box::new(stream); + if scheme == "https" { + tls_connect(stream, host, tls_config, "proxy").await + } else { + Ok(stream) + } +} + +fn websocket_destination(url: &str) -> Result<(String, String), CodexError> { + let destination = url::Url::parse(url) + .map_err(|_| websocket_protocol_error("WebSocket destination URL was invalid"))?; + let host = destination + .host_str() + .ok_or_else(|| websocket_protocol_error("WebSocket destination URL had no host"))?; + let port = destination + .port_or_known_default() + .ok_or_else(|| websocket_protocol_error("WebSocket destination URL had no usable port"))?; + let authority = match destination.host() { + Some(url::Host::Ipv6(address)) => format!("[{address}]:{port}"), + Some(_) => format!("{host}:{port}"), + None => { + return Err(websocket_protocol_error( + "WebSocket destination URL had no host", + )); + } }; - parsed[url::Position::BeforeHost..url::Position::AfterPort].to_string() + Ok((host.to_string(), authority)) +} + +fn tungstenite_handshake_error(error: tokio_tungstenite::tungstenite::Error) -> CodexError { + if let tokio_tungstenite::tungstenite::Error::Http(response) = error { + let status = response.status().as_u16(); + let retry_after = response + .headers() + .get(http::header::RETRY_AFTER) + .and_then(|value| value.to_str().ok()) + .map(str::to_string); + return CodexError { + status, + message: format!("WebSocket upgrade rejected with status {status}"), + detail: Some(GENERIC_HANDSHAKE_ERROR_DETAIL.to_string()), + retry_after, + origin: CodexErrorOrigin::WebSocketHandshake, + }; + } + CodexError { + status: 0, + message: "WebSocket upgrade request failed".to_string(), + detail: None, + retry_after: None, + origin: CodexErrorOrigin::WebSocketHandshake, + } +} + +async fn connect_via_http_proxy_tunnel( + proxy_config: &WebSocketProxyConfig, + route: WebSocketProxyRoute, + url: &str, + headers: &HeaderMap, +) -> Result { + let (host, authority) = websocket_destination(url)?; + let stream = connect_to_http_proxy(&route, proxy_config.tls_config.clone()).await?; + let stream = establish_connect_tunnel(stream, &authority, route.basic_auth.as_ref()).await?; + let stream = tls_connect( + stream, + &host, + proxy_config.tls_config.clone(), + "destination", + ) + .await?; + let websocket_key = generate_key(); + let subprotocols = requested_subprotocols(headers); + let request = tunneled_websocket_request(url, headers, &websocket_key)?; + let (websocket, response) = tokio_tungstenite::client_async(request, stream) + .await + .map_err(tungstenite_handshake_error)?; + validate_websocket_upgrade( + response.version(), + response.headers(), + &websocket_key, + &subprotocols, + )?; + Ok(websocket) +} + +async fn connect_via_http_upgrade( + websocket_client: &reqwest::Client, + url: &str, + headers: &HeaderMap, +) -> Result { + let http_url = to_http_upgrade_url(url).map_err(|error| CodexError { + status: 0, + message: error.message, + detail: None, + retry_after: None, + origin: CodexErrorOrigin::WebSocketHandshake, + })?; + let websocket_key = generate_key(); + let subprotocols = requested_subprotocols(headers); + let mut request = websocket_client + .get(http_url) + .version(http::Version::HTTP_11) + .header(http::header::CONNECTION, "Upgrade") + .header(http::header::UPGRADE, "websocket") + .header(http::header::SEC_WEBSOCKET_VERSION, "13") + .header(http::header::SEC_WEBSOCKET_KEY, &websocket_key); + + for (key, value) in headers { + if skip_websocket_request_header(key) { + continue; + } + request = request.header(key.clone(), value.clone()); + } + + let response = request.send().await.map_err(reqwest_handshake_error)?; + if response.status() != http::StatusCode::SWITCHING_PROTOCOLS { + let status = response.status().as_u16(); + let retry_after = response + .headers() + .get(http::header::RETRY_AFTER) + .and_then(|value| value.to_str().ok()) + .map(str::to_string); + let detail = if status == http::StatusCode::PROXY_AUTHENTICATION_REQUIRED.as_u16() { + GENERIC_HANDSHAKE_ERROR_DETAIL.to_string() + } else { + let body = bounded_handshake_error_body(response).await; + handshake_error_detail(Some(&body)) + }; + return Err(CodexError { + status, + message: format!("WebSocket upgrade rejected with status {status}"), + detail: Some(detail), + retry_after, + origin: CodexErrorOrigin::WebSocketHandshake, + }); + } + + validate_websocket_upgrade( + response.version(), + response.headers(), + &websocket_key, + &subprotocols, + )?; + let upgraded = response + .upgrade() + .await + .map_err(|_| websocket_protocol_error("WebSocket upgrade stream was not available"))?; + let upgraded: BoxedWebSocketIo = Box::new(upgraded); + Ok(WebSocketStream::from_raw_socket(upgraded, Role::Client, None).await) +} + +async fn connect_with_timeout( + websocket_client: &reqwest::Client, + proxy_config: &WebSocketProxyConfig, + url: &str, + headers: &HeaderMap, + connect_timeout_ms: u64, +) -> Result { + let connect = async { + if let Some(route) = proxy_config.http_connect_route(url)? { + connect_via_http_proxy_tunnel(proxy_config, route, url, headers).await + } else { + connect_via_http_upgrade(websocket_client, url, headers).await + } + }; + tokio::time::timeout(Duration::from_millis(connect_timeout_ms), connect) + .await + .map_err(|_| CodexError { + status: 0, + message: format!("WebSocket connect timeout after {connect_timeout_ms}ms"), + detail: None, + retry_after: None, + origin: CodexErrorOrigin::WebSocketHandshake, + })? } // --------------------------------------------------------------------------- @@ -891,13 +1580,16 @@ struct WsEvent { payload: serde_json::Value, } -async fn collect_ws_events( - ws: &mut WebSocketStream>, +async fn collect_ws_events( + ws: &mut WebSocketStream, idle_timeout_ms: u64, pool_key: Option<&str>, pool_entry: Option<&Arc>, traffic: Option<&TrafficCapture>, -) -> Result<(Vec, Option), CodexError> { +) -> Result<(Vec, Option), CodexError> +where + S: AsyncRead + AsyncWrite + Unpin, +{ let mut sse_body: Vec = Vec::new(); let mut terminal_event: Option = None; let response_event_budget = Duration::from_millis(idle_timeout_ms); @@ -1047,14 +1739,17 @@ async fn collect_ws_events( Ok((sse_body, terminal_event)) } -async fn stream_ws_events( - ws: &mut WebSocketStream>, +async fn stream_ws_events( + ws: &mut WebSocketStream, idle_timeout_ms: u64, pool_key: Option<&str>, pool_entry: Option<&Arc>, traffic: Option>, tx: mpsc::Sender>, -) -> bool { +) -> bool +where + S: AsyncRead + AsyncWrite + Unpin, +{ let started_at = Instant::now(); let mut sse_body: Vec = Vec::new(); let response_event_budget = Duration::from_millis(idle_timeout_ms); @@ -1248,6 +1943,28 @@ fn summarize_json_request_size(body: &serde_json::Value, body_json: &str) -> ser mod tests { use super::*; + static WS_POOL_TEST_LOCK: AsyncMutex<()> = AsyncMutex::const_new(()); + + async fn lock_ws_pool_tests() -> tokio::sync::MutexGuard<'static, ()> { + WS_POOL_TEST_LOCK.lock().await + } + + fn test_websocket_client() -> reqwest::Client { + reqwest::Client::builder() + .http1_only() + .redirect(reqwest::redirect::Policy::none()) + .no_proxy() + .build() + .unwrap() + } + + #[test] + fn websocket_tls_configuration_is_shared() { + let first = websocket_tls_config(); + let second = websocket_tls_config(); + assert!(Arc::ptr_eq(&first, &second)); + } + #[test] fn event_error_status_requires_error_event_and_checks_numeric_fallbacks() { assert_eq!( @@ -1292,16 +2009,105 @@ mod tests { } #[test] - fn websocket_host_header_preserves_explicit_port() { + fn websocket_upgrade_url_preserves_authority_path_and_query() { + assert_eq!( + to_http_upgrade_url("wss://chatgpt.com/backend-api/codex/responses?mode=live").unwrap(), + "https://chatgpt.com/backend-api/codex/responses?mode=live" + ); assert_eq!( - websocket_host_header("wss://chatgpt.com/backend-api/codex/responses"), - "chatgpt.com" + to_http_upgrade_url("ws://127.0.0.1:4141/backend-api/codex/responses").unwrap(), + "http://127.0.0.1:4141/backend-api/codex/responses" ); assert_eq!( - websocket_host_header("ws://127.0.0.1:4141/backend-api/codex/responses"), - "127.0.0.1:4141" + to_http_upgrade_url("ws://[::1]:4141/path").unwrap(), + "http://[::1]:4141/path" + ); + } + + #[test] + fn validates_tokenized_websocket_upgrade_headers() { + let key = "dGhlIHNhbXBsZSBub25jZQ=="; + let mut headers = HeaderMap::new(); + headers.insert(http::header::UPGRADE, "h2c, WebSocket".parse().unwrap()); + headers.insert( + http::header::CONNECTION, + "keep-alive, Upgrade".parse().unwrap(), + ); + headers.insert( + http::header::SEC_WEBSOCKET_ACCEPT, + derive_accept_key(key.as_bytes()).parse().unwrap(), + ); + headers.insert( + http::header::SEC_WEBSOCKET_PROTOCOL, + "responses".parse().unwrap(), + ); + + validate_websocket_upgrade( + http::Version::HTTP_11, + &headers, + key, + &["responses".to_string()], + ) + .unwrap(); + } + + #[test] + fn rejects_invalid_websocket_upgrade_headers() { + let key = "dGhlIHNhbXBsZSBub25jZQ=="; + let valid = || { + let mut headers = HeaderMap::new(); + headers.insert(http::header::UPGRADE, "websocket".parse().unwrap()); + headers.insert(http::header::CONNECTION, "Upgrade".parse().unwrap()); + headers.insert( + http::header::SEC_WEBSOCKET_ACCEPT, + derive_accept_key(key.as_bytes()).parse().unwrap(), + ); + headers + }; + + let mut missing_upgrade = valid(); + missing_upgrade.remove(http::header::UPGRADE); + assert!( + validate_websocket_upgrade(http::Version::HTTP_11, &missing_upgrade, key, &[]).is_err() + ); + + let mut missing_connection = valid(); + missing_connection.remove(http::header::CONNECTION); + assert!( + validate_websocket_upgrade(http::Version::HTTP_11, &missing_connection, key, &[]) + .is_err() + ); + + let mut wrong_accept = valid(); + wrong_accept.insert(http::header::SEC_WEBSOCKET_ACCEPT, "wrong".parse().unwrap()); + assert!( + validate_websocket_upgrade(http::Version::HTTP_11, &wrong_accept, key, &[]).is_err() + ); + + let unsolicited_extension = { + let mut headers = valid(); + headers.insert( + http::header::SEC_WEBSOCKET_EXTENSIONS, + "permessage-deflate".parse().unwrap(), + ); + headers + }; + assert!( + validate_websocket_upgrade(http::Version::HTTP_11, &unsolicited_extension, key, &[],) + .is_err() + ); + + assert!(validate_websocket_upgrade(http::Version::HTTP_10, &valid(), key, &[]).is_err()); + + let mut unsolicited_protocol = valid(); + unsolicited_protocol.insert( + http::header::SEC_WEBSOCKET_PROTOCOL, + "unexpected".parse().unwrap(), + ); + assert!( + validate_websocket_upgrade(http::Version::HTTP_11, &unsolicited_protocol, key, &[]) + .is_err() ); - assert_eq!(websocket_host_header("ws://[::1]:4141/path"), "[::1]:4141"); } #[test] @@ -1310,9 +2116,14 @@ mod tests { headers.insert("openai-beta", "responses=experimental".parse().unwrap()); headers.insert("content-length", "10".parse().unwrap()); headers.insert("authorization", "Bearer tok".parse().unwrap()); + headers.insert( + http::header::PROXY_AUTHORIZATION, + "Basic dXNlcjpwYXNz".parse().unwrap(), + ); let ws = codex_websocket_headers(&headers); assert_eq!(ws.get("openai-beta").unwrap(), WEBSOCKET_PROTOCOL_HEADER); assert!(!ws.contains_key("content-length")); + assert!(!ws.contains_key(http::header::PROXY_AUTHORIZATION)); assert_eq!(ws.get("authorization").unwrap(), "Bearer tok"); } @@ -1391,6 +2202,7 @@ mod tests { #[tokio::test] async fn pool_checkout_is_exclusive_and_removal_is_identity_safe() { + let _pool_test_guard = lock_ws_pool_tests().await; clear_codex_websocket_pool_for_tests(); let first = Arc::new(PoolEntry { ws: Arc::new(AsyncMutex::new(create_dummy_stream_async().await)), @@ -1416,17 +2228,19 @@ mod tests { clear_codex_websocket_pool_for_tests(); } - #[test] - fn pool_invalidation() { + #[tokio::test] + async fn pool_invalidation() { + let _pool_test_guard = lock_ws_pool_tests().await; clear_codex_websocket_pool_for_tests(); // Verify pool operations work through the public API // We insert an entry directly into the pool, then invalidate it + let stream = create_dummy_stream_async().await; { let mut guard = WS_POOL.lock().unwrap(); guard.insert( "test-session".to_string(), Arc::new(PoolEntry { - ws: Arc::new(AsyncMutex::new(create_dummy_stream())), + ws: Arc::new(AsyncMutex::new(stream)), created_at: now_ms(), }), ); @@ -1454,7 +2268,10 @@ mod tests { .unwrap(); }); + let client = test_websocket_client(); let err = match connect_with_timeout( + &client, + &WebSocketProxyConfig::direct(), &format!("ws://{addr}/backend-api/codex/responses"), &HeaderMap::new(), 1_000, @@ -1489,7 +2306,10 @@ mod tests { .unwrap(); }); + let client = test_websocket_client(); let err = match connect_with_timeout( + &client, + &WebSocketProxyConfig::direct(), &format!("ws://{addr}/backend-api/codex/responses"), &HeaderMap::new(), 1_000, @@ -1506,8 +2326,245 @@ mod tests { assert_eq!(err.origin, CodexErrorOrigin::WebSocketHandshake); } + #[tokio::test] + async fn websocket_connects_through_explicit_http_proxy() { + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + use tokio::net::TcpListener; + + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let proxy_addr = listener.local_addr().unwrap(); + let (captured_tx, captured_rx) = tokio::sync::oneshot::channel(); + let proxy = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.unwrap(); + let mut request = Vec::new(); + let mut buffer = [0_u8; 1024]; + while !request.windows(4).any(|window| window == b"\r\n\r\n") { + let read = stream.read(&mut buffer).await.unwrap(); + assert!(read > 0); + request.extend_from_slice(&buffer[..read]); + } + let request_text = String::from_utf8(request).unwrap(); + let key = request_text + .lines() + .find_map(|line| { + let (name, value) = line.split_once(':')?; + name.eq_ignore_ascii_case("sec-websocket-key") + .then(|| value.trim().to_string()) + }) + .unwrap(); + let _ = captured_tx.send(request_text); + let response = format!( + "HTTP/1.1 101 Switching Protocols\r\nConnection: Upgrade\r\nUpgrade: websocket\r\nSec-WebSocket-Accept: {}\r\n\r\n", + derive_accept_key(key.as_bytes()) + ); + stream.write_all(response.as_bytes()).await.unwrap(); + + let mut websocket = WebSocketStream::from_raw_socket(stream, Role::Server, None).await; + assert_eq!( + websocket.next().await.unwrap().unwrap(), + Message::Text("hello".to_string()) + ); + websocket + .send(Message::Text("proxy-ok".to_string())) + .await + .unwrap(); + }); + + let client = reqwest::Client::builder() + .http1_only() + .redirect(reqwest::redirect::Policy::none()) + .proxy( + reqwest::Proxy::http(format!("http://proxy-user:proxy-pass@{proxy_addr}")).unwrap(), + ) + .build() + .unwrap(); + let mut websocket = connect_with_timeout( + &client, + &WebSocketProxyConfig::direct(), + "ws://codex.invalid/backend-api/codex/responses", + &HeaderMap::new(), + 2_000, + ) + .await + .unwrap(); + websocket + .send(Message::Text("hello".to_string())) + .await + .unwrap(); + assert_eq!( + websocket.next().await.unwrap().unwrap(), + Message::Text("proxy-ok".to_string()) + ); + + let captured = captured_rx.await.unwrap(); + assert!( + captured.starts_with("GET http://codex.invalid/backend-api/codex/responses HTTP/1.1") + ); + assert!( + captured + .to_ascii_lowercase() + .contains("proxy-authorization: basic ") + ); + proxy.await.unwrap(); + } + + #[tokio::test] + async fn connect_tunnel_accepts_fragmented_non_200_success() { + let (client, mut proxy) = tokio::io::duplex(4096); + let proxy_task = tokio::spawn(async move { + let mut request = Vec::new(); + let mut buffer = [0_u8; 256]; + while !request.windows(4).any(|window| window == b"\r\n\r\n") { + let read = proxy.read(&mut buffer).await.unwrap(); + assert!(read > 0); + request.extend_from_slice(&buffer[..read]); + } + assert!(request.starts_with(b"CONNECT codex.invalid:4443 HTTP/1.1\r\n")); + proxy.write_all(b"HTTP/1.").await.unwrap(); + tokio::task::yield_now().await; + proxy + .write_all(b"1 204 No Content\r\nX-Proxy: ok\r\n\r\nprefixed") + .await + .unwrap(); + }); + + let stream: BoxedWebSocketIo = Box::new(client); + let mut stream = establish_connect_tunnel(stream, "codex.invalid:4443", None) + .await + .unwrap(); + let mut prefix = [0_u8; 8]; + stream.read_exact(&mut prefix).await.unwrap(); + assert_eq!(&prefix, b"prefixed"); + proxy_task.await.unwrap(); + } + + #[tokio::test] + async fn connect_tunnel_classifies_fragmented_proxy_authentication() { + let (client, mut proxy) = tokio::io::duplex(4096); + tokio::spawn(async move { + let mut request = [0_u8; 512]; + let _ = proxy.read(&mut request).await.unwrap(); + proxy.write_all(b"HTTP/1.1 4").await.unwrap(); + tokio::task::yield_now().await; + proxy + .write_all(b"07 Proxy Authentication Required\r\n\r\n") + .await + .unwrap(); + }); + + let stream: BoxedWebSocketIo = Box::new(client); + let error = match establish_connect_tunnel(stream, "codex.invalid:4443", None).await { + Ok(_) => panic!("proxy authentication should be rejected"), + Err(error) => error, + }; + assert_eq!( + error.status, + http::StatusCode::PROXY_AUTHENTICATION_REQUIRED.as_u16() + ); + assert_eq!(error.message, "WebSocket proxy authentication failed"); + } + + #[test] + fn websocket_proxy_routing_honors_no_proxy_and_leaves_socks_to_reqwest() { + let http_proxy = "http://proxy.example:8080"; + let config = WebSocketProxyConfig::new(None, Some(http_proxy), None, None); + assert!(config.uses_proxy_for("wss://codex.invalid/responses")); + assert!( + config + .http_connect_route("wss://codex.invalid/responses") + .unwrap() + .is_some() + ); + + let bypass = WebSocketProxyConfig::new(None, Some(http_proxy), None, Some("codex.invalid")); + assert!(!bypass.uses_proxy_for("wss://codex.invalid/responses")); + assert!( + bypass + .http_connect_route("wss://codex.invalid/responses") + .unwrap() + .is_none() + ); + + let socks = + WebSocketProxyConfig::new(None, Some("socks5h://proxy.example:1080"), None, None); + assert!(socks.uses_proxy_for("wss://codex.invalid/responses")); + assert!( + socks + .http_connect_route("wss://codex.invalid/responses") + .unwrap() + .is_none() + ); + } + + #[tokio::test] + async fn websocket_wss_uses_http_connect_without_leaking_proxy_credentials() { + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + use tokio::net::TcpListener; + + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let proxy_addr = listener.local_addr().unwrap(); + let (captured_tx, captured_rx) = tokio::sync::oneshot::channel(); + let proxy = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.unwrap(); + let mut request = Vec::new(); + let mut buffer = [0_u8; 1024]; + while !request.windows(4).any(|window| window == b"\r\n\r\n") { + let read = stream.read(&mut buffer).await.unwrap(); + assert!(read > 0); + request.extend_from_slice(&buffer[..read]); + } + let _ = captured_tx.send(String::from_utf8(request).unwrap()); + stream + .write_all( + b"HTTP/1.1 407 Proxy Authentication Required\r\nContent-Length: 0\r\n\r\n", + ) + .await + .unwrap(); + }); + + let proxy_url = format!("http://secret-user:secret-pass@{proxy_addr}"); + let client = reqwest::Client::builder() + .http1_only() + .redirect(reqwest::redirect::Policy::none()) + .proxy(reqwest::Proxy::https(&proxy_url).unwrap()) + .build() + .unwrap(); + let proxy_config = WebSocketProxyConfig::new(None, Some(&proxy_url), None, None); + let error = match connect_with_timeout( + &client, + &proxy_config, + "wss://codex.invalid:4443/backend-api/codex/responses", + &HeaderMap::new(), + 2_000, + ) + .await + { + Ok(_) => panic!("proxy rejection should fail the WebSocket connection"), + Err(error) => error, + }; + + let captured = captured_rx.await.unwrap(); + assert!(captured.starts_with("CONNECT codex.invalid:4443 HTTP/1.1")); + assert!( + captured + .to_ascii_lowercase() + .contains("proxy-authorization: basic ") + ); + assert!(!error.message.contains("secret-user")); + assert!(!error.message.contains("secret-pass")); + assert!( + !error + .detail + .as_deref() + .unwrap_or_default() + .contains("secret") + ); + proxy.await.unwrap(); + } + #[tokio::test] async fn binary_frame_invalidates_pool_key() { + let _pool_test_guard = lock_ws_pool_tests().await; clear_codex_websocket_pool_for_tests(); let pooled_stream = create_dummy_stream_async().await; { @@ -1544,6 +2601,7 @@ mod tests { #[tokio::test] async fn response_start_timeout_ignores_rate_limits_and_pings() { + let _pool_test_guard = lock_ws_pool_tests().await; clear_codex_websocket_pool_for_tests(); let pooled_stream = create_dummy_stream_async().await; { @@ -1598,6 +2656,7 @@ mod tests { #[tokio::test] async fn response_idle_timeout_ignores_pings_after_response_event() { + let _pool_test_guard = lock_ws_pool_tests().await; clear_codex_websocket_pool_for_tests(); let pooled_stream = create_dummy_stream_async().await; { @@ -1649,7 +2708,7 @@ mod tests { ); } - async fn create_dummy_stream_async() -> WebSocketStream> { + async fn create_dummy_stream_async() -> CodexWebSocketStream { use tokio::net::TcpListener; let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); @@ -1660,14 +2719,16 @@ mod tests { futures_util::future::pending::<()>().await; }); let url = format!("ws://{addr}/"); - let (ws, _) = tokio::time::timeout( - Duration::from_millis(1000), - tokio_tungstenite::connect_async(&url), + let client = test_websocket_client(); + connect_with_timeout( + &client, + &WebSocketProxyConfig::direct(), + &url, + &HeaderMap::new(), + 1_000, ) .await .unwrap() - .unwrap(); - ws } #[test] @@ -1696,32 +2757,4 @@ mod tests { ); } } - - fn create_dummy_stream() -> WebSocketStream> { - // Use a connected TcpStream pair with connect_async which returns - // WebSocketStream> - use tokio::net::TcpListener; - let rt = tokio::runtime::Runtime::new().unwrap(); - rt.block_on(async { - let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let addr = listener.local_addr().unwrap(); - let _conn = tokio::spawn(async move { - let (socket, _) = listener.accept().await.unwrap(); - // Accept WebSocket handshake - let _ = tokio_tungstenite::accept_async(socket).await; - // Keep alive - futures_util::future::pending::<()>().await; - }); - // Use connect_async to get MaybeTlsStream - let url = format!("ws://{}/", addr); - let (ws, _) = tokio::time::timeout( - Duration::from_millis(1000), - tokio_tungstenite::connect_async(&url), - ) - .await - .unwrap() - .unwrap(); - ws - }) - } } diff --git a/src/traffic.rs b/src/traffic.rs index 36d60efa..571d7ee4 100644 --- a/src/traffic.rs +++ b/src/traffic.rs @@ -634,4 +634,20 @@ mod tests { "nested payload leaked: {rendered}" ); } + + #[test] + fn traffic_redacts_proxy_authorization() { + let redacted = redact_traffic(&serde_json::json!({ + "headers": { + "proxy-authorization": "Basic dXNlcjpwYXNz", + "x-safe": "kept" + } + })); + + assert_eq!(redacted["headers"]["x-safe"], "kept"); + assert_eq!( + redacted["headers"]["proxy-authorization"], + "[redacted len=18]" + ); + } } diff --git a/tests/codex_websocket_proxy.rs b/tests/codex_websocket_proxy.rs new file mode 100644 index 00000000..2f68e917 --- /dev/null +++ b/tests/codex_websocket_proxy.rs @@ -0,0 +1,478 @@ +use axum::body::Body; +use axum::http::{Method, Request, StatusCode}; +use claude_code_proxy::providers::codex::websocket::clear_codex_websocket_pool_for_tests; +use claude_code_proxy::{registry::Registry, server::app}; +use futures_util::{SinkExt, StreamExt}; +use http_body_util::BodyExt; +use serde_json::json; +use std::path::Path; +use std::sync::{Arc, Mutex}; +use std::time::Duration; +use tempfile::TempDir; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::net::{TcpListener, TcpStream}; +use tokio::sync::oneshot; +use tokio_tungstenite::tungstenite::Message; +use tower::util::ServiceExt; + +struct EnvGuard { + key: &'static str, + previous: Option, +} + +impl EnvGuard { + fn set(key: &'static str, value: impl AsRef) -> Self { + let previous = std::env::var_os(key); + unsafe { + std::env::set_var(key, value); + } + Self { key, previous } + } + + fn unset(key: &'static str) -> Self { + let previous = std::env::var_os(key); + unsafe { + std::env::remove_var(key); + } + Self { key, previous } + } +} + +impl Drop for EnvGuard { + fn drop(&mut self) { + unsafe { + match self.previous.take() { + Some(value) => std::env::set_var(self.key, value), + None => std::env::remove_var(self.key), + } + } + } +} + +struct ZeroRetryDelayGuard; + +impl ZeroRetryDelayGuard { + fn new() -> Self { + claude_code_proxy::retry::set_zero_retry_delay_for_tests(true); + Self + } +} + +impl Drop for ZeroRetryDelayGuard { + fn drop(&mut self) { + claude_code_proxy::retry::set_zero_retry_delay_for_tests(false); + } +} + +fn clear_proxy_environment() -> Vec { + [ + "HTTP_PROXY", + "http_proxy", + "HTTPS_PROXY", + "https_proxy", + "ALL_PROXY", + "all_proxy", + "NO_PROXY", + "no_proxy", + "REQUEST_METHOD", + ] + .into_iter() + .map(EnvGuard::unset) + .collect() +} + +fn configure_codex(config_dir: &Path, base_url: &str) -> Vec { + vec![ + EnvGuard::set("CCP_CONFIG_DIR", config_dir), + EnvGuard::set("CCP_ALIAS_PROVIDER", "codex"), + EnvGuard::set("CCP_CODEX_TRANSPORT", "websocket"), + EnvGuard::set("CCP_CODEX_BASE_URL", base_url), + EnvGuard::set("CCP_CODEX_PREVIOUS_RESPONSE_ID", "0"), + EnvGuard::set("CCP_CODEX_SERVER_COMPACTION", "0"), + ] +} + +fn write_codex_auth(config_dir: &Path) { + let dir = config_dir.join("codex"); + std::fs::create_dir_all(&dir).unwrap(); + std::fs::write( + dir.join("auth.json"), + serde_json::to_vec(&json!({ + "access": "test-access", + "refresh": "test-refresh", + "expires": 4_102_444_800_000_i64, + "account_id": "acct_test" + })) + .unwrap(), + ) + .unwrap(); +} + +async fn call_messages(session_id: &str) -> (StatusCode, String) { + call_messages_with_stream(session_id, false).await +} + +async fn call_messages_with_stream(session_id: &str, stream: bool) -> (StatusCode, String) { + let response = app(Arc::new(Registry::with_default_alias())) + .oneshot( + Request::builder() + .method(Method::POST) + .uri("/v1/messages") + .header("content-type", "application/json") + .header("x-claude-code-session-id", session_id) + .body(Body::from( + json!({ + "model": "gpt-5.6-sol", + "max_tokens": 64, + "stream": stream, + "messages": [{"role": "user", "content": "hello"}] + }) + .to_string(), + )) + .unwrap(), + ) + .await + .unwrap(); + let status = response.status(); + let body = response.into_body().collect().await.unwrap().to_bytes(); + (status, String::from_utf8_lossy(&body).into_owned()) +} + +async fn send_codex_response(websocket: &mut tokio_tungstenite::WebSocketStream) { + let request = websocket.next().await.unwrap().unwrap(); + assert!(matches!(request, Message::Text(_))); + for event in [ + r#"{"type":"response.output_item.added","output_index":0,"item":{"type":"message","id":"msg_proxy"}}"#, + r#"{"type":"response.output_text.delta","output_index":0,"delta":"proxy env ok"}"#, + r#"{"type":"response.output_item.done","output_index":0,"item":{"type":"message"}}"#, + r#"{"type":"response.completed","response":{"id":"resp_proxy","usage":{"input_tokens":5,"output_tokens":2}}}"#, + ] { + websocket + .send(Message::Text(event.to_string())) + .await + .unwrap(); + } +} + +#[derive(Debug)] +struct CapturedProxyRequest { + target: String, + has_proxy_authorization: bool, +} + +#[allow(clippy::result_large_err)] +async fn spawn_websocket_proxy() -> ( + String, + oneshot::Receiver, + tokio::task::JoinHandle<()>, +) { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let (captured_tx, captured_rx) = oneshot::channel(); + let task = tokio::spawn(async move { + let (stream, _) = listener.accept().await.unwrap(); + let websocket = tokio_tungstenite::accept_hdr_async( + stream, + move |request: &tokio_tungstenite::tungstenite::handshake::server::Request, + response| { + let _ = captured_tx.send(CapturedProxyRequest { + target: request.uri().to_string(), + has_proxy_authorization: request + .headers() + .contains_key(http::header::PROXY_AUTHORIZATION), + }); + Ok(response) + }, + ) + .await + .unwrap(); + let mut websocket = websocket; + send_codex_response(&mut websocket).await; + }); + ( + format!("http://proxy-user:proxy-pass@{addr}"), + captured_rx, + task, + ) +} + +async fn spawn_direct_websocket() -> (String, tokio::task::JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let task = tokio::spawn(async move { + let (stream, _) = listener.accept().await.unwrap(); + let mut websocket = tokio_tungstenite::accept_async(stream).await.unwrap(); + send_codex_response(&mut websocket).await; + }); + (format!("http://{addr}/responses"), task) +} + +async fn read_http_head(stream: &mut TcpStream) -> String { + let mut request = Vec::new(); + let mut buffer = [0_u8; 1024]; + while !request.windows(4).any(|window| window == b"\r\n\r\n") { + let read = stream.read(&mut buffer).await.unwrap(); + if read == 0 { + break; + } + request.extend_from_slice(&buffer[..read]); + assert!(request.len() <= 16 * 1024); + } + String::from_utf8(request).unwrap() +} + +async fn spawn_rejecting_proxy( + response: &'static [u8], +) -> ( + String, + Arc>>, + oneshot::Sender<()>, + tokio::task::JoinHandle<()>, +) { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let captured = Arc::new(Mutex::new(Vec::new())); + let task_captured = captured.clone(); + let (stop_tx, mut stop_rx) = oneshot::channel(); + let task = tokio::spawn(async move { + loop { + let accepted = tokio::select! { + _ = &mut stop_rx => break, + accepted = listener.accept() => accepted, + }; + let Ok((mut stream, _)) = accepted else { + break; + }; + let request = read_http_head(&mut stream).await; + let mut lines = request.lines(); + let target = lines.next().unwrap_or_default().to_string(); + let has_proxy_authorization = lines.any(|line| { + line.split_once(':') + .is_some_and(|(name, _)| name.eq_ignore_ascii_case("proxy-authorization")) + }); + task_captured.lock().unwrap().push(CapturedProxyRequest { + target, + has_proxy_authorization, + }); + stream.write_all(response).await.unwrap(); + } + }); + ( + format!("http://secret-user:secret-pass@{addr}"), + captured, + stop_tx, + task, + ) +} + +#[tokio::test] +async fn codex_websocket_inherits_environment_proxy_configuration() { + let config_dir = TempDir::new().unwrap(); + write_codex_auth(config_dir.path()); + + clear_codex_websocket_pool_for_tests(); + let (proxy_url, captured, proxy_task) = spawn_websocket_proxy().await; + let all_proxy_trap = TcpListener::bind("127.0.0.1:0").await.unwrap(); + { + let mut guards = clear_proxy_environment(); + guards.extend(configure_codex( + config_dir.path(), + "http://codex.invalid/backend-api/codex/responses", + )); + guards.push(EnvGuard::set("HTTP_PROXY", &proxy_url)); + guards.push(EnvGuard::set( + "ALL_PROXY", + format!("http://{}", all_proxy_trap.local_addr().unwrap()), + )); + let (status, body) = call_messages("proxy-env-http").await; + assert_eq!(status, StatusCode::OK); + assert!(body.contains("proxy env ok")); + } + let captured = captured.await.unwrap(); + assert_eq!( + captured.target, + "http://codex.invalid/backend-api/codex/responses" + ); + assert!(captured.has_proxy_authorization); + proxy_task.await.unwrap(); + assert!( + tokio::time::timeout(Duration::from_millis(100), all_proxy_trap.accept()) + .await + .is_err() + ); + + clear_codex_websocket_pool_for_tests(); + let (all_proxy_url, captured, proxy_task) = spawn_websocket_proxy().await; + { + let mut guards = clear_proxy_environment(); + guards.extend(configure_codex( + config_dir.path(), + "http://all-proxy.invalid/backend-api/codex/responses", + )); + guards.push(EnvGuard::set("ALL_PROXY", &all_proxy_url)); + let (status, body) = call_messages("proxy-env-all").await; + assert_eq!(status, StatusCode::OK); + assert!(body.contains("proxy env ok")); + } + let captured = captured.await.unwrap(); + assert_eq!( + captured.target, + "http://all-proxy.invalid/backend-api/codex/responses" + ); + proxy_task.await.unwrap(); + + clear_codex_websocket_pool_for_tests(); + let (direct_url, direct_task) = spawn_direct_websocket().await; + let cgi_proxy_trap = TcpListener::bind("127.0.0.1:0").await.unwrap(); + { + let mut guards = clear_proxy_environment(); + guards.extend(configure_codex(config_dir.path(), &direct_url)); + guards.push(EnvGuard::set( + "HTTP_PROXY", + format!("http://{}", cgi_proxy_trap.local_addr().unwrap()), + )); + guards.push(EnvGuard::set("REQUEST_METHOD", "GET")); + let (status, body) = call_messages("proxy-env-cgi-safe").await; + assert_eq!(status, StatusCode::OK); + assert!(body.contains("proxy env ok")); + } + direct_task.await.unwrap(); + assert!( + tokio::time::timeout(Duration::from_millis(100), cgi_proxy_trap.accept()) + .await + .is_err() + ); + + clear_codex_websocket_pool_for_tests(); + let (direct_url, direct_task) = spawn_direct_websocket().await; + let trap_proxy = TcpListener::bind("127.0.0.1:0").await.unwrap(); + { + let mut guards = clear_proxy_environment(); + guards.extend(configure_codex(config_dir.path(), &direct_url)); + guards.push(EnvGuard::set( + "HTTP_PROXY", + format!("http://{}", trap_proxy.local_addr().unwrap()), + )); + guards.push(EnvGuard::set("NO_PROXY", "127.0.0.1")); + let (status, body) = call_messages("proxy-env-no-proxy").await; + assert_eq!(status, StatusCode::OK); + assert!(body.contains("proxy env ok")); + } + direct_task.await.unwrap(); + assert!( + tokio::time::timeout(Duration::from_millis(100), trap_proxy.accept()) + .await + .is_err() + ); + + clear_codex_websocket_pool_for_tests(); + let (https_proxy_url, captured, stop, proxy_task) = spawn_rejecting_proxy( + b"HTTP/1.1 407 Proxy Authentication Required\r\nProxy-Authenticate: Basic\r\nContent-Length: 0\r\n\r\n", + ) + .await; + let (status, body) = { + let mut guards = clear_proxy_environment(); + guards.extend(configure_codex( + config_dir.path(), + "https://codex.invalid:4443/backend-api/codex/responses", + )); + guards.push(EnvGuard::set("HTTPS_PROXY", &https_proxy_url)); + call_messages("proxy-env-https").await + }; + let _ = stop.send(()); + proxy_task.await.unwrap(); + assert!(!status.is_success()); + assert!(!body.contains("secret-user")); + assert!(!body.contains("secret-pass")); + { + let captured = captured.lock().unwrap(); + assert_eq!(captured.len(), 1); + assert!( + captured + .iter() + .all(|request| request.target.starts_with("CONNECT codex.invalid:4443 ")) + ); + assert!( + captured + .iter() + .all(|request| request.has_proxy_authorization) + ); + } + + clear_codex_websocket_pool_for_tests(); + let (https_proxy_url, captured, stop, proxy_task) = + spawn_rejecting_proxy(b"HTTP/1.1 403 Forbidden\r\nContent-Length: 0\r\n\r\n").await; + let status = { + let mut guards = clear_proxy_environment(); + guards.extend(configure_codex( + config_dir.path(), + "https://codex.invalid:4443/backend-api/codex/responses", + )); + guards.push(EnvGuard::set("HTTPS_PROXY", &https_proxy_url)); + call_messages("proxy-env-connect-rejected").await.0 + }; + let _ = stop.send(()); + proxy_task.await.unwrap(); + assert!(!status.is_success()); + assert_eq!(captured.lock().unwrap().len(), 1); + + clear_codex_websocket_pool_for_tests(); + let (https_proxy_url, captured, stop, proxy_task) = + spawn_rejecting_proxy(b"HTTP/1.1 403 Forbidden\r\nContent-Length: 0\r\n\r\n").await; + { + let _retry_delay = ZeroRetryDelayGuard::new(); + let mut guards = clear_proxy_environment(); + guards.extend(configure_codex( + config_dir.path(), + "https://codex.invalid:4443/backend-api/codex/responses", + )); + guards.push(EnvGuard::set("HTTPS_PROXY", &https_proxy_url)); + let _ = call_messages_with_stream("proxy-env-live-connect-rejected", true).await; + } + let _ = stop.send(()); + proxy_task.await.unwrap(); + assert_eq!(captured.lock().unwrap().len(), 1); + + clear_codex_websocket_pool_for_tests(); + let origin = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let origin_url = format!("http://{}/responses", origin.local_addr().unwrap()); + let (failing_proxy_url, captured, stop, proxy_task) = spawn_rejecting_proxy( + b"HTTP/1.1 502 Bad Gateway\r\nRetry-After: 120\r\nContent-Length: 0\r\n\r\n", + ) + .await; + let (status, _) = { + let mut guards = clear_proxy_environment(); + guards.extend(configure_codex(config_dir.path(), &origin_url)); + guards.push(EnvGuard::set("HTTP_PROXY", &failing_proxy_url)); + guards.push(EnvGuard::set("CCP_CODEX_TRANSPORT", "auto")); + call_messages("proxy-env-no-direct-fallback").await + }; + let _ = stop.send(()); + proxy_task.await.unwrap(); + assert!(!status.is_success()); + { + let captured = captured.lock().unwrap(); + assert!( + captured + .iter() + .any(|request| request.target.starts_with("GET http://")) + ); + assert!( + captured + .iter() + .any(|request| request.target.starts_with("POST http://")) + ); + assert!( + captured + .iter() + .all(|request| request.has_proxy_authorization) + ); + } + assert!( + tokio::time::timeout(Duration::from_millis(100), origin.accept()) + .await + .is_err() + ); + + clear_codex_websocket_pool_for_tests(); +} diff --git a/tests/smoke_cutover.rs b/tests/smoke_cutover.rs index 35b91da6..955a3f11 100644 --- a/tests/smoke_cutover.rs +++ b/tests/smoke_cutover.rs @@ -85,6 +85,7 @@ async fn call_messages(model: &str) -> Response { } async fn call_messages_body(body: Value) -> Response { + let _no_proxy_env = EnvGuard::set("NO_PROXY", "127.0.0.1,localhost"); app(Arc::new(Registry::with_default_alias())) .oneshot( Request::builder()