use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use bytes::Bytes;
use http_body_util::{BodyExt, Empty, Limited};
use hyper::body::Body;
use hyper_util::rt::TokioIo;
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use url::Url;
use crate::dns::Resolver;
use crate::proxy::{OutboundProxies, ProxyTarget};
const MAX_PROXY_ERROR_BYTES: usize = 512;
pub(crate) const MAX_RESPONSE_BYTES: usize = 1024 * 1024;
pub(crate) const MAX_ERROR_BODY_CHARS: usize = 200;
pub(crate) fn error_excerpt(body: &[u8]) -> String {
String::from_utf8_lossy(body)
.chars()
.take(MAX_ERROR_BODY_CHARS)
.collect()
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct Endpoint {
pub host: String,
pub port: u16,
pub https: bool,
}
impl Endpoint {
pub(crate) fn from_url(url: &Url) -> Result<Self, String> {
let host = url
.host_str()
.ok_or_else(|| format!("{url} has no host"))?
.to_string();
let https = match url.scheme() {
"https" => true,
"http" => false,
other => return Err(format!("unsupported scheme: {other}")),
};
let port = url
.port_or_known_default()
.unwrap_or(if https { 443 } else { 80 });
Ok(Self { host, port, https })
}
pub(crate) fn tls(host: &str, port: u16) -> Self {
Self {
host: host.to_string(),
port,
https: true,
}
}
pub(crate) fn host_for_lookup(&self) -> &str {
self.host
.strip_prefix('[')
.and_then(|rest| rest.strip_suffix(']'))
.unwrap_or(&self.host)
}
pub(crate) fn authority(&self) -> String {
let default = if self.https { 443 } else { 80 };
if self.port == default {
self.host.clone()
} else {
format!("{}:{}", self.host, self.port)
}
}
pub(crate) fn connect_authority(&self) -> String {
format!("{}:{}", self.host, self.port)
}
}
pub(crate) fn webpki_tls_config() -> rustls::ClientConfig {
let roots = rustls::RootCertStore {
roots: webpki_roots::TLS_SERVER_ROOTS.to_vec(),
};
rustls::ClientConfig::builder_with_provider(Arc::new(rustls::crypto::ring::default_provider()))
.with_safe_default_protocol_versions()
.expect("ring provider supports the default protocol versions")
.with_root_certificates(roots)
.with_no_client_auth()
}
pub(crate) enum ClientStream {
Direct(tokio::net::TcpStream),
Tunnelled(TokioIo<hyper::upgrade::Upgraded>),
}
impl AsyncRead for ClientStream {
fn poll_read(
self: Pin<&mut Self>,
context: &mut Context<'_>,
buffer: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
match self.get_mut() {
Self::Direct(stream) => Pin::new(stream).poll_read(context, buffer),
Self::Tunnelled(stream) => Pin::new(stream).poll_read(context, buffer),
}
}
}
impl AsyncWrite for ClientStream {
fn poll_write(
self: Pin<&mut Self>,
context: &mut Context<'_>,
buffer: &[u8],
) -> Poll<std::io::Result<usize>> {
match self.get_mut() {
Self::Direct(stream) => Pin::new(stream).poll_write(context, buffer),
Self::Tunnelled(stream) => Pin::new(stream).poll_write(context, buffer),
}
}
fn poll_flush(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<std::io::Result<()>> {
match self.get_mut() {
Self::Direct(stream) => Pin::new(stream).poll_flush(context),
Self::Tunnelled(stream) => Pin::new(stream).poll_flush(context),
}
}
fn poll_shutdown(self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<std::io::Result<()>> {
match self.get_mut() {
Self::Direct(stream) => Pin::new(stream).poll_shutdown(context),
Self::Tunnelled(stream) => Pin::new(stream).poll_shutdown(context),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum RequestForm {
Origin,
Absolute,
}
pub(crate) struct Connection<B> {
sender: hyper::client::conn::http1::SendRequest<B>,
form: RequestForm,
proxy_authorization: Option<String>,
}
impl<B> Connection<B>
where
B: Body + 'static,
{
pub(crate) fn request_target(&self, url: &Url) -> String {
match self.form {
RequestForm::Absolute => url.as_str().to_string(),
RequestForm::Origin => {
let mut target = url.path().to_string();
if let Some(query) = url.query() {
target.push('?');
target.push_str(query);
}
target
}
}
}
pub(crate) async fn send_request(
&mut self,
mut request: hyper::Request<B>,
) -> hyper::Result<hyper::Response<hyper::body::Incoming>> {
if let Some(credential) = &self.proxy_authorization
&& let Ok(value) = hyper::header::HeaderValue::from_str(credential)
{
request
.headers_mut()
.insert(hyper::header::PROXY_AUTHORIZATION, value);
}
self.sender.send_request(request).await
}
}
#[derive(Clone)]
pub struct Outbound {
resolver: Arc<dyn Resolver>,
proxies: Arc<OutboundProxies>,
}
impl Outbound {
pub fn new(resolver: Arc<dyn Resolver>, proxies: Arc<OutboundProxies>) -> Self {
Self { resolver, proxies }
}
pub(crate) async fn connect<B>(
&self,
endpoint: &Endpoint,
tls: &Arc<rustls::ClientConfig>,
) -> Result<Connection<B>, String>
where
B: Body + Send + 'static,
B::Data: Send,
B::Error: Into<Box<dyn std::error::Error + Send + Sync>>,
{
connect(self.resolver.as_ref(), &self.proxies, endpoint, tls).await
}
pub(crate) async fn connect_stream(&self, endpoint: &Endpoint) -> Result<ClientStream, String> {
connect_stream(self.resolver.as_ref(), &self.proxies, endpoint).await
}
}
pub(crate) async fn connect_stream(
resolver: &dyn Resolver,
proxies: &OutboundProxies,
endpoint: &Endpoint,
) -> Result<ClientStream, String> {
match proxies.select(endpoint) {
Some(proxy) => tunnel(resolver, proxy, endpoint)
.await
.map(ClientStream::Tunnelled),
None => dial(resolver, endpoint).await.map(ClientStream::Direct),
}
}
pub(crate) async fn connect<B>(
resolver: &dyn Resolver,
proxies: &OutboundProxies,
endpoint: &Endpoint,
tls: &Arc<rustls::ClientConfig>,
) -> Result<Connection<B>, String>
where
B: Body + Send + 'static,
B::Data: Send,
B::Error: Into<Box<dyn std::error::Error + Send + Sync>>,
{
let proxy = proxies.select(endpoint);
if let Some(proxy) = proxy
&& !endpoint.https
{
let stream = dial(resolver, proxy.endpoint())
.await
.map_err(|error| format!("connecting to proxy {}: {error}", proxy.redacted()))?;
let sender = spawn_handshake(TokioIo::new(stream)).await?;
return Ok(Connection {
sender,
form: RequestForm::Absolute,
proxy_authorization: proxy.authorization().map(str::to_string),
});
}
let stream = match proxy {
Some(proxy) => ClientStream::Tunnelled(tunnel(resolver, proxy, endpoint).await?),
None => ClientStream::Direct(dial(resolver, endpoint).await?),
};
let sender = if endpoint.https {
let server_name =
rustls_pki_types::ServerName::try_from(endpoint.host_for_lookup().to_string())
.map_err(|error| format!("{}: {error}", endpoint.host))?;
let stream = tokio_rustls::TlsConnector::from(tls.clone())
.connect(server_name, stream)
.await
.map_err(|error| format!("TLS handshake with {}: {error}", endpoint.host))?;
spawn_handshake(TokioIo::new(stream)).await?
} else {
spawn_handshake(TokioIo::new(stream)).await?
};
Ok(Connection {
sender,
form: RequestForm::Origin,
proxy_authorization: None,
})
}
async fn dial(
resolver: &dyn Resolver,
endpoint: &Endpoint,
) -> Result<tokio::net::TcpStream, String> {
crate::dns::connect(resolver, endpoint.host_for_lookup(), endpoint.port)
.await
.map_err(|error| format!("connecting to {}:{}: {error}", endpoint.host, endpoint.port))
}
async fn tunnel(
resolver: &dyn Resolver,
proxy: &ProxyTarget,
endpoint: &Endpoint,
) -> Result<TokioIo<hyper::upgrade::Upgraded>, String> {
let socket = dial(resolver, proxy.endpoint())
.await
.map_err(|error| format!("connecting to proxy {}: {error}", proxy.redacted()))?;
let (mut sender, connection) = hyper::client::conn::http1::handshake(TokioIo::new(socket))
.await
.map_err(|error| format!("HTTP handshake with proxy {}: {error}", proxy.redacted()))?;
tokio::spawn(async move {
let _ = connection.with_upgrades().await;
});
let authority = endpoint.connect_authority();
let mut builder = hyper::Request::connect(&authority)
.header(hyper::header::HOST, &authority)
.header(hyper::header::USER_AGENT, "acme-proxy");
if let Some(credential) = proxy.authorization() {
builder = builder.header(hyper::header::PROXY_AUTHORIZATION, credential);
}
let request = builder
.body(Empty::<Bytes>::new())
.map_err(|error| format!("building the CONNECT request: {error}"))?;
let response = sender
.send_request(request)
.await
.map_err(|error| format!("CONNECT {authority} via {}: {error}", proxy.redacted()))?;
if !response.status().is_success() {
let status = response.status();
let excerpt = Limited::new(response.into_body(), MAX_PROXY_ERROR_BYTES)
.collect()
.await
.map(|body| {
String::from_utf8_lossy(&body.to_bytes())
.split_whitespace()
.collect::<Vec<_>>()
.join(" ")
.chars()
.take(200)
.collect::<String>()
})
.unwrap_or_default();
return Err(format!(
"proxy {} refused CONNECT {authority}: {status} {excerpt}",
proxy.redacted()
));
}
hyper::upgrade::on(response)
.await
.map(TokioIo::new)
.map_err(|error| {
format!(
"proxy {} did not hand over the tunnel to {authority}: {error}",
proxy.redacted()
)
})
}
async fn spawn_handshake<B, I>(io: I) -> Result<hyper::client::conn::http1::SendRequest<B>, String>
where
B: Body + Send + 'static,
B::Data: Send,
B::Error: Into<Box<dyn std::error::Error + Send + Sync>>,
I: hyper::rt::Read + hyper::rt::Write + Unpin + Send + 'static,
{
let (sender, connection) = hyper::client::conn::http1::handshake(io)
.await
.map_err(|error| format!("HTTP handshake: {error}"))?;
tokio::spawn(async move {
let _ = connection.await;
});
Ok(sender)
}
#[cfg(test)]
mod tests {
use super::*;
fn url(value: &str) -> Url {
Url::parse(value).unwrap()
}
#[test]
fn default_ports_follow_the_scheme() {
let http = Endpoint::from_url(&url("http://example.com/x")).unwrap();
assert_eq!(http.port, 80);
assert!(!http.https);
assert_eq!(http.host, "example.com");
let https = Endpoint::from_url(&url("https://example.com/x")).unwrap();
assert_eq!(https.port, 443);
assert!(https.https);
}
#[test]
fn an_explicit_port_wins() {
let endpoint = Endpoint::from_url(&url("https://example.com:8443/x")).unwrap();
assert_eq!(endpoint.port, 8443);
assert!(endpoint.https);
}
#[test]
fn the_authority_omits_a_default_port() {
assert_eq!(
Endpoint::from_url(&url("https://example.com/x"))
.unwrap()
.authority(),
"example.com"
);
assert_eq!(
Endpoint::from_url(&url("http://example.com/x"))
.unwrap()
.authority(),
"example.com"
);
assert_eq!(
Endpoint::from_url(&url("https://example.com:8443/x"))
.unwrap()
.authority(),
"example.com:8443"
);
}
#[test]
fn a_url_with_no_host_is_refused() {
let error = Endpoint::from_url(&url("file:///etc/passwd")).unwrap_err();
assert!(
error.contains("no host") || error.contains("unsupported scheme"),
"{error}"
);
}
#[test]
fn an_unsupported_scheme_is_refused() {
let error = Endpoint::from_url(&url("ftp://example.com/x")).unwrap_err();
assert!(error.contains("unsupported scheme"), "{error}");
}
#[test]
fn an_ipv6_literal_survives_the_round_trip() {
let endpoint = Endpoint::from_url(&url("https://[2001:db8::1]:8443/x")).unwrap();
assert_eq!(endpoint.port, 8443);
assert_eq!(endpoint.authority(), "[2001:db8::1]:8443");
}
#[test]
fn an_ipv6_literal_loses_its_brackets_for_a_lookup() {
let endpoint = Endpoint::from_url(&url("https://[2001:db8::1]:8443/x")).unwrap();
assert_eq!(endpoint.host_for_lookup(), "2001:db8::1");
assert!(
endpoint
.host_for_lookup()
.parse::<std::net::IpAddr>()
.is_ok()
);
let named = Endpoint::from_url(&url("https://example.com/x")).unwrap();
assert_eq!(named.host_for_lookup(), "example.com");
}
#[test]
fn the_webpki_config_builds() {
let config = webpki_tls_config();
assert!(config.alpn_protocols.is_empty());
}
#[test]
fn a_connect_authority_always_carries_the_port() {
let https = Endpoint::from_url(&url("https://example.com/x")).unwrap();
assert_eq!(https.authority(), "example.com");
assert_eq!(https.connect_authority(), "example.com:443");
let http = Endpoint::from_url(&url("http://example.com/x")).unwrap();
assert_eq!(http.connect_authority(), "example.com:80");
let literal = Endpoint::from_url(&url("https://[2001:db8::1]/x")).unwrap();
assert_eq!(literal.connect_authority(), "[2001:db8::1]:443");
}
#[test]
fn an_endpoint_can_be_built_without_a_url() {
let endpoint = Endpoint::tls("example.com", 8443);
assert!(endpoint.https);
assert_eq!(endpoint.connect_authority(), "example.com:8443");
}
#[test]
fn the_request_target_follows_the_form() {
let target = url("http://example.com/a/b?c=d&e=f");
for (form, expected) in [
(RequestForm::Origin, "/a/b?c=d&e=f"),
(RequestForm::Absolute, "http://example.com/a/b?c=d&e=f"),
] {
let connection = Connection::<Empty<Bytes>> {
sender: unreachable_sender(),
form,
proxy_authorization: None,
};
assert_eq!(connection.request_target(&target), expected);
}
let no_query = url("http://example.com/a");
let connection = Connection::<Empty<Bytes>> {
sender: unreachable_sender(),
form: RequestForm::Origin,
proxy_authorization: None,
};
assert_eq!(connection.request_target(&no_query), "/a");
}
fn unreachable_sender() -> hyper::client::conn::http1::SendRequest<Empty<Bytes>> {
let (sender, connection) = futures_lite_block_on(async {
let (client, _server) = tokio::io::duplex(64);
hyper::client::conn::http1::handshake(TokioIo::new(client))
.await
.unwrap()
});
drop(connection);
sender
}
fn futures_lite_block_on<F: Future>(future: F) -> F::Output {
tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap()
.block_on(future)
}
mod loopback {
use super::*;
use crate::proxy::{OutboundProxies, ProxyTarget};
use crate::testutil::{FakeProxy, ProxyBehaviour};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
struct UnreachableResolver;
#[async_trait::async_trait]
impl Resolver for UnreachableResolver {
async fn reverse(&self, _ip: std::net::IpAddr) -> Result<Vec<String>, String> {
unreachable!()
}
async fn forward(&self, _name: &str) -> Result<Vec<std::net::IpAddr>, String> {
unreachable!("a literal 127.0.0.1 must short-circuit before this is called")
}
async fn txt(&self, _name: &str) -> Result<Vec<String>, String> {
unreachable!()
}
}
struct LoopbackResolver;
#[async_trait::async_trait]
impl Resolver for LoopbackResolver {
async fn reverse(&self, _ip: std::net::IpAddr) -> Result<Vec<String>, String> {
unreachable!()
}
async fn forward(&self, _name: &str) -> Result<Vec<std::net::IpAddr>, String> {
Ok(vec![std::net::IpAddr::from([127, 0, 0, 1])])
}
async fn txt(&self, _name: &str) -> Result<Vec<String>, String> {
unreachable!()
}
}
fn through(proxy: &FakeProxy) -> OutboundProxies {
OutboundProxies::always(ProxyTarget::for_test(&proxy.url()))
}
fn tunnelling(port: u16) -> ProxyBehaviour {
ProxyBehaviour::Tunnel {
status: "HTTP/1.1 200 Connection established\r\n",
force_port: Some(port),
}
}
async fn origin(response: &'static str) -> u16 {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut buffer = vec![0u8; 1024];
let _ = stream.read(&mut buffer).await;
let _ = stream.write_all(response.as_bytes()).await;
let _ = stream.shutdown().await;
});
port
}
#[tokio::test]
async fn connect_stream_tunnels_to_the_origin() {
let port = origin("pong").await;
let proxy = FakeProxy::start(tunnelling(port)).await;
let mut stream = connect_stream(
&LoopbackResolver,
&through(&proxy),
&Endpoint::tls("origin.example", 443),
)
.await
.expect("the tunnel must open");
stream.write_all(b"ping").await.unwrap();
let mut answer = String::new();
stream.read_to_string(&mut answer).await.unwrap();
assert_eq!(answer, "pong");
assert_eq!(proxy.connections(), 1);
let request = proxy.requests().remove(0);
assert!(
request.starts_with("CONNECT origin.example:443 HTTP/1.1"),
"{request}"
);
assert!(
!request.to_lowercase().contains("connection: close"),
"{request}"
);
assert!(
!request.to_lowercase().contains("proxy-connection"),
"{request}"
);
}
#[tokio::test]
async fn a_squid_shaped_reply_still_opens_the_tunnel() {
let port = origin("pong").await;
let proxy = FakeProxy::start(ProxyBehaviour::Tunnel {
status: "HTTP/1.0 200 Connection established\r\nProxy-Agent: squid/6.10\r\n",
force_port: Some(port),
})
.await;
let mut stream = connect_stream(
&LoopbackResolver,
&through(&proxy),
&Endpoint::tls("origin.example", 443),
)
.await
.expect("a 1.0 reply is still a tunnel");
stream.write_all(b"ping").await.unwrap();
let mut answer = String::new();
stream.read_to_string(&mut answer).await.unwrap();
assert_eq!(answer, "pong");
}
#[tokio::test]
async fn a_cleartext_target_is_forwarded_with_its_credentials() {
let proxy = FakeProxy::start(ProxyBehaviour::Forward(
"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok",
))
.await;
let proxies = OutboundProxies::always(ProxyTarget::for_test(&format!(
"http://user:pass@127.0.0.1:{}",
proxy.port
)));
let target = Url::parse("http://origin.example/a?b=c").unwrap();
let endpoint = Endpoint::from_url(&target).unwrap();
let mut connection = connect::<Empty<Bytes>>(
&UnreachableResolver,
&proxies,
&endpoint,
&Arc::new(webpki_tls_config()),
)
.await
.expect("the proxy is the peer, so the origin need not exist");
let request = hyper::Request::builder()
.uri(connection.request_target(&target))
.header(hyper::header::HOST, endpoint.authority())
.body(Empty::<Bytes>::new())
.unwrap();
assert_eq!(
connection.send_request(request).await.unwrap().status(),
200
);
let seen = proxy.requests().remove(0);
assert!(
seen.starts_with("GET http://origin.example/a?b=c HTTP/1.1"),
"{seen}"
);
assert!(
seen.to_lowercase()
.contains("proxy-authorization: basic dxnlcjpwyxnz"),
"{seen}"
);
}
#[tokio::test]
async fn https_is_tunnelled_with_the_origin_s_own_sni() {
use rustls::server::{ClientHello, ResolvesServerCert};
use rustls::sign::CertifiedKey;
use std::sync::Mutex;
#[derive(Debug)]
struct RecordingCert {
key: Arc<CertifiedKey>,
names: Arc<Mutex<Vec<String>>>,
}
impl ResolvesServerCert for RecordingCert {
fn resolve(&self, hello: ClientHello<'_>) -> Option<Arc<CertifiedKey>> {
self.names
.lock()
.unwrap()
.push(hello.server_name().unwrap_or_default().to_string());
Some(self.key.clone())
}
}
let key_pair = rcgen::KeyPair::generate().unwrap();
let certificate = rcgen::CertificateParams::new(vec!["origin.example".to_string()])
.unwrap()
.self_signed(&key_pair)
.unwrap();
let provider = rustls::crypto::ring::default_provider();
let signing_key = provider
.key_provider
.load_private_key(
rustls_pki_types::PrivatePkcs8KeyDer::from(key_pair.serialize_der()).into(),
)
.unwrap();
let names = Arc::new(Mutex::new(Vec::new()));
let resolver = RecordingCert {
key: Arc::new(CertifiedKey::new(
vec![certificate.der().clone()],
signing_key,
)),
names: names.clone(),
};
let server_config = rustls::ServerConfig::builder_with_provider(Arc::new(provider))
.with_safe_default_protocol_versions()
.unwrap()
.with_no_client_auth()
.with_cert_resolver(Arc::new(resolver));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let origin_port = listener.local_addr().unwrap().port();
let seen_inside = Arc::new(Mutex::new(String::new()));
let recorder = seen_inside.clone();
tokio::spawn(async move {
let acceptor = tokio_rustls::TlsAcceptor::from(Arc::new(server_config));
let (stream, _) = listener.accept().await.unwrap();
let mut stream = acceptor.accept(stream).await.unwrap();
let mut buffer = vec![0u8; 2048];
let read = stream.read(&mut buffer).await.unwrap();
*recorder.lock().unwrap() = String::from_utf8_lossy(&buffer[..read]).into_owned();
let _ = stream
.write_all(b"HTTP/1.1 204 No Content\r\nConnection: close\r\n\r\n")
.await;
let _ = stream.shutdown().await;
});
let proxy = FakeProxy::start(ProxyBehaviour::Tunnel {
status: "HTTP/1.1 200 Connection established\r\n",
force_port: Some(origin_port),
})
.await;
let proxies = OutboundProxies::always(ProxyTarget::for_test(&format!(
"http://user:pass@127.0.0.1:{}",
proxy.port
)));
let target = Url::parse("https://origin.example/x").unwrap();
let endpoint = Endpoint::from_url(&target).unwrap();
let mut connection = connect::<Empty<Bytes>>(
&LoopbackResolver,
&proxies,
&endpoint,
&crate::challenge::tls_alpn_01::accept_any_client_config(&[]).unwrap(),
)
.await
.expect("the tunnel must carry the TLS session");
let request = hyper::Request::builder()
.uri(connection.request_target(&target))
.header(hyper::header::HOST, endpoint.authority())
.body(Empty::<Bytes>::new())
.unwrap();
assert_eq!(
connection.send_request(request).await.unwrap().status(),
204
);
let connect_request = proxy.requests().remove(0);
assert!(
connect_request.starts_with("CONNECT origin.example:443 HTTP/1.1"),
"{connect_request}"
);
assert!(
connect_request
.to_lowercase()
.contains("proxy-authorization"),
"{connect_request}"
);
assert_eq!(names.lock().unwrap().as_slice(), ["origin.example"]);
let inside = seen_inside.lock().unwrap().clone();
assert!(inside.starts_with("GET /x HTTP/1.1"), "{inside}");
assert!(
!inside.to_lowercase().contains("proxy-authorization"),
"{inside}"
);
}
#[tokio::test]
async fn a_refused_connect_reports_the_status_and_the_body() {
let proxy = FakeProxy::start(ProxyBehaviour::Refuse(
"HTTP/1.1 407 Proxy Authentication Required\r\nContent-Length: 20\r\n\r\n\
credentials required",
))
.await;
let proxies = OutboundProxies::always(ProxyTarget::for_test(&format!(
"http://user:hunter2@127.0.0.1:{}",
proxy.port
)));
let Err(error) = connect_stream(
&LoopbackResolver,
&proxies,
&Endpoint::tls("origin.example", 443),
)
.await
else {
panic!("a 407 is not a tunnel");
};
assert!(error.contains("407"), "{error}");
assert!(error.contains("credentials required"), "{error}");
assert!(!error.contains("hunter2"), "{error}");
}
#[tokio::test]
async fn an_unreachable_proxy_names_the_proxy() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
drop(listener);
let proxies =
OutboundProxies::always(ProxyTarget::for_test(&format!("http://127.0.0.1:{port}")));
let Err(error) = connect_stream(
&LoopbackResolver,
&proxies,
&Endpoint::tls("origin.example", 443),
)
.await
else {
panic!("a dead proxy is not a tunnel");
};
assert!(error.contains("proxy"), "{error}");
assert!(error.contains(&port.to_string()), "{error}");
assert!(!error.contains("origin.example"), "{error}");
}
#[tokio::test]
async fn a_bypassed_target_never_reaches_the_proxy() {
let port = origin("pong").await;
let proxy = FakeProxy::start(tunnelling(port)).await;
let proxies = through(&proxy).with_bypass(&["bypassed.example"]).unwrap();
let mut stream = connect_stream(
&LoopbackResolver,
&proxies,
&Endpoint::tls("bypassed.example", port),
)
.await
.expect("a bypassed target still connects, just directly");
stream.write_all(b"ping").await.unwrap();
let mut answer = String::new();
stream.read_to_string(&mut answer).await.unwrap();
assert_eq!(answer, "pong");
assert_eq!(proxy.connections(), 0, "the proxy must not be dialled");
}
}
}