#[macro_use]
mod common;
use std::net::SocketAddr;
use std::sync::{Arc, Mutex};
use apollo_opentelemetry_test::{TelemetryContext, assert_metrics_snapshot};
use axum::Router;
use axum::routing::get;
use bytes::Bytes;
use http::Method;
use http_body_util::{BodyExt as _, Empty};
use hyper::body::Incoming;
use hyper::client::conn::http1 as http1_client;
use hyper::server::conn::http1;
use hyper::service::service_fn;
use hyper_util::rt::TokioIo;
use indoc::indoc;
use tower::BoxError;
use tower::ServiceExt as _;
use common::*;
type BoxBody = http_body_util::combinators::BoxBody<Bytes, BoxError>;
fn empty_boxbody() -> BoxBody {
Empty::<Bytes>::new()
.map_err(|e: std::convert::Infallible| match e {})
.boxed()
}
async fn spawn_proxy() -> (SocketAddr, Arc<Mutex<Vec<String>>>) {
let captured = Arc::new(Mutex::new(Vec::<String>::new()));
let captured_clone = captured.clone();
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
loop {
let Ok((stream, _)) = listener.accept().await else {
break;
};
let captured = captured_clone.clone();
tokio::spawn(async move {
http1::Builder::new()
.preserve_header_case(true)
.serve_connection(
TokioIo::new(stream),
service_fn(move |req| handle_proxy_request(req, captured.clone())),
)
.with_upgrades()
.await
.ok();
});
}
});
(addr, captured)
}
async fn handle_proxy_request(
req: hyper::Request<Incoming>,
captured_auth: Arc<Mutex<Vec<String>>>,
) -> Result<hyper::Response<BoxBody>, BoxError> {
use tokio::net::TcpStream;
if let Some(auth) = req.headers().get(http::header::PROXY_AUTHORIZATION) {
captured_auth
.lock()
.unwrap()
.push(auth.to_str().unwrap_or("").to_string());
}
if req.method() == Method::CONNECT {
let authority = req.uri().authority().unwrap().to_string();
tokio::spawn(async move {
let Ok(mut target) = TcpStream::connect(&authority).await else {
return;
};
let Ok(upgraded) = hyper::upgrade::on(req).await else {
return;
};
tokio::io::copy_bidirectional(&mut TokioIo::new(upgraded), &mut target)
.await
.ok();
});
Ok(hyper::Response::builder()
.status(200)
.body(empty_boxbody())
.unwrap())
} else {
let uri = req.uri().clone();
let host = uri.host().ok_or("missing host in proxy request")?;
let port = uri.port_u16().unwrap_or(80);
let tcp = TcpStream::connect(format!("{host}:{port}")).await?;
let (mut sender, conn) = http1_client::handshake::<_, Incoming>(TokioIo::new(tcp)).await?;
tokio::spawn(async move {
conn.await.expect("plain http proxy connection errored");
});
let (mut parts, body) = req.into_parts();
let path_and_query = uri.path_and_query().map(|pq| pq.as_str()).unwrap_or("/");
parts.uri = path_and_query.parse()?;
let resp = sender
.send_request(hyper::Request::from_parts(parts, body))
.await?;
Ok(resp.map(|b| b.map_err(BoxError::from).boxed()))
}
}
async fn spawn_tls_h1_server(router: Router) -> SocketAddr {
use hyper_util::service::TowerToHyperService;
use rcgen::{CertificateParams, KeyPair};
use rustls::ServerConfig;
use rustls::crypto::aws_lc_rs as aws_lc_rs_crypto;
use tokio_rustls::TlsAcceptor;
let key_pair = KeyPair::generate().unwrap();
let params = CertificateParams::new(vec!["127.0.0.1".to_string()]).unwrap();
let cert = params.self_signed(&key_pair).unwrap();
let cert_der = cert.der().clone();
let key_der = rustls::pki_types::PrivateKeyDer::Pkcs8(key_pair.serialize_der().into());
let mut server_config =
ServerConfig::builder_with_provider(Arc::new(aws_lc_rs_crypto::default_provider()))
.with_safe_default_protocol_versions()
.unwrap()
.with_no_client_auth()
.with_single_cert(vec![cert_der], key_der)
.unwrap();
server_config.alpn_protocols = vec![b"http/1.1".to_vec()];
let acceptor = TlsAcceptor::from(Arc::new(server_config));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
loop {
let Ok((stream, _)) = listener.accept().await else {
break;
};
let acceptor = acceptor.clone();
let router = router.clone();
tokio::spawn(async move {
let Ok(tls_stream) = acceptor.accept(stream).await else {
return;
};
let svc = TowerToHyperService::new(router);
http1::Builder::new()
.serve_connection(TokioIo::new(tls_stream), svc)
.await
.ok();
});
}
});
addr
}
#[tokio::test]
async fn proxy_routes_http_request_through_proxy() {
let backend =
spawn_server(Router::new().route("/", get(|| async { "hello from backend" }))).await;
let (proxy_addr, _) = spawn_proxy().await;
let config = parse_config_with_vars(
indoc! {r#"
proxy:
url: ${env.PROXY_URL}
"#},
&[("PROXY_URL", &format!("http://{proxy_addr}"))],
);
let req = http::Request::builder()
.method(http::Method::GET)
.uri(format!("http://{backend}/"))
.body(empty_body())
.unwrap();
let resp = new_client(&config).oneshot(req).await.expect("request ok");
let body = resp.into_body().collect().await.unwrap().to_bytes();
assert_eq!(body, "hello from backend");
}
#[tokio::test]
async fn proxy_routes_https_request_via_connect_tunnel() {
let backend =
spawn_tls_h1_server(Router::new().route("/", get(|| async { "hello via tunnel" }))).await;
let (proxy_addr, _) = spawn_proxy().await;
let config = parse_config_with_vars(
indoc! {r#"
proxy:
url: ${env.PROXY_URL}
tls:
danger_accept_invalid_certs: true
"#},
&[("PROXY_URL", &format!("http://{proxy_addr}"))],
);
let req = http::Request::builder()
.method(http::Method::GET)
.uri(format!("https://{backend}/"))
.body(empty_body())
.unwrap();
let resp = new_client(&config).oneshot(req).await.expect("request ok");
let body = resp.into_body().collect().await.unwrap().to_bytes();
assert_eq!(body, "hello via tunnel");
}
#[tokio::test]
async fn proxy_sends_proxy_authorization_header_when_credentials_present() {
let backend = spawn_server(Router::new().route("/", get(|| async { "ok" }))).await;
let (proxy_addr, captured_auth) = spawn_proxy().await;
let config = parse_config_with_vars(
indoc! {r#"
proxy:
url: ${env.PROXY_URL}
"#},
&[("PROXY_URL", &format!("http://alice:s3cr3t@{proxy_addr}"))],
);
let req = http::Request::builder()
.method(http::Method::GET)
.uri(format!("http://{backend}/"))
.body(empty_body())
.unwrap();
new_client(&config).oneshot(req).await.expect("request ok");
let headers = captured_auth.lock().unwrap();
assert_eq!(
headers.len(),
1,
"proxy should have received one auth header"
);
assert_eq!(headers[0], "Basic YWxpY2U6czNjcjN0");
}
#[tokio::test]
async fn proxy_sends_proxy_authorization_header_on_https_connect_request() {
let backend = spawn_tls_h1_server(Router::new().route("/", get(|| async { "ok" }))).await;
let (proxy_addr, captured_auth) = spawn_proxy().await;
let config = parse_config_with_vars(
indoc! {r#"
proxy:
url: ${env.PROXY_URL}
tls:
danger_accept_invalid_certs: true
"#},
&[("PROXY_URL", &format!("http://alice:s3cr3t@{proxy_addr}"))],
);
let req = http::Request::builder()
.method(http::Method::GET)
.uri(format!("https://{backend}/"))
.body(empty_body())
.unwrap();
new_client(&config).oneshot(req).await.expect("request ok");
let headers = captured_auth.lock().unwrap();
assert_eq!(
headers.len(),
1,
"proxy should have received one auth header on CONNECT"
);
assert_eq!(headers[0], "Basic YWxpY2U6czNjcjN0");
}
#[tokio::test]
async fn proxy_https_tunnelled_request_does_not_leak_proxy_authorization_to_backend() {
use axum::http::HeaderMap;
let backend_headers: Arc<Mutex<Option<HeaderMap>>> = Arc::new(Mutex::new(None));
let captured = backend_headers.clone();
let backend = spawn_tls_h1_server(Router::new().route(
"/",
get(move |headers: HeaderMap| {
let captured = captured.clone();
async move {
*captured.lock().unwrap() = Some(headers);
"ok"
}
}),
))
.await;
let (proxy_addr, _) = spawn_proxy().await;
let config = parse_config_with_vars(
indoc! {r#"
proxy:
url: ${env.PROXY_URL}
tls:
danger_accept_invalid_certs: true
"#},
&[("PROXY_URL", &format!("http://alice:s3cr3t@{proxy_addr}"))],
);
let req = http::Request::builder()
.method(http::Method::GET)
.uri(format!("https://{backend}/"))
.body(empty_body())
.unwrap();
new_client(&config).oneshot(req).await.expect("request ok");
let received = backend_headers
.lock()
.unwrap()
.clone()
.expect("backend must have received the request");
assert!(
!received.contains_key(http::header::PROXY_AUTHORIZATION),
"backend received Proxy-Authorization on tunnelled request: {received:?}"
);
}
async fn reject_with_status(
status: http::StatusCode,
_req: hyper::Request<Incoming>,
) -> Result<hyper::Response<BoxBody>, std::convert::Infallible> {
Ok(hyper::Response::builder()
.status(status)
.body(empty_boxbody())
.unwrap())
}
async fn spawn_rejecting_connect_proxy(status: http::StatusCode) -> SocketAddr {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
loop {
let Ok((stream, _)) = listener.accept().await else {
break;
};
tokio::spawn(async move {
http1::Builder::new()
.serve_connection(
TokioIo::new(stream),
service_fn(move |req| reject_with_status(status, req)),
)
.await
.ok();
});
}
});
addr
}
#[tokio::test]
async fn proxy_returns_proxy_tunnel_error_when_connect_rejected() {
let proxy_addr =
spawn_rejecting_connect_proxy(http::StatusCode::PROXY_AUTHENTICATION_REQUIRED).await;
let config = parse_config_with_vars(
indoc! {r#"
proxy:
url: ${env.PROXY_URL}
tls:
danger_accept_invalid_certs: true
"#},
&[("PROXY_URL", &format!("http://{proxy_addr}"))],
);
let req = http::Request::builder()
.method(http::Method::GET)
.uri("https://example.com/")
.body(empty_body())
.unwrap();
let err = new_client(&config)
.oneshot(req)
.await
.expect_err("should fail");
match err {
apollo_http_client::HttpClientError::ProxyTunnel { status } => {
assert_eq!(status, http::StatusCode::PROXY_AUTHENTICATION_REQUIRED);
}
other => panic!("expected ProxyTunnel error, got {other}"),
}
}
#[tokio::test]
async fn proxy_connect_timeout_fires_when_proxy_unresponsive() {
let proxy_addr = spawn_stall_server().await;
let config = parse_config_with_vars(
indoc! {r#"
proxy:
url: ${env.PROXY_URL}
connect_timeout: 100ms
tls:
danger_accept_invalid_certs: true
"#},
&[("PROXY_URL", &format!("http://{proxy_addr}"))],
);
let req = http::Request::builder()
.method(http::Method::GET)
.uri("https://example.com/")
.body(empty_body())
.unwrap();
let err = new_client(&config)
.oneshot(req)
.await
.expect_err("should fail");
assert!(
matches!(err, apollo_http_client::HttpClientError::ConnectionTimeout),
"expected ConnectionTimeout, got: {err}"
);
}
#[tokio::test]
async fn proxy_https_tunnel_fails_when_target_not_tls() {
let backend = spawn_drop_server().await;
let (proxy_addr, _) = spawn_proxy().await;
let config = parse_config_with_vars(
indoc! {r#"
proxy:
url: ${env.PROXY_URL}
tls:
danger_accept_invalid_certs: true
"#},
&[("PROXY_URL", &format!("http://{proxy_addr}"))],
);
let req = http::Request::builder()
.method(http::Method::GET)
.uri(format!("https://{backend}/"))
.body(empty_body())
.unwrap();
let err = new_client(&config)
.oneshot(req)
.await
.expect_err("should fail");
assert!(
matches!(err, apollo_http_client::HttpClientError::Transport { .. }),
"expected Transport error from failed tunnel TLS handshake, got: {err}"
);
}
#[tokio::test]
async fn proxy_routes_h2_request_via_connect_tunnel() {
let backend =
spawn_tls_h2_server(Router::new().route("/", get(|| async { "hello via h2 tunnel" })))
.await;
let (proxy_addr, _) = spawn_proxy().await;
let config = parse_config_with_vars(
indoc! {r#"
protocol: http2
proxy:
url: ${env.PROXY_URL}
tls:
danger_accept_invalid_certs: true
"#},
&[("PROXY_URL", &format!("http://{proxy_addr}"))],
);
let req = http::Request::builder()
.method(http::Method::GET)
.uri(format!("https://{backend}/"))
.body(empty_body())
.unwrap();
let resp = new_client(&config).oneshot(req).await.expect("request ok");
let body = resp.into_body().collect().await.unwrap().to_bytes();
assert_eq!(body, "hello via h2 tunnel");
}
#[tokio::test]
async fn proxy_alpn_negotiates_h2_via_connect_tunnel() {
let backend =
spawn_tls_h2_server(Router::new().route("/", get(|| async { "hello via alpn-h2 tunnel" })))
.await;
let (proxy_addr, _) = spawn_proxy().await;
let config = parse_config_with_vars(
indoc! {r#"
protocol: alpn
proxy:
url: ${env.PROXY_URL}
tls:
danger_accept_invalid_certs: true
"#},
&[("PROXY_URL", &format!("http://{proxy_addr}"))],
);
let req = http::Request::builder()
.method(http::Method::GET)
.uri(format!("https://{backend}/"))
.body(empty_body())
.unwrap();
let resp = new_client(&config).oneshot(req).await.expect("request ok");
let body = resp.into_body().collect().await.unwrap().to_bytes();
assert_eq!(body, "hello via alpn-h2 tunnel");
}
#[tokio::test]
async fn proxy_connect_timeout_fires_during_tunnel_tls_handshake() {
let backend = spawn_stall_server().await;
let (proxy_addr, _) = spawn_proxy().await;
let config = parse_config_with_vars(
indoc! {r#"
proxy:
url: ${env.PROXY_URL}
connect_timeout: 200ms
tls:
danger_accept_invalid_certs: true
"#},
&[("PROXY_URL", &format!("http://{proxy_addr}"))],
);
let req = http::Request::builder()
.method(http::Method::GET)
.uri(format!("https://{backend}/"))
.body(empty_body())
.unwrap();
let err = new_client(&config)
.oneshot(req)
.await
.expect_err("should fail");
assert!(
matches!(err, apollo_http_client::HttpClientError::ConnectionTimeout),
"expected ConnectionTimeout from stalled tunnel TLS handshake, got: {err}"
);
}
#[tokio::test]
async fn proxy_request_records_connection_state_metrics() {
let ctx = integration_context();
let backend = spawn_server(Router::new().route("/", get(|| async { "ok" }))).await;
let (proxy_addr, _) = spawn_proxy().await;
let config = parse_config_with_vars(
indoc! {r#"
proxy:
url: ${env.PROXY_URL}
"#},
&[("PROXY_URL", &format!("http://{proxy_addr}"))],
);
let req = http::Request::builder()
.method(http::Method::GET)
.uri(format!("http://{backend}/"))
.body(empty_body())
.unwrap();
let resp = new_client(&config).oneshot(req).await.expect("request ok");
resp.into_body().collect().await.unwrap();
snapshot!(ctx, @r#"
- name: http.client.active_requests
description: Number of HTTP requests currently in flight
unit: "{request}"
data:
type: Sum
data_points:
- attributes:
http.request.method: GET
server.address: 127.0.0.1
server.port: "<port>"
value: 0
is_monotonic: false
temporality: Cumulative
- name: http.client.open_connections
description: Number of open connections in the HTTP client pool
unit: "{connection}"
data:
type: Sum
data_points:
- attributes:
http.connection.state: active
network.peer.address: 127.0.0.1
network.peer.port: "<peer_port>"
network.protocol.version: "1.1"
server.address: 127.0.0.1
server.port: "<port>"
value: 0
- attributes:
http.connection.state: connecting
network.peer.address: 127.0.0.1
network.peer.port: "<peer_port>"
server.address: 127.0.0.1
server.port: "<port>"
value: 0
- attributes:
http.connection.state: idle
network.peer.address: 127.0.0.1
network.peer.port: "<peer_port>"
network.protocol.version: "1.1"
server.address: 127.0.0.1
server.port: "<port>"
value: 0
is_monotonic: false
temporality: Cumulative
- name: http.client.request.duration
description: Duration of HTTP client requests
unit: s
data:
type: Histogram
data_points:
- attributes:
http.request.method: GET
http.response.status_code: "200"
network.protocol.version: "1.1"
server.address: 127.0.0.1
server.port: "<port>"
count: 1
sum: 0
min: 0
max: 0
bounds:
- 0.005
- 0.01
- 0.025
- 0.05
- 0.075
- 0.1
- 0.25
- 0.5
- 0.75
- 1
- 2.5
- 5
- 7.5
- 10
bucket_counts:
- 1
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
temporality: Cumulative
"#);
}
#[tokio::test]
async fn proxy_unreachable_records_error_and_resets_counters() {
let ctx = TelemetryContext::new();
let config = parse_config(indoc! {r#"
proxy:
url: http://127.0.0.1:1
connect_timeout: 1s
"#});
let req = http::Request::builder()
.method(http::Method::GET)
.uri("http://example.com/")
.body(empty_body())
.unwrap();
let err = new_client(&config)
.oneshot(req)
.await
.expect_err("should fail");
assert!(
matches!(
err,
apollo_http_client::HttpClientError::Transport { .. }
| apollo_http_client::HttpClientError::ConnectionTimeout
),
"unexpected error: {err}"
);
assert_metrics_snapshot!(ctx, @r#"
- name: http.client.active_requests
description: Number of HTTP requests currently in flight
unit: "{request}"
data:
type: Sum
data_points:
- attributes:
http.request.method: GET
server.address: example.com
server.port: "80"
value: 0
is_monotonic: false
temporality: Cumulative
- name: http.client.open_connections
description: Number of open connections in the HTTP client pool
unit: "{connection}"
data:
type: Sum
data_points:
- attributes:
http.connection.state: connecting
network.peer.address: 127.0.0.1
network.peer.port: "1"
server.address: example.com
server.port: "80"
value: 0
is_monotonic: false
temporality: Cumulative
- name: http.client.request.duration
description: Duration of HTTP client requests
unit: s
data:
type: Histogram
data_points:
- attributes:
error.type: _OTHER
http.request.method: GET
network.protocol.version: "1.1"
server.address: example.com
server.port: "80"
count: 1
sum: 0
min: 0
max: 0
bounds:
- 0.005
- 0.01
- 0.025
- 0.05
- 0.075
- 0.1
- 0.25
- 0.5
- 0.75
- 1
- 2.5
- 5
- 7.5
- 10
bucket_counts:
- 1
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
- 0
temporality: Cumulative
"#);
}