#![allow(non_snake_case)]
use std::sync::Arc;
use windows_sys::Win32::Foundation::{
SEC_E_OK, SEC_I_COMPLETE_AND_CONTINUE, SEC_I_COMPLETE_NEEDED, SEC_I_CONTINUE_NEEDED,
};
use windows_sys::Win32::Security::Authentication::Identity::{
AcquireCredentialsHandleW, CompleteAuthToken, DeleteSecurityContext, FreeContextBuffer,
FreeCredentialsHandle, InitializeSecurityContextW, QueryContextAttributesW, SecBuffer,
SecBufferDesc, SecPkgContext_NegotiationInfoW, ISC_REQ_MUTUAL_AUTH, SECBUFFER_TOKEN,
SECBUFFER_VERSION, SECPKG_ATTR_NEGOTIATION_INFO, SECPKG_CRED_OUTBOUND, SECURITY_NATIVE_DREP,
};
use windows_sys::Win32::Security::Credentials::SecHandle;
use super::{NegotiateError, ProviderFactory, StepResult, TokenProvider};
const KERBEROS: &str = "Kerberos";
const NEGOTIATE_PACKAGE: &str = "Negotiate";
fn wide(value: &str) -> Vec<u16> {
value.encode_utf16().chain(std::iter::once(0)).collect()
}
unsafe fn from_wide(ptr: *const u16) -> String {
if ptr.is_null() {
return String::new();
}
let mut len = 0usize;
while *ptr.add(len) != 0 {
len += 1;
}
String::from_utf16_lossy(std::slice::from_raw_parts(ptr, len))
}
fn win_error(what: &str, status: i32) -> String {
format!("{what} failed (0x{:08X})", status as u32)
}
struct Credentials {
handle: SecHandle,
}
impl Credentials {
fn acquire() -> Result<Self, NegotiateError> {
let mut package = wide(NEGOTIATE_PACKAGE);
let mut handle = SecHandle {
dwLower: 0,
dwUpper: 0,
};
let mut expiry: i64 = 0;
let status = unsafe {
AcquireCredentialsHandleW(
std::ptr::null(),
package.as_mut_ptr(),
SECPKG_CRED_OUTBOUND,
std::ptr::null(),
std::ptr::null(),
None,
std::ptr::null(),
&mut handle,
&mut expiry,
)
};
if status != SEC_E_OK {
return Err(NegotiateError::NoTicket(win_error(
"AcquireCredentialsHandleW",
status,
)));
}
Ok(Self { handle })
}
}
impl Drop for Credentials {
fn drop(&mut self) {
unsafe { FreeCredentialsHandle(&self.handle) };
}
}
pub struct SspiProvider;
impl SspiProvider {
pub fn new() -> Self {
Self
}
}
impl Default for SspiProvider {
fn default() -> Self {
Self::new()
}
}
impl ProviderFactory for SspiProvider {
fn new_provider(&self, spn: &str) -> Result<Box<dyn TokenProvider>, NegotiateError> {
Ok(Box::new(SspiContext::new(spn)?))
}
fn name(&self) -> &'static str {
"sspi"
}
}
struct SspiContext {
credentials: Arc<Credentials>,
target: Vec<u16>,
context: SecHandle,
established: bool,
done: bool,
package_checked: bool,
}
impl SspiContext {
fn new(spn: &str) -> Result<Self, NegotiateError> {
Ok(Self {
credentials: Arc::new(Credentials::acquire()?),
target: wide(spn),
context: SecHandle {
dwLower: 0,
dwUpper: 0,
},
established: false,
done: false,
package_checked: false,
})
}
fn assert_kerberos(&mut self) -> Result<(), NegotiateError> {
if self.package_checked {
return Ok(());
}
let mut info = SecPkgContext_NegotiationInfoW {
PackageInfo: std::ptr::null_mut(),
NegotiationState: 0,
};
let status = unsafe {
QueryContextAttributesW(
&self.context,
SECPKG_ATTR_NEGOTIATION_INFO,
&mut info as *mut _ as *mut std::ffi::c_void,
)
};
if status != SEC_E_OK {
return Err(NegotiateError::Provider(win_error(
"QueryContextAttributesW(SECPKG_ATTR_NEGOTIATION_INFO)",
status,
)));
}
let package = if info.PackageInfo.is_null() {
String::new()
} else {
unsafe { from_wide((*info.PackageInfo).Name) }
};
if !info.PackageInfo.is_null() {
unsafe { FreeContextBuffer(info.PackageInfo as *mut std::ffi::c_void) };
}
self.package_checked = true;
if package.eq_ignore_ascii_case(KERBEROS) {
Ok(())
} else {
Err(NegotiateError::NtlmSelected(format!(
"Negotiate selected \"{package}\" on this host, not Kerberos"
)))
}
}
}
impl TokenProvider for SspiContext {
fn step(&mut self, peer: Option<&[u8]>) -> StepResult {
if self.done {
return StepResult::Done(None);
}
let mut input_buffer = SecBuffer {
cbBuffer: peer.map_or(0, |t| t.len() as u32),
BufferType: SECBUFFER_TOKEN,
pvBuffer: peer.map_or(std::ptr::null_mut(), |t| {
t.as_ptr() as *mut std::ffi::c_void
}),
};
let mut input_desc = SecBufferDesc {
ulVersion: SECBUFFER_VERSION,
cBuffers: 1,
pBuffers: &mut input_buffer,
};
let mut output_buffer = SecBuffer {
cbBuffer: 0,
BufferType: SECBUFFER_TOKEN,
pvBuffer: std::ptr::null_mut(),
};
let mut output_desc = SecBufferDesc {
ulVersion: SECBUFFER_VERSION,
cBuffers: 1,
pBuffers: &mut output_buffer,
};
let mut new_context = SecHandle {
dwLower: 0,
dwUpper: 0,
};
let mut attrs: u32 = 0;
let mut expiry: i64 = 0;
let status = unsafe {
InitializeSecurityContextW(
&self.credentials.handle,
if self.established {
&self.context
} else {
std::ptr::null()
},
self.target.as_ptr(),
ISC_REQ_MUTUAL_AUTH,
0,
SECURITY_NATIVE_DREP,
if peer.is_some() {
&input_desc
} else {
std::ptr::null()
},
0,
&mut new_context,
&mut output_desc,
&mut attrs,
&mut expiry,
)
};
let _ = &mut input_desc;
if status < 0 {
self.done = true;
return StepResult::Failed(NegotiateError::NoTicket(win_error(
"InitializeSecurityContextW",
status,
)));
}
self.context = new_context;
self.established = true;
let token = if output_buffer.pvBuffer.is_null() || output_buffer.cbBuffer == 0 {
Vec::new()
} else {
unsafe {
std::slice::from_raw_parts(
output_buffer.pvBuffer as *const u8,
output_buffer.cbBuffer as usize,
)
.to_vec()
}
};
if !output_buffer.pvBuffer.is_null() {
unsafe { FreeContextBuffer(output_buffer.pvBuffer) };
}
if let Err(e) = self.assert_kerberos() {
self.done = true;
return StepResult::Failed(e);
}
if status == SEC_I_COMPLETE_NEEDED || status == SEC_I_COMPLETE_AND_CONTINUE {
let mut complete_buffer = SecBuffer {
cbBuffer: token.len() as u32,
BufferType: SECBUFFER_TOKEN,
pvBuffer: token.as_ptr() as *mut std::ffi::c_void,
};
let complete_desc = SecBufferDesc {
ulVersion: SECBUFFER_VERSION,
cBuffers: 1,
pBuffers: &mut complete_buffer,
};
let complete = unsafe { CompleteAuthToken(&self.context, &complete_desc) };
if complete < 0 {
self.done = true;
return StepResult::Failed(NegotiateError::Provider(win_error(
"CompleteAuthToken",
complete,
)));
}
}
if status == SEC_I_CONTINUE_NEEDED || status == SEC_I_COMPLETE_AND_CONTINUE {
return StepResult::Continue(token);
}
self.done = true;
StepResult::Done((!token.is_empty()).then_some(token))
}
}
impl Drop for SspiContext {
fn drop(&mut self) {
if self.established {
unsafe { DeleteSecurityContext(&self.context) };
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn wide_strings_are_nul_terminated() {
let encoded = wide("HTTP/proxy.corp");
assert_eq!(*encoded.last().expect("non-empty"), 0);
assert_eq!(encoded.len(), "HTTP/proxy.corp".len() + 1);
}
#[test]
fn a_wide_string_round_trips() {
let encoded = wide("HTTP/pröxy.corp");
let decoded = unsafe { from_wide(encoded.as_ptr()) };
assert_eq!(decoded, "HTTP/pröxy.corp");
}
#[test]
fn a_null_package_name_reads_as_empty_rather_than_dereferencing() {
assert_eq!(unsafe { from_wide(std::ptr::null()) }, "");
}
#[test]
fn only_kerberos_satisfies_the_package_check() {
assert!(KERBEROS.eq_ignore_ascii_case("kerberos"));
assert!(!"NTLM".eq_ignore_ascii_case(KERBEROS));
assert!(!"Negotiate".eq_ignore_ascii_case(KERBEROS));
}
#[test]
fn an_outbound_credential_can_be_acquired_on_this_host() {
Credentials::acquire().expect("Negotiate credentials must be acquirable on Windows");
}
#[test]
fn minting_a_first_leg_either_produces_kerberos_or_refuses_ntlm() {
let provider = SspiProvider::new();
let mut context = match provider.new_provider("HTTP/localhost") {
Ok(context) => context,
Err(NegotiateError::NoTicket(_)) => return,
Err(other) => panic!("unexpected provider error: {other}"),
};
match context.step(None) {
StepResult::Continue(token) | StepResult::Done(Some(token)) => {
assert!(!token.is_empty(), "a produced token must not be empty");
}
StepResult::Done(None) => {}
StepResult::Failed(NegotiateError::NtlmSelected(detail)) => {
assert!(detail.contains("not Kerberos"), "{detail}");
}
StepResult::Failed(NegotiateError::NoTicket(_)) => {}
StepResult::Failed(other) => panic!("unexpected failure: {other}"),
}
}
}