use std::fmt;
use windows_impersonation_token_sys::{
ApplyError as ImpersonationApplyError, CaptureError as ImpersonationCaptureError,
ImpersonationToken,
};
use crate::capture_set::{CapturableAspect, CaptureSet};
use crate::captured::Captured;
use crate::declared::{Declared, DeclaredError};
use crate::error_mode::{
ApplyError as ErrorModeApplyError, ErrorModeGuard, RestoreError as ErrorModeRestoreError,
ThreadErrorMode, UnsupportedBits,
};
use crate::transaction::{TransactionContext, TransactionError};
use crate::{impersonation, transaction};
#[derive(Debug)]
#[non_exhaustive]
pub enum CaptureFailure {
Impersonation(ImpersonationCaptureError),
ErrorMode(UnsupportedBits),
Transaction(TransactionError),
}
#[derive(Debug)]
pub struct CaptureError {
failure: CaptureFailure,
}
impl CaptureError {
#[must_use]
pub const fn aspect(&self) -> CapturableAspect {
match self.failure {
CaptureFailure::Impersonation(_) => CapturableAspect::Impersonation,
CaptureFailure::ErrorMode(_) => CapturableAspect::ErrorMode,
CaptureFailure::Transaction(_) => CapturableAspect::Transaction,
}
}
#[must_use]
pub const fn failure(&self) -> &CaptureFailure {
&self.failure
}
#[must_use]
pub fn raw_os_error(&self) -> Option<i32> {
match &self.failure {
CaptureFailure::Impersonation(error) => error.raw_os_error(),
CaptureFailure::ErrorMode(_) => None,
CaptureFailure::Transaction(error) => error.raw_os_error(),
}
}
}
impl fmt::Display for CaptureError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "capturing the {} aspect failed: ", self.aspect())?;
match &self.failure {
CaptureFailure::Impersonation(error) => write!(f, "{error}"),
CaptureFailure::ErrorMode(error) => write!(f, "{error}"),
CaptureFailure::Transaction(error) => write!(f, "{error}"),
}
}
}
impl std::error::Error for CaptureError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match &self.failure {
CaptureFailure::Impersonation(error) => Some(error),
CaptureFailure::ErrorMode(error) => Some(error),
CaptureFailure::Transaction(error) => Some(error),
}
}
}
#[derive(Debug)]
#[must_use = "an ambient state that is never applied captured a context for nothing"]
pub struct AmbientState {
impersonation: Captured<ImpersonationToken>,
error_mode: Captured<ThreadErrorMode>,
transaction: Captured<TransactionContext>,
declared: Declared,
}
impl AmbientState {
pub fn capture(set: CaptureSet) -> Result<Self, CaptureError> {
let impersonation = if set.contains(CaptureSet::IMPERSONATION) {
impersonation::capture().map_err(|error| CaptureError {
failure: CaptureFailure::Impersonation(error),
})?
} else {
Captured::NotCaptured
};
let error_mode = if set.contains(CaptureSet::ERROR_MODE) {
Captured::Present(ThreadErrorMode::capture().map_err(|error| CaptureError {
failure: CaptureFailure::ErrorMode(error),
})?)
} else {
Captured::NotCaptured
};
let transaction = if set.contains(CaptureSet::TRANSACTION) {
transaction::capture().map_err(|error| CaptureError {
failure: CaptureFailure::Transaction(error),
})?
} else {
Captured::NotCaptured
};
Ok(Self {
impersonation,
error_mode,
transaction,
declared: Declared::none(),
})
}
pub fn with_declared(mut self, declared: Declared) -> Self {
self.declared = declared;
self
}
#[must_use]
pub fn captured_set(&self) -> CaptureSet {
let mut set = CaptureSet::NONE;
if self.impersonation.was_captured() {
set = set.union(CaptureSet::IMPERSONATION);
}
if self.error_mode.was_captured() {
set = set.union(CaptureSet::ERROR_MODE);
}
if self.transaction.was_captured() {
set = set.union(CaptureSet::TRANSACTION);
}
set
}
#[must_use]
pub const fn impersonation(&self) -> &Captured<ImpersonationToken> {
&self.impersonation
}
#[must_use]
pub const fn error_mode(&self) -> &Captured<ThreadErrorMode> {
&self.error_mode
}
#[must_use]
pub const fn transaction(&self) -> &Captured<TransactionContext> {
&self.transaction
}
#[must_use]
pub const fn declared(&self) -> &Declared {
&self.declared
}
pub fn with_applied<F, T>(&self, operation: F) -> Result<Applied<T>, ApplyError>
where
F: FnOnce() -> T,
{
let error_mode_guard = match self.error_mode.present() {
Some(mode) => Some(mode.apply().map_err(|error| ApplyError {
failure: ApplyFailure::ErrorMode(error),
})?),
None => None,
};
let declared_guard = match self.declared.install() {
Ok(guard) => guard,
Err(error) => {
release_error_mode(error_mode_guard);
return Err(ApplyError {
failure: ApplyFailure::Declared(error),
});
}
};
let transaction_guard = match transaction::install(&self.transaction) {
Ok(guard) => guard,
Err(error) => {
drop(declared_guard);
release_error_mode(error_mode_guard);
return Err(ApplyError {
failure: ApplyFailure::Transaction(error),
});
}
};
let outcome = impersonation::with_applied(&self.impersonation, operation);
let value = match outcome {
Ok(value) => value,
Err(error) => {
drop(transaction_guard);
drop(declared_guard);
release_error_mode(error_mode_guard);
return Err(ApplyError {
failure: ApplyFailure::Impersonation(error),
});
}
};
let transaction = transaction_guard.release().err();
let declared = declared_guard.release().err();
let error_mode = match error_mode_guard {
Some(guard) => guard.release().err(),
None => None,
};
Ok(Applied {
value,
restore: RestoreReport {
error_mode,
declared,
transaction,
},
})
}
}
fn release_error_mode(guard: Option<ErrorModeGuard>) {
if let Some(guard) = guard {
let _ = guard.release();
}
}
#[derive(Debug)]
#[must_use = "ignoring the restore report discards evidence that the thread is contaminated"]
pub struct Applied<T> {
value: T,
restore: RestoreReport,
}
impl<T> Applied<T> {
pub const fn value(&self) -> &T {
&self.value
}
pub fn into_value(self) -> T {
self.value
}
pub const fn restore(&self) -> &RestoreReport {
&self.restore
}
pub fn into_clean_value(self) -> Result<T, RestoreReport> {
if self.restore.is_clean() {
Ok(self.value)
} else {
Err(self.restore)
}
}
}
#[derive(Debug, Default)]
pub struct RestoreReport {
error_mode: Option<ErrorModeRestoreError>,
declared: Option<DeclaredError>,
transaction: Option<TransactionError>,
}
impl RestoreReport {
#[must_use]
pub const fn is_clean(&self) -> bool {
self.error_mode.is_none() && self.declared.is_none() && self.transaction.is_none()
}
#[must_use]
pub const fn error_mode(&self) -> Option<&ErrorModeRestoreError> {
self.error_mode.as_ref()
}
#[must_use]
pub const fn declared(&self) -> Option<&DeclaredError> {
self.declared.as_ref()
}
#[must_use]
pub const fn transaction(&self) -> Option<&TransactionError> {
self.transaction.as_ref()
}
}
impl fmt::Display for RestoreReport {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
if self.is_clean() {
return f.write_str("the thread was restored cleanly");
}
f.write_str("the thread is contaminated:")?;
if let Some(error) = &self.error_mode {
write!(f, " error mode: {error};")?;
}
if let Some(error) = &self.declared {
write!(f, " declared: {error};")?;
}
if let Some(error) = &self.transaction {
write!(f, " transaction: {error};")?;
}
Ok(())
}
}
impl std::error::Error for RestoreReport {}
#[derive(Debug)]
#[non_exhaustive]
pub enum ApplyFailure {
ErrorMode(ErrorModeApplyError),
Declared(DeclaredError),
Transaction(TransactionError),
Impersonation(ImpersonationApplyError),
}
#[derive(Debug)]
pub struct ApplyError {
failure: ApplyFailure,
}
impl ApplyError {
#[must_use]
pub const fn failure(&self) -> &ApplyFailure {
&self.failure
}
#[must_use]
pub fn raw_os_error(&self) -> Option<i32> {
match &self.failure {
ApplyFailure::ErrorMode(error) => error.raw_os_error(),
ApplyFailure::Declared(error) => error.raw_os_error(),
ApplyFailure::Transaction(error) => error.raw_os_error(),
ApplyFailure::Impersonation(error) => error.raw_os_error(),
}
}
}
impl fmt::Display for ApplyError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("applying the ambient state failed: ")?;
match &self.failure {
ApplyFailure::ErrorMode(error) => write!(f, "{error}"),
ApplyFailure::Declared(error) => write!(f, "{error}"),
ApplyFailure::Transaction(error) => write!(f, "{error}"),
ApplyFailure::Impersonation(error) => write!(f, "{error}"),
}
}
}
impl std::error::Error for ApplyError {
fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
match &self.failure {
ApplyFailure::ErrorMode(error) => Some(error),
ApplyFailure::Declared(error) => Some(error),
ApplyFailure::Transaction(error) => Some(error),
ApplyFailure::Impersonation(error) => Some(error),
}
}
}
#[cfg(test)]
mod tests;