diff --git a/codex-rs/network-proxy/src/certs.rs b/codex-rs/network-proxy/src/certs.rs index d3ee637a72..3ef9cb8d96 100644 --- a/codex-rs/network-proxy/src/certs.rs +++ b/codex-rs/network-proxy/src/certs.rs @@ -23,6 +23,7 @@ use rama_tls_rustls::server::TlsAcceptorData; use sha2::Digest as _; use sha2::Sha256; use std::collections::HashMap; +use std::collections::HashSet; use std::ffi::OsStr; #[cfg(windows)] use std::ffi::OsString; @@ -176,7 +177,7 @@ fn managed_ca_trust_bundle_for_cert_path( .map(|value| (key, value.clone())) }) .collect(); - let trust_bundle = build_managed_ca_trust_bundle(cert_path)?; + let trust_bundle = build_managed_ca_trust_bundle(cert_path, &startup_env_values)?; let path = persist_managed_ca_trust_bundle(cert_path, &trust_bundle)?; Ok(ManagedMitmCaTrustBundle { @@ -186,7 +187,10 @@ fn managed_ca_trust_bundle_for_cert_path( }) } -fn build_managed_ca_trust_bundle(managed_ca_cert_path: &Path) -> Result { +fn build_managed_ca_trust_bundle( + managed_ca_cert_path: &Path, + startup_env_values: &HashMap<&'static str, String>, +) -> Result { let mut trust_bundle = String::new(); let rustls_native_certs::CertificateResult { certs, errors, .. } = crate::native_certs::load_platform_native_certs(); @@ -199,6 +203,16 @@ fn build_managed_ca_trust_bundle(managed_ca_cert_path: &Path) -> Result for cert in certs { push_certificate_pem(&mut trust_bundle, cert.as_ref()); } + let mut appended_startup_paths = HashSet::new(); + for path in CUSTOM_CA_ENV_KEYS + .into_iter() + .filter_map(|key| startup_env_values.get(key)) + .map(Path::new) + { + if path != managed_ca_cert_path && appended_startup_paths.insert(path) { + append_pem_file(&mut trust_bundle, path)?; + } + } append_pem_file(&mut trust_bundle, managed_ca_cert_path)?; Ok(trust_bundle) } @@ -866,24 +880,30 @@ mod tests { fn managed_ca_trust_bundle_records_startup_ca_env_values() { let dir = tempdir().unwrap(); let managed_ca_cert_path = dir.path().join("ca.pem"); + let startup_ca_bundle_path = dir.path().join("startup-ca.pem"); + let startup_cert_dir = dir.path().join("startup-certs"); fs::write(&managed_ca_cert_path, "managed ca\n").unwrap(); + fs::write(&startup_ca_bundle_path, "startup ca\n").unwrap(); + fs::create_dir(&startup_cert_dir).unwrap(); + let startup_ca_bundle_path = startup_ca_bundle_path.display().to_string(); + let startup_cert_dir = startup_cert_dir.display().to_string(); let env = HashMap::from([ - ("SSL_CERT_FILE", "/tmp/startup-ca.pem".to_string()), - (SSL_CERT_DIR_ENV_KEY, "/tmp/startup-certs".to_string()), + ("SSL_CERT_FILE", startup_ca_bundle_path.clone()), + (SSL_CERT_DIR_ENV_KEY, startup_cert_dir.clone()), ]); let trust_bundle = managed_ca_trust_bundle_for_cert_path(&managed_ca_cert_path, &env).unwrap(); assert_eq!( trust_bundle.startup_env_values, HashMap::from([ - ("SSL_CERT_FILE", "/tmp/startup-ca.pem".to_string()), - (SSL_CERT_DIR_ENV_KEY, "/tmp/startup-certs".to_string()), + ("SSL_CERT_FILE", startup_ca_bundle_path), + (SSL_CERT_DIR_ENV_KEY, startup_cert_dir), ]) ); } #[test] - fn managed_ca_trust_bundle_does_not_append_startup_ca_override_to_baseline() { + fn managed_ca_trust_bundle_appends_startup_ca_override_to_baseline() { let dir = tempdir().unwrap(); let managed_ca_cert_path = dir.path().join("ca.pem"); let startup_ca_bundle_path = dir.path().join("startup-ca.pem"); @@ -898,7 +918,7 @@ mod tests { managed_ca_trust_bundle_for_cert_path(&managed_ca_cert_path, &env).unwrap(); let baseline_bundle = fs::read_to_string(trust_bundle.path).unwrap(); - assert!(!baseline_bundle.contains("startup ca")); + assert!(baseline_bundle.contains("startup ca")); assert!(baseline_bundle.contains("managed ca")); }