use std::ptr;
use super::sspi_ffi::{
CompleteAuthToken, CredHandle, CtxtHandle, DeleteSecurityContext, FreeCredentialsHandle,
InitializeSecurityContextW, PVOID, SEC_E_NO_CREDENTIALS, SEC_E_OK, SECBUFFER_CHANNEL_BINDINGS,
SECBUFFER_TOKEN, SECBUFFER_VERSION, SECURITY_STATUS, STANDARD_CONTEXT_REQ, SecBuffer,
SecBufferDesc, TimeStamp, acquire_credentials, get_sspi_error_message, is_success_status,
needs_complete, needs_continue, to_wide_string,
};
use crate::security::{
IntegratedAuthConfig, SecurityContext, SecurityError, SecurityPackage, SspiAuthToken,
spn::make_spn_canonicalized,
};
const DEFAULT_MAX_TOKEN_SIZE: usize = 12288;
struct CredHandleWrapper(CredHandle);
unsafe impl Send for CredHandleWrapper {}
unsafe impl Sync for CredHandleWrapper {}
struct CtxtHandleWrapper(CtxtHandle);
unsafe impl Send for CtxtHandleWrapper {}
unsafe impl Sync for CtxtHandleWrapper {}
pub struct WindowsSspiContext {
spn: String,
package: SecurityPackage,
is_complete: bool,
is_loopback: bool,
tried_empty_spn: bool,
channel_bindings: Option<Vec<u8>>,
cred_handle: Option<CredHandleWrapper>,
ctx_handle: Option<CtxtHandleWrapper>,
max_token_size: usize,
}
impl WindowsSspiContext {
pub fn new(
config: &IntegratedAuthConfig,
server: &str,
port: u16,
) -> Result<Self, SecurityError> {
let spn = config
.server_spn
.clone()
.unwrap_or_else(|| make_spn_canonicalized(server, None, port));
let package_name = config.security_package.as_sspi_name();
let (cred_handle, _expiry) = acquire_credentials(package_name).map_err(|status| {
SecurityError::AcquireCredentialsFailed {
code: status as u32,
message: get_sspi_error_message(status),
}
})?;
let max_token_size = super::sspi_ffi::get_max_token_size(package_name)
.unwrap_or(DEFAULT_MAX_TOKEN_SIZE as u32) as usize;
Ok(Self {
spn,
package: config.security_package,
is_complete: false,
is_loopback: config.is_loopback,
tried_empty_spn: false,
channel_bindings: config.channel_bindings.clone(),
cred_handle: Some(CredHandleWrapper(cred_handle)),
ctx_handle: None,
max_token_size,
})
}
pub fn check_availability() -> Result<(), SecurityError> {
match super::sspi_ffi::get_max_token_size("Negotiate") {
Ok(_) => Ok(()),
Err(status) => Err(SecurityError::LoadLibraryFailed(format!(
"Failed to query SSPI: {}",
get_sspi_error_message(status)
))),
}
}
fn generate_token_impl(
&mut self,
server_token: Option<&[u8]>,
target_spn: &str,
) -> Result<(Vec<u8>, SECURITY_STATUS), SecurityError> {
let cred_handle = self
.cred_handle
.as_ref()
.ok_or_else(|| SecurityError::InternalError("No credential handle".to_string()))?;
let spn_wide = to_wide_string(target_spn);
let mut input_buffer = SecBuffer::default();
let mut input_buffer_desc = SecBufferDesc::default();
let mut channel_bindings_buffer = SecBuffer::default();
let mut input_buffers: Vec<SecBuffer>;
if let Some(server_data) = server_token {
input_buffer.cbBuffer = server_data.len() as u32;
input_buffer.BufferType = SECBUFFER_TOKEN;
input_buffer.pvBuffer = server_data.as_ptr() as PVOID;
if let Some(ref cb) = self.channel_bindings {
channel_bindings_buffer.cbBuffer = cb.len() as u32;
channel_bindings_buffer.BufferType = SECBUFFER_CHANNEL_BINDINGS;
channel_bindings_buffer.pvBuffer = cb.as_ptr() as PVOID;
input_buffers = vec![input_buffer, channel_bindings_buffer];
input_buffer_desc.cBuffers = 2;
} else {
input_buffers = vec![input_buffer];
input_buffer_desc.cBuffers = 1;
}
input_buffer_desc.pBuffers = input_buffers.as_mut_ptr();
} else if let Some(ref cb) = self.channel_bindings {
channel_bindings_buffer.cbBuffer = cb.len() as u32;
channel_bindings_buffer.BufferType = SECBUFFER_CHANNEL_BINDINGS;
channel_bindings_buffer.pvBuffer = cb.as_ptr() as PVOID;
input_buffers = vec![channel_bindings_buffer];
input_buffer_desc.cBuffers = 1;
input_buffer_desc.pBuffers = input_buffers.as_mut_ptr();
}
let mut output_token_data = vec![0u8; self.max_token_size];
let mut output_buffer = SecBuffer {
cbBuffer: output_token_data.len() as u32,
BufferType: SECBUFFER_TOKEN,
pvBuffer: output_token_data.as_mut_ptr() as PVOID,
};
let mut output_buffer_desc = SecBufferDesc {
ulVersion: SECBUFFER_VERSION,
cBuffers: 1,
pBuffers: &mut output_buffer,
};
let (ctx_in_ptr, mut new_ctx_handle) = match &self.ctx_handle {
Some(ctx) => (&ctx.0 as *const CtxtHandle, ctx.0),
None => (ptr::null(), CtxtHandle::default()),
};
let mut context_attr: u32 = 0;
let mut expiry = TimeStamp::default();
let status = unsafe {
InitializeSecurityContextW(
&cred_handle.0,
if ctx_in_ptr.is_null() {
ptr::null()
} else {
ctx_in_ptr
},
spn_wide.as_ptr(),
STANDARD_CONTEXT_REQ,
0, 0, if server_token.is_some() || self.channel_bindings.is_some() {
&input_buffer_desc
} else {
ptr::null()
},
0, &mut new_ctx_handle,
&mut output_buffer_desc,
&mut context_attr,
&mut expiry,
)
};
if new_ctx_handle.is_valid() {
self.ctx_handle = Some(CtxtHandleWrapper(new_ctx_handle));
}
if !is_success_status(status) {
return Err(SecurityError::InitContextFailed {
code: status as u32,
message: get_sspi_error_message(status),
});
}
if let Some(ctx) = self.ctx_handle.as_ref().filter(|_| needs_complete(status)) {
let complete_status = unsafe { CompleteAuthToken(&ctx.0, &output_buffer_desc) };
if complete_status != SEC_E_OK {
return Err(SecurityError::InitContextFailed {
code: complete_status as u32,
message: format!(
"CompleteAuthToken failed: {}",
get_sspi_error_message(complete_status)
),
});
}
}
let output_size = output_buffer.cbBuffer as usize;
output_token_data.truncate(output_size);
Ok((output_token_data, status))
}
}
impl SecurityContext for WindowsSspiContext {
fn package_name(&self) -> &str {
self.package.as_sspi_name()
}
fn generate_token(
&mut self,
server_token: Option<&[u8]>,
) -> Result<SspiAuthToken, SecurityError> {
let result = self.generate_token_impl(server_token, &self.spn.clone());
match result {
Ok((token, status)) => {
if !needs_continue(status) {
self.is_complete = true;
}
Ok(SspiAuthToken {
data: token,
is_complete: self.is_complete,
})
}
Err(SecurityError::InitContextFailed { code, .. })
if code == SEC_E_NO_CREDENTIALS as u32 =>
{
Err(SecurityError::NoCredentials)
}
Err(ref e) if self.is_loopback && !self.tried_empty_spn => {
tracing::debug!(
spn = %self.spn,
error = %e,
has_ctx = self.ctx_handle.is_some(),
has_server_token = server_token.is_some(),
"Loopback SPN retry: retrying with empty SPN"
);
self.tried_empty_spn = true;
self.spn = String::new();
let (token, status) = self.generate_token_impl(server_token, "")?;
if !needs_continue(status) {
self.is_complete = true;
}
Ok(SspiAuthToken {
data: token,
is_complete: self.is_complete,
})
}
Err(e) => Err(e),
}
}
fn is_complete(&self) -> bool {
self.is_complete
}
fn spn(&self) -> &str {
&self.spn
}
}
impl Drop for WindowsSspiContext {
fn drop(&mut self) {
if let Some(ref mut ctx) = self.ctx_handle {
unsafe {
DeleteSecurityContext(&mut ctx.0);
}
}
if let Some(ref mut cred) = self.cred_handle {
unsafe {
FreeCredentialsHandle(&mut cred.0);
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_sspi_package_name() {
assert_eq!(SecurityPackage::Negotiate.as_sspi_name(), "Negotiate");
assert_eq!(SecurityPackage::Kerberos.as_sspi_name(), "Kerberos");
assert_eq!(SecurityPackage::Ntlm.as_sspi_name(), "NTLM");
}
}