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::{
pem_wallet_path, read_ewallet_pem, sso_wallet_path, WalletError,
};
if pem_wallet_path(dir).exists() {
return read_ewallet_pem(dir, password).map_err(Error::from);
}
let sso = sso_wallet_path(dir);
if sso.exists() {
let bytes = std::fs::read(&sso).map_err(|source| WalletError::Io {
path: sso.display().to_string(),
source,
})?;
return oracledb_protocol::tls::sso::parse_cwallet_sso(&bytes).map_err(Error::from);
}
if dir.join("ewallet.p12").exists() {
return Err(WalletError::UnsupportedFormat {
format: "ewallet.p12",
}
.into());
}
Err(WalletError::FileMissing("ewallet.pem 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 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 = if params.use_sni {
let sni = build_sni(&descriptor.service_name, server_type);
if sni_is_rustls_valid(&sni) {
config.enable_sni = true;
sni
} else {
config.enable_sni = false;
SNI_PLACEHOLDER.to_string()
}
} else {
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 ewallet_p12_only_wallet_is_typed_unsupported_format() {
let dir = std::path::PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("tests")
.join("fixtures")
.join("tls")
.join("p12_only");
let err = load_wallet(&dir, None).expect_err("p12-only wallet must be unsupported");
let wallet_err = if let Error::Wallet(wallet_err) = err {
wallet_err
} else {
assert!(matches!(err, Error::Wallet(_)), "expected wallet error");
return;
};
let format =
if let oracledb_protocol::tls::wallet::WalletError::UnsupportedFormat { format } =
&wallet_err
{
format
} else {
assert!(
matches!(
&wallet_err,
oracledb_protocol::tls::wallet::WalletError::UnsupportedFormat { .. }
),
"expected UnsupportedFormat, got {wallet_err:?}"
);
return;
};
assert_eq!(*format, "ewallet.p12");
let sensitive_path = dir.display().to_string();
assert!(!format!("{wallet_err}").contains(&sensitive_path));
assert!(!format!("{wallet_err:?}").contains(&sensitive_path));
}
}