use std::mem;
use std::ptr;
use std::string::String;
use std::ffi::OsStr;
use std::os::raw::{c_void, c_char};
use std::os::windows::ffi::OsStrExt;
use crypt32::{CertOpenStore, CertCloseStore, CertAddEncodedCertificateToStore,
CertFreeCertificateContext, CertGetCertificateChain,
CertFreeCertificateChain, CertVerifyCertificateChainPolicy};
use winapi::minwindef::DWORD;
use winapi::wincrypt::{PCCERT_CHAIN_CONTEXT, CERT_STORE_PROV_MEMORY, HCERTSTORE,
CERT_STORE_DEFER_CLOSE_UNTIL_LAST_FREE_FLAG, PCCERT_CONTEXT,
X509_ASN_ENCODING, CERT_STORE_ADD_ALWAYS, CERT_CHAIN_PARA,
CERT_CHAIN_POLICY_PARA, CERT_CHAIN_POLICY_STATUS,
CERT_CHAIN_POLICY_SSL, szOID_PKIX_KP_SERVER_AUTH,
szOID_SERVER_GATED_CRYPTO, szOID_SGC_NETSCAPE};
use winapi::winnt::LPWSTR;
use ValidationResult;
pub fn validate_cert_chain(encoded_certs: &[&[u8]], hostname: &str) -> ValidationResult {
let context = match build_cert_context(encoded_certs) {
Ok(context) => context,
Err(e) => return e,
};
let chain = match build_chain(context) {
Ok(chain) => chain,
Err(e) => return e,
};
verify_chain_against_policy(chain, hostname)
}
macro_rules! as_cchar_vec {
($e:expr) => {
{
let mut bytes = String::from($e).into_bytes();
bytes.push(0);
bytes.iter().map(|&b| b as c_char).collect()
}
}
}
#[repr(C)]
#[allow(non_snake_case)]
struct SSL_EXTRA_CERT_CHAIN_POLICY_PARA {
cbSize: DWORD,
dwAuthType: DWORD,
fdwChecks: DWORD,
pwszServerName: LPWSTR
}
impl Copy for SSL_EXTRA_CERT_CHAIN_POLICY_PARA {}
impl Clone for SSL_EXTRA_CERT_CHAIN_POLICY_PARA { fn clone(&self) -> SSL_EXTRA_CERT_CHAIN_POLICY_PARA {*self} }
struct CertStore(HCERTSTORE);
impl Drop for CertStore {
fn drop(&mut self) {
unsafe {
CertCloseStore(self.0 as *mut _, 0);
}
}
}
struct CertContext(PCCERT_CONTEXT);
impl Drop for CertContext {
fn drop(&mut self) {
unsafe {
CertFreeCertificateContext(self.0 as *mut _);
}
}
}
struct CertChainContext(PCCERT_CHAIN_CONTEXT);
impl Drop for CertChainContext {
fn drop(&mut self) {
unsafe {
CertFreeCertificateChain(self.0 as *mut _);
}
}
}
fn verify_chain_against_policy(chain_context: CertChainContext, hostname: &str) -> ValidationResult {
let mut encoded_host = OsStr::new(hostname).encode_wide().chain(Some(0).into_iter()).collect::<Vec<_>>();
let mut extra_policy = SSL_EXTRA_CERT_CHAIN_POLICY_PARA {
cbSize: mem::size_of::<SSL_EXTRA_CERT_CHAIN_POLICY_PARA>() as u32,
dwAuthType: 2, fdwChecks: 0,
pwszServerName: encoded_host.as_mut_ptr(), };
let mut policy = CERT_CHAIN_POLICY_PARA {
cbSize: mem::size_of::<CERT_CHAIN_POLICY_PARA>() as u32,
dwFlags: 0,
pvExtraPolicyPara: &mut extra_policy as *mut _ as *mut c_void,
};
let mut result = CERT_CHAIN_POLICY_STATUS {
cbSize: mem::size_of::<CERT_CHAIN_POLICY_STATUS>() as u32,
dwError: 0,
lChainIndex: 0,
lElementIndex: 0,
pvExtraPolicyStatus: ptr::null_mut(),
};
unsafe {
let verified = CertVerifyCertificateChainPolicy(
CERT_CHAIN_POLICY_SSL as *const i8,
chain_context.0,
&mut policy,
&mut result,
);
if verified == 0 {
return ValidationResult::ErrorDuringValidation;
}
}
match result.dwError {
0 => ValidationResult::Trusted,
_ => ValidationResult::NotTrusted,
}
}
fn build_chain(cert_context: CertContext) -> Result<CertChainContext, ValidationResult> {
let mut server_auth: Vec<c_char> = as_cchar_vec!(szOID_PKIX_KP_SERVER_AUTH);
let mut server_gated_crypto: Vec<c_char> = as_cchar_vec!(szOID_SERVER_GATED_CRYPTO);
let mut sgc_netscape: Vec<c_char> = as_cchar_vec!(szOID_SGC_NETSCAPE);
let mut usage = [
server_auth.as_mut_ptr(),
server_gated_crypto.as_mut_ptr(),
sgc_netscape.as_mut_ptr(),
];
let mut chain_parameters: CERT_CHAIN_PARA = unsafe{ mem::zeroed() };
chain_parameters.RequestedUsage.dwType = 1; chain_parameters.RequestedUsage.Usage.cUsageIdentifier = usage.len() as u32;
chain_parameters.RequestedUsage.Usage.rgpszUsageIdentifier = usage.as_mut_ptr();
chain_parameters.cbSize = mem::size_of::<CERT_CHAIN_PARA>() as u32;
let mut chain_context_ptr = ptr::null();
unsafe {
let got_chain = CertGetCertificateChain(
ptr::null_mut(), cert_context.0, ptr::null_mut(), (*(cert_context.0)).hCertStore, &mut chain_parameters, 0, ptr::null_mut(), &mut chain_context_ptr
);
if got_chain == 0 {
return Err(ValidationResult::NotTrusted);
}
if chain_context_ptr.is_null() {
return Err(ValidationResult::ErrorDuringValidation);
}
}
let context = CertChainContext(chain_context_ptr as PCCERT_CHAIN_CONTEXT);
Ok(context)
}
fn build_cert_context(encoded_certs: &[&[u8]]) -> Result<CertContext, ValidationResult> {
let store_ptr = unsafe {
let backing_store = CertOpenStore(
CERT_STORE_PROV_MEMORY as *const i8,
0,
0,
CERT_STORE_DEFER_CLOSE_UNTIL_LAST_FREE_FLAG,
ptr::null(),
);
if backing_store.is_null() {
return Err(ValidationResult::ErrorDuringValidation);
}
backing_store
};
let store = CertStore(store_ptr);
let mut primary_cert_ptr = ptr::null();
unsafe {
let ok = CertAddEncodedCertificateToStore(
store.0,
X509_ASN_ENCODING,
encoded_certs[0].as_ptr(),
encoded_certs[0].len() as u32,
CERT_STORE_ADD_ALWAYS,
&mut primary_cert_ptr,
);
if ok == 0 {
return Err(ValidationResult::MalformedCertificateInChain);
}
}
let primary_cert = CertContext(primary_cert_ptr);
for cert in &encoded_certs[1..] {
unsafe {
let ok = CertAddEncodedCertificateToStore(
store.0,
X509_ASN_ENCODING,
cert.as_ptr(),
cert.len() as u32,
CERT_STORE_ADD_ALWAYS,
ptr::null_mut(),
);
if ok == 0 {
return Err(ValidationResult::MalformedCertificateInChain);
}
}
}
Ok(primary_cert)
}
#[cfg(test)]
mod test {
use windows::validate_cert_chain;
use test::{expired_chain, certifi_chain, self_signed_chain};
use ValidationResult;
#[test]
fn can_validate_good_chain() {
let chain = certifi_chain();
let valid = validate_cert_chain(&chain, "certifi.io");
assert_eq!(valid, ValidationResult::Trusted);
}
#[test]
fn fails_on_bad_hostname() {
let chain = certifi_chain();
let valid = validate_cert_chain(&chain, "lukasa.co.uk");
assert_eq!(valid, ValidationResult::NotTrusted);
}
#[test]
fn fails_on_bad_cert() {
let mut good_chain = certifi_chain();
let originals = good_chain.split_first_mut().unwrap();
let leaf = originals.0;
let intermediates = originals.1;
let mut certs = vec![&leaf[1..50]];
certs.extend(intermediates.iter());
let valid = validate_cert_chain(&certs, "certifi.io");
assert_eq!(valid, ValidationResult::MalformedCertificateInChain);
}
#[test]
fn fails_on_expired_cert() {
let chain = expired_chain();
let valid = validate_cert_chain(&chain, "expired.badssl.com");
assert_eq!(valid, ValidationResult::NotTrusted);
}
#[test]
fn test_fails_on_self_signed() {
let chain = self_signed_chain();
let valid = validate_cert_chain(&chain, "self-signed.badssl.com");
assert_eq!(valid, ValidationResult::NotTrusted);
}
#[test]
fn test_fails_on_invalid_asn1() {
let mut chain = certifi_chain();
let mut first_chain = chain[0].to_vec();
first_chain[0] = 0xff;
let mut chain_builder = vec![first_chain.as_slice()];
chain_builder.append(&mut chain[1..].to_vec());
let new_chain = chain_builder.as_slice();
let valid = validate_cert_chain(&new_chain, "certifi.io");
assert_eq!(valid, ValidationResult::MalformedCertificateInChain);
}
}