use std::io;
use std::ptr;
use std::sync::{Arc, OnceLock};
use tracing::debug;
use windows_sys::Win32::Foundation;
use windows_sys::Win32::Security::Authentication::Identity;
use windows_sys::Win32::Security::Credentials;
use super::errors::sec_status_to_io_error;
pub(crate) struct CredHandle(Credentials::SecHandle);
unsafe impl Send for CredHandle {}
unsafe impl Sync for CredHandle {}
impl Drop for CredHandle {
fn drop(&mut self) {
unsafe {
Identity::FreeCredentialsHandle(&self.0);
}
}
}
impl CredHandle {
pub(crate) fn raw(&self) -> &Credentials::SecHandle {
&self.0
}
}
#[derive(Debug, Copy, Clone, PartialEq, Eq, Hash)]
#[allow(clippy::enum_variant_names)] pub(crate) enum CredKind {
NoValidate,
ManualValidate,
AutoValidate,
}
impl CredKind {
pub(crate) fn manual_validation_isc_bit(self) -> u32 {
match self {
CredKind::NoValidate | CredKind::ManualValidate => {
Identity::ISC_REQ_MANUAL_CRED_VALIDATION
}
CredKind::AutoValidate => 0,
}
}
}
struct CredCache {
no_validate: OnceLock<Arc<CredHandle>>,
manual_validate: OnceLock<Arc<CredHandle>>,
auto_validate: OnceLock<Arc<CredHandle>>,
}
static CACHE: CredCache = CredCache {
no_validate: OnceLock::new(),
manual_validate: OnceLock::new(),
auto_validate: OnceLock::new(),
};
pub(crate) fn get_or_acquire(kind: CredKind) -> io::Result<Arc<CredHandle>> {
let slot = match kind {
CredKind::NoValidate => &CACHE.no_validate,
CredKind::ManualValidate => &CACHE.manual_validate,
CredKind::AutoValidate => &CACHE.auto_validate,
};
if let Some(existing) = slot.get() {
debug!(kind = ?kind, "win_tls: credential cache HIT");
return Ok(existing.clone());
}
debug!(kind = ?kind, "win_tls: credential cache MISS, acquiring");
let acquired = Arc::new(acquire_client_cred(kind)?);
Ok(slot.get_or_init(move || acquired).clone())
}
fn cred_flags_for(kind: CredKind) -> u32 {
let base = Identity::SCH_USE_STRONG_CRYPTO | Identity::SCH_CRED_NO_DEFAULT_CREDS;
match kind {
CredKind::NoValidate => base | Identity::SCH_CRED_NO_SERVERNAME_CHECK,
CredKind::ManualValidate => base | Identity::SCH_CRED_NO_SERVERNAME_CHECK,
CredKind::AutoValidate => base | Identity::SCH_CRED_AUTO_CRED_VALIDATION,
}
}
fn acquire_client_cred(kind: CredKind) -> io::Result<CredHandle> {
let cred_flags = cred_flags_for(kind);
unsafe {
let mut cred_data: Identity::SCH_CREDENTIALS = std::mem::zeroed();
cred_data.dwVersion = Identity::SCH_CREDENTIALS_VERSION;
cred_data.dwFlags = cred_flags;
let mut handle: Credentials::SecHandle = std::mem::zeroed();
let status = Identity::AcquireCredentialsHandleW(
ptr::null(),
Identity::UNISP_NAME_W,
Identity::SECPKG_CRED_OUTBOUND,
ptr::null_mut(),
&cred_data as *const _ as *const _,
None,
ptr::null_mut(),
&mut handle,
ptr::null_mut(),
);
if status == Foundation::SEC_E_OK {
debug!(
kind = ?kind,
flags = format!("0x{:08x}", cred_flags),
"win_tls: AcquireCredentialsHandle OK"
);
Ok(CredHandle(handle))
} else {
debug!(
kind = ?kind,
flags = format!("0x{:08x}", cred_flags),
status = format!("0x{:08x}", status as u32),
"win_tls: AcquireCredentialsHandle FAILED"
);
Err(sec_status_to_io_error(
status,
"AcquireCredentialsHandleW(UNISP_NAME) failed",
))
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn acquire_each_kind_succeeds() {
for kind in [
CredKind::NoValidate,
CredKind::ManualValidate,
CredKind::AutoValidate,
] {
let h = get_or_acquire(kind).unwrap_or_else(|e| panic!("acquire {kind:?} failed: {e}"));
assert!(h.raw().dwLower != 0 || h.raw().dwUpper != 0);
}
}
#[test]
fn same_kind_returns_same_handle() {
let a = get_or_acquire(CredKind::NoValidate).unwrap();
let b = get_or_acquire(CredKind::NoValidate).unwrap();
assert_eq!(a.raw().dwLower, b.raw().dwLower);
assert_eq!(a.raw().dwUpper, b.raw().dwUpper);
}
#[test]
fn different_kinds_return_different_handles() {
let no = get_or_acquire(CredKind::NoValidate).unwrap();
let auto = get_or_acquire(CredKind::AutoValidate).unwrap();
assert!(
no.raw().dwLower != auto.raw().dwLower || no.raw().dwUpper != auto.raw().dwUpper,
"expected distinct SecHandles for different CredKinds"
);
}
#[test]
fn manual_validation_isc_bit_matches_kind() {
assert_ne!(CredKind::NoValidate.manual_validation_isc_bit(), 0);
assert_ne!(CredKind::ManualValidate.manual_validation_isc_bit(), 0);
assert_eq!(CredKind::AutoValidate.manual_validation_isc_bit(), 0);
}
#[test]
fn cred_flags_for_each_kind() {
let base = Identity::SCH_USE_STRONG_CRYPTO | Identity::SCH_CRED_NO_DEFAULT_CREDS;
let no = cred_flags_for(CredKind::NoValidate);
assert_eq!(no, base | Identity::SCH_CRED_NO_SERVERNAME_CHECK);
let manual = cred_flags_for(CredKind::ManualValidate);
assert_eq!(manual, base | Identity::SCH_CRED_NO_SERVERNAME_CHECK);
let auto = cred_flags_for(CredKind::AutoValidate);
assert_eq!(auto, base | Identity::SCH_CRED_AUTO_CRED_VALIDATION);
assert_eq!(auto & Identity::SCH_CRED_NO_SERVERNAME_CHECK, 0);
for f in [no, manual, auto] {
assert_eq!(f & base, base);
}
}
#[test]
fn drop_frees_zeroed_handle_without_panicking() {
let h = CredHandle(unsafe { std::mem::zeroed() });
drop(h);
}
}