#![allow(unused_imports)]
#[macro_use]
mod common;
use std::net::SocketAddr;
use std::sync::Arc;
use apollo_http_client::{HttpBody, HttpClient, HttpClientConfig, HttpClientError};
use axum::Router;
use axum::routing::get;
use bytes::Bytes;
use http_body_util::{BodyExt as _, Full};
use hyper::server::conn::http1;
use hyper_util::rt::TokioIo;
use hyper_util::service::TowerToHyperService;
use indoc::indoc;
use rcgen::{BasicConstraints, CertificateParams, CertifiedIssuer, DnType, IsCa, KeyPair};
use rustls::ServerConfig;
use rustls::crypto::aws_lc_rs as aws_lc_rs_crypto;
use rustls_pki_types::pem::PemObject;
use tokio_rustls::TlsAcceptor;
use tower::ServiceExt as _;
use common::*;
struct Pki {
ca_pem: String,
leaf_cert_pem: String,
leaf_cert_der: rustls::pki_types::CertificateDer<'static>,
leaf_key_pem: String,
leaf_key_der: rustls::pki_types::PrivateKeyDer<'static>,
}
fn generate_pki(sans: &[&str], label: &str) -> Pki {
let ca_key = KeyPair::generate().unwrap();
let mut ca_params = CertificateParams::new(Vec::<String>::new()).unwrap();
ca_params.is_ca = IsCa::Ca(BasicConstraints::Unconstrained);
ca_params
.distinguished_name
.push(DnType::CommonName, format!("apollo-http-client {label} CA"));
let ca = CertifiedIssuer::self_signed(ca_params, ca_key).unwrap();
let leaf_key = KeyPair::generate().unwrap();
let san_strings: Vec<String> = sans.iter().map(|s| (*s).to_string()).collect();
let mut leaf_params = CertificateParams::new(san_strings).unwrap();
leaf_params.distinguished_name.push(
DnType::CommonName,
format!("apollo-http-client {label} leaf"),
);
let leaf_cert = leaf_params.signed_by(&leaf_key, &ca).unwrap();
Pki {
ca_pem: ca.pem(),
leaf_cert_pem: leaf_cert.pem(),
leaf_cert_der: leaf_cert.der().clone(),
leaf_key_pem: leaf_key.serialize_pem(),
leaf_key_der: rustls::pki_types::PrivateKeyDer::Pkcs8(leaf_key.serialize_der().into()),
}
}
fn generate_server_pki() -> Pki {
generate_pki(&["127.0.0.1"], "server")
}
fn generate_client_pki() -> Pki {
generate_pki(&[], "client")
}
async fn spawn_tls_server(
cert: rustls::pki_types::CertificateDer<'static>,
key: rustls::pki_types::PrivateKeyDer<'static>,
) -> SocketAddr {
spawn_tls_server_with_client_auth(cert, key, None).await
}
async fn spawn_tls_server_with_client_auth(
cert: rustls::pki_types::CertificateDer<'static>,
key: rustls::pki_types::PrivateKeyDer<'static>,
client_ca_pem: Option<&str>,
) -> SocketAddr {
let provider = Arc::new(aws_lc_rs_crypto::default_provider());
let builder = ServerConfig::builder_with_provider(provider.clone())
.with_safe_default_protocol_versions()
.unwrap();
let builder = match client_ca_pem {
None => builder.with_no_client_auth(),
Some(pem) => {
let mut roots = rustls::RootCertStore::empty();
for cert in rustls::pki_types::CertificateDer::pem_slice_iter(pem.as_bytes()) {
roots.add(cert.unwrap()).unwrap();
}
let verifier = rustls::server::WebPkiClientVerifier::builder_with_provider(
Arc::new(roots),
provider,
)
.build()
.unwrap();
builder.with_client_cert_verifier(verifier)
}
};
let server_config = builder.with_single_cert(vec![cert], key).unwrap();
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();
let router: Router = Router::new().route("/", get(|| async { "hello" }));
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 io = TokioIo::new(tls_stream);
let svc = TowerToHyperService::new(router);
http1::Builder::new().serve_connection(io, svc).await.ok();
});
}
});
addr
}
async fn get_https(client: &HttpClient, addr: SocketAddr) -> Result<String, HttpClientError> {
let req = http::Request::builder()
.method(http::Method::GET)
.uri(format!("https://{addr}/"))
.body(empty_body())
.unwrap();
let resp = client.clone().oneshot(req).await?;
let bytes = resp.into_body().collect().await.unwrap().to_bytes();
Ok(String::from_utf8(bytes.to_vec()).unwrap())
}
#[tokio::test]
async fn bundle_only_trust_accepts_signed_server() {
let pki = generate_server_pki();
let addr = spawn_tls_server(pki.leaf_cert_der, pki.leaf_key_der).await;
let client = new_client(&parse_config_with_vars(
indoc! {r#"
tls:
use_native_certificate_store: false
certificate_authorities: ${env.CA_PEM}
"#},
&[("CA_PEM", &pki.ca_pem)],
));
let body = get_https(&client, addr).await.expect("request succeeds");
assert_eq!(body, "hello");
}
#[tokio::test]
async fn bundle_plus_native_trust_accepts_signed_server() {
let pki = generate_server_pki();
let addr = spawn_tls_server(pki.leaf_cert_der, pki.leaf_key_der).await;
let client = new_client(&parse_config_with_vars(
indoc! {r#"
tls:
use_native_certificate_store: true
certificate_authorities: ${env.CA_PEM}
"#},
&[("CA_PEM", &pki.ca_pem)],
));
let body = get_https(&client, addr).await.expect("request succeeds");
assert_eq!(body, "hello");
}
#[tokio::test]
async fn multi_cert_pem_is_supported() {
let pki_a = generate_server_pki();
let pki_b = generate_server_pki();
let addr_b = spawn_tls_server(pki_b.leaf_cert_der, pki_b.leaf_key_der).await;
let bundle = format!("{}{}", pki_a.ca_pem, pki_b.ca_pem);
let client = new_client(&parse_config_with_vars(
indoc! {r#"
tls:
use_native_certificate_store: false
certificate_authorities: ${env.CA_PEM}
"#},
&[("CA_PEM", &bundle)],
));
let body = get_https(&client, addr_b).await.expect("request succeeds");
assert_eq!(body, "hello");
}
#[tokio::test]
async fn native_only_rejects_unknown_ca() {
let pki = generate_server_pki();
let addr = spawn_tls_server(pki.leaf_cert_der, pki.leaf_key_der).await;
let client = new_client(&parse_config("{}"));
let err = get_https(&client, addr)
.await
.expect_err("untrusted CA should be rejected");
assert!(
matches!(err, HttpClientError::Transport { .. }),
"expected Transport error, got {err:?}"
);
}
#[tokio::test]
async fn bundle_only_rejects_server_signed_by_native_ca() {
let server_pki = generate_server_pki();
let other_pki = generate_server_pki();
let addr = spawn_tls_server(server_pki.leaf_cert_der, server_pki.leaf_key_der).await;
let client = new_client(&parse_config_with_vars(
indoc! {r#"
tls:
use_native_certificate_store: false
certificate_authorities: ${env.CA_PEM}
"#},
&[("CA_PEM", &other_pki.ca_pem)],
));
let err = get_https(&client, addr)
.await
.expect_err("server signed by an out-of-bundle CA should be rejected");
assert!(
matches!(err, HttpClientError::Transport { .. }),
"expected Transport error, got {err:?}"
);
}
#[test]
fn invalid_pem_returns_error_at_construction() {
let config = parse_config(indoc! {"
tls:
certificate_authorities: not pem
"});
let err = expect_build_err(&config);
assert!(
matches!(err, HttpClientError::TlsCertificateAuthoritiesEmpty),
"expected TlsCertificateAuthoritiesEmpty, got {err:?}"
);
}
#[test]
fn pem_with_no_certificates_returns_error() {
let config = parse_config(indoc! {"
tls:
certificate_authorities: |
-----BEGIN PRIVATE KEY-----
MIGHAgEAMBMGByqGSM49AgEGCCqGSM49AwEHBG0wawIBAQQg
-----END PRIVATE KEY-----
"});
let err = expect_build_err(&config);
assert!(
matches!(err, HttpClientError::TlsCertificateAuthoritiesEmpty),
"expected TlsCertificateAuthoritiesEmpty, got {err:?}"
);
}
#[test]
fn no_trust_anchors_returns_error() {
let config = parse_config(indoc! {"
tls:
use_native_certificate_store: false
"});
let err = expect_build_err(&config);
assert!(
matches!(err, HttpClientError::TlsNoTrustAnchors),
"expected TlsNoTrustAnchors, got {err:?}"
);
}
#[tokio::test]
async fn danger_accept_invalid_certs_bypasses_verification() {
let pki = generate_server_pki();
let addr = spawn_tls_server(pki.leaf_cert_der, pki.leaf_key_der).await;
let client = new_client(&parse_config(indoc! {"
tls:
use_native_certificate_store: false
danger_accept_invalid_certs: true
"}));
let body = get_https(&client, addr).await.expect("request succeeds");
assert_eq!(body, "hello");
}
#[tokio::test]
async fn mtls_handshake_succeeds_when_client_cert_is_trusted() {
let server_pki = generate_server_pki();
let client_pki = generate_client_pki();
let addr = spawn_tls_server_with_client_auth(
server_pki.leaf_cert_der,
server_pki.leaf_key_der,
Some(&client_pki.ca_pem),
)
.await;
let client = new_client(&parse_mtls_config(
&server_pki.ca_pem,
&client_pki.leaf_cert_pem,
&client_pki.leaf_key_pem,
));
let body = get_https(&client, addr).await.expect("request succeeds");
assert_eq!(body, "hello");
}
#[tokio::test]
async fn mtls_client_works_against_server_without_client_auth() {
let server_pki = generate_server_pki();
let client_pki = generate_client_pki();
let addr = spawn_tls_server(server_pki.leaf_cert_der, server_pki.leaf_key_der).await;
let client = new_client(&parse_mtls_config(
&server_pki.ca_pem,
&client_pki.leaf_cert_pem,
&client_pki.leaf_key_pem,
));
let body = get_https(&client, addr).await.expect("request succeeds");
assert_eq!(body, "hello");
}
#[test]
fn invalid_client_cert_chain_returns_error() {
let client_pki = generate_client_pki();
let err = expect_build_err(&parse_mtls_config(
&client_pki.ca_pem,
"not pem",
&client_pki.leaf_key_pem,
));
assert!(
matches!(err, HttpClientError::TlsClientCertificateChainEmpty),
"expected TlsClientCertificateChainEmpty, got {err:?}"
);
}
#[test]
fn invalid_client_key_returns_error() {
let client_pki = generate_client_pki();
let err = expect_build_err(&parse_mtls_config(
&client_pki.ca_pem,
&client_pki.leaf_cert_pem,
"not a key",
));
assert!(
matches!(err, HttpClientError::TlsClientKeyMissing),
"expected TlsClientKeyMissing, got {err:?}"
);
}
#[test]
fn mismatched_client_cert_and_key_return_client_auth_error() {
let pki_a = generate_client_pki();
let pki_b = generate_client_pki();
let err = expect_build_err(&parse_mtls_config(
&pki_a.ca_pem,
&pki_a.leaf_cert_pem,
&pki_b.leaf_key_pem,
));
assert!(
matches!(err, HttpClientError::TlsClientAuth { .. }),
"expected TlsClientAuth, got {err:?}"
);
}
#[test]
fn encrypted_client_key_returns_targeted_error() {
let client_pki = generate_client_pki();
let encrypted_key =
"-----BEGIN ENCRYPTED PRIVATE KEY-----\nQUFBQUFBQUFB\n-----END ENCRYPTED PRIVATE KEY-----";
let err = expect_build_err(&parse_mtls_config(
&client_pki.ca_pem,
&client_pki.leaf_cert_pem,
encrypted_key,
));
assert!(
matches!(err, HttpClientError::TlsClientKeyEncrypted),
"expected TlsClientKeyEncrypted, got {err:?}"
);
}
#[tokio::test]
async fn builder_with_tls_config_uses_injected_ca() {
let pki = generate_server_pki();
let addr = spawn_tls_server(pki.leaf_cert_der, pki.leaf_key_der).await;
let provider = Arc::new(aws_lc_rs_crypto::default_provider());
let mut root_store = rustls::RootCertStore::empty();
for cert in rustls::pki_types::CertificateDer::pem_slice_iter(pki.ca_pem.as_bytes()) {
root_store.add(cert.unwrap()).unwrap();
}
let client_config = rustls::ClientConfig::builder_with_provider(provider)
.with_safe_default_protocol_versions()
.unwrap()
.with_root_certificates(root_store)
.with_no_client_auth();
let config = parse_config(indoc! {"
tls:
use_native_certificate_store: false
"});
let client = HttpClient::builder(config)
.with_tls_config(Arc::new(client_config))
.build()
.expect("builder with injected TLS config succeeds");
let body = get_https(&client, addr)
.await
.expect("request succeeds with injected CA");
assert_eq!(body, "hello");
}
fn parse_mtls_config(
server_ca_pem: &str,
client_cert_pem: &str,
client_key_pem: &str,
) -> HttpClientConfig {
parse_config_with_vars(
indoc! {r#"
tls:
use_native_certificate_store: false
certificate_authorities: ${env.CA_PEM}
client_authentication:
certificate_chain: ${env.CLIENT_CERT}
key: ${env.CLIENT_KEY}
"#},
&[
("CA_PEM", server_ca_pem),
("CLIENT_CERT", client_cert_pem),
("CLIENT_KEY", client_key_pem),
],
)
}
fn expect_build_err(config: &HttpClientConfig) -> HttpClientError {
match HttpClient::new(config) {
Ok(_) => panic!("expected HttpClient::new to fail"),
Err(e) => e,
}
}