use std::sync::Arc;
use asupersync::net::TcpStream;
use asupersync::tls::{TlsConnector, TlsStream};
use oracledb_protocol::net::EasyConnect;
use oracledb_protocol::tls::dn::{check_cert_dn, check_server_name, DnMatchError};
use oracledb_protocol::tls::sni::build_sni;
use oracledb_protocol::tls::wallet::{resolve_wallet_dir, WalletContents};
use rustls::client::danger::{HandshakeSignatureValid, ServerCertVerified, ServerCertVerifier};
use rustls::crypto::{verify_tls12_signature, verify_tls13_signature, WebPkiSupportedAlgorithms};
use rustls::pki_types::{CertificateDer, ServerName, UnixTime};
use rustls::{ClientConfig, DigitallySignedStruct, Error as RustlsError, SignatureScheme};
use crate::Error;
#[derive(Clone, Debug, Default)]
pub struct TlsParams {
pub wallet: Option<WalletContents>,
pub dn_match: bool,
pub server_cert_dn: Option<String>,
pub expected_host: String,
pub use_sni: bool,
}
#[derive(Debug)]
pub(crate) struct OracleServerCertVerifier {
trust_anchor_ders: Vec<Vec<u8>>,
supported_algs: WebPkiSupportedAlgorithms,
dn_match: bool,
server_cert_dn: Option<String>,
expected_host: String,
}
impl OracleServerCertVerifier {
fn run_dn_match(&self, end_entity: &CertificateDer<'_>) -> Result<(), RustlsError> {
if !self.dn_match {
return Ok(());
}
let (subject_dn, san_dns, common_names) = parse_cert_identity(end_entity)?;
let result = if let Some(expected_dn) = self.server_cert_dn.as_deref() {
check_cert_dn(expected_dn, &subject_dn)
} else {
check_server_name(&self.expected_host, &san_dns, &common_names)
};
result.map_err(dn_error_to_rustls)
}
}
impl ServerCertVerifier for OracleServerCertVerifier {
fn verify_server_cert(
&self,
end_entity: &CertificateDer<'_>,
intermediates: &[CertificateDer<'_>],
_server_name: &ServerName<'_>,
_ocsp_response: &[u8],
now: UnixTime,
) -> Result<ServerCertVerified, RustlsError> {
let owned_certs: Vec<CertificateDer<'static>> = self
.trust_anchor_ders
.iter()
.map(|der| CertificateDer::from(der.clone()))
.collect();
let anchors: Vec<rustls_pki_types::TrustAnchor<'_>> = owned_certs
.iter()
.filter_map(|c| webpki::anchor_from_trusted_cert(c).ok())
.collect();
if anchors.is_empty() {
return Err(RustlsError::General(
"wallet contained no usable CA trust anchors".to_string(),
));
}
let ee = webpki::EndEntityCert::try_from(end_entity)
.map_err(|e| RustlsError::General(format!("invalid server certificate: {e}")))?;
ee.verify_for_usage(
self.supported_algs.all,
&anchors,
intermediates,
now,
webpki::KeyUsage::server_auth(),
None,
None,
)
.map_err(|e| {
RustlsError::General(format!("TCPS server certificate chain is not trusted: {e}"))
})?;
self.run_dn_match(end_entity)?;
Ok(ServerCertVerified::assertion())
}
fn verify_tls12_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &DigitallySignedStruct,
) -> Result<HandshakeSignatureValid, RustlsError> {
verify_tls12_signature(message, cert, dss, &self.supported_algs)
}
fn verify_tls13_signature(
&self,
message: &[u8],
cert: &CertificateDer<'_>,
dss: &DigitallySignedStruct,
) -> Result<HandshakeSignatureValid, RustlsError> {
verify_tls13_signature(message, cert, dss, &self.supported_algs)
}
fn supported_verify_schemes(&self) -> Vec<SignatureScheme> {
self.supported_algs.supported_schemes()
}
}
fn dn_error_to_rustls(err: DnMatchError) -> RustlsError {
RustlsError::General(err.to_string())
}
fn parse_cert_identity(
cert: &CertificateDer<'_>,
) -> Result<(String, Vec<String>, Vec<String>), RustlsError> {
use x509_cert::der::Decode;
let parsed = x509_cert::Certificate::from_der(cert.as_ref())
.map_err(|e| RustlsError::General(format!("server certificate parse error: {e}")))?;
let subject_dn = parsed.tbs_certificate.subject.to_string();
let mut common_names = Vec::new();
for rdn in parsed.tbs_certificate.subject.0.iter() {
for atv in rdn.0.iter() {
if atv.oid.to_string() == "2.5.4.3" {
if let Ok(s) = std::str::from_utf8(atv.value.value()) {
common_names.push(s.to_string());
} else if let Ok(s) = atv.value.decode_as::<x509_cert::der::asn1::Utf8StringRef>() {
common_names.push(s.as_str().to_string());
}
}
}
}
let mut san_dns = Vec::new();
if let Some(extensions) = parsed.tbs_certificate.extensions.as_ref() {
for ext in extensions.iter() {
if ext.extn_id.to_string() == "2.5.29.17" {
if let Ok(san) =
x509_cert::ext::pkix::SubjectAltName::from_der(ext.extn_value.as_bytes())
{
for name in san.0.iter() {
if let x509_cert::ext::pkix::name::GeneralName::DnsName(dns) = name {
san_dns.push(dns.as_str().to_string());
}
}
}
}
}
}
Ok((subject_dn, san_dns, common_names))
}
pub(crate) fn build_client_config(params: &TlsParams) -> Result<ClientConfig, Error> {
use rustls::crypto::ring::default_provider;
let provider = Arc::new(default_provider());
let supported_algs = provider.signature_verification_algorithms;
let trust_anchor_ders: Vec<Vec<u8>> = match ¶ms.wallet {
Some(w) if !w.ca_certificates.is_empty() => w.ca_certificates.clone(),
_ => load_system_roots(),
};
if trust_anchor_ders.is_empty() {
return Err(Error::Tls(
"no trust anchors available for TCPS: supply a wallet (ewallet.pem) \
or install system root certificates"
.to_string(),
));
}
let verifier = Arc::new(OracleServerCertVerifier {
trust_anchor_ders,
supported_algs,
dn_match: params.dn_match,
server_cert_dn: params.server_cert_dn.clone(),
expected_host: params.expected_host.clone(),
});
let builder = ClientConfig::builder_with_provider(provider)
.with_safe_default_protocol_versions()
.map_err(|e| Error::Tls(format!("TLS protocol setup failed: {e}")))?
.dangerous()
.with_custom_certificate_verifier(verifier);
let mut config = if let Some(w) = ¶ms.wallet {
if w.has_client_identity() {
let chain: Vec<CertificateDer<'static>> = w
.client_cert_chain
.iter()
.map(|der| CertificateDer::from(der.clone()))
.collect();
let key = client_private_key(w)?;
builder
.with_client_auth_cert(chain, key)
.map_err(|e| Error::Tls(format!("client certificate setup failed: {e}")))?
} else {
builder.with_no_client_auth()
}
} else {
builder.with_no_client_auth()
};
config.enable_sni = true;
Ok(config)
}
fn client_private_key(
w: &WalletContents,
) -> Result<rustls::pki_types::PrivateKeyDer<'static>, Error> {
let der = w
.client_private_key
.as_ref()
.ok_or_else(|| Error::Tls("wallet has a client cert but no private key".to_string()))?;
rustls::pki_types::PrivateKeyDer::try_from(der.clone())
.map_err(|e| Error::Tls(format!("client private key parse failed: {e}")))
}
fn load_system_roots() -> Vec<Vec<u8>> {
const BUNDLES: &[&str] = &[
"/etc/ssl/certs/ca-certificates.crt",
"/etc/pki/tls/certs/ca-bundle.crt",
"/etc/ssl/ca-bundle.pem",
"/etc/ssl/cert.pem",
];
for path in BUNDLES {
if let Ok(bytes) = std::fs::read(path) {
let mut reader = std::io::BufReader::new(&bytes[..]);
let certs: Vec<Vec<u8>> = rustls_pemfile_certs(&mut reader);
if !certs.is_empty() {
return certs;
}
}
}
Vec::new()
}
fn rustls_pemfile_certs(reader: &mut dyn std::io::BufRead) -> Vec<Vec<u8>> {
oracledb_protocol::tls::wallet::parse_pem_certificates(reader)
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn resolve_tls_params(
descriptor: &EasyConnect,
wallet_location: Option<&str>,
wallet_password: Option<&str>,
ssl_server_dn_match: bool,
ssl_server_cert_dn: Option<&str>,
use_sni: bool,
) -> Result<TlsParams, Error> {
let tns_admin = std::env::var("TNS_ADMIN").ok();
let wallet = match resolve_wallet_dir(wallet_location, tns_admin.as_deref()) {
Some(dir) => Some(load_wallet(&dir, wallet_password)?),
None => None,
};
Ok(TlsParams {
wallet,
dn_match: ssl_server_dn_match,
server_cert_dn: ssl_server_cert_dn.map(str::to_string),
expected_host: descriptor.host.clone(),
use_sni,
})
}
fn load_wallet(dir: &std::path::Path, password: Option<&str>) -> Result<WalletContents, Error> {
use oracledb_protocol::tls::wallet::{
p12_wallet_path, pem_wallet_path, read_ewallet_p12, read_ewallet_pem, sso_wallet_path,
WalletError,
};
let read_sso = || -> Result<Option<WalletContents>, WalletError> {
let sso = sso_wallet_path(dir);
if !sso.exists() {
return Ok(None);
}
let bytes = std::fs::read(&sso).map_err(|source| WalletError::Io {
path: sso.display().to_string(),
source,
})?;
oracledb_protocol::tls::sso::parse_cwallet_sso(&bytes).map(Some)
};
let falls_through_to_autologin = |e: &WalletError| {
matches!(
e,
WalletError::KeyDecrypt(_)
| WalletError::Pkcs12(_)
| WalletError::PasswordRequired { .. }
| WalletError::UnsupportedFormat { .. }
)
};
let have_p12 = p12_wallet_path(dir).exists();
let primary: Option<(&'static str, Result<WalletContents, WalletError>)> =
if pem_wallet_path(dir).exists() {
Some(("ewallet.pem", read_ewallet_pem(dir, password)))
} else if have_p12 && password.is_some() {
Some(("ewallet.p12", read_ewallet_p12(dir, password)))
} else {
None
};
match primary {
Some((_, Ok(contents))) => Ok(contents),
Some((name, Err(primary_err))) => {
if falls_through_to_autologin(&primary_err) {
if let Ok(Some(sso)) = read_sso() {
obs_warn!(
skipped_wallet = name,
"wallet {name} could not be used ({primary_err}); \
falling back to auto-login cwallet.sso"
);
let _ = name;
return Ok(sso);
}
}
Err(primary_err.into())
}
None => {
if let Some(sso) = read_sso()? {
return Ok(sso);
}
if have_p12 {
return read_ewallet_p12(dir, password).map_err(Error::from);
}
Err(
WalletError::FileMissing("ewallet.pem, ewallet.p12, or cwallet.sso".to_string())
.into(),
)
}
}
}
const SNI_PLACEHOLDER: &str = "oracle.invalid";
fn sni_is_rustls_valid(sni: &str) -> bool {
rustls::pki_types::ServerName::try_from(sni.to_string()).is_ok()
}
pub(crate) fn decide_sni(
use_sni: bool,
service_name: &str,
server_type: Option<&str>,
) -> Result<Option<String>, Error> {
if !use_sni {
return Ok(None);
}
let sni = build_sni(service_name, server_type);
if sni_is_rustls_valid(&sni) {
Ok(Some(sni))
} else {
Err(Error::UnsupportedSni(sni))
}
}
pub async fn tls_handshake(
descriptor: &EasyConnect,
server_type: Option<&str>,
params: &TlsParams,
tcp: TcpStream,
) -> Result<TlsStream<TcpStream>, Error> {
let mut config = build_client_config(params)?;
let server_name = match decide_sni(params.use_sni, &descriptor.service_name, server_type)? {
Some(sni) => {
config.enable_sni = true;
sni
}
None => {
config.enable_sni = false;
SNI_PLACEHOLDER.to_string()
}
};
let connector =
TlsConnector::new(config).with_handshake_timeout(std::time::Duration::from_secs(20));
connector
.connect(&server_name, tcp)
.await
.map_err(|e| Error::Tls(format!("TCPS handshake failed: {e}")))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn build_config_requires_trust_anchors() {
let params = TlsParams {
wallet: Some(WalletContents::default()),
dn_match: true,
server_cert_dn: None,
expected_host: "db.example.com".to_string(),
use_sni: false,
};
let _ = build_client_config(¶ms);
}
#[test]
fn oracle_sni_is_rejected_by_rustls_servername() {
assert!(!sni_is_rustls_valid("S8.FREEPDB1.V3.319"));
assert!(sni_is_rustls_valid("db.example.com"));
}
#[test]
fn decide_sni_without_use_sni_sends_no_sni() {
let decided = decide_sni(false, "FREEPDB1", None).expect("no-SNI must be Ok");
assert!(decided.is_none(), "use_sni=false must not send an SNI");
}
#[test]
fn decide_sni_with_use_sni_fails_closed_not_silent() {
let err = decide_sni(true, "FREEPDB1", None)
.expect_err("use_sni=true with an un-encodable Oracle SNI must fail closed");
match err {
Error::UnsupportedSni(sni) => {
assert!(
sni.starts_with('S') && sni.contains("FREEPDB1"),
"error must name the Oracle SNI string, got {sni:?}"
);
}
other => panic!("expected Error::UnsupportedSni, got {other:?}"),
}
}
fn fixture_tls_dir() -> std::path::PathBuf {
std::path::PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("tests")
.join("fixtures")
.join("tls")
}
#[test]
fn ewallet_p12_only_wallet_without_password_is_password_required() {
let dir = fixture_tls_dir().join("p12_only");
let err = load_wallet(&dir, None).expect_err("p12-only wallet without password");
let wallet_err = if let Error::Wallet(wallet_err) = err {
wallet_err
} else {
panic!("expected wallet error, got {err:?}");
};
assert!(
matches!(
&wallet_err,
oracledb_protocol::tls::wallet::WalletError::PasswordRequired { format }
if *format == "ewallet.p12"
),
"expected PasswordRequired, got {wallet_err:?}"
);
let sensitive_path = dir.display().to_string();
assert!(!format!("{wallet_err}").contains(&sensitive_path));
assert!(!format!("{wallet_err:?}").contains(&sensitive_path));
}
#[test]
fn ewallet_p12_only_wallet_with_password_garbage_is_typed_pkcs12_error() {
let dir = fixture_tls_dir().join("p12_only");
let err = load_wallet(&dir, Some("any-password")).expect_err("dummy p12 must not parse");
let wallet_err = if let Error::Wallet(wallet_err) = err {
wallet_err
} else {
panic!("expected wallet error, got {err:?}");
};
assert!(
matches!(
&wallet_err,
oracledb_protocol::tls::wallet::WalletError::Pkcs12(_)
),
"expected Pkcs12, got {wallet_err:?}"
);
assert!(!format!("{wallet_err}").contains("any-password"));
assert!(!format!("{wallet_err:?}").contains("any-password"));
}
fn temp_wallet_dir(label: &str, files: &[&str]) -> std::path::PathBuf {
let dir = std::env::temp_dir().join(format!(
"oracledb-wallet-test-{label}-{}",
std::process::id()
));
std::fs::create_dir_all(&dir).expect("create temp wallet dir");
for name in files {
std::fs::copy(
fixture_tls_dir().join(name),
dir.join(wallet_file_name(name)),
)
.expect("copy fixture");
}
dir
}
fn wallet_file_name(fixture: &str) -> &'static str {
match fixture {
"ewallet_orapki.p12" => "ewallet.p12",
"cwallet_orapki.sso" => "cwallet.sso",
"ewallet.pem" => "ewallet.pem",
other => panic!("unmapped fixture {other}"),
}
}
#[test]
fn adb_style_wallet_dir_prefers_p12_with_password_and_sso_without() {
let dir = temp_wallet_dir("adb", &["ewallet_orapki.p12", "cwallet_orapki.sso"]);
let with_pw =
load_wallet(&dir, Some("WalletPass123")).expect("p12 path must load with password");
assert!(with_pw.has_client_identity());
let without_pw = load_wallet(&dir, None).expect("sso path must load without password");
assert!(without_pw.has_client_identity());
assert_eq!(with_pw.ca_certificates, without_pw.ca_certificates);
}
#[test]
fn wallet_dir_prefers_pem_over_p12_and_sso() {
let dir = temp_wallet_dir(
"pem-first",
&["ewallet.pem", "ewallet_orapki.p12", "cwallet_orapki.sso"],
);
let wallet = load_wallet(&dir, None).expect("pem path must load");
assert!(wallet.has_client_identity());
use oracledb_protocol::tls::wallet::parse_ewallet_pem;
let pem_bytes =
std::fs::read(fixture_tls_dir().join("ewallet.pem")).expect("read pem fixture");
let direct = parse_ewallet_pem(&pem_bytes, None).expect("parse pem fixture");
assert_eq!(wallet.ca_certificates, direct.ca_certificates);
}
#[test]
fn unusable_p12_falls_through_to_auto_login_sso() {
let dir = temp_wallet_dir("fallthrough", &["ewallet_orapki.p12", "cwallet_orapki.sso"]);
let fell_through = load_wallet(&dir, Some("not-the-password!"))
.expect("wrong p12 password must fall through to the auto-login cwallet.sso");
assert!(fell_through.has_client_identity());
let sso_only = temp_wallet_dir("fallthrough-sso", &["cwallet_orapki.sso"]);
let direct = load_wallet(&sso_only, None).expect("sso path must load");
assert_eq!(fell_through.ca_certificates, direct.ca_certificates);
assert_eq!(fell_through.client_private_key, direct.client_private_key);
}
#[test]
fn unusable_p12_without_sso_preserves_original_typed_error() {
let dir = temp_wallet_dir("no-sso", &["ewallet_orapki.p12"]);
let err = load_wallet(&dir, Some("not-the-password!"))
.expect_err("wrong p12 password with no sso must fail closed");
let wallet_err = if let Error::Wallet(wallet_err) = err {
wallet_err
} else {
panic!("expected wallet error, got {err:?}");
};
assert!(
matches!(
&wallet_err,
oracledb_protocol::tls::wallet::WalletError::Pkcs12(_)
),
"expected the original typed Pkcs12 error, got {wallet_err:?}"
);
for rendered in [format!("{wallet_err}"), format!("{wallet_err:?}")] {
assert!(!rendered.contains("not-the-password!"), "password leaked");
let lower = rendered.to_ascii_lowercase();
assert!(
!lower.contains("fall") && !lower.contains("auto-login") && !lower.contains("sso"),
"preserved error must not mention the fallthrough, got {rendered:?}"
);
}
}
}