#[cfg(any(test, feature = "test-support"))]
use std::sync::{
Condvar, Mutex,
atomic::{AtomicBool, Ordering},
};
use std::{net::SocketAddr, sync::Arc, time::Instant};
use pingora_core::{
Result,
upstreams::peer::{HttpPeer, HttpUpstreamRequestPolicy},
};
use praxis_core::connectivity::{Upstream, peer as peer_utils};
use tracing::debug;
use super::super::context::PingoraRequestCtx;
#[cfg(any(test, feature = "test-support"))]
static UPSTREAM_RETRY_GATE_ARMED: AtomicBool = AtomicBool::new(false);
#[cfg(any(test, feature = "test-support"))]
static UPSTREAM_RETRY_GATE_PARK: Mutex<()> = Mutex::new(());
#[cfg(any(test, feature = "test-support"))]
static UPSTREAM_RETRY_GATE_CV: Condvar = Condvar::new();
#[cfg(any(test, feature = "test-support"))]
static UPSTREAM_RETRY_GATE_TEST_LOCK: Mutex<()> = Mutex::new(());
#[cfg(any(test, feature = "test-support"))]
#[doc(hidden)]
pub fn lock_upstream_retry_gate_tests() -> std::sync::MutexGuard<'static, ()> {
#[expect(clippy::expect_used, reason = "poisoned mutex is unrecoverable")]
UPSTREAM_RETRY_GATE_TEST_LOCK
.lock()
.expect("upstream retry gate test lock")
}
#[cfg(any(test, feature = "test-support"))]
#[doc(hidden)]
pub struct UpstreamRetryGateRelease;
#[cfg(any(test, feature = "test-support"))]
impl Drop for UpstreamRetryGateRelease {
fn drop(&mut self) {
clear_upstream_retry_gate_wait();
}
}
#[cfg(any(test, feature = "test-support"))]
fn clear_upstream_retry_gate_wait() {
UPSTREAM_RETRY_GATE_ARMED.store(false, Ordering::SeqCst);
UPSTREAM_RETRY_GATE_CV.notify_all();
}
#[cfg(any(test, feature = "test-support"))]
#[doc(hidden)]
pub fn arm_upstream_retry_gate() -> (std::sync::MutexGuard<'static, ()>, UpstreamRetryGateRelease) {
let guard = lock_upstream_retry_gate_tests();
UPSTREAM_RETRY_GATE_ARMED.store(true, Ordering::SeqCst);
(guard, UpstreamRetryGateRelease)
}
#[cfg(any(test, feature = "test-support"))]
fn wait_for_upstream_retry_gate(retries: u32) {
if retries > 0 && UPSTREAM_RETRY_GATE_ARMED.load(Ordering::SeqCst) {
#[expect(clippy::expect_used, reason = "poisoned mutex/condvar is unrecoverable")]
{
let mut park = UPSTREAM_RETRY_GATE_PARK.lock().expect("upstream retry gate park lock");
while UPSTREAM_RETRY_GATE_ARMED.load(Ordering::SeqCst) {
park = UPSTREAM_RETRY_GATE_CV.wait(park).expect("upstream retry gate wait");
}
drop(park);
}
}
}
#[expect(clippy::too_many_lines, reason = "retry orchestration reads clearer as one function")]
pub(super) async fn execute(ctx: &mut PingoraRequestCtx) -> Result<Box<HttpPeer>> {
#[cfg(any(test, feature = "test-support"))]
wait_for_upstream_retry_gate(ctx.retries);
if let Some(backoff) = ctx.pending_backoff.take()
&& !backoff.is_zero()
{
debug!(?backoff, "applying retry backoff");
tokio::time::sleep(backoff).await;
}
if ctx.reselect_on_retry {
ctx.reselect_on_retry = false;
if let Some(reselector) = ctx.endpoint_reselector.clone() {
let health = ctx
.pinned_pipeline
.as_ref()
.and_then(|p| p.health_registry())
.and_then(|reg| ctx.cluster.as_deref().and_then(|c| reg.get(c)));
if let Some(addr) = reselector.select_address(health, &ctx.attempted_endpoints) {
debug!(upstream = %addr, "selected alternate host for retry");
if let Some(prev) = ctx.upstream_for_retry.as_ref() {
reselector.release(&prev.address);
}
if !ctx.attempted_endpoints.iter().any(|e| e.as_ref() == addr.as_ref()) {
ctx.attempted_endpoints.push(Arc::clone(&addr));
}
ctx.selected_endpoint_index = Some(reselected_endpoint_index(health, &addr));
let mut upstream = reselector.build_upstream(addr);
carry_forward_sni(&mut upstream, ctx.upstream_for_retry.as_ref());
apply_per_try_timeout(ctx, &mut upstream);
ctx.upstream_for_retry = Some(upstream);
} else {
debug!("no alternate host available; reusing previous upstream if present");
}
}
}
ctx.upstream_connect_start = Some(Instant::now());
if ctx.upstream_for_retry.is_none() {
let mut upstream = ctx.upstream.take();
if let Some(ref mut u) = upstream {
apply_per_try_timeout(ctx, u);
}
ctx.upstream_for_retry = upstream;
}
if ctx.upstream_for_retry.is_some() {
ctx.upstream_contacted = true;
}
if let Some(deadline) = ctx.extensions.get::<praxis_core::grpc::GrpcDeadline>().copied()
&& let Some(upstream) = ctx.upstream_for_retry.as_mut()
{
apply_grpc_deadline(deadline, upstream);
}
let upstream = ctx.upstream_for_retry.as_ref().ok_or_else(|| {
let cluster = &ctx.cluster;
pingora_core::Error::explain(
pingora_core::ErrorType::InternalError,
format!("no upstream selected (cluster: {cluster:?}); is a load_balancer configured?"),
)
})?;
let allow_private = ctx
.pinned_pipeline
.as_ref()
.is_some_and(|pipeline| pipeline.allow_private_upstreams());
build_peer(upstream, allow_private).await
}
fn reselected_endpoint_index(health: Option<&praxis_core::health::ClusterHealthState>, addr: &str) -> usize {
health.and_then(|h| h.endpoint_index(addr)).unwrap_or(usize::MAX)
}
fn apply_grpc_deadline(deadline: praxis_core::grpc::GrpcDeadline, upstream: &mut Upstream) {
let Some(remaining) = deadline.remaining() else {
return;
};
let opts = Arc::make_mut(&mut upstream.connection);
opts.connection_timeout = Some(shorter(opts.connection_timeout, remaining));
opts.total_connection_timeout = Some(shorter(opts.total_connection_timeout, remaining));
opts.read_timeout = Some(shorter(opts.read_timeout, remaining));
opts.write_timeout = Some(shorter(opts.write_timeout, remaining));
}
fn shorter(configured: Option<std::time::Duration>, remaining: std::time::Duration) -> std::time::Duration {
configured.map_or(remaining, |configured| configured.min(remaining))
}
fn carry_forward_sni(upstream: &mut Upstream, previous: Option<&Upstream>) {
let Some(sni) = previous.and_then(|p| p.tls.as_ref()).and_then(|t| t.sni()) else {
return;
};
if let Some(tls) = upstream.tls.as_mut()
&& tls.sni().is_none()
{
tls.set_sni(Arc::<str>::from(sni));
}
}
fn apply_per_try_timeout(ctx: &PingoraRequestCtx, upstream: &mut Upstream) {
let Some(policy) = ctx.retry_policy.as_ref() else {
return;
};
let Some(per_try_ms) = policy.per_try_timeout_ms else {
return;
};
let opts = Arc::make_mut(&mut upstream.connection);
let timeout = std::time::Duration::from_millis(per_try_ms);
opts.connection_timeout = Some(timeout);
opts.total_connection_timeout = Some(timeout);
opts.read_timeout = Some(timeout);
opts.write_timeout = Some(timeout);
}
async fn build_peer(upstream: &Upstream, allow_private: bool) -> Result<Box<HttpPeer>> {
let addr: SocketAddr = resolve_upstream(upstream, allow_private).await?;
let tls_enabled = upstream.tls.is_some();
let sni = upstream
.tls
.as_ref()
.and_then(|t| t.sni().map(str::to_owned))
.unwrap_or_else(|| {
if tls_enabled {
peer_utils::derive_sni(&upstream.address)
} else {
String::new()
}
});
let mut peer = HttpPeer::new(addr, tls_enabled, sni);
peer.options.http_upstream_request_policy = HttpUpstreamRequestPolicy::preserve();
peer_utils::apply_connection_options(&mut peer, &upstream.connection);
if let Some(tls) = &upstream.tls {
peer_utils::apply_cached_tls(&mut peer, tls, &upstream.address);
}
Ok(Box::new(peer))
}
async fn resolve_upstream(upstream: &Upstream, allow_private: bool) -> Result<SocketAddr> {
peer_utils::resolve_upstream_checked(upstream, allow_private)
.await
.map_err(|error| {
let etype = match error {
peer_utils::AddressResolutionError::PrivateAddress { .. }
| peer_utils::AddressResolutionError::UntrustedRange { .. } => pingora_core::ErrorType::ConnectError,
_ => pingora_core::ErrorType::InternalError,
};
pingora_core::Error::explain(etype, error.to_string())
})
}
#[cfg(test)]
#[expect(clippy::allow_attributes, reason = "blanket test suppressions")]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::indexing_slicing,
clippy::field_reassign_with_default,
clippy::too_many_lines,
clippy::significant_drop_tightening,
clippy::print_stderr,
reason = "tests"
)]
mod tests {
use praxis_core::connectivity::ConnectionOptions;
use praxis_tls::{CachedClusterTls, ClusterTls};
use super::*;
#[test]
fn reselected_upstream_keeps_the_previous_attempts_sni() {
let mut previous = tls_upstream("10.0.0.1:443", None);
if let Some(tls) = previous.tls.as_mut() {
tls.set_sni("api.example.com");
}
let mut reselected = tls_upstream("10.0.0.2:443", None);
carry_forward_sni(&mut reselected, Some(&previous));
assert_eq!(
reselected.tls.as_ref().and_then(|t| t.sni()),
Some("api.example.com"),
"a retry must present the same SNI as the first attempt"
);
}
#[test]
fn reselected_upstream_keeps_an_explicit_cluster_sni() {
let previous = tls_upstream("10.0.0.1:443", Some("from-host.example.com"));
let mut reselected = tls_upstream("10.0.0.2:443", Some("configured.example.com"));
carry_forward_sni(&mut reselected, Some(&previous));
assert_eq!(
reselected.tls.as_ref().and_then(|t| t.sni()),
Some("configured.example.com"),
"an SNI set in the cluster config must win over the carried one"
);
}
#[tokio::test]
async fn reselected_peer_presents_its_own_sni_when_authority_follows_the_endpoint() {
for host in ["alpha.reselect-sni.test", "beta.reselect-sni.test"] {
peer_utils::seed_dns(host, &[std::net::IpAddr::from([127, 0, 0, 1])]);
}
let config: serde_yaml::Value = serde_yaml::from_str(
r#"
clusters:
- name: mixed
endpoints: ["alpha.reselect-sni.test:443", "beta.reselect-sni.test:443"]
http:
authority: { from: endpoint }
tls: {}
"#,
)
.unwrap();
let lb = praxis_filter::LoadBalancerFilter::from_config(&config).unwrap();
let mut pipeline =
praxis_filter::FilterPipeline::build(&mut [], &praxis_filter::FilterRegistry::with_builtins()).unwrap();
pipeline.set_allow_private_upstreams(true);
let pipeline = Arc::new(pipeline);
let mut headers = http::HeaderMap::new();
headers.insert(http::header::HOST, http::HeaderValue::from_static("client.example.com"));
let request = praxis_filter::Request {
method: http::Method::GET,
uri: http::Uri::from_static("/"),
headers,
};
let mut ctx = PingoraRequestCtx::default();
ctx.cluster = Some(Arc::from("mixed"));
let mut filter_ctx = ctx.build_filter_context(&pipeline, &request, None);
drop(lb.on_request(&mut filter_ctx).await.unwrap());
ctx.cluster = filter_ctx.cluster.take();
ctx.upstream = filter_ctx.upstream.take();
ctx.endpoint_reselector = filter_ctx.endpoint_reselector.take();
ctx.attempted_endpoints = std::mem::take(&mut filter_ctx.attempted_endpoints);
drop(filter_ctx);
ctx.pinned_pipeline = Some(Arc::clone(&pipeline));
let first = execute(&mut ctx).await.expect("first attempt should build a peer");
let first_address = Arc::clone(&ctx.upstream_for_retry.as_ref().unwrap().address);
assert_eq!(
first.sni,
peer_utils::derive_sni(&first_address),
"the first attempt should present its endpoint's name, not the downstream Host"
);
ctx.reselect_on_retry = true;
let retry = execute(&mut ctx).await.expect("the retry should build a peer");
let retry_address = Arc::clone(&ctx.upstream_for_retry.as_ref().unwrap().address);
assert_ne!(
retry_address, first_address,
"the retry should reselect the other endpoint"
);
assert_eq!(
retry.sni,
peer_utils::derive_sni(&retry_address),
"a reselected endpoint must present its own name, not carry the first attempt's SNI forward"
);
}
#[tokio::test]
async fn valid_address_builds_peer() {
assert!(
build_peer(&make_upstream("127.0.0.1:8080"), false).await.is_ok(),
"valid address should build peer"
);
}
#[tokio::test]
async fn build_peer_with_tls_enabled() {
let tls = ClusterTls {
sni: Some("api.example.com".to_owned()),
..ClusterTls::default()
};
let upstream = Upstream {
address: Arc::from("127.0.0.1:8443"),
authority: None,
connection: Arc::new(ConnectionOptions::default()),
tls: Some(CachedClusterTls::try_from_config(&tls).unwrap()),
};
let peer = build_peer(&upstream, false).await.expect("should build TLS peer");
assert!(!peer.sni.is_empty(), "TLS peer should have a non-empty SNI");
assert_eq!(peer.sni, "api.example.com", "peer SNI should match configured value");
}
#[test]
fn sni_not_set_with_hostname_address_derives_sni() {
let sni = peer_utils::derive_sni("backend.example.com:8443");
assert_eq!(
sni, "backend.example.com",
"SNI should be derived from hostname address"
);
}
#[test]
fn sni_not_set_with_ip_address_leaves_sni_empty() {
let sni = peer_utils::derive_sni("127.0.0.1:8443");
assert_eq!(sni, "", "SNI should be empty for IP address");
}
#[tokio::test]
async fn build_peer_without_tls() {
let upstream = Upstream {
address: Arc::from("127.0.0.1:8080"),
authority: None,
connection: Arc::new(ConnectionOptions::default()),
tls: None,
};
let peer = build_peer(&upstream, false).await.expect("should build plain peer");
assert_eq!(peer.sni, "", "plain peer should have empty SNI");
}
#[tokio::test]
async fn build_peer_with_tls_verify_disabled() {
let tls = ClusterTls {
sni: Some("self-signed.local".to_owned()),
verify: false,
..ClusterTls::default()
};
let upstream = Upstream {
address: Arc::from("127.0.0.1:8443"),
authority: None,
connection: Arc::new(ConnectionOptions::default()),
tls: Some(CachedClusterTls::try_from_config(&tls).unwrap()),
};
let peer = build_peer(&upstream, false)
.await
.expect("should build peer with verification disabled");
assert!(
!peer.options.verify_cert,
"verify_cert should be false when verify is disabled"
);
assert!(
!peer.options.verify_hostname,
"verify_hostname should be false when verify is disabled"
);
}
#[tokio::test]
async fn build_peer_with_tls_verify_enabled() {
let tls = ClusterTls {
sni: Some("api.example.com".to_owned()),
..ClusterTls::default()
};
let upstream = Upstream {
address: Arc::from("127.0.0.1:8443"),
authority: None,
connection: Arc::new(ConnectionOptions::default()),
tls: Some(CachedClusterTls::try_from_config(&tls).unwrap()),
};
let peer = build_peer(&upstream, false)
.await
.expect("should build peer with verification enabled");
assert!(
peer.options.verify_cert,
"verify_cert should be true (default) when verify is enabled"
);
assert!(
peer.options.verify_hostname,
"verify_hostname should be true (default) when verify is enabled"
);
}
#[tokio::test]
async fn resolve_address_parses_socket_addr() {
let addr = resolve_upstream(&make_upstream("127.0.0.1:8080"), false)
.await
.expect("socket addr should parse");
assert_eq!(addr.port(), 8080, "port should match");
}
#[tokio::test]
async fn resolve_address_resolves_localhost() {
if !localhost_resolution_available() {
eprintln!("skipping: localhost did not resolve in this environment");
return;
}
let addr = resolve_upstream(&make_upstream("localhost:8080"), true)
.await
.expect("localhost should resolve");
assert_eq!(addr.port(), 8080, "port should match");
}
#[tokio::test]
async fn resolve_address_fails_for_no_port() {
assert!(
resolve_upstream(&make_upstream("127.0.0.1"), false).await.is_err(),
"address without port should return error"
);
}
#[tokio::test]
async fn hostname_address_builds_peer() {
if !localhost_resolution_available() {
eprintln!("skipping: localhost did not resolve in this environment");
return;
}
assert!(
build_peer(&make_upstream("localhost:8080"), true).await.is_ok(),
"hostname address should build peer via DNS resolution"
);
}
#[tokio::test]
async fn build_peer_rejects_hostname_resolving_to_private_address() {
if !localhost_resolution_available() {
eprintln!("skipping: localhost did not resolve in this environment");
return;
}
let err = build_peer(&make_upstream("localhost:8080"), false)
.await
.expect_err("a hostname resolving to loopback must be refused by default");
let message = err.to_string();
assert!(
message.contains("private/reserved IP address"),
"the error must explain the SSRF rejection: {message}"
);
assert!(
message.contains("allow_private_upstreams"),
"the error must name the override flag: {message}"
);
}
#[tokio::test]
async fn execute_rejects_private_upstream_without_pinned_override() {
if !localhost_resolution_available() {
eprintln!("skipping: localhost did not resolve in this environment");
return;
}
let mut ctx = PingoraRequestCtx::default();
ctx.upstream = Some(make_upstream("localhost:8080"));
let err = execute(&mut ctx)
.await
.expect_err("an unconfigured pipeline must not permit private upstreams");
assert!(
err.to_string().contains("private/reserved IP address"),
"unexpected error: {err}"
);
}
#[tokio::test]
async fn execute_allows_private_upstream_when_pipeline_permits_it() {
if !localhost_resolution_available() {
eprintln!("skipping: localhost did not resolve in this environment");
return;
}
let mut pipeline =
praxis_filter::FilterPipeline::build(&mut [], &praxis_filter::FilterRegistry::with_builtins())
.expect("an empty pipeline should build");
pipeline.set_allow_private_upstreams(true);
let mut ctx = PingoraRequestCtx::default();
ctx.pinned_pipeline = Some(Arc::new(pipeline));
ctx.upstream = Some(make_upstream("localhost:8080"));
execute(&mut ctx)
.await
.expect("allow_private_upstreams must permit a loopback resolution");
}
#[tokio::test]
async fn invalid_address_returns_error() {
assert!(
build_peer(&make_upstream("invalid host:8080"), false).await.is_err(),
"syntactically invalid address should return error"
);
}
#[tokio::test]
async fn missing_port_returns_error() {
assert!(
build_peer(&make_upstream("127.0.0.1"), false).await.is_err(),
"address without port should return error"
);
}
#[test]
fn reselected_endpoint_index_points_at_new_address() {
use praxis_core::health::{ClusterHealthEntry, ClusterHealthState, EndpointHealth};
let entry = ClusterHealthEntry::new(
vec![EndpointHealth::default(), EndpointHealth::default()],
vec![Arc::from("127.0.0.1:3001"), Arc::from("127.0.0.1:3002")],
Some(3),
Some(2),
);
let health: ClusterHealthState = Arc::new(entry);
assert_eq!(
reselected_endpoint_index(Some(&health), "127.0.0.1:3002"),
1,
"reselection must resolve the reselected address's index"
);
assert_eq!(
reselected_endpoint_index(Some(&health), "10.0.0.9:9999"),
usize::MAX,
"unknown address is a no-op index"
);
assert_eq!(
reselected_endpoint_index(None, "127.0.0.1:3001"),
usize::MAX,
"missing registry is a no-op index"
);
}
#[tokio::test]
async fn execute_first_call_moves_upstream_to_retry() {
let mut ctx = PingoraRequestCtx::default();
ctx.upstream = Some(make_upstream("127.0.0.1:8080"));
let result = execute(&mut ctx).await;
assert!(result.is_ok(), "first execute should succeed");
assert!(ctx.upstream.is_none(), "upstream should be consumed");
assert!(ctx.upstream_for_retry.is_some(), "should save for retry");
assert_eq!(
&*ctx.upstream_for_retry.as_ref().unwrap().address,
"127.0.0.1:8080",
"saved retry address should match original"
);
}
#[tokio::test]
async fn execute_retry_reuses_saved_upstream() {
let mut ctx = PingoraRequestCtx::default();
ctx.upstream = None;
ctx.upstream_for_retry = Some(make_upstream("127.0.0.1:9090"));
let result = execute(&mut ctx).await;
assert!(result.is_ok(), "retry execute should succeed");
assert!(
ctx.upstream_for_retry.is_some(),
"retry upstream should remain for further retries"
);
}
#[tokio::test]
async fn execute_reapplies_grpc_deadline_on_retry() {
let mut ctx = PingoraRequestCtx::default();
ctx.upstream_for_retry = Some(make_upstream("127.0.0.1:9090"));
ctx.extensions.insert(praxis_core::grpc::GrpcDeadline::new(
Instant::now() + std::time::Duration::from_millis(50),
false,
true,
));
let _peer = execute(&mut ctx).await.expect("retry execute should succeed");
let read = ctx
.upstream_for_retry
.as_ref()
.unwrap()
.connection
.read_timeout
.expect("a retry's read timeout must be bounded by the remaining deadline");
assert!(
read <= std::time::Duration::from_millis(50),
"a retry must inherit the remaining deadline, not attempt one's budget: {read:?}"
);
}
#[tokio::test]
async fn execute_enforces_the_deadline_even_when_propagation_is_off() {
let mut ctx = PingoraRequestCtx::default();
ctx.upstream_for_retry = Some(make_upstream("127.0.0.1:9090"));
ctx.extensions.insert(praxis_core::grpc::GrpcDeadline::new(
Instant::now() + std::time::Duration::from_millis(50),
false,
false,
));
let _peer = execute(&mut ctx).await.expect("execute should succeed");
let read = ctx
.upstream_for_retry
.as_ref()
.unwrap()
.connection
.read_timeout
.expect("propagate: false suppresses only the header, not transport enforcement");
assert!(
read <= std::time::Duration::from_millis(50),
"the transport budget must still shrink to the deadline: {read:?}"
);
}
#[tokio::test]
async fn execute_no_upstream_no_retry_returns_error() {
let mut ctx = PingoraRequestCtx::default();
ctx.upstream = None;
ctx.upstream_for_retry = None;
let result = execute(&mut ctx).await;
assert!(result.is_err(), "execute with no upstream should return error");
let err = result.unwrap_err().to_string();
assert!(err.contains("no upstream selected"), "unexpected error message: {err}");
assert!(
err.contains("is a load_balancer configured?"),
"error should mention load_balancer: {err}"
);
}
#[tokio::test]
async fn execute_no_upstream_error_includes_cluster_name() {
let mut ctx = PingoraRequestCtx::default();
ctx.cluster = Some(Arc::from("my-api"));
ctx.upstream = None;
ctx.upstream_for_retry = None;
let result = execute(&mut ctx).await;
assert!(result.is_err(), "execute with no upstream should return error");
let err = result.unwrap_err().to_string();
assert!(err.contains("my-api"), "error should include cluster name: {err}");
}
#[tokio::test]
async fn build_peer_with_cached_ca() {
let ca = gen_ca_file();
let ca_path = ca.ca_path.to_str().expect("ca path should be valid UTF-8");
let tls = ClusterTls {
ca: Some(praxis_tls::CaConfig {
ca_path: ca_path.to_owned(),
crl_paths: Vec::new(),
}),
sni: Some("api.example.com".to_owned()),
..ClusterTls::default()
};
let upstream = Upstream {
address: Arc::from("127.0.0.1:8443"),
authority: None,
connection: Arc::new(ConnectionOptions::default()),
tls: Some(CachedClusterTls::try_from_config(&tls).unwrap()),
};
let peer = build_peer(&upstream, false)
.await
.expect("should build peer with cached CA");
assert!(peer.options.ca.is_some(), "peer should have custom CA set from cache");
}
#[test]
fn ca_from_cached_produces_wrapped_x509() {
let ca = gen_ca_file();
let ca_path = ca.ca_path.to_str().expect("ca path should be valid UTF-8");
let cached = praxis_tls::CachedCaCerts::from_pem_file(ca_path).expect("valid CA should parse");
let wrapped = peer_utils::ca_from_cached(&cached);
assert_eq!(wrapped.len(), 1, "should produce one WrappedX509");
}
#[test]
fn client_cert_from_cached_produces_cert_key() {
praxis_tls::provider::install();
let pair = gen_cert_key_files();
let cert_path = pair.cert_path.to_str().expect("cert path should be valid UTF-8");
let key_path = pair.key_path.to_str().expect("key path should be valid UTF-8");
let cached =
praxis_tls::CachedClientCert::from_pem_files(cert_path, key_path).expect("valid cert+key should parse");
let _cert_key = peer_utils::client_cert_from_cached(&cached);
}
fn localhost_resolution_available() -> bool {
use std::net::ToSocketAddrs as _;
"localhost:8080"
.to_socket_addrs()
.is_ok_and(|mut addrs| addrs.next().is_some())
}
fn make_upstream(address: &str) -> Upstream {
Upstream {
address: Arc::from(address),
authority: None,
connection: Arc::new(ConnectionOptions::default()),
tls: None,
}
}
struct TestCa {
ca_path: std::path::PathBuf,
_temp_dir: tempfile::TempDir,
}
struct TestCertKey {
cert_path: std::path::PathBuf,
key_path: std::path::PathBuf,
_temp_dir: tempfile::TempDir,
}
fn gen_ca_file() -> TestCa {
use rcgen::{CertificateParams, DnType, IsCa, KeyPair};
let ca_key = KeyPair::generate().expect("CA key generation should succeed");
let mut ca_params = CertificateParams::new(Vec::<String>::new()).expect("CA params should be valid");
ca_params.is_ca = IsCa::Ca(rcgen::BasicConstraints::Unconstrained);
ca_params.distinguished_name.push(DnType::CommonName, "Test CA");
let ca_cert = ca_params.self_signed(&ca_key).expect("CA self-sign should succeed");
let temp_dir = tempfile::TempDir::new().expect("tempdir creation should succeed");
let ca_path = temp_dir.path().join("ca.pem");
std::fs::write(&ca_path, ca_cert.pem()).expect("write CA PEM should succeed");
TestCa {
ca_path,
_temp_dir: temp_dir,
}
}
fn gen_cert_key_files() -> TestCertKey {
use rcgen::{CertificateParams, DnType, KeyPair};
let key = KeyPair::generate().expect("key generation should succeed");
let mut params = CertificateParams::new(Vec::<String>::new()).expect("params should be valid");
params.distinguished_name.push(DnType::CommonName, "Test Cert");
let cert = params.self_signed(&key).expect("self-sign should succeed");
let temp_dir = tempfile::TempDir::new().expect("tempdir creation should succeed");
let cert_path = temp_dir.path().join("cert.pem");
let key_path = temp_dir.path().join("key.pem");
std::fs::write(&cert_path, cert.pem()).expect("write cert PEM should succeed");
std::fs::write(&key_path, key.serialize_pem()).expect("write key PEM should succeed");
TestCertKey {
cert_path,
key_path,
_temp_dir: temp_dir,
}
}
fn tls_upstream(address: &str, sni: Option<&str>) -> Upstream {
let tls = ClusterTls {
sni: sni.map(str::to_owned),
..ClusterTls::default()
};
Upstream {
address: Arc::from(address),
authority: None,
connection: Arc::new(ConnectionOptions::default()),
tls: Some(CachedClusterTls::try_from_config(&tls).unwrap()),
}
}
}