From 2aff7208fe95f331d9bb966bbd265c36ee5ecebf Mon Sep 17 00:00:00 2001 From: richardopenai Date: Wed, 16 Sep 2026 08:54:12 +0000 Subject: [PATCH] Add a hidden HTTP/3 TCP tunnel command (#45900) ## What changed Add `codex tcp-tunnel` and the `codex-tcp-tunnel` crate to forward loopback TCP connections to an explicit target through a TLS-verified HTTP/3 CONNECT proxy. - Require the proxy origin to match an approved HTTPS origin in a supplied policy file. - Read bearer tokens and optional bounded, non-forwarding `x-` headers from stdin. Support token updates for new connections, `LISTENING` and `AUTH_UPDATED` notifications, and shutdown when the control pipe closes in token-update mode. - Preserve the listener across proxy reconnects without replaying TCP streams, and let accepted streams continue while a proxy drains. ## Testing Add tests for hidden CLI parsing, proxy and target validation, credential renewal, control-pipe closure, and invalid input without secret disclosure. A local HTTP/3 proxy test covers token replacement, transport recovery without stream replay, and graceful draining. GitOrigin-RevId: c6af3025301c61e9fe90940cf5a6df0039465e69 --- MODULE.bazel.lock | 2 + codex-rs/Cargo.lock | 49 +++ codex-rs/Cargo.toml | 2 + codex-rs/cli/Cargo.toml | 1 + codex-rs/cli/src/main.rs | 37 ++ codex-rs/tcp-tunnel/BUILD.bazel | 6 + codex-rs/tcp-tunnel/Cargo.toml | 30 ++ codex-rs/tcp-tunnel/README.md | 7 + codex-rs/tcp-tunnel/src/control.rs | 77 +++++ codex-rs/tcp-tunnel/src/lib.rs | 483 +++++++++++++++++++++++++++ codex-rs/tcp-tunnel/src/lib_tests.rs | 379 +++++++++++++++++++++ 11 files changed, 1073 insertions(+) create mode 100644 codex-rs/tcp-tunnel/BUILD.bazel create mode 100644 codex-rs/tcp-tunnel/Cargo.toml create mode 100644 codex-rs/tcp-tunnel/README.md create mode 100644 codex-rs/tcp-tunnel/src/control.rs create mode 100644 codex-rs/tcp-tunnel/src/lib.rs create mode 100644 codex-rs/tcp-tunnel/src/lib_tests.rs diff --git a/MODULE.bazel.lock b/MODULE.bazel.lock index 55bde0a88c..bd2e8f7316 100644 --- a/MODULE.bazel.lock +++ b/MODULE.bazel.lock @@ -1084,6 +1084,8 @@ "git+https://github.com/dzbarsky/rules_rust?rev=b56cbaa8465e74127f1ea216f813cd377295ad81#b56cbaa8465e74127f1ea216f813cd377295ad81_runfiles": "{\"dependencies\":[],\"features\":{},\"strip_prefix\":\"\"}", "git+https://github.com/helix-editor/nucleo.git?rev=4253de9faabb4e5c6d81d946a5e35a90f87347ee#4253de9faabb4e5c6d81d946a5e35a90f87347ee_nucleo": "{\"dependencies\":[{\"default_features\":true,\"features\":[],\"name\":\"nucleo-matcher\",\"optional\":false},{\"default_features\":true,\"features\":[\"send_guard\",\"arc_lock\"],\"name\":\"parking_lot\",\"optional\":false},{\"name\":\"rayon\"}],\"features\":{},\"strip_prefix\":\"\"}", "git+https://github.com/helix-editor/nucleo.git?rev=4253de9faabb4e5c6d81d946a5e35a90f87347ee#4253de9faabb4e5c6d81d946a5e35a90f87347ee_nucleo-matcher": "{\"dependencies\":[{\"name\":\"memchr\"},{\"default_features\":true,\"features\":[],\"name\":\"unicode-segmentation\",\"optional\":true}],\"features\":{\"default\":[\"unicode-normalization\",\"unicode-casefold\",\"unicode-segmentation\"],\"unicode-casefold\":[],\"unicode-normalization\":[],\"unicode-segmentation\":[\"dep:unicode-segmentation\"]},\"strip_prefix\":\"matcher\"}", + "git+https://github.com/hyperium/h3?rev=e07e69412876f7e26f026bd75a48b2704d8c8283#e07e69412876f7e26f026bd75a48b2704d8c8283_h3": "{\"dependencies\":[{\"name\":\"bytes\",\"req\":\"1\"},{\"name\":\"fastrand\",\"req\":\"2.0.1\"},{\"default_features\":false,\"features\":[\"io\"],\"name\":\"futures-util\",\"optional\":false,\"req\":\"0.3\"},{\"name\":\"http\",\"req\":\"1\"},{\"default_features\":false,\"features\":[],\"name\":\"pin-project-lite\",\"optional\":false,\"req\":\"0.2\"},{\"default_features\":true,\"features\":[\"sync\"],\"name\":\"tokio\",\"optional\":false,\"req\":\"1\"},{\"default_features\":true,\"features\":[],\"name\":\"tracing\",\"optional\":true,\"req\":\"0.1.40\"}],\"features\":{\"i-implement-a-third-party-backend-and-opt-into-breaking-changes\":[],\"tracing\":[\"dep:tracing\"]},\"strip_prefix\":\"h3\"}", + "git+https://github.com/hyperium/h3?rev=e07e69412876f7e26f026bd75a48b2704d8c8283#e07e69412876f7e26f026bd75a48b2704d8c8283_h3-quinn": "{\"dependencies\":[{\"name\":\"bytes\",\"req\":\"1\"},{\"default_features\":true,\"features\":[],\"name\":\"futures\",\"optional\":false,\"req\":\"0.3.28\"},{\"default_features\":true,\"features\":[],\"name\":\"h3\",\"optional\":false,\"req\":\"0.0.8\"},{\"default_features\":true,\"features\":[],\"name\":\"h3-datagram\",\"optional\":true,\"req\":\"0.0.2\"},{\"default_features\":false,\"features\":[\"futures-io\"],\"name\":\"quinn\",\"optional\":false,\"req\":\"0.11.7\"},{\"default_features\":false,\"features\":[\"io-util\"],\"name\":\"tokio\",\"optional\":false,\"req\":\"1\"},{\"default_features\":true,\"features\":[],\"name\":\"tokio-util\",\"optional\":false,\"req\":\"0.7.9\"},{\"default_features\":true,\"features\":[],\"name\":\"tracing\",\"optional\":true,\"req\":\"0.1.40\"}],\"features\":{\"datagram\":[\"dep:h3-datagram\"],\"tracing\":[\"dep:tracing\"]},\"strip_prefix\":\"h3-quinn\"}", "git+https://github.com/microsoft/mxc?rev=6cd3d58f05d3447e67109cfb75e042803b843ca4#6cd3d58f05d3447e67109cfb75e042803b843ca4_appcontainer_common": "{\"dependencies\":[{\"name\":\"getrandom\",\"req\":\"0.2\"},{\"default_features\":true,\"features\":[\"derive\"],\"name\":\"serde\",\"optional\":false,\"req\":\"1\"},{\"name\":\"serde_json\",\"req\":\"1\"},{\"name\":\"thiserror\",\"req\":\"2\"},{\"default_features\":true,\"features\":[],\"name\":\"wxc_common\",\"optional\":false},{\"name\":\"flatbuffers\",\"req\":\"25\",\"target\":\"cfg(target_os = \\\"windows\\\")\"},{\"default_features\":true,\"features\":[],\"name\":\"learning_mode_core\",\"optional\":false,\"target\":\"cfg(target_os = \\\"windows\\\")\"},{\"default_features\":true,\"features\":[],\"name\":\"learning_mode_windows\",\"optional\":false,\"target\":\"cfg(target_os = \\\"windows\\\")\"},{\"default_features\":true,\"features\":[],\"name\":\"process_security_environment_spec\",\"optional\":false,\"target\":\"cfg(target_os = \\\"windows\\\")\"},{\"default_features\":true,\"features\":[],\"name\":\"sandbox_spec\",\"optional\":false,\"target\":\"cfg(target_os = \\\"windows\\\")\"},{\"name\":\"widestring\",\"req\":\"1\",\"target\":\"cfg(target_os = \\\"windows\\\")\"},{\"default_features\":true,\"features\":[\"Win32_Foundation\",\"Win32_Networking_WinSock\",\"Win32_NetworkManagement_WindowsFirewall\",\"Win32_Security\",\"Win32_Security_Authorization\",\"Win32_Security_Credentials\",\"Win32_Security_Isolation\",\"Win32_Storage_FileSystem\",\"Win32_System_Com\",\"Win32_System_Console\",\"Win32_System_Diagnostics_Debug\",\"Win32_System_Diagnostics_Etw\",\"Win32_System_Diagnostics_ToolHelp\",\"Win32_System_Environment\",\"Win32_System_IO\",\"Win32_System_JobObjects\",\"Win32_System_LibraryLoader\",\"Win32_System_Memory\",\"Win32_System_Ole\",\"Win32_System_Pipes\",\"Win32_System_ProcessStatus\",\"Win32_System_Registry\",\"Win32_System_SystemInformation\",\"Win32_System_SystemServices\",\"Win32_System_Threading\",\"Win32_System_Time\",\"Win32_System_Variant\",\"Win32_System_WindowsProgramming\",\"Win32_UI_Shell\"],\"name\":\"windows\",\"optional\":false,\"req\":\"0.62\",\"target\":\"cfg(target_os = \\\"windows\\\")\"},{\"name\":\"windows-core\",\"req\":\"0.62\",\"target\":\"cfg(target_os = \\\"windows\\\")\"},{\"name\":\"winreg\",\"req\":\"0.55\",\"target\":\"cfg(target_os = \\\"windows\\\")\"}],\"features\":{\"tier2_bfs\":[]},\"strip_prefix\":\"backends/appcontainer/common\"}", "git+https://github.com/microsoft/mxc?rev=6cd3d58f05d3447e67109cfb75e042803b843ca4#6cd3d58f05d3447e67109cfb75e042803b843ca4_learning_mode_core": "{\"dependencies\":[{\"name\":\"same-file\",\"req\":\"1\"},{\"default_features\":true,\"features\":[\"derive\"],\"name\":\"serde\",\"optional\":false,\"req\":\"1\"},{\"name\":\"serde_json\",\"req\":\"1\"},{\"name\":\"sha2\",\"req\":\"0.10\"},{\"name\":\"tempfile\",\"req\":\"3\"},{\"name\":\"thiserror\",\"req\":\"2\"}],\"features\":{},\"strip_prefix\":\"core/learning_mode_core\"}", "git+https://github.com/microsoft/mxc?rev=6cd3d58f05d3447e67109cfb75e042803b843ca4#6cd3d58f05d3447e67109cfb75e042803b843ca4_learning_mode_windows": "{\"dependencies\":[{\"default_features\":true,\"features\":[],\"name\":\"learning_mode_core\",\"optional\":false},{\"name\":\"thiserror\",\"req\":\"2\"},{\"name\":\"sha2\",\"req\":\"0.10\",\"target\":\"cfg(target_os = \\\"windows\\\")\"},{\"default_features\":true,\"features\":[\"Win32_Foundation\",\"Win32_Networking_WinSock\",\"Win32_NetworkManagement_WindowsFirewall\",\"Win32_Security\",\"Win32_Security_Authorization\",\"Win32_Security_Credentials\",\"Win32_Security_Isolation\",\"Win32_Storage_FileSystem\",\"Win32_System_Com\",\"Win32_System_Console\",\"Win32_System_Diagnostics_Debug\",\"Win32_System_Diagnostics_Etw\",\"Win32_System_Diagnostics_ToolHelp\",\"Win32_System_Environment\",\"Win32_System_IO\",\"Win32_System_JobObjects\",\"Win32_System_LibraryLoader\",\"Win32_System_Memory\",\"Win32_System_Ole\",\"Win32_System_Pipes\",\"Win32_System_ProcessStatus\",\"Win32_System_Registry\",\"Win32_System_SystemInformation\",\"Win32_System_SystemServices\",\"Win32_System_Threading\",\"Win32_System_Time\",\"Win32_System_Variant\",\"Win32_System_WindowsProgramming\",\"Win32_UI_Shell\"],\"name\":\"windows\",\"optional\":false,\"req\":\"0.62\",\"target\":\"cfg(target_os = \\\"windows\\\")\"},{\"name\":\"windows-core\",\"req\":\"0.62\",\"target\":\"cfg(target_os = \\\"windows\\\")\"},{\"default_features\":true,\"features\":[],\"name\":\"wxc_common\",\"optional\":false,\"target\":\"cfg(target_os = \\\"windows\\\")\"}],\"features\":{},\"strip_prefix\":\"backends/learning_mode/windows\"}", diff --git a/codex-rs/Cargo.lock b/codex-rs/Cargo.lock index 757328f641..d2f5babd5b 100644 --- a/codex-rs/Cargo.lock +++ b/codex-rs/Cargo.lock @@ -2628,6 +2628,7 @@ dependencies = [ "codex-skills-extension", "codex-state", "codex-stdio-to-uds", + "codex-tcp-tunnel", "codex-terminal-detection", "codex-thread-store", "codex-tui", @@ -4634,6 +4635,27 @@ dependencies = [ "tokio", ] +[[package]] +name = "codex-tcp-tunnel" +version = "0.0.0" +dependencies = [ + "anyhow", + "bytes", + "clap", + "h3", + "h3-quinn", + "http 1.4.0", + "pretty_assertions", + "quinn", + "rand 0.9.3", + "rcgen", + "rustls", + "rustls-native-certs", + "serde_json", + "tokio", + "url", +] + [[package]] name = "codex-terminal-detection" version = "0.0.0" @@ -8468,6 +8490,32 @@ dependencies = [ "tracing", ] +[[package]] +name = "h3" +version = "0.0.8" +source = "git+https://github.com/hyperium/h3?rev=e07e69412876f7e26f026bd75a48b2704d8c8283#e07e69412876f7e26f026bd75a48b2704d8c8283" +dependencies = [ + "bytes", + "fastrand", + "futures-util", + "http 1.4.0", + "pin-project-lite", + "tokio", +] + +[[package]] +name = "h3-quinn" +version = "0.0.10" +source = "git+https://github.com/hyperium/h3?rev=e07e69412876f7e26f026bd75a48b2704d8c8283#e07e69412876f7e26f026bd75a48b2704d8c8283" +dependencies = [ + "bytes", + "futures", + "h3", + "quinn", + "tokio", + "tokio-util", +] + [[package]] name = "half" version = "2.7.1" @@ -12045,6 +12093,7 @@ checksum = "b9e20a958963c291dc322d98411f541009df2ced7b5a4f2bd52337638cfccf20" dependencies = [ "bytes", "cfg_aliases 0.2.1", + "futures-io", "pin-project-lite", "quinn-proto", "quinn-udp 0.5.14", diff --git a/codex-rs/Cargo.toml b/codex-rs/Cargo.toml index f4ecb1ce04..5963eaba18 100644 --- a/codex-rs/Cargo.toml +++ b/codex-rs/Cargo.toml @@ -102,6 +102,7 @@ members = [ "stdio-to-uds", "otel", "otel-trace-websocket", + "tcp-tunnel", "tui", "user-verification", "tools", @@ -271,6 +272,7 @@ codex-terminal-detection = { path = "terminal-detection" } codex-test-binary-support = { path = "test-binary-support" } codex-thread-store = { path = "thread-store" } codex-tools = { path = "tools" } +codex-tcp-tunnel = { path = "tcp-tunnel" } codex-tui = { path = "tui" } codex-uds = { path = "uds" } codex-utils-absolute-path = { path = "utils/absolute-path" } diff --git a/codex-rs/cli/Cargo.toml b/codex-rs/cli/Cargo.toml index 8668aee897..51181994e5 100644 --- a/codex-rs/cli/Cargo.toml +++ b/codex-rs/cli/Cargo.toml @@ -61,6 +61,7 @@ codex-models-manager = { workspace = true } codex-plugin = { workspace = true } codex-protocol = { workspace = true } codex-responses-api-proxy = { workspace = true } +codex-tcp-tunnel = { workspace = true } codex-rmcp-client = { workspace = true } codex-rollout = { workspace = true } codex-rollout-trace = { workspace = true } diff --git a/codex-rs/cli/src/main.rs b/codex-rs/cli/src/main.rs index 39048c90f6..339ab94573 100644 --- a/codex-rs/cli/src/main.rs +++ b/codex-rs/cli/src/main.rs @@ -150,6 +150,9 @@ enum Subcommand { /// Browse all agent sessions on the shared local app-server daemon. Agents(AgentsCommand), + /// Internal: forward a local TCP socket through an HTTP/3 CONNECT proxy. + #[clap(hide = true)] + TcpTunnel(codex_tcp_tunnel::Args), /// Run Codex non-interactively. #[clap(visible_alias = "e")] Exec(ExecCli), @@ -1276,6 +1279,9 @@ async fn cli_main( .await?; handle_app_exit(exit_info, daemon_cli_executable.as_deref())?; } + Some(Subcommand::TcpTunnel(args)) => { + return codex_tcp_tunnel::run(args).await; + } Some(Subcommand::Exec(mut exec_cli)) => { reject_remote_mode_for_subcommand( root_remote.as_deref(), @@ -2674,6 +2680,7 @@ fn unsupported_subcommand_name_for_strict_config( Some(Subcommand::ResponsesApiProxy(_)) => Some("responses-api-proxy"), Some(Subcommand::StdioToUds(_)) => Some("stdio-to-uds"), Some(Subcommand::Features(_)) => Some("features"), + Some(Subcommand::TcpTunnel(_)) => Some("tcp-tunnel"), } } @@ -3780,6 +3787,36 @@ mod tests { err.to_string() } + #[test] + fn tcp_tunnel_is_hidden_and_accepts_its_control_input_contract() { + let root_help = help_from_args(&["codex", "--help"]); + assert!(!root_help.contains("tcp-tunnel"), "{root_help}"); + let help = help_from_args(&["codex", "tcp-tunnel", "--help"]); + for option in [ + "--target", + "--auth-token-stdin", + "--auth-token-updates-stdin", + "--connect-headers-stdin", + ] { + assert!(help.contains(option), "{help}"); + } + let cli = MultitoolCli::try_parse_from([ + "codex", + "tcp-tunnel", + "--proxy-url", + "https://proxy.example.org", + "--proxy-origins-file", + "/tmp/approved-origins", + "--target", + "[::1]:22", + "--auth-token-stdin", + "--auth-token-updates-stdin", + "--connect-headers-stdin", + ]) + .expect("generic tunnel should parse"); + assert!(matches!(cli.subcommand, Some(Subcommand::TcpTunnel(_)))); + } + #[test] fn plugin_marketplace_help_uses_plugin_namespace() { let help = help_from_args(&["codex", "plugin", "marketplace", "--help"]); diff --git a/codex-rs/tcp-tunnel/BUILD.bazel b/codex-rs/tcp-tunnel/BUILD.bazel new file mode 100644 index 0000000000..4cab3ab94f --- /dev/null +++ b/codex-rs/tcp-tunnel/BUILD.bazel @@ -0,0 +1,6 @@ +load("//:defs.bzl", "codex_rust_crate") + +codex_rust_crate( + name = "tcp-tunnel", + crate_name = "codex_tcp_tunnel", +) diff --git a/codex-rs/tcp-tunnel/Cargo.toml b/codex-rs/tcp-tunnel/Cargo.toml new file mode 100644 index 0000000000..617ed7e067 --- /dev/null +++ b/codex-rs/tcp-tunnel/Cargo.toml @@ -0,0 +1,30 @@ +[package] +name = "codex-tcp-tunnel" +version.workspace = true +edition.workspace = true +license.workspace = true + +[lib] +doctest = false + +[dependencies] +anyhow = { workspace = true } +bytes = { workspace = true } +clap = { workspace = true, features = ["derive"] } +h3 = { git = "https://github.com/hyperium/h3", rev = "e07e69412876f7e26f026bd75a48b2704d8c8283" } +h3-quinn = { git = "https://github.com/hyperium/h3", rev = "e07e69412876f7e26f026bd75a48b2704d8c8283" } +http = { workspace = true } +quinn = { version = "0.11", default-features = false, features = ["runtime-tokio", "rustls-ring"] } +rand = { workspace = true } +rustls = { workspace = true, features = ["ring", "std"] } +rustls-native-certs = { workspace = true } +serde_json = { workspace = true } +tokio = { workspace = true, features = ["io-util", "macros", "net", "rt-multi-thread", "sync", "time"] } +url = { workspace = true } + +[dev-dependencies] +pretty_assertions = { workspace = true } +rcgen = { workspace = true } + +[lints] +workspace = true diff --git a/codex-rs/tcp-tunnel/README.md b/codex-rs/tcp-tunnel/README.md new file mode 100644 index 0000000000..0e49499f05 --- /dev/null +++ b/codex-rs/tcp-tunnel/README.md @@ -0,0 +1,7 @@ +# HTTP/3 TCP tunnel + +The ordinary Codex binary includes the hidden `codex tcp-tunnel` command. It forwards a loopback TCP listener to an explicitly selected host and port through a TLS-verified HTTP/3 CONNECT proxy. The proxy URL must be an exact HTTPS origin in the supplied policy file; the target is provided separately as `host:port` or `[IPv6]:port`. + +The initial bearer must be provided as a single line on standard input (`--auth-token-stdin`). With `--auth-token-updates-stdin`, later lines replace the bearer for new connections, emit `AUTH_UPDATED`, and closure of the controlling pipe stops the tunnel. With `--connect-headers-stdin`, the first line is instead a JSON list of `["x-name","value"]` pairs; token lines follow. Only non-forwarding `x-` extension headers are accepted. No header values or credentials belong in the command arguments. + +The process emits `LISTENING 127.0.0.1:` after connecting to the proxy. A proxy handshake does not authorize the target; each local connection makes its own authenticated CONNECT. The listener survives transport reconnects, but TCP streams are never replayed. diff --git a/codex-rs/tcp-tunnel/src/control.rs b/codex-rs/tcp-tunnel/src/control.rs new file mode 100644 index 0000000000..801658aa59 --- /dev/null +++ b/codex-rs/tcp-tunnel/src/control.rs @@ -0,0 +1,77 @@ +//! Bounded parsing of HTTP/3 CONNECT metadata and bearer input. +//! Extension values remain sensitive and may never replace authorization or forwarding headers. +use std::io::BufRead; + +use anyhow::Context; +use anyhow::Result; +use anyhow::anyhow; +use anyhow::ensure; +use http::HeaderMap; +use http::HeaderValue; +use http::header; + +pub(super) const MAX_TOKEN_BYTES: usize = 64 * 1024; +pub(super) const MAX_CONNECT_METADATA_BYTES: usize = 16 * 1024; +pub(super) const MAX_CONNECT_HEADERS: usize = 16; +pub(super) const MAX_CONNECT_HEADER_VALUE_BYTES: usize = 4 * 1024; + +pub(super) fn read_connect_headers(reader: &mut impl BufRead) -> Result { + let mut line = Vec::new(); + std::io::Read::take(&mut *reader, (MAX_CONNECT_METADATA_BYTES + 1) as u64) + .read_until(b'\n', &mut line) + .context("reading CONNECT metadata")?; + ensure!( + !line.is_empty() && line.len() <= MAX_CONNECT_METADATA_BYTES && line.ends_with(b"\n"), + "invalid or oversized CONNECT metadata" + ); + let pairs: Vec<(String, String)> = + serde_json::from_slice(&line).map_err(|_| anyhow!("invalid CONNECT metadata"))?; + ensure!( + pairs.len() <= MAX_CONNECT_HEADERS, + "too many CONNECT metadata headers" + ); + let mut headers = HeaderMap::new(); + for (name, value) in pairs { + let name = header::HeaderName::from_bytes(name.as_bytes()) + .map_err(|_| anyhow!("invalid CONNECT metadata header name"))?; + ensure!( + !headers.contains_key(&name), + "duplicate CONNECT metadata header" + ); + let extension = name.as_str(); + ensure!( + extension.starts_with("x-") + && !extension.starts_with("x-forwarded-") + && extension != "x-real-ip", + "CONNECT metadata must use non-forwarding extension headers" + ); + let mut value = HeaderValue::from_str(&value) + .map_err(|_| anyhow!("invalid CONNECT metadata header value"))?; + ensure!( + value.as_bytes().len() <= MAX_CONNECT_HEADER_VALUE_BYTES, + "CONNECT metadata header value is too long" + ); + value.set_sensitive(true); + headers.insert(name, value); + } + Ok(headers) +} + +pub(super) fn read_auth_token(reader: &mut impl BufRead) -> Result> { + let mut line = Vec::new(); + std::io::Read::take(&mut *reader, (MAX_TOKEN_BYTES + 1) as u64) + .read_until(b'\n', &mut line) + .context("reading MASQUE token")?; + if line.is_empty() { + return Ok(None); + } + ensure!(line.len() <= MAX_TOKEN_BYTES, "MASQUE token is too long"); + let token = std::str::from_utf8(&line) + .context("invalid MASQUE token encoding")? + .trim(); + ensure!(!token.is_empty(), "empty MASQUE token"); + let mut auth = + HeaderValue::from_str(&format!("Bearer {token}")).context("invalid MASQUE token header")?; + auth.set_sensitive(true); + Ok(Some(auth)) +} diff --git a/codex-rs/tcp-tunnel/src/lib.rs b/codex-rs/tcp-tunnel/src/lib.rs new file mode 100644 index 0000000000..2908a504fd --- /dev/null +++ b/codex-rs/tcp-tunnel/src/lib.rs @@ -0,0 +1,483 @@ +//! Loopback TCP forwarding through an HTTP/3 CONNECT proxy. +//! Bearer credentials and optional proxy-specific metadata enter only via stdin. +//! The listener survives transport loss; individual TCP streams are never replayed. +use std::io::BufRead; +use std::io::Write; +use std::net::IpAddr; +use std::net::SocketAddr; +use std::path::PathBuf; +use std::sync::Arc; +use std::time::Duration; + +use anyhow::Context; +use anyhow::Result; +use anyhow::anyhow; +use anyhow::bail; +use anyhow::ensure; +use bytes::Buf; +use bytes::Bytes; +use clap::Args as ClapArgs; +use http::HeaderMap; +use http::HeaderValue; +use http::Method; +use http::Request; +use http::Uri; +use http::header; +use http::uri::Authority; +use quinn::crypto::rustls::QuicClientConfig; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; +use tokio::net::TcpListener; +use tokio::net::TcpStream; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio::sync::watch; +use url::Host; +use url::Url; + +mod control; + +use control::read_auth_token; +use control::read_connect_headers; + +const PROXY_CONNECT_TIMEOUT: Duration = Duration::from_secs(10); + +#[derive(Debug, ClapArgs)] +pub struct Args { + /// HTTPS origin of the HTTP/3 proxy. + #[arg(long)] + proxy_url: Url, + /// File containing one exact approved HTTPS proxy origin per line. + #[arg(long)] + proxy_origins_file: PathBuf, + /// CONNECT target authority, including a nonzero port. + #[arg(long)] + target: String, + #[arg(long, default_value = "127.0.0.1:0")] + listen_addr: SocketAddr, + /// Read the initial bearer from stdin. + #[arg(long, default_value_t = false)] + auth_token_stdin: bool, + /// Read replacement bearers from stdin and stop when the controlling pipe closes. + #[arg(long, requires = "auth_token_stdin")] + auth_token_updates_stdin: bool, + /// Read a JSON list of extension-header name/value pairs before the first bearer. + #[arg(long, requires = "auth_token_stdin")] + connect_headers_stdin: bool, +} + +#[derive(Debug, PartialEq, Eq)] +struct ProxyTarget { + host: String, + port: u16, + authority: String, +} + +struct TunnelHeaders { + auth: watch::Receiver, + connect: HeaderMap, +} + +struct ProxyConnection { + endpoint: quinn::Endpoint, + quic: quinn::Connection, + connection: h3::client::Connection, + sender: h3::client::SendRequest, +} + +enum ProxyClosure { + Draining, + Closed(anyhow::Error), +} + +impl ProxyTarget { + fn parse(proxy_url: &Url, trusted_origins: &[String], target: &str) -> Result { + ensure!(proxy_url.scheme() == "https", "proxy URL must use HTTPS"); + ensure!( + proxy_url.username().is_empty() + && proxy_url.password().is_none() + && proxy_url.query().is_none() + && proxy_url.fragment().is_none() + && proxy_url.path() == "/", + "proxy URL must be an HTTPS origin" + ); + ensure!( + trusted_origins.contains(&proxy_url.origin().ascii_serialization()), + "proxy URL origin is not trusted" + ); + let host = match proxy_url.host().context("proxy URL has no host")? { + Host::Domain(host) => host.to_owned(), + Host::Ipv4(address) => address.to_string(), + Host::Ipv6(address) => address.to_string(), + }; + ensure!( + !target.contains('@'), + "CONNECT target must not contain credentials" + ); + let authority: Authority = target.parse().context("invalid CONNECT target")?; + ensure!( + !authority.host().is_empty() && authority.port_u16().is_some_and(|port| port != 0), + "CONNECT target must include a nonzero port" + ); + Ok(Self { + host, + port: proxy_url.port_or_known_default().unwrap_or(443), + authority: authority.to_string(), + }) + } +} + +fn parse_trusted_origins(origins: &str) -> Result> { + let trusted_origins = origins + .lines() + .map(str::trim) + .filter(|line| !line.is_empty()) + .map(|line| { + let origin = Url::parse(line).context("invalid trusted proxy origin")?; + ensure!( + origin.scheme() == "https" + && origin.username().is_empty() + && origin.password().is_none() + && origin.query().is_none() + && origin.fragment().is_none() + && origin.path() == "/" + && line == origin.origin().ascii_serialization(), + "trusted proxy origins must be exact HTTPS origins" + ); + Ok(line.to_owned()) + }) + .collect::>>()?; + ensure!(!trusted_origins.is_empty(), "no trusted proxy origins"); + Ok(trusted_origins) +} + +#[cfg(test)] +#[path = "lib_tests.rs"] +mod tests; + +pub async fn run(args: Args) -> Result<()> { + ensure!(args.auth_token_stdin, "--auth-token-stdin is required"); + ensure!( + args.listen_addr.ip().is_loopback(), + "listener must be loopback" + ); + let origins = std::fs::read_to_string(&args.proxy_origins_file) + .context("reading trusted proxy origins")?; + let trusted_origins = parse_trusted_origins(&origins)?; + let target = ProxyTarget::parse(&args.proxy_url, &trusted_origins, &args.target)?; + let (metadata, mut tokens) = control_input( + std::io::BufReader::new(std::io::stdin()), + args.connect_headers_stdin, + )?; + let connect = metadata.await.context("CONNECT metadata input closed")??; + let initial_auth = tokens + .recv() + .await + .context("MASQUE credential input closed")?? + .context("empty MASQUE token")?; + let (updates, auth) = watch::channel(initial_auth); + let headers = TunnelHeaders { auth, connect }; + + let listener = TcpListener::bind(args.listen_addr) + .await + .context("binding loopback listener")?; + let native = rustls_native_certs::load_native_certs(); + let mut roots = rustls::RootCertStore::empty(); + roots.add_parsable_certificates(native.certs); + ensure!( + !roots.is_empty(), + "no native TLS roots: {:?}", + native.errors + ); + let mut tls = rustls::ClientConfig::builder_with_provider(Arc::new( + rustls::crypto::ring::default_provider(), + )) + .with_safe_default_protocol_versions()? + .with_root_certificates(roots) + .with_no_client_auth(); + tls.alpn_protocols = vec![b"h3".to_vec()]; + let mut config = quinn::ClientConfig::new(Arc::new(QuicClientConfig::try_from(tls)?)); + let mut transport = quinn::TransportConfig::default(); + transport.max_idle_timeout(Some(Duration::from_secs(30).try_into()?)); + transport.keep_alive_interval(Some(Duration::from_secs(10))); + config.transport_config(Arc::new(transport)); + let (ready, readiness) = oneshot::channel(); + let tunnel = async move { + let connected = ProxyConnection::connect(&target, &config).await?; + // The line is the caller readiness contract, including the assigned port. + writeln!(std::io::stdout(), "LISTENING {}", listener.local_addr()?)?; + std::io::stdout().flush()?; + let _ = ready.send(()); + serve(listener, target, config, headers, connected).await + }; + if args.auth_token_updates_stdin { + tokio::select! { + result = tunnel => result, + result = update_auth_tokens(&mut tokens, updates, readiness, std::io::stdout()) => result, + } + } else { + tunnel.await + } +} + +type TokenStream = mpsc::Receiver>>; + +async fn update_auth_tokens( + tokens: &mut TokenStream, + updates: watch::Sender, + readiness: oneshot::Receiver<()>, + mut output: impl Write, +) -> Result<()> { + let mut readiness = Some(readiness); + while let Some(result) = tokens.recv().await { + let Some(token) = result? else { + break; + }; + updates.send_replace(token); + if let Some(readiness) = readiness.take() { + readiness.await.context("MASQUE tunnel did not start")?; + } + writeln!(output, "AUTH_UPDATED")?; + output.flush()?; + } + Ok(()) +} + +fn control_input( + mut reader: impl BufRead + Send + 'static, + connect_headers_stdin: bool, +) -> Result<(oneshot::Receiver>, TokenStream)> { + let (metadata, initial_metadata) = oneshot::channel(); + let (send, tokens) = mpsc::channel(1); + // Tokio's blocking stdin can prevent runtime shutdown after a proxy startup failure. + std::thread::Builder::new() + .name("masque-credentials".to_owned()) + .spawn(move || { + let headers = if connect_headers_stdin { + read_connect_headers(&mut reader) + } else { + Ok(HeaderMap::new()) + }; + let valid = headers.is_ok(); + let _ = metadata.send(headers); + if !valid { + return; + } + loop { + let token = read_auth_token(&mut reader); + let finished = !matches!(&token, Ok(Some(_))); + if send.blocking_send(token).is_err() || finished { + break; + } + } + }) + .context("starting MASQUE credential input")?; + Ok((initial_metadata, tokens)) +} + +impl ProxyConnection { + async fn connect(target: &ProxyTarget, config: &quinn::ClientConfig) -> Result { + tokio::time::timeout(PROXY_CONNECT_TIMEOUT, async { + let mut last_error = "MASQUE proxy resolved to no addresses".to_owned(); + for peer in tokio::net::lookup_host((target.host.as_str(), target.port)).await? { + let bind = match peer.ip() { + IpAddr::V4(_) => "0.0.0.0:0", + IpAddr::V6(_) => "[::]:0", + }; + let mut endpoint = match quinn::Endpoint::client(bind.parse()?) { + Ok(endpoint) => endpoint, + Err(error) => { + last_error = error.to_string(); + continue; + } + }; + endpoint.set_default_client_config(config.clone()); + let connecting = match endpoint.connect(peer, &target.host) { + Ok(connecting) => connecting, + Err(error) => { + last_error = error.to_string(); + continue; + } + }; + let quic = match tokio::time::timeout(Duration::from_secs(5), connecting).await { + Ok(Ok(quic)) => quic, + Ok(Err(error)) => { + last_error = error.to_string(); + continue; + } + Err(_) => { + last_error = "QUIC handshake timed out".to_owned(); + continue; + } + }; + let (connection, sender) = match h3::client::builder() + .enable_extended_connect(true) + .build(h3_quinn::Connection::new(quic.clone())) + .await + { + Ok(connection) => connection, + Err(error) => { + last_error = error.to_string(); + continue; + } + }; + return Ok(Self { + endpoint, + quic, + connection, + sender, + }); + } + bail!("QUIC handshake failed: {last_error}"); + }) + .await + .context("connecting to MASQUE proxy timed out")? + } +} + +async fn serve( + listener: TcpListener, + target: ProxyTarget, + config: quinn::ClientConfig, + headers: TunnelHeaders, + mut connected: ProxyConnection, +) -> Result<()> { + loop { + let ProxyConnection { + endpoint, + quic, + mut connection, + sender, + } = connected; + let mut driver = tokio::spawn(async move { connection.wait_idle().await }); + let (draining, mut drain_started) = mpsc::channel(1); + let failure = loop { + tokio::select! { + incoming = listener.accept() => { + let (socket, _) = incoming?; + let mut sender = sender.clone(); + let draining = draining.clone(); + let mut request = Request::builder().method(Method::CONNECT) + .uri(Uri::builder().authority(target.authority.as_str()).build()?) + .body(())?; + *request.headers_mut() = headers.connect.clone(); + request.headers_mut().insert(header::AUTHORIZATION, headers.auth.borrow().clone()); + tokio::spawn(async move { + if let Err(error) = bridge(socket, &mut sender, request, draining).await { + eprintln!("MASQUE TCP connection failed: {error:#}"); + } + }); + } + _ = drain_started.recv() => break ProxyClosure::Draining, + closed = &mut driver => break ProxyClosure::Closed(match closed { + Ok(error) => anyhow!("MASQUE HTTP/3 connection closed: {error}"), + Err(error) => anyhow!("MASQUE HTTP/3 driver stopped: {error}"), + }), + closed = quic.closed() => break ProxyClosure::Closed(anyhow!(closed).context("MASQUE QUIC connection closed")), + } + }; + drop(sender); + match failure { + ProxyClosure::Draining => { + eprintln!("MASQUE HTTP/3 proxy is draining; reconnecting"); + // A refused new CONNECT is not replayed. Accepted streams keep their old transport. + tokio::spawn(async move { + tokio::select! { + _ = &mut driver => {}, + _ = quic.closed() => driver.abort(), + } + drop(endpoint); + }); + } + ProxyClosure::Closed(error) => { + driver.abort(); + quic.close(/*error_code*/ 0_u8.into(), b"reconnecting"); + drop(endpoint); + eprintln!("MASQUE proxy connection lost; reconnecting: {error:#}"); + } + } + + let mut retry_delay = Duration::from_secs(1); + loop { + let result = { + let attempt = async { + tokio::time::sleep(rand::random_range(retry_delay / 2..=retry_delay)).await; + ProxyConnection::connect(&target, &config).await + }; + tokio::pin!(attempt); + loop { + tokio::select! { + incoming = listener.accept() => { + // Offline local connections fail promptly; their traffic is never queued or replayed. + let (socket, _) = incoming?; + drop(socket); + } + result = &mut attempt => break result, + } + } + }; + match result { + Ok(reconnected) => { + connected = reconnected; + break; + } + Err(error) => { + eprintln!("MASQUE proxy reconnect failed: {error:#}"); + retry_delay = (retry_delay * 2).min(Duration::from_secs(15)); + } + } + } + } +} + +async fn bridge( + mut socket: TcpStream, + sender: &mut h3::client::SendRequest, + request: Request<()>, + draining: mpsc::Sender<()>, +) -> Result<()> { + let mut stream = match sender.send_request(request).await { + Ok(stream) => stream, + Err(error) => { + if matches!(error, h3::error::StreamError::RemoteClosing { .. }) { + let _ = draining.try_send(()); + } + return Err(error).context("sending CONNECT"); + } + }; + let status = stream + .recv_response() + .await + .context("receiving CONNECT response")? + .status(); + if !status.is_success() { + bail!("MASQUE CONNECT rejected with status {}", status.as_u16()); + } + let (mut send, mut recv) = stream.split(); + let (mut local_read, mut local_write) = socket.split(); + let upload = async { + let mut buffer = vec![0; 64 * 1024]; + loop { + let count = local_read.read(&mut buffer).await?; + if count == 0 { + send.finish().await?; + return Ok::<_, anyhow::Error>(()); + } + send.send_data(Bytes::copy_from_slice(&buffer[..count])) + .await?; + } + }; + let download = async { + while let Some(mut data) = recv.recv_data().await? { + while data.has_remaining() { + let chunk = data.chunk(); + local_write.write_all(chunk).await?; + data.advance(chunk.len()); + } + } + local_write.shutdown().await?; + Ok::<_, anyhow::Error>(()) + }; + tokio::try_join!(upload, download)?; + Ok(()) +} diff --git a/codex-rs/tcp-tunnel/src/lib_tests.rs b/codex-rs/tcp-tunnel/src/lib_tests.rs new file mode 100644 index 0000000000..1a3d83699a --- /dev/null +++ b/codex-rs/tcp-tunnel/src/lib_tests.rs @@ -0,0 +1,379 @@ +use std::io::Cursor; +use std::sync::Arc; +use std::sync::atomic::AtomicUsize; +use std::sync::atomic::Ordering; +use std::time::Duration; + +use anyhow::Result; +use bytes::Buf; +use bytes::Bytes; +use http::Response; +use pretty_assertions::assert_eq; +use rcgen::CertifiedKey; +use rcgen::generate_simple_self_signed; +use rustls::pki_types::PrivatePkcs8KeyDer; +use tokio::io::AsyncReadExt; +use tokio::io::AsyncWriteExt; +use tokio::net::TcpListener; +use tokio::net::TcpStream; +use tokio::sync::mpsc; +use tokio::sync::oneshot; +use tokio::sync::watch; +use tokio::time::timeout; +use url::Url; + +use super::ProxyConnection; +use super::ProxyTarget; +use super::TunnelHeaders; +use super::control::MAX_CONNECT_HEADER_VALUE_BYTES; +use super::control::MAX_CONNECT_HEADERS; +use super::control::MAX_CONNECT_METADATA_BYTES; +use super::control::MAX_TOKEN_BYTES; +use super::control::read_connect_headers; +use super::control_input; +use super::parse_trusted_origins; +use super::serve; +use super::update_auth_tokens; + +#[test] +fn proxy_policy_and_target_admission() { + let origins = parse_trusted_origins("\nhttps://proxy.example.org\n").unwrap(); + assert_eq!(origins, vec!["https://proxy.example.org".to_string()]); + let proxy = Url::parse("https://proxy.example.org").unwrap(); + for target in ["127.0.0.1:22", "[::1]:2222", "target.example:443"] { + assert_eq!( + ProxyTarget::parse(&proxy, &origins, target).unwrap(), + ProxyTarget { + host: "proxy.example.org".into(), + port: 443, + authority: target.into() + }, + ); + } + let ipv6 = "https://[::1]:8443"; + assert_eq!( + ProxyTarget::parse( + &Url::parse(ipv6).unwrap(), + &parse_trusted_origins(ipv6).unwrap(), + "target.example:22" + ) + .unwrap(), + ProxyTarget { + host: "::1".into(), + port: 8443, + authority: "target.example:22".into() + }, + ); + for origin in [ + "http://proxy.example.org", + "https://proxy.example.org/path", + "https://user@proxy.example.org", + "https://proxy.example.org?query", + "https://proxy.example.org#fragment", + ] { + assert!(parse_trusted_origins(origin).is_err(), "{origin}"); + assert!( + ProxyTarget::parse(&Url::parse(origin).unwrap(), &origins, "127.0.0.1:22").is_err(), + "{origin}" + ); + } + for origin in [ + "", + "https://proxy.example.org:443", + "https://proxy.example.org/", + ] { + assert!(parse_trusted_origins(origin).is_err(), "{origin}"); + } + for (proxy, target) in [ + ("https://evil.example", "127.0.0.1:22"), + ("https://proxy.example.org.attacker.net", "127.0.0.1:22"), + ("https://proxy.example.org", "127.0.0.1:0"), + ("https://proxy.example.org", "target.example"), + ("https://proxy.example.org", "user@target.example:22"), + ("https://proxy.example.org", "target.example:65536"), + ("https://proxy.example.org", "target.example:22/path"), + ] { + assert!( + ProxyTarget::parse(&Url::parse(proxy).unwrap(), &origins, target).is_err(), + "{proxy} {target}" + ); + } +} + +#[tokio::test] +async fn control_input_acknowledges_bearer_renewal_with_or_without_metadata() -> Result<()> { + for (input, includes_metadata) in [ + (&b"first-secret\nsecond-secret\n"[..], false), + ( + &b"[[\"x-test-route\",\"private-value\"]]\nfirst-secret\nsecond-secret\n"[..], + true, + ), + ] { + let (metadata, mut tokens) = control_input(Cursor::new(input), includes_metadata)?; + let metadata = metadata.await??; + if includes_metadata { + assert_eq!(metadata["x-test-route"], "private-value"); + assert!(metadata["x-test-route"].is_sensitive()); + } else { + assert!(metadata.is_empty()); + } + let initial = tokens.recv().await.unwrap()?.unwrap(); + assert_eq!(initial, "Bearer first-secret"); + assert!(initial.is_sensitive()); + let (updates, latest) = watch::channel(initial); + let (ready, readiness) = oneshot::channel(); + let mut output = Vec::new(); + ready.send(()).unwrap(); + update_auth_tokens(&mut tokens, updates, readiness, &mut output).await?; + assert_eq!(*latest.borrow(), "Bearer second-secret"); + assert!(latest.borrow().is_sensitive()); + assert_eq!(output, b"AUTH_UPDATED\n"); + } + Ok(()) +} + +#[tokio::test] +async fn closed_parent_pipe_finishes_even_before_proxy_readiness() -> Result<()> { + let (_, mut tokens) = control_input(Cursor::new([]), /*connect_headers_stdin*/ false)?; + let (updates, _) = watch::channel("Bearer initial".parse()?); + let (_ready, readiness) = oneshot::channel(); + let mut output = Vec::new(); + timeout( + Duration::from_secs(1), + update_auth_tokens(&mut tokens, updates, readiness, &mut output), + ) + .await??; + assert!(output.is_empty()); + Ok(()) +} + +#[tokio::test] +async fn invalid_control_input_stops_credentials_without_exposing_secrets() -> Result<()> { + let oversized_header = format!( + "[[\"x-test\",\"private-value{}\"]]\n", + "x".repeat(MAX_CONNECT_HEADER_VALUE_BYTES) + ); + let oversized_metadata = vec![b'x'; MAX_CONNECT_METADATA_BYTES + 1]; + let mut oversized_count = serde_json::to_vec( + &(0..=MAX_CONNECT_HEADERS) + .map(|index| (format!("x-test-{index}"), "private-value")) + .collect::>(), + )?; + oversized_count.push(b'\n'); + let invalid_metadata: &[(&str, &[u8])] = &[ + ( + "authorization", + b"[[\"authorization\",\"private-value\"]]\n", + ), + ("forwarding", b"[[\"x-forwarded-for\",\"private-value\"]]\n"), + ("real IP", b"[[\"x-real-ip\",\"private-value\"]]\n"), + ( + "duplicate", + b"[[\"x-test\",\"one\"],[\"X-Test\",\"private-value\"]]\n", + ), + ( + "invalid value", + b"[[\"x-test\",\"private-value\\nmore\"]]\n", + ), + ("invalid JSON", b"\"private-value\"\n"), + ("missing newline", b"[[\"x-test\",\"private-value\"]]"), + ("value limit", oversized_header.as_bytes()), + ("metadata limit", &oversized_metadata), + ("header count", &oversized_count), + ]; + for &(case, input) in invalid_metadata { + let mut input = input.to_vec(); + input.extend_from_slice(b"private-secret\n"); + let (metadata, mut tokens) = + control_input(Cursor::new(input), /*connect_headers_stdin*/ true)?; + let error = metadata.await?.expect_err(case); + assert!(!format!("{error:#}").contains("private-"), "{case}"); + assert!(tokens.recv().await.is_none(), "{case}"); + } + for (case, input) in [ + ("invalid bearer", b" private-secret\0 \n".to_vec()), + ("empty bearer", b"\n".to_vec()), + ("bearer limit", vec![b'x'; MAX_TOKEN_BYTES + 1]), + ] { + let (metadata, mut tokens) = + control_input(Cursor::new(input), /*connect_headers_stdin*/ false)?; + assert!(metadata.await??.is_empty(), "{case}"); + let error = tokens.recv().await.expect(case).expect_err(case); + assert!(!format!("{error:#}").contains("private-"), "{case}"); + assert!(tokens.recv().await.is_none(), "{case}"); + } + Ok(()) +} + +#[tokio::test] +async fn transport_reconnect_keeps_listener_and_does_not_replay_old_tcp_streams() -> Result<()> { + let CertifiedKey { cert, signing_key } = + generate_simple_self_signed(vec!["127.0.0.1".to_string()])?; + let provider = Arc::new(rustls::crypto::ring::default_provider()); + let mut server_tls = rustls::ServerConfig::builder_with_provider(provider.clone()) + .with_safe_default_protocol_versions()? + .with_no_client_auth() + .with_single_cert( + vec![cert.der().clone()], + PrivatePkcs8KeyDer::from(signing_key.serialize_der()).into(), + )?; + server_tls.alpn_protocols = vec![b"h3".to_vec()]; + let server_config = quinn::ServerConfig::with_crypto(Arc::new( + quinn::crypto::rustls::QuicServerConfig::try_from(server_tls)?, + )); + let endpoint = quinn::Endpoint::server(server_config, "127.0.0.1:0".parse()?)?; + let proxy_address = endpoint.local_addr()?; + + let mut roots = rustls::RootCertStore::empty(); + roots.add(cert.der().clone())?; + let mut client_tls = rustls::ClientConfig::builder_with_provider(provider) + .with_safe_default_protocol_versions()? + .with_root_certificates(roots) + .with_no_client_auth(); + client_tls.alpn_protocols = vec![b"h3".to_vec()]; + let config = quinn::ClientConfig::new(Arc::new( + quinn::crypto::rustls::QuicClientConfig::try_from(client_tls)?, + )); + + let request_count = Arc::new(AtomicUsize::new(0)); + let server_request_count = request_count.clone(); + let (accepted, mut connections) = mpsc::unbounded_channel(); + let (authorizations, mut tokens) = mpsc::unbounded_channel(); + let server = tokio::spawn(async move { + while let Some(incoming) = endpoint.accept().await { + let accepted = accepted.clone(); + let request_count = server_request_count.clone(); + let authorizations = authorizations.clone(); + tokio::spawn(async move { + let quic = incoming.await.unwrap(); + let (shutdown, mut shutdown_requests) = mpsc::channel::>(1); + accepted.send((quic.clone(), shutdown)).unwrap(); + let mut http = h3::server::builder() + .enable_extended_connect(true) + .build(h3_quinn::Connection::new(quic)) + .await + .unwrap(); + loop { + let request = tokio::select! { + Some(finished) = shutdown_requests.recv() => { + http.shutdown(/*max_requests*/ 0).await.unwrap(); + let _ = finished.send(()); + continue; + } + request = http.accept() => match request { + Ok(Some(request)) => request, + _ => break, + }, + }; + let (request, mut stream) = request.resolve_request().await.unwrap(); + authorizations + .send(( + request.headers()[http::header::AUTHORIZATION].clone(), + request.headers()["x-test-route"].clone(), + )) + .unwrap(); + request_count.fetch_add(1, Ordering::SeqCst); + tokio::spawn(async move { + if stream + .send_response(Response::builder().status(200).body(()).unwrap()) + .await + .is_ok() + && stream.send_data(Bytes::from_static(b"ready")).await.is_ok() + { + while let Ok(Some(mut data)) = stream.recv_data().await { + let count = data.remaining(); + if stream.send_data(data.copy_to_bytes(count)).await.is_err() { + break; + } + } + } + }); + } + }); + } + }); + + let listener = TcpListener::bind("127.0.0.1:0").await?; + let address = listener.local_addr()?; + let target = ProxyTarget { + host: proxy_address.ip().to_string(), + port: proxy_address.port(), + authority: "127.0.0.1:22".to_string(), + }; + let (updates, auth) = watch::channel("Bearer local-test".parse()?); + let headers = TunnelHeaders { + auth, + connect: read_connect_headers(&mut &b"[[\"x-test-route\",\"custom-route\"]]\n"[..])?, + }; + let initial = ProxyConnection::connect(&target, &config).await?; + let client = tokio::spawn(serve(listener, target, config, headers, initial)); + let (first_connection, _) = timeout(Duration::from_secs(5), connections.recv()) + .await? + .unwrap(); + let mut first = TcpStream::connect(address).await?; + let mut data = [0; 5]; + timeout(Duration::from_secs(5), first.read_exact(&mut data)).await??; + assert_eq!(data, *b"ready"); + assert_eq!( + tokens.recv().await.unwrap(), + ("Bearer local-test".parse()?, "custom-route".parse()?) + ); + + updates.send_replace("Bearer replacement-test".parse()?); + let mut refreshed = TcpStream::connect(address).await?; + timeout(Duration::from_secs(5), refreshed.read_exact(&mut data)).await??; + assert_eq!( + tokens.recv().await.unwrap(), + ("Bearer replacement-test".parse()?, "custom-route".parse()?) + ); + + first_connection.close(/*error_code*/ 0_u8.into(), b"test interruption"); + let old_stream = timeout(Duration::from_secs(5), first.read(&mut data)).await?; + assert!(matches!(old_stream, Ok(0) | Err(_))); + let (_, graceful_shutdown) = timeout(Duration::from_secs(10), connections.recv()) + .await? + .unwrap(); + assert_eq!(request_count.load(Ordering::SeqCst), 2); + + let connect_ready_listener = move || async move { + loop { + let mut candidate = TcpStream::connect(address).await?; + let mut data = [0; 5]; + if timeout(Duration::from_millis(250), candidate.read_exact(&mut data)) + .await + .is_ok_and(|result| result.is_ok()) + { + return Ok::<_, anyhow::Error>((data, candidate)); + } + tokio::time::sleep(Duration::from_millis(50)).await; + } + }; + let (recovered, mut persistent) = + timeout(Duration::from_secs(5), connect_ready_listener()).await??; + assert_eq!(recovered, *b"ready"); + assert_eq!( + tokens.recv().await.unwrap(), + ("Bearer replacement-test".parse()?, "custom-route".parse()?) + ); + assert_eq!(request_count.load(Ordering::SeqCst), 3); + + let (sent, goaway_sent) = oneshot::channel(); + graceful_shutdown.send(sent).await?; + goaway_sent.await?; + persistent.write_all(b"still").await?; + timeout(Duration::from_secs(5), persistent.read_exact(&mut data)).await??; + assert_eq!(data, *b"still"); + + let new_request = tokio::spawn(connect_ready_listener()); + timeout(Duration::from_secs(10), connections.recv()) + .await? + .unwrap(); + let (ready, _) = timeout(Duration::from_secs(5), new_request).await???; + assert_eq!(ready, *b"ready"); + persistent.write_all(b"alive").await?; + timeout(Duration::from_secs(5), persistent.read_exact(&mut data)).await??; + assert_eq!(data, *b"alive"); + client.abort(); + server.abort(); + Ok(()) +}