use reqwest::Certificate;
use std::env;
use std::fs;
use std::path::Path;
use std::sync::OnceLock;
use std::time::Duration;
use tracing::{debug, trace, warn};
pub const DEFAULT_CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
pub const DEFAULT_REQUEST_TIMEOUT: Duration = Duration::from_secs(30);
static CA_BUNDLE_CACHE: OnceLock<Option<Vec<Certificate>>> = OnceLock::new();
pub fn load_ca_certificates() -> Option<&'static Vec<Certificate>> {
CA_BUNDLE_CACHE
.get_or_init(|| {
let ca_path = env::var("AWS_CA_BUNDLE")
.or_else(|_| env::var("SSL_CERT_FILE"))
.ok();
let path = match ca_path {
Some(p) => p,
None => {
trace!("No custom CA bundle configured (AWS_CA_BUNDLE/SSL_CERT_FILE not set)");
return None;
}
};
debug!("Loading custom CA bundle from: {}", path);
load_certificates_from_file(&path)
})
.as_ref()
}
fn load_certificates_from_file(path: &str) -> Option<Vec<Certificate>> {
let path = Path::new(path);
if !path.exists() {
warn!(
"CA bundle file does not exist: {}. Using default certificate roots.",
path.display()
);
return None;
}
let pem_data = match fs::read(path) {
Ok(data) => data,
Err(e) => {
warn!(
"Failed to read CA bundle file '{}': {}. Using default certificate roots.",
path.display(),
e
);
return None;
}
};
let certs = parse_pem_certificates(&pem_data);
if certs.is_empty() {
warn!(
"No valid certificates found in CA bundle file '{}'. Using default certificate roots.",
path.display()
);
return None;
}
debug!(
"Loaded {} certificate(s) from CA bundle: {}",
certs.len(),
path.display()
);
Some(certs)
}
fn parse_pem_certificates(pem_data: &[u8]) -> Vec<Certificate> {
let pem_str = match std::str::from_utf8(pem_data) {
Ok(s) => s,
Err(e) => {
warn!("CA bundle is not valid UTF-8: {}", e);
return vec![];
}
};
let cert_marker_begin = "-----BEGIN CERTIFICATE-----";
let cert_marker_end = "-----END CERTIFICATE-----";
let mut all_certs = Vec::new();
let mut pos = 0;
while let Some(start) = pem_str[pos..].find(cert_marker_begin) {
let abs_start = pos + start;
if let Some(end) = pem_str[abs_start..].find(cert_marker_end) {
let abs_end = abs_start + end + cert_marker_end.len();
let cert_pem = &pem_str[abs_start..abs_end];
if let Ok(cert) = Certificate::from_pem(cert_pem.as_bytes()) {
all_certs.push(cert);
}
pos = abs_end;
} else {
warn!("Malformed PEM: found BEGIN but no END marker");
break;
}
}
if all_certs.is_empty() {
return vec![];
}
if validate_certificates(&all_certs) {
debug!(
"All {} certificate(s) from CA bundle are valid",
all_certs.len()
);
return all_certs;
}
debug!(
"Some certificates have unsupported features, filtering {} certificates...",
all_certs.len()
);
let valid_certs = filter_valid_certificates(all_certs);
if valid_certs.is_empty() {
warn!("No valid certificates found in CA bundle after filtering");
} else {
debug!(
"Filtered to {} valid certificate(s) (rustls compatible)",
valid_certs.len()
);
}
valid_certs
}
fn validate_certificates(certs: &[Certificate]) -> bool {
let mut builder = reqwest::blocking::Client::builder();
for cert in certs {
builder = builder.add_root_certificate(cert.clone());
}
builder.build().is_ok()
}
fn filter_valid_certificates(certs: Vec<Certificate>) -> Vec<Certificate> {
if certs.is_empty() {
return vec![];
}
if certs.len() == 1 {
if validate_certificates(&certs) {
return certs;
} else {
return vec![];
}
}
if validate_certificates(&certs) {
return certs;
}
let mid = certs.len() / 2;
let (left, right) = certs.split_at(mid);
let mut valid = filter_valid_certificates(left.to_vec());
valid.extend(filter_valid_certificates(right.to_vec()));
valid
}
pub fn configure_tls_blocking(
mut builder: reqwest::blocking::ClientBuilder,
) -> reqwest::blocking::ClientBuilder {
builder = builder
.connect_timeout(DEFAULT_CONNECT_TIMEOUT)
.timeout(DEFAULT_REQUEST_TIMEOUT);
if let Some(certs) = load_ca_certificates() {
builder = builder.tls_built_in_root_certs(true);
for cert in certs {
builder = builder.add_root_certificate(cert.clone());
}
}
builder
}
#[allow(dead_code)]
pub fn create_blocking_client() -> Result<reqwest::blocking::Client, reqwest::Error> {
configure_tls_blocking(reqwest::blocking::Client::builder()).build()
}
pub fn create_blocking_client_with_timeout(
timeout: Duration,
) -> Result<reqwest::blocking::Client, reqwest::Error> {
configure_tls_blocking(reqwest::blocking::Client::builder())
.timeout(timeout)
.build()
}
pub fn create_async_client() -> Result<reqwest::Client, reqwest::Error> {
configure_tls_async(reqwest::Client::builder()).build()
}
pub fn configure_tls_async(mut builder: reqwest::ClientBuilder) -> reqwest::ClientBuilder {
builder = builder
.connect_timeout(DEFAULT_CONNECT_TIMEOUT)
.timeout(DEFAULT_REQUEST_TIMEOUT);
if let Some(certs) = load_ca_certificates() {
builder = builder.tls_built_in_root_certs(true);
for cert in certs {
builder = builder.add_root_certificate(cert.clone());
}
}
builder
}
#[cfg(test)]
mod tests {
use super::*;
const DIGICERT_ROOT_CA: &str = r#"-----BEGIN CERTIFICATE-----
MIIDrzCCApegAwIBAgIQCDvgVpBCRrGhdWrJWZHHSjANBgkqhkiG9w0BAQUFADBh
MQswCQYDVQQGEwJVUzEVMBMGA1UEChMMRGlnaUNlcnQgSW5jMRkwFwYDVQQLExB3
d3cuZGlnaWNlcnQuY29tMSAwHgYDVQQDExdEaWdpQ2VydCBHbG9iYWwgUm9vdCBD
QTAeFw0wNjExMTAwMDAwMDBaFw0zMTExMTAwMDAwMDBaMGExCzAJBgNVBAYTAlVT
MRUwEwYDVQQKEwxEaWdpQ2VydCBJbmMxGTAXBgNVBAsTEHd3dy5kaWdpY2VydC5j
b20xIDAeBgNVBAMTF0RpZ2lDZXJ0IEdsb2JhbCBSb290IENBMIIBIjANBgkqhkiG
9w0BAQEFAAOCAQ8AMIIBCgKCAQEA4jvhEXLeqKTTo1eqUKKPC3eQyaKl7hLOllsB
CSDMAZOnTjC3U/dDxGkAV53ijSLdhwZAAIEJzs4bg7/fzTtxRuLWZscFs3YnFo97
nh6Vfe63SKMI2tavegw5BmV/Sl0fvBf4q77uKNd0f3p4mVmFaG5cIzJLv07A6Fpt
43C/dxC//AH2hdmoRBBYMql1GNXRor5H4idq9Joz+EkIYIvUX7Q6hL+hqkpMfT7P
T19sdl6gSzeRntwi5m3OFBqOasv+zbMUZBfHWymeMr/y7vrTC0LUq7dBMtoM1O/4
gdW7jVg/tRvoSSiicNoxBN33shbyTApOB6jtSj1etX+jkMOvJwIDAQABo2MwYTAO
BgNVHQ8BAf8EBAMCAYYwDwYDVR0TAQH/BAUwAwEB/zAdBgNVHQ4EFgQUA95QNVbR
TLtm8KPiGxvDl7I90VUwHwYDVR0jBBgwFoAUA95QNVbRTLtm8KPiGxvDl7I90VUw
DQYJKoZIhvcNAQEFBQADggEBAMucN6pIExIK+t1EnE9SsPTfrgT1eXkIoyQY/Esr
hMAtudXH/vTBH1jLuG2cenTnmCmrEbXjcKChzUyImZOMkXDiqw8cvpOp/2PV5Adg
06O/nVsJ8dWO41P0jmP6P6fbtGbfYmbW0W5BjfIttep3Sp+dWOIrWcBAI+0tKIJF
PnlUkiaY4IBIqDfv8NZ5YBberOgOzW6sRBc4L0na4UU+Krk2U886UAb3LujEV0ls
YSEY1QSteDwsOoBrp+uvFRTp2InBuThs4pFsiv9kuXclVzDAGySj4dzp30d8tbQk
CAUw7C29C79Fv1C5qfPrmAESrciIxpg0X40KPMbp1ZWVbd4=
-----END CERTIFICATE-----"#;
#[test]
fn test_parse_valid_certificate() {
let certs = parse_pem_certificates(DIGICERT_ROOT_CA.as_bytes());
assert_eq!(certs.len(), 1, "Should parse valid certificate");
}
#[test]
fn test_parse_certificate_bundle() {
let pem = format!("{}\n{}", DIGICERT_ROOT_CA, DIGICERT_ROOT_CA);
let certs = parse_pem_certificates(pem.as_bytes());
assert_eq!(certs.len(), 2, "Should parse each certificate individually");
}
#[test]
fn test_load_ca_certificates_not_set() {
}
}