use std::{
collections::HashMap,
net::SocketAddr,
sync::{Arc, Mutex, OnceLock},
time::Instant,
};
use pingora_core::upstreams::peer::HttpPeer;
use super::ConnectionOptions;
const DNS_TTL_SECS: u64 = 60;
const MAX_DNS_ENTRIES: usize = 1_024;
#[derive(Debug, thiserror::Error)]
pub enum AddressResolutionError {
#[error("DNS resolution task failed for '{address}': {message}")]
Task {
address: String,
message: String,
},
#[error("upstream address resolution failed for '{address}': {source}")]
Resolve {
address: String,
#[source]
source: std::io::Error,
},
#[error("upstream address '{0}' resolved to zero addresses")]
Empty(String),
}
struct DnsCacheEntry {
addrs: Vec<SocketAddr>,
resolved_at: Instant,
}
fn dns_cache() -> &'static Mutex<HashMap<String, DnsCacheEntry>> {
static CACHE: OnceLock<Mutex<HashMap<String, DnsCacheEntry>>> = OnceLock::new();
CACHE.get_or_init(|| Mutex::new(HashMap::new()))
}
#[expect(
clippy::too_many_lines,
reason = "cache lookup, resolution, and insertion form one operation"
)]
pub async fn resolve_address(address: &str) -> Result<SocketAddr, AddressResolutionError> {
if let Ok(addr) = address.parse::<SocketAddr>() {
return Ok(addr);
}
if let Some(addr) = lookup_cached(address) {
return Ok(addr);
}
let owned = address.to_owned();
let task_address = owned.clone();
let addrs = tokio::task::spawn_blocking(move || {
use std::net::ToSocketAddrs as _;
task_address.to_socket_addrs().map(Iterator::collect::<Vec<_>>)
})
.await
.map_err(|error| AddressResolutionError::Task {
address: owned.clone(),
message: error.to_string(),
})?
.map_err(|source| AddressResolutionError::Resolve {
address: owned.clone(),
source,
})?;
let preferred = select_preferred_address(&addrs, address)?;
let mut cache = dns_cache().lock().unwrap_or_else(std::sync::PoisonError::into_inner);
if cache.len() >= MAX_DNS_ENTRIES && !cache.contains_key(address) {
cache.retain(|_, entry| entry.resolved_at.elapsed().as_secs() < DNS_TTL_SECS);
if cache.len() >= MAX_DNS_ENTRIES
&& let Some(oldest) = cache
.iter()
.min_by_key(|(_, entry)| entry.resolved_at)
.map(|(key, _)| key.clone())
{
cache.remove(&oldest);
}
}
cache.insert(
owned,
DnsCacheEntry {
addrs,
resolved_at: Instant::now(),
},
);
drop(cache);
Ok(preferred)
}
fn lookup_cached(address: &str) -> Option<SocketAddr> {
let cache = dns_cache().lock().unwrap_or_else(std::sync::PoisonError::into_inner);
cache.get(address).and_then(|entry| {
(entry.resolved_at.elapsed().as_secs() < DNS_TTL_SECS)
.then(|| {
entry
.addrs
.iter()
.find(|addr| addr.is_ipv4())
.or_else(|| entry.addrs.first())
.copied()
})
.flatten()
})
}
fn select_preferred_address(addrs: &[SocketAddr], address: &str) -> Result<SocketAddr, AddressResolutionError> {
addrs
.iter()
.find(|addr| addr.is_ipv4())
.or_else(|| addrs.first())
.copied()
.ok_or_else(|| AddressResolutionError::Empty(address.to_owned()))
}
#[inline]
pub fn apply_connection_options(peer: &mut HttpPeer, opts: &ConnectionOptions) {
peer.options.connection_timeout = opts.connection_timeout;
peer.options.total_connection_timeout = opts.total_connection_timeout;
peer.options.idle_timeout = opts.idle_timeout;
peer.options.read_timeout = opts.read_timeout;
peer.options.write_timeout = opts.write_timeout;
}
pub fn apply_cached_tls(peer: &mut HttpPeer, tls: &praxis_tls::CachedClusterTls, address: &str) {
if !tls.verify() {
tracing::debug!(upstream = %address, "upstream TLS verification disabled for this peer");
peer.options.verify_cert = false;
peer.options.verify_hostname = false;
}
if let Some(ca) = tls.ca() {
peer.options.ca = Some(Arc::from(ca_from_cached(ca)));
}
if let Some(client) = tls.client_cert() {
peer.client_cert_key = Some(Arc::new(client_cert_from_cached(client)));
}
}
pub fn ca_from_cached(cached: &praxis_tls::CachedCaCerts) -> Vec<pingora_core::utils::tls::WrappedX509> {
cached
.der_certs()
.iter()
.filter_map(|der| {
pingora_core::utils::tls::WrappedX509::parse(der.clone())
.inspect_err(|e| tracing::warn!("failed to parse cached CA cert: {e}"))
.ok()
})
.collect()
}
pub fn client_cert_from_cached(cached: &praxis_tls::CachedClientCert) -> pingora_core::utils::tls::CertKey {
pingora_core::utils::tls::CertKey::new(cached.cert_der().to_vec(), cached.key_der().to_vec())
}
pub fn derive_sni(address: &str) -> String {
let host = address.rsplit_once(':').map_or(address, |(h, _)| h);
let host_bare = host.strip_prefix('[').and_then(|h| h.strip_suffix(']')).unwrap_or(host);
if host_bare.parse::<std::net::IpAddr>().is_ok() {
tracing::warn!(
address,
"upstream is an IP without explicit SNI; TLS hostname verification is meaningless"
);
return String::new();
}
tracing::debug!(address, sni = host, "derived SNI from upstream address");
host.to_owned()
}
#[cfg(test)]
#[expect(clippy::allow_attributes, reason = "blanket test suppressions")]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::indexing_slicing, reason = "tests")]
mod tests {
use super::*;
#[tokio::test]
async fn resolve_address_parses_literal_without_dns() {
let address = resolve_address("127.0.0.1:8080").await.unwrap();
assert_eq!(address, "127.0.0.1:8080".parse().unwrap());
}
#[tokio::test]
async fn resolve_address_rejects_missing_port() {
resolve_address("127.0.0.1").await.unwrap_err();
}
#[test]
fn preferred_address_favors_ipv4() {
let ipv6 = "[::1]:8080".parse().unwrap();
let ipv4 = "127.0.0.1:8080".parse().unwrap();
assert_eq!(select_preferred_address(&[ipv6, ipv4], "example:8080").unwrap(), ipv4);
}
#[test]
fn derive_sni_extracts_hostname() {
assert_eq!(
derive_sni("backend.example.com:8443"),
"backend.example.com",
"should extract hostname from host:port"
);
}
#[test]
fn derive_sni_returns_empty_for_ip() {
assert_eq!(derive_sni("127.0.0.1:8443"), "", "should return empty for IP address");
}
#[test]
fn derive_sni_returns_empty_for_ipv6() {
assert_eq!(derive_sni("[::1]:8443"), "", "should return empty for IPv6 address");
}
#[test]
fn apply_connection_options_sets_timeouts() {
use std::time::Duration;
let opts = ConnectionOptions {
connection_timeout: Some(Duration::from_secs(1)),
read_timeout: Some(Duration::from_secs(2)),
write_timeout: Some(Duration::from_secs(3)),
idle_timeout: Some(Duration::from_secs(4)),
total_connection_timeout: Some(Duration::from_secs(5)),
};
let mut peer = HttpPeer::new("127.0.0.1:80", false, String::new());
apply_connection_options(&mut peer, &opts);
assert_eq!(peer.options.connection_timeout, Some(Duration::from_secs(1)));
assert_eq!(peer.options.read_timeout, Some(Duration::from_secs(2)));
assert_eq!(peer.options.write_timeout, Some(Duration::from_secs(3)));
assert_eq!(peer.options.idle_timeout, Some(Duration::from_secs(4)));
assert_eq!(peer.options.total_connection_timeout, Some(Duration::from_secs(5)));
}
}