use windows_impersonation_token_sys::{CaptureError, ImpersonationToken};
use crate::buffer::NativeBuffer;
use crate::completion::EnumerationId;
use crate::engine::EngineState;
use crate::error::{BeginError, BeginFailure};
use crate::request::EnumerationRequest;
use crate::session::SessionShared;
use crate::submission_ring::{
BeginMessage, CancelSlot, ControlMessage, SubmitRejection, release_cancel_slot,
release_retire_slot,
};
use std::sync::Arc;
#[must_use = "dropping the handle cancels the enumeration; use `detach` to let it run"]
pub struct EnumerationHandle {
enumeration: EnumerationId,
shared: Arc<SessionShared>,
cancel: Option<CancelSlot>,
}
impl EnumerationHandle {
pub(crate) fn new(
enumeration: EnumerationId,
shared: Arc<SessionShared>,
cancel: CancelSlot,
) -> Self {
Self {
enumeration,
shared,
cancel: Some(cancel),
}
}
#[must_use]
pub fn id(&self) -> EnumerationId {
self.enumeration
}
pub fn cancel(mut self) {
self.enqueue_cancel();
}
pub fn detach(mut self) {
if let Some(slot) = self.cancel.take() {
release_cancel_slot(&self.shared.submissions, slot);
}
}
fn enqueue_cancel(&mut self) {
let Some(slot) = self.cancel.take() else {
return;
};
let pushed = self.shared.submissions.push_cancel(slot, self.enumeration);
self.shared.ring_servicer(pushed);
}
}
impl Drop for EnumerationHandle {
fn drop(&mut self) {
self.enqueue_cancel();
}
}
impl std::fmt::Debug for EnumerationHandle {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("EnumerationHandle")
.field("enumeration", &self.enumeration)
.field("cancellable", &self.cancel.is_some())
.finish_non_exhaustive()
}
}
pub(crate) fn try_begin(
shared: &Arc<SessionShared>,
request: EnumerationRequest,
) -> Result<EnumerationHandle, BeginError> {
let token = match ImpersonationToken::capture() {
Ok(token) => token,
Err(error) => return Err(BeginError::capture(request, error)),
};
try_begin_with_token(shared, request, token)
}
pub(crate) fn try_begin_with_token(
shared: &Arc<SessionShared>,
request: EnumerationRequest,
token: ImpersonationToken,
) -> Result<EnumerationHandle, BeginError> {
if shared.submissions.is_abandoned() {
return Err(BeginError::rejected(
BeginFailure::Abandoned,
request,
Some(token),
));
}
let Some(cancel) = shared.submissions.reserve_cancel() else {
return Err(BeginError::rejected(
BeginFailure::SubmissionRingFull,
request,
Some(token),
));
};
let Some(retire) = shared.submissions.reserve_retire() else {
release_cancel_slot(&shared.submissions, cancel);
return Err(BeginError::rejected(
BeginFailure::SubmissionRingFull,
request,
Some(token),
));
};
let enumeration = shared.next_enumeration_id();
let Some(terminal) = shared.completions.reserve_terminal(enumeration) else {
release_cancel_slot(&shared.submissions, cancel);
release_retire_slot(&shared.submissions, retire);
return Err(BeginError::rejected(
BeginFailure::CompletionRingFull,
request,
Some(token),
));
};
let Some(buffer) = NativeBuffer::try_new(request.buffer_capacity()) else {
release_cancel_slot(&shared.submissions, cancel);
release_retire_slot(&shared.submissions, retire);
drop(terminal);
return Err(BeginError::rejected(
BeginFailure::BufferAllocation,
request,
Some(token),
));
};
let message = ControlMessage::Begin(Box::new(BeginMessage {
enumeration,
engine: EngineState::new(request, token, buffer),
terminal,
retire,
}));
match shared.submissions.try_push(message) {
Ok(pushed) => {
shared.ring_servicer(pushed);
Ok(EnumerationHandle::new(
enumeration,
Arc::clone(shared),
cancel,
))
}
Err((message, rejection)) => {
release_cancel_slot(&shared.submissions, cancel);
let (request, token) = match message {
ControlMessage::Begin(begin) => {
let begin = *begin;
release_retire_slot(&shared.submissions, begin.retire);
begin.engine.into_parts()
}
_ => unreachable!("the message pushed above is always a begin"),
};
let failure = match rejection {
SubmitRejection::Full => BeginFailure::SubmissionRingFull,
SubmitRejection::Abandoned => BeginFailure::Abandoned,
};
Err(BeginError::rejected(failure, request, Some(token)))
}
}
}
pub type TokenCaptureError = CaptureError;
#[cfg(test)]
mod tests;