use std::net::SocketAddr;
use std::path::PathBuf;
use std::sync::atomic::{AtomicU32, Ordering};
use std::sync::Arc;
use std::time::Duration;
use alkhttp::client::{ClientCertConfig, HttpClientBuildError, HttpClientConfig, SharedHttpClient};
struct TestPki {
ca_pem: Vec<u8>,
server_pem: Vec<u8>,
server_key_pem: Vec<u8>,
client_pem: Vec<u8>,
client_key_pem: Vec<u8>,
}
impl TestPki {
fn generate() -> Self {
let mut ca_params =
rcgen::CertificateParams::new(vec!["alkhttp test CA".to_string()]).expect("CA params");
ca_params.is_ca = rcgen::IsCa::Ca(rcgen::BasicConstraints::Unconstrained);
let ca_key = rcgen::KeyPair::generate().expect("CA key");
let ca_cert = ca_params.self_signed(&ca_key).expect("self-signed CA");
let issuer = rcgen::Issuer::from_params(&ca_params, &ca_key);
let mut server_params =
rcgen::CertificateParams::new(vec!["127.0.0.1".to_string(), "localhost".to_string()])
.expect("server params");
server_params.is_ca = rcgen::IsCa::NoCa;
let server_key = rcgen::KeyPair::generate().expect("server key");
let server_cert = server_params
.signed_by(&server_key, &issuer)
.expect("server leaf");
let mut client_params =
rcgen::CertificateParams::new(vec!["alkhttp test client".to_string()])
.expect("client params");
client_params.is_ca = rcgen::IsCa::NoCa;
client_params.extended_key_usages = vec![rcgen::ExtendedKeyUsagePurpose::ClientAuth];
let client_key = rcgen::KeyPair::generate().expect("client key");
let client_cert = client_params
.signed_by(&client_key, &issuer)
.expect("client leaf");
Self {
ca_pem: ca_cert.pem().into_bytes(),
server_pem: server_cert.pem().into_bytes(),
server_key_pem: server_key.serialize_pem().into_bytes(),
client_pem: client_cert.pem().into_bytes(),
client_key_pem: client_key.serialize_pem().into_bytes(),
}
}
fn write_config_files(
&self,
with_client_cert: bool,
) -> (Option<PathBuf>, Option<ClientCertConfig>, PathBuf) {
let dir = std::env::temp_dir().join(format!(
"alkhttp-tls-test-{}-{}",
std::process::id(),
uuid::Uuid::new_v4()
));
std::fs::create_dir_all(&dir).expect("temp dir");
let write = |name: &str, bytes: &[u8]| {
let path = dir.join(name);
std::fs::write(&path, bytes).expect("write pem");
path
};
let ca = Some(write("ca.pem", &self.ca_pem));
let client = if with_client_cert {
Some(ClientCertConfig {
cert_pem: write("client-cert.pem", &self.client_pem),
key_pem: write("client-key.pem", &self.client_key_pem),
})
} else {
None
};
(ca, client, dir)
}
}
struct TlsTestServer {
origin: String,
shutdown: Option<tokio::sync::oneshot::Sender<()>>,
handshakes: Arc<AtomicU32>,
}
impl TlsTestServer {
fn handshakes(&self) -> u32 {
self.handshakes.load(Ordering::SeqCst)
}
async fn spawn(pki: &TestPki, require_client_cert: bool) -> Self {
use rustls_pki_types::pem::PemObject;
let server_certs: Vec<rustls_pki_types::CertificateDer<'static>> =
rustls_pki_types::pem::PemObject::pem_slice_iter(&pki.server_pem)
.map(|c: Result<rustls_pki_types::CertificateDer<'_>, _>| {
c.expect("server cert parses")
})
.collect();
let server_key = rustls_pki_types::PrivateKeyDer::from_pem_slice(&pki.server_key_pem)
.expect("server key parses");
let server_trust = if require_client_cert {
let mut trust = rustls::RootCertStore::empty();
let ca_iter = rustls_pki_types::pem::PemObject::pem_slice_iter(&pki.ca_pem).map(
|c: Result<rustls_pki_types::CertificateDer<'_>, _>| c.expect("CA cert parses"),
);
for ca in ca_iter {
trust.add(ca).expect("CA added to server trust store");
}
Some(trust)
} else {
None
};
let config = match &server_trust {
Some(trust) => {
let verifier =
rustls::server::WebPkiClientVerifier::builder(Arc::new(trust.clone()))
.build()
.expect("client verifier");
rustls::ServerConfig::builder()
.with_client_cert_verifier(verifier)
.with_single_cert(server_certs, server_key)
.expect("server config with client auth")
}
None => rustls::ServerConfig::builder()
.with_no_client_auth()
.with_single_cert(server_certs, server_key)
.expect("server config"),
};
let tls_config = Arc::new(config);
let handshakes = Arc::new(AtomicU32::new(0));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("bind 127.0.0.1:0");
let addr: SocketAddr = listener.local_addr().expect("local addr");
let (shutdown_tx, shutdown_rx) = tokio::sync::oneshot::channel::<()>();
let hs_counter = Arc::clone(&handshakes);
tokio::spawn(async move {
let acceptor = tokio_rustls::TlsAcceptor::from(tls_config);
let mut shutdown = std::pin::pin!(shutdown_rx);
loop {
let accept = tokio::select! {
_ = &mut shutdown => break,
accepted = listener.accept() => match accepted {
Ok((sock, _)) => sock,
Err(_) => break,
},
};
let acceptor = acceptor.clone();
let hs = Arc::clone(&hs_counter);
tokio::spawn(async move {
let Ok(mut tls_stream) = acceptor.accept(accept).await else {
return;
};
hs.fetch_add(1, Ordering::SeqCst);
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let mut buf = [0u8; 4096];
loop {
let n = tls_stream.read(&mut buf).await.unwrap_or(0);
if n == 0 {
break;
}
if String::from_utf8_lossy(&buf[..n]).contains("\r\n\r\n") {
break;
}
}
let body = b"ok";
let response = format!(
"HTTP/1.1 200 OK\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
body.len(),
String::from_utf8_lossy(body),
);
let _ = tls_stream.write_all(response.as_bytes()).await;
let _ = tls_stream.shutdown().await;
});
}
});
Self {
origin: format!("https://127.0.0.1:{}", addr.port()),
shutdown: Some(shutdown_tx),
handshakes,
}
}
}
impl Drop for TlsTestServer {
fn drop(&mut self) {
if let Some(shutdown) = self.shutdown.take() {
let _ = shutdown.send(());
}
}
}
fn client_config(ca: Option<PathBuf>, cert: Option<ClientCertConfig>) -> HttpClientConfig {
HttpClientConfig {
ca_bundle: ca,
client_cert: cert,
..HttpClientConfig::default()
}
}
fn cleanup_dir(dir: &PathBuf) {
let _ = std::fs::remove_dir_all(dir);
}
#[tokio::test]
async fn client_with_ca_bundle_connects_to_private_roots_server() {
let pki = TestPki::generate();
let server = TlsTestServer::spawn(&pki, false).await;
let (ca, _cert, dir) = pki.write_config_files(false);
let http = SharedHttpClient::new(client_config(ca, None)).expect("client builds with CA");
let response = http
.client()
.get(format!("{}/ping", server.origin))
.send()
.await
.expect("request over private roots succeeds");
assert_eq!(response.status(), 200, "server answers over TLS");
assert_eq!(
response.text().await.unwrap(),
"ok",
"the TLS-secured body arrives intact"
);
assert_eq!(
server.handshakes(),
1,
"exactly one TLS handshake was completed"
);
cleanup_dir(&dir);
}
#[tokio::test]
async fn client_without_ca_bundle_rejects_private_roots_server() {
let pki = TestPki::generate();
let server = TlsTestServer::spawn(&pki, false).await;
let http = SharedHttpClient::new(HttpClientConfig::default())
.expect("client builds with default (webpki) roots");
let result = http
.client()
.get(format!("{}/ping", server.origin))
.send()
.await;
let error = result
.expect_err("a private-roots server must be rejected by a client without the CA bundle");
let text = error_chain_text(&error);
assert!(
text.contains("certificate"),
"the chain names the TLS verification failure, got: {text}"
);
}
fn error_chain_text(error: &reqwest_middleware::Error) -> String {
let mut text = error.to_string().to_lowercase();
let mut source = std::error::Error::source(error);
while let Some(err) = source {
text.push(' ');
text.push_str(&err.to_string().to_lowercase());
source = err.source();
}
text
}
#[tokio::test]
async fn mtls_client_cert_is_presented_and_accepted_end_to_end() {
let pki = TestPki::generate();
let server = TlsTestServer::spawn(&pki, true).await;
let (ca, cert, dir) = pki.write_config_files(true);
let http = SharedHttpClient::new(client_config(ca, cert))
.expect("client builds with CA bundle + client identity");
let response = http
.client()
.get(format!("{}/ping", server.origin))
.send()
.await
.expect("mTLS handshake with client identity succeeds");
assert_eq!(response.status(), 200, "server answers the mTLS client");
assert_eq!(response.text().await.unwrap(), "ok");
assert_eq!(
server.handshakes(),
1,
"the client-cert handshake completed through the full middleware stack"
);
cleanup_dir(&dir);
}
#[tokio::test]
async fn mtls_server_rejects_client_without_identity() {
let pki = TestPki::generate();
let server = TlsTestServer::spawn(&pki, true).await;
let (ca, _cert, dir) = pki.write_config_files(false);
let http = SharedHttpClient::new(client_config(ca, None))
.expect("client builds with CA bundle but no client identity");
let result = http
.client()
.get(format!("{}/ping", server.origin))
.send()
.await;
let error = result
.expect_err("an mTLS-requiring server must reject a client that presents no certificate");
let text = error_chain_text(&error);
assert!(
text.contains("certificate") || text.contains("alert") || text.contains("handshake"),
"the chain names a TLS/certificate-level rejection, got: {text}"
);
cleanup_dir(&dir);
}
#[tokio::test]
async fn reload_to_a_ca_bundle_backed_client_succeeds() {
let pki = TestPki::generate();
let server = TlsTestServer::spawn(&pki, false).await;
let (ca, _cert, dir) = pki.write_config_files(false);
let http = SharedHttpClient::new(HttpClientConfig::default()).expect("initial client");
assert!(
http.client()
.get(format!("{}/ping", server.origin))
.send()
.await
.is_err(),
"before the reload the private-roots server is unreachable"
);
let reloaded = client_config(ca, None);
http.reload(reloaded)
.await
.expect("reload with a valid CA bundle succeeds");
let response = http
.client()
.get(format!("{}/ping", server.origin))
.send()
.await
.expect("after the reload the CA bundle is trusted");
assert_eq!(response.status(), 200);
cleanup_dir(&dir);
tokio::time::sleep(Duration::from_millis(1)).await;
}
#[test]
fn nonexistent_ca_bundle_path_fails_ca_bundle_read_with_path() {
let path = std::env::temp_dir().join(format!(
"alkhttp-pem-read-{}-{}-missing.pem",
std::process::id(),
uuid::Uuid::new_v4()
));
let error = SharedHttpClient::new(client_config(Some(path.clone()), None))
.expect_err("a nonexistent CA path must fail the build");
match error {
HttpClientBuildError::CaBundleRead { path: p, .. } => {
assert_eq!(p, path, "the error names the unreadable path");
}
other => panic!("expected CaBundleRead, got {other:?}"),
}
}
#[tokio::test]
async fn reload_with_nonexistent_client_cert_path_fails_client_cert_read() {
let pki = TestPki::generate();
let (_ca, _cert, dir) = pki.write_config_files(false);
let http = SharedHttpClient::new(HttpClientConfig::default()).expect("initial client");
let missing_key = std::env::temp_dir().join(format!(
"alkhttp-pem-read-{}-{}-missing-key.pem",
std::process::id(),
uuid::Uuid::new_v4()
));
let missing_cert = std::env::temp_dir().join(format!(
"alkhttp-pem-read-{}-{}-missing-cert.pem",
std::process::id(),
uuid::Uuid::new_v4()
));
let error = http
.reload(client_config(
None,
Some(ClientCertConfig {
cert_pem: missing_cert.clone(),
key_pem: missing_key.clone(),
}),
))
.await
.expect_err("a nonexistent client-cert key path must fail the reload");
match error {
HttpClientBuildError::ClientCertRead { path: p, .. } => {
assert_eq!(
p, missing_cert,
"the error names the unreadable cert path (read before the key)"
);
}
other => panic!("expected ClientCertRead, got {other:?}"),
}
assert!(
http.config().client_cert.is_none(),
"the reload failure keeps the previous generation's clients+config (FWD-12)"
);
cleanup_dir(&dir);
}
#[test]
fn corrupt_ca_bundle_fails_ca_bundle_parse_with_path() {
let dir = std::env::temp_dir().join(format!(
"alkhttp-pem-parse-{}-{}",
std::process::id(),
uuid::Uuid::new_v4()
));
std::fs::create_dir_all(&dir).expect("temp dir");
let ca_path = dir.join("ca.pem");
std::fs::write(
&ca_path,
b"-----BEGIN CERTIFICATE-----\n!!not-base64!!\n-----END CERTIFICATE-----\n",
)
.expect("write corrupt pem");
let error = SharedHttpClient::new(client_config(Some(ca_path.clone()), None))
.expect_err("a corrupt CA bundle must fail the build");
match error {
HttpClientBuildError::CaBundleParse { path, .. } => {
assert_eq!(path, ca_path, "the error names the unparseable path");
}
other => panic!("expected CaBundleParse, got {other:?}"),
}
cleanup_dir(&dir);
}
#[test]
fn garbage_client_cert_fails_client_cert_parse_with_path_and_no_key_material() {
let dir = std::env::temp_dir().join(format!(
"alkhttp-pem-parse-{}-{}",
std::process::id(),
uuid::Uuid::new_v4()
));
std::fs::create_dir_all(&dir).expect("temp dir");
let cert_path = dir.join("client-cert.pem");
let key_path = dir.join("client-key.pem");
std::fs::write(
&cert_path,
b"-----BEGIN GARBAGE-----\nnope\n-----END GARBAGE-----\n",
)
.expect("write garbage cert");
std::fs::write(
&key_path,
b"-----BEGIN PRIVATE KEY-----\nnot-a-key\n-----END PRIVATE KEY-----\n",
)
.expect("write garbage key");
let key_marker = "not-a-key";
let error = SharedHttpClient::new(client_config(
None,
Some(ClientCertConfig {
cert_pem: cert_path.clone(),
key_pem: key_path,
}),
))
.expect_err("a non-parseable client identity must fail the build");
let rendered = format!("{error}");
match error {
HttpClientBuildError::ClientCertParse { path, .. } => {
assert_eq!(path, cert_path, "the error names the identity's cert path");
}
other => panic!("expected ClientCertParse, got {other:?}"),
}
assert!(
!rendered.contains(key_marker),
"the error must never echo key material: {rendered}"
);
cleanup_dir(&dir);
}
#[tokio::test]
async fn config_accessor_reflects_a_tls_config_reload() {
let pki = TestPki::generate();
let (ca, _cert, dir) = pki.write_config_files(false);
let http = SharedHttpClient::new(HttpClientConfig::default()).expect("initial client");
assert!(
http.config().ca_bundle.is_none(),
"the initial config has no CA bundle"
);
http.reload(client_config(ca, None))
.await
.expect("reload with the CA bundle succeeds");
let visible = http.config();
assert!(
visible.ca_bundle.is_some(),
"config() reflects the reloaded generation's CA bundle (FWD-12)"
);
cleanup_dir(&dir);
}