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
This commit is contained in:
richardopenai
2026-09-16 08:54:12 +00:00
committed by copyberry
parent 50d77959bf
commit 2aff7208fe
11 changed files with 1073 additions and 0 deletions

2
MODULE.bazel.lock generated
View File

@@ -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\"}",

49
codex-rs/Cargo.lock generated
View File

@@ -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",

View File

@@ -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" }

View File

@@ -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 }

View File

@@ -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"]);

View File

@@ -0,0 +1,6 @@
load("//:defs.bzl", "codex_rust_crate")
codex_rust_crate(
name = "tcp-tunnel",
crate_name = "codex_tcp_tunnel",
)

View File

@@ -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

View File

@@ -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:<port>` 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.

View File

@@ -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<HeaderMap> {
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<Option<HeaderValue>> {
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))
}

View File

@@ -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<HeaderValue>,
connect: HeaderMap,
}
struct ProxyConnection {
endpoint: quinn::Endpoint,
quic: quinn::Connection,
connection: h3::client::Connection<h3_quinn::Connection, Bytes>,
sender: h3::client::SendRequest<h3_quinn::OpenStreams, Bytes>,
}
enum ProxyClosure {
Draining,
Closed(anyhow::Error),
}
impl ProxyTarget {
fn parse(proxy_url: &Url, trusted_origins: &[String], target: &str) -> Result<Self> {
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<Vec<String>> {
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::<Result<Vec<_>>>()?;
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<Result<Option<HeaderValue>>>;
async fn update_auth_tokens(
tokens: &mut TokenStream,
updates: watch::Sender<HeaderValue>,
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<Result<HeaderMap>>, 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<Self> {
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<h3_quinn::OpenStreams, Bytes>,
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(())
}

View File

@@ -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::<Vec<_>>(),
)?;
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::<oneshot::Sender<()>>(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(())
}