#![cfg(feature = "tokio")]
#[path = "proxy_connect_coverage/observer_events.rs"]
mod observer_events;
use std::convert::Infallible;
use std::sync::Arc;
use std::sync::Mutex;
use std::time::Duration;
use bytes::Bytes;
use http_body_util::Full;
use hyper::Response;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
use aioduct::HttpEngineSend;
use aioduct::observer::{
ConnectionEvent, ConnectionPhase, RequestEvent, RequestObserver, RequestPhase, RetryKind,
};
use aioduct::runtime::TokioRuntime;
use aioduct::runtime::tokio_rt::TcpConnector;
use aioduct_test_server::h1::h1_server_with;
use aioduct_test_server::h2::h2_server_with;
#[derive(Default, Clone)]
struct RecordingObserver {
events: Arc<Mutex<Vec<RequestPhase>>>,
connection_events: Arc<Mutex<Vec<ConnectionPhase>>>,
}
impl RequestObserver for RecordingObserver {
fn on_event(&self, event: &RequestEvent) {
self.events.lock().unwrap().push(event.phase.clone());
}
fn on_connection_event(&self, event: &ConnectionEvent) {
self.connection_events
.lock()
.unwrap()
.push(event.phase.clone());
}
}
impl RecordingObserver {
fn phases(&self) -> Vec<String> {
self.events
.lock()
.unwrap()
.iter()
.map(|p| match p {
RequestPhase::Started => "Started".into(),
RequestPhase::PoolCheckoutComplete { outcome, .. } => {
format!("PoolCheckoutComplete({outcome:?})")
}
RequestPhase::DnsResolved { .. } => "DnsResolved".into(),
RequestPhase::TcpConnected { .. } => "TcpConnected".into(),
RequestPhase::TlsHandshakeComplete { .. } => "TlsHandshakeComplete".into(),
RequestPhase::RequestSent { .. } => "RequestSent".into(),
RequestPhase::ResponseStarted { .. } => "ResponseStarted".into(),
RequestPhase::ResponseComplete { .. } => "ResponseComplete".into(),
RequestPhase::Failed { .. } => "Failed".into(),
RequestPhase::BytesTransferred { .. } => "BytesTransferred".into(),
RequestPhase::TransferComplete { .. } => "TransferComplete".into(),
RequestPhase::TransferAborted { .. } => "TransferAborted".into(),
RequestPhase::TrailersReceived { .. } => "TrailersReceived".into(),
RequestPhase::Redirected { .. } => "Redirected".into(),
RequestPhase::Retrying { .. } => "Retrying".into(),
})
.collect()
}
}
#[tokio::test]
async fn concurrent_h2_multiplex_exercises_checkout_path() {
let (addr, counter) = h2_server_with(|_req| async move {
tokio::time::sleep(Duration::from_millis(50)).await;
Ok::<_, Infallible>(Response::new(Full::new(Bytes::from("h2-ok"))))
})
.await;
let obs = RecordingObserver::default();
let client = Arc::new(
HttpEngineSend::<TokioRuntime, TcpConnector>::builder()
.request_observer(obs.clone())
.timeout(Duration::from_secs(5))
.build()
.unwrap(),
);
let resp = client
.get(&format!("http://{addr}/first"))
.unwrap()
.h2c_prior_knowledge()
.send()
.await
.unwrap();
assert_eq!(resp.status(), 200);
assert_eq!(resp.text().await.unwrap(), "h2-ok");
let mut handles = vec![];
for i in 0..5 {
let c = client.clone();
let a = addr;
handles.push(tokio::spawn(async move {
let resp = c
.get(&format!("http://{a}/concurrent-{i}"))
.unwrap()
.h2c_prior_knowledge()
.send()
.await
.unwrap();
assert_eq!(resp.status(), 200);
resp.text().await.unwrap()
}));
}
for h in handles {
let body = h.await.unwrap();
assert_eq!(body, "h2-ok");
}
assert_eq!(
counter.connections(),
1,
"all H2 requests should multiplex over one connection"
);
assert_eq!(counter.requests(), 6);
let phases = obs.phases();
let miss_count = phases.iter().filter(|p| p.contains("Miss")).count();
let hit_count = phases.iter().filter(|p| p.contains("Hit")).count();
assert!(
miss_count >= 1,
"should have at least one pool miss (initial connection), got: {phases:?}"
);
assert!(
hit_count >= 1,
"concurrent H2 requests should see pool hits for multiplexed connection, got: {phases:?}"
);
}
#[cfg(feature = "rustls")]
#[tokio::test]
async fn connect_tunnel_succeeds_through_proxy() {
aioduct_test_server::tls::install_crypto_provider();
let (target_addr, cert_der, _counter) =
aioduct_test_server::tls::tls_h1_server(&[b"http/1.1"]).await;
let proxy_listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let proxy_addr = proxy_listener.local_addr().unwrap();
tokio::spawn(async move {
loop {
let (mut client, _) = proxy_listener.accept().await.unwrap();
tokio::spawn(async move {
let mut buf = [0u8; 4096];
let n = client.read(&mut buf).await.unwrap();
let req_str = String::from_utf8_lossy(&buf[..n]);
if !req_str.starts_with("CONNECT") {
let _ = client.write_all(b"HTTP/1.1 400 Bad Request\r\n\r\n").await;
return;
}
let target = req_str.split_whitespace().nth(1).unwrap_or("").to_string();
let _ = client
.write_all(b"HTTP/1.1 200 Connection Established\r\n\r\n")
.await;
let target_connect = if target.contains(':') {
target.clone()
} else {
format!("{target}:443")
};
let actual_target = format!(
"127.0.0.1:{}",
target_connect.rsplit(':').next().unwrap_or("443")
);
let mut upstream = match tokio::net::TcpStream::connect(&actual_target).await {
Ok(s) => s,
Err(_) => return,
};
let _ = tokio::io::copy_bidirectional(&mut client, &mut upstream).await;
});
}
});
let cert = aioduct::tls::Certificate::from_der(cert_der.to_vec());
let client = HttpEngineSend::<TokioRuntime, TcpConnector>::builder()
.proxy(aioduct::ProxyConfig::http(&format!("http://{proxy_addr}")).unwrap())
.add_root_certificates(&[cert])
.danger_accept_invalid_hostnames(true)
.timeout(Duration::from_secs(5))
.build()
.unwrap();
let resp = client
.get(&format!("https://localhost:{}/", target_addr.port()))
.unwrap()
.send()
.await
.unwrap();
assert_eq!(resp.status(), 200);
assert_eq!(resp.text().await.unwrap(), "hello tls");
}
#[tokio::test]
async fn proxy_connection_with_keepalive() {
let (target_addr, _counter) = aioduct_test_server::h1::h1_server().await;
let proxy_listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let proxy_addr = proxy_listener.local_addr().unwrap();
tokio::spawn(async move {
loop {
let (mut client, _) = proxy_listener.accept().await.unwrap();
tokio::spawn(async move {
let mut buf = [0u8; 4096];
let n = client.read(&mut buf).await.unwrap();
let req_str = String::from_utf8_lossy(&buf[..n]);
if !req_str.starts_with("CONNECT") {
return;
}
let target = req_str.split_whitespace().nth(1).unwrap_or("");
let _ = client
.write_all(b"HTTP/1.1 200 Connection Established\r\n\r\n")
.await;
let mut upstream = match tokio::net::TcpStream::connect(target).await {
Ok(s) => s,
Err(_) => return,
};
let _ = tokio::io::copy_bidirectional(&mut client, &mut upstream).await;
});
}
});
let client = HttpEngineSend::<TokioRuntime, TcpConnector>::builder()
.proxy(aioduct::ProxyConfig::http(&format!("http://{proxy_addr}")).unwrap())
.tcp_keepalive(Duration::from_secs(30))
.tcp_keepalive_interval(Duration::from_secs(10))
.tcp_keepalive_retries(3)
.timeout(Duration::from_secs(5))
.build()
.unwrap();
let resp = client
.get(&format!("http://{target_addr}/keepalive-test"))
.unwrap()
.send()
.await
.unwrap();
assert_eq!(resp.status(), 200);
let body = resp.text().await.unwrap();
assert!(
body.contains("hello aioduct"),
"request through proxy with keepalive should succeed, got: {body}"
);
}
#[tokio::test]
async fn proxy_connection_with_fast_open() {
let (target_addr, _counter) = aioduct_test_server::h1::h1_server().await;
let proxy_listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let proxy_addr = proxy_listener.local_addr().unwrap();
tokio::spawn(async move {
loop {
let (mut client, _) = proxy_listener.accept().await.unwrap();
tokio::spawn(async move {
let mut buf = [0u8; 4096];
let n = client.read(&mut buf).await.unwrap();
let req_str = String::from_utf8_lossy(&buf[..n]);
if !req_str.starts_with("CONNECT") {
return;
}
let target = req_str.split_whitespace().nth(1).unwrap_or("");
let _ = client
.write_all(b"HTTP/1.1 200 Connection Established\r\n\r\n")
.await;
let mut upstream = match tokio::net::TcpStream::connect(target).await {
Ok(s) => s,
Err(_) => return,
};
let _ = tokio::io::copy_bidirectional(&mut client, &mut upstream).await;
});
}
});
let client = HttpEngineSend::<TokioRuntime, TcpConnector>::builder()
.proxy(aioduct::ProxyConfig::http(&format!("http://{proxy_addr}")).unwrap())
.tcp_fast_open(true)
.timeout(Duration::from_secs(5))
.build()
.unwrap();
let resp = client
.get(&format!("http://{target_addr}/fast-open-test"))
.unwrap()
.send()
.await
.unwrap();
assert_eq!(resp.status(), 200);
let body = resp.text().await.unwrap();
assert!(
body.contains("hello aioduct"),
"request through proxy with fast_open should succeed, got: {body}"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn socks4_proxy_connection() {
let (target_addr, _counter) = aioduct_test_server::h1::h1_server().await;
let socks_listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let socks_addr = socks_listener.local_addr().unwrap();
tokio::spawn(async move {
loop {
let (mut client, _) = socks_listener.accept().await.unwrap();
tokio::spawn(async move {
let mut buf = [0u8; 1024];
let n = client.read(&mut buf).await.unwrap();
if n < 8 {
return;
}
assert_eq!(buf[0], 0x04); assert_eq!(buf[1], 0x01);
let port = ((buf[2] as u16) << 8) | (buf[3] as u16);
let is_socks4a = buf[4] == 0 && buf[5] == 0 && buf[6] == 0 && buf[7] != 0;
if is_socks4a {
let _userid_start = 8;
}
client
.write_all(&[0x00, 0x5a, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00])
.await
.unwrap();
let target = format!("127.0.0.1:{port}");
let mut upstream = match tokio::net::TcpStream::connect(target).await {
Ok(s) => s,
Err(_) => return,
};
let _ = tokio::io::copy_bidirectional(&mut client, &mut upstream).await;
});
}
});
let client = HttpEngineSend::<TokioRuntime, TcpConnector>::builder()
.proxy(aioduct::ProxyConfig::socks4(&format!("socks4://{socks_addr}")).unwrap())
.timeout(Duration::from_secs(5))
.build()
.unwrap();
let resp = client
.get(&format!("http://localhost:{}/", target_addr.port()))
.unwrap()
.send()
.await
.unwrap();
assert_eq!(resp.status(), 200);
assert_eq!(resp.text().await.unwrap(), "hello aioduct");
}
#[tokio::test]
async fn observer_reports_stale_retry_on_rst() {
use std::sync::atomic::{AtomicU32, Ordering};
let request_count = Arc::new(AtomicU32::new(0));
let request_count2 = request_count.clone();
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
loop {
let (mut stream, _) = listener.accept().await.unwrap();
let count = request_count2.clone();
tokio::spawn(async move {
let n = count.fetch_add(1, Ordering::SeqCst);
if n == 0 {
let mut buf = [0u8; 4096];
let _ = stream.read(&mut buf).await;
let response = b"HTTP/1.1 200 OK\r\nContent-Length: 5\r\nConnection: keep-alive\r\n\r\nfirst";
let _ = stream.write_all(response).await;
let _ = stream.flush().await;
let mut peek = [0u8; 1];
match stream.read(&mut peek).await {
Ok(0) | Err(_) => return,
Ok(_) => {}
}
let raw = stream.into_std().unwrap();
let sock = socket2::SockRef::from(&raw);
let _ = sock.set_linger(Some(Duration::from_secs(0)));
drop(raw);
} else {
let mut buf = [0u8; 4096];
let _ = stream.read(&mut buf).await;
let response =
b"HTTP/1.1 200 OK\r\nContent-Length: 7\r\nConnection: close\r\n\r\nretried";
let _ = stream.write_all(response).await;
let _ = stream.flush().await;
}
});
}
});
let obs = RecordingObserver::default();
let client = HttpEngineSend::<TokioRuntime, TcpConnector>::builder()
.request_observer(obs.clone())
.timeout(Duration::from_secs(5))
.build()
.unwrap();
let resp = client
.get(&format!("http://{addr}/"))
.unwrap()
.send()
.await
.unwrap();
assert_eq!(resp.status(), 200);
assert_eq!(resp.text().await.unwrap(), "first");
let resp = client
.get(&format!("http://{addr}/"))
.unwrap()
.send()
.await
.unwrap();
assert_eq!(resp.status(), 200);
assert_eq!(resp.text().await.unwrap(), "retried");
let phases = obs.phases();
let has_failed_retry = obs.events.lock().unwrap().iter().any(|p| {
matches!(
p,
RequestPhase::Failed {
retry: RetryKind::StaleConnection,
..
}
)
});
assert!(
has_failed_retry,
"observer should report Failed with retry: StaleConnection on stale connection, got: {phases:?}"
);
let has_stale_retry = phases.iter().any(|p| p.contains("StaleRetry"));
assert!(
has_stale_retry,
"observer should report PoolCheckoutComplete(StaleRetry), got: {phases:?}"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn socks5_proxy_with_keepalive_and_fast_open() {
let (target_addr, _counter) = aioduct_test_server::h1::h1_server().await;
let socks_listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let socks_addr = socks_listener.local_addr().unwrap();
tokio::spawn(async move {
loop {
let (mut client, _) = socks_listener.accept().await.unwrap();
tokio::spawn(async move {
let mut buf = [0u8; 256];
let n = client.read(&mut buf).await.unwrap();
if n < 3 || buf[0] != 0x05 {
return;
}
client.write_all(&[0x05, 0x00]).await.unwrap();
let n = client.read(&mut buf).await.unwrap();
if n < 7 {
return;
}
let port = match buf[3] {
0x01 => u16::from_be_bytes([buf[8], buf[9]]),
0x03 => {
let domain_len = buf[4] as usize;
let port_offset = 5 + domain_len;
u16::from_be_bytes([buf[port_offset], buf[port_offset + 1]])
}
0x04 => u16::from_be_bytes([buf[20], buf[21]]),
_ => return,
};
let target = format!("127.0.0.1:{port}");
let mut upstream = match tokio::net::TcpStream::connect(target).await {
Ok(s) => s,
Err(_) => return,
};
client
.write_all(&[0x05, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00])
.await
.unwrap();
let _ = tokio::io::copy_bidirectional(&mut client, &mut upstream).await;
});
}
});
let client = HttpEngineSend::<TokioRuntime, TcpConnector>::builder()
.proxy(aioduct::ProxyConfig::socks5(&format!("socks5://{socks_addr}")).unwrap())
.tcp_keepalive(Duration::from_secs(15))
.tcp_keepalive_interval(Duration::from_secs(5))
.tcp_fast_open(true)
.timeout(Duration::from_secs(5))
.build()
.unwrap();
let resp = client
.get(&format!("http://localhost:{}/", target_addr.port()))
.unwrap()
.send()
.await
.unwrap();
assert_eq!(resp.status(), 200);
assert_eq!(
resp.text().await.unwrap(),
"hello aioduct",
"SOCKS5 proxy with keepalive+fast_open should succeed"
);
}
#[cfg(feature = "rustls")]
#[tokio::test]
async fn connect_tunnel_with_auth_succeeds() {
use std::sync::atomic::{AtomicBool, Ordering as AtomicOrdering};
aioduct_test_server::tls::install_crypto_provider();
let (target_addr, cert_der, _counter) =
aioduct_test_server::tls::tls_h1_server(&[b"http/1.1"]).await;
let auth_received = Arc::new(AtomicBool::new(false));
let auth_received_clone = auth_received.clone();
let proxy_listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let proxy_addr = proxy_listener.local_addr().unwrap();
tokio::spawn(async move {
loop {
let (mut client, _) = proxy_listener.accept().await.unwrap();
let auth_flag = auth_received_clone.clone();
tokio::spawn(async move {
let mut buf = [0u8; 4096];
let n = client.read(&mut buf).await.unwrap();
let req_str = String::from_utf8_lossy(&buf[..n]);
if !req_str.starts_with("CONNECT") {
let _ = client.write_all(b"HTTP/1.1 400 Bad Request\r\n\r\n").await;
return;
}
for line in req_str.lines() {
if line.to_lowercase().starts_with("proxy-authorization:") {
auth_flag.store(true, AtomicOrdering::SeqCst);
}
}
let target = req_str.split_whitespace().nth(1).unwrap_or("").to_string();
let port_str = target.rsplit(':').next().unwrap_or("443");
let _ = client
.write_all(b"HTTP/1.1 200 Connection Established\r\n\r\n")
.await;
let actual_target = format!("127.0.0.1:{port_str}");
let mut upstream = match tokio::net::TcpStream::connect(&actual_target).await {
Ok(s) => s,
Err(_) => return,
};
let _ = tokio::io::copy_bidirectional(&mut client, &mut upstream).await;
});
}
});
let cert = aioduct::tls::Certificate::from_der(cert_der.to_vec());
let client = HttpEngineSend::<TokioRuntime, TcpConnector>::builder()
.proxy(
aioduct::ProxyConfig::http(&format!("http://{proxy_addr}"))
.unwrap()
.basic_auth("user", "pass"),
)
.add_root_certificates(&[cert])
.danger_accept_invalid_hostnames(true)
.timeout(Duration::from_secs(5))
.build()
.unwrap();
let resp = client
.get(&format!("https://localhost:{}/", target_addr.port()))
.unwrap()
.send()
.await
.unwrap();
assert_eq!(resp.status(), 200);
assert_eq!(resp.text().await.unwrap(), "hello tls");
assert!(
auth_received.load(AtomicOrdering::SeqCst),
"CONNECT tunnel should include Proxy-Authorization header"
);
}
#[tokio::test]
async fn direct_connection_keepalive_and_fast_open() {
let (addr, _counter) = h1_server_with(|_req| async move {
Ok::<_, Infallible>(Response::new(Full::new(Bytes::from("keepalive-ok"))))
})
.await;
let client = HttpEngineSend::<TokioRuntime, TcpConnector>::builder()
.tcp_keepalive(Duration::from_secs(30))
.tcp_keepalive_interval(Duration::from_secs(10))
.tcp_keepalive_retries(3)
.tcp_fast_open(true)
.timeout(Duration::from_secs(5))
.build()
.unwrap();
let resp = client
.get(&format!("http://{addr}/"))
.unwrap()
.send()
.await
.unwrap();
assert_eq!(resp.status(), 200);
assert_eq!(resp.text().await.unwrap(), "keepalive-ok");
}
#[cfg(feature = "rustls")]
#[tokio::test]
async fn connection_coalescing_reuses_h2_with_sans() {
use std::sync::Arc;
aioduct_test_server::tls::install_crypto_provider();
let cert_params =
rcgen::generate_simple_self_signed(vec!["localhost".into(), "alt.localhost".into()])
.unwrap();
let cert_der = rustls::pki_types::CertificateDer::from(cert_params.cert.der().to_vec());
let key_der =
rustls::pki_types::PrivateKeyDer::Pkcs8(cert_params.signing_key.serialize_der().into());
let mut server_tls_config =
rustls::ServerConfig::builder_with_provider(aioduct_test_server::tls::crypto_provider())
.with_safe_default_protocol_versions()
.unwrap()
.with_no_client_auth()
.with_single_cert(vec![cert_der.clone()], key_der)
.unwrap();
server_tls_config.alpn_protocols = vec![b"h2".to_vec()];
let server_tls_config = Arc::new(server_tls_config);
let tls_acceptor = tokio_rustls::TlsAcceptor::from(server_tls_config);
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn({
let tls_acceptor = tls_acceptor.clone();
async move {
loop {
let (stream, _) = listener.accept().await.unwrap();
let acceptor = tls_acceptor.clone();
tokio::spawn(async move {
let tls_stream = match acceptor.accept(stream).await {
Ok(s) => s,
Err(_) => return,
};
let io = aioduct::runtime::tokio_rt::TokioIo::new(tls_stream);
let _ =
hyper::server::conn::http2::Builder::new(aioduct_test_server::TokioExec)
.serve_connection(
io,
hyper::service::service_fn(|_req| async {
Ok::<_, Infallible>(Response::new(Full::new(Bytes::from(
"coalesced-ok",
))))
}),
)
.await;
});
}
}
});
let mut root_store = rustls::RootCertStore::empty();
root_store.add(cert_der.clone()).unwrap();
let mut client_tls_config =
rustls::ClientConfig::builder_with_provider(aioduct_test_server::tls::crypto_provider())
.with_safe_default_protocol_versions()
.unwrap()
.with_root_certificates(root_store)
.with_no_client_auth();
client_tls_config.alpn_protocols = vec![b"h2".to_vec()];
let connector = aioduct::tls::RustlsConnector::new(Arc::new(client_tls_config));
let obs = RecordingObserver::default();
let client = HttpEngineSend::<TokioRuntime, TcpConnector>::builder()
.tls(connector)
.connection_coalescing(true)
.request_observer(obs.clone())
.timeout(Duration::from_secs(5))
.build()
.unwrap();
let resp = client
.get(&format!("https://localhost:{}/", addr.port()))
.unwrap()
.send()
.await
.unwrap();
assert_eq!(resp.status(), 200);
assert_eq!(resp.text().await.unwrap(), "coalesced-ok");
let port = addr.port();
let client_with_resolver = HttpEngineSend::<TokioRuntime, TcpConnector>::builder()
.tls(aioduct::tls::RustlsConnector::new({
let mut root_store2 = rustls::RootCertStore::empty();
root_store2.add(cert_der).unwrap();
let mut cfg2 = rustls::ClientConfig::builder_with_provider(
aioduct_test_server::tls::crypto_provider(),
)
.with_safe_default_protocol_versions()
.unwrap()
.with_root_certificates(root_store2)
.with_no_client_auth();
cfg2.alpn_protocols = vec![b"h2".to_vec()];
Arc::new(cfg2)
}))
.connection_coalescing(true)
.request_observer(obs.clone())
.resolver(move |host: &str, _port: u16| {
let port = port;
let _ = host;
Box::pin(async move { Ok(std::net::SocketAddr::from(([127, 0, 0, 1], port))) })
as std::pin::Pin<
Box<
dyn std::future::Future<Output = std::io::Result<std::net::SocketAddr>>
+ Send,
>,
>
})
.timeout(Duration::from_secs(5))
.build()
.unwrap();
let resp = client_with_resolver
.get(&format!("https://localhost:{port}/setup"))
.unwrap()
.send()
.await
.unwrap();
assert_eq!(resp.status(), 200);
let _ = resp.text().await.unwrap();
let resp = client_with_resolver
.get(&format!("https://alt.localhost:{port}/coalesced"))
.unwrap()
.send()
.await
.unwrap();
assert_eq!(resp.status(), 200);
assert_eq!(resp.text().await.unwrap(), "coalesced-ok");
let phases = obs.phases();
let _has_coalesced = phases.iter().any(|p| p.contains("Coalesced"));
}
#[tokio::test]
async fn proxy_connect_sends_custom_headers() {
let (target_addr, _counter) = h1_server_with(|_req| async move {
Ok::<_, Infallible>(Response::new(Full::new(Bytes::from("through proxy"))))
})
.await;
let captured = Arc::new(Mutex::new(String::new()));
let cap = captured.clone();
let proxy_listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let proxy_addr = proxy_listener.local_addr().unwrap();
tokio::spawn(async move {
let (mut client, _) = proxy_listener.accept().await.unwrap();
let mut buf = Vec::new();
let mut tmp = [0u8; 512];
loop {
let n = client.read(&mut tmp).await.unwrap();
buf.extend_from_slice(&tmp[..n]);
if buf.windows(4).any(|w| w == b"\r\n\r\n") {
break;
}
}
let req_str = String::from_utf8_lossy(&buf).to_string();
*cap.lock().unwrap() = req_str.clone();
let _ = client
.write_all(b"HTTP/1.1 200 Connection Established\r\n\r\n")
.await;
let actual = format!("127.0.0.1:{}", target_addr.port());
if let Ok(mut upstream) = tokio::net::TcpStream::connect(&actual).await {
let _ = tokio::io::copy_bidirectional(&mut client, &mut upstream).await;
}
});
let proxy = aioduct::ProxyConfig::http(&format!("http://{proxy_addr}"))
.unwrap()
.header(
http::header::HeaderName::from_static("x-proxy-token"),
http::HeaderValue::from_static("secret-123"),
);
let client = HttpEngineSend::<TokioRuntime, TcpConnector>::builder()
.proxy(proxy)
.timeout(Duration::from_secs(5))
.build()
.unwrap();
let resp = client
.get(&format!("http://localhost:{}/", target_addr.port()))
.unwrap()
.send()
.await
.unwrap();
assert_eq!(resp.status(), 200);
let _ = resp.text().await.unwrap();
let connect_req = captured.lock().unwrap().clone();
assert!(
connect_req.starts_with("CONNECT "),
"expected a CONNECT request, got: {connect_req}"
);
assert!(
connect_req
.to_lowercase()
.contains("x-proxy-token: secret-123"),
"custom CONNECT header missing, got: {connect_req}"
);
}