use std::cell::Cell;
use std::collections::HashMap;
use std::fmt;
use std::io;
use std::os::windows::io::{AsHandle, AsRawHandle, BorrowedHandle, FromRawHandle, OwnedHandle};
use std::panic::Location;
use std::sync::{Arc, Mutex, MutexGuard};
use windows_sys::Win32::Foundation::{HANDLE, INVALID_HANDLE_VALUE, WAIT_TIMEOUT};
use windows_sys::Win32::System::IO::{
CancelIoEx, CreateIoCompletionPort, GetQueuedCompletionStatus, OVERLAPPED,
PostQueuedCompletionStatus,
};
use crate::identity::{OperationId, OperationRegistry};
use crate::{Operation, OperationState, UnassociatedEndpoint};
const RUN_DOWN_POLL_MS: u32 = 5;
struct Track {
location: &'static Location<'static>,
#[cfg(feature = "operation-backtrace")]
backtrace: std::backtrace::Backtrace,
}
struct PortState {
live: OperationRegistry,
tracked: Mutex<HashMap<usize, Track>>,
}
impl PortState {
fn new() -> Self {
Self {
live: OperationRegistry::new(),
tracked: Mutex::new(HashMap::new()),
}
}
}
fn lock<T>(mutex: &Mutex<T>) -> MutexGuard<'_, T> {
mutex.lock().unwrap_or_else(|poison| poison.into_inner())
}
pub struct CompletionPort {
handle: OwnedHandle,
state: Arc<PortState>,
}
impl CompletionPort {
pub fn new(concurrency: u32) -> io::Result<Self> {
let handle = unsafe {
CreateIoCompletionPort(INVALID_HANDLE_VALUE, std::ptr::null_mut(), 0, concurrency)
};
if handle.is_null() {
return Err(io::Error::last_os_error());
}
let handle = unsafe { OwnedHandle::from_raw_handle(handle) };
Ok(Self {
handle,
state: Arc::new(PortState::new()),
})
}
pub fn associate(
&self,
endpoint: UnassociatedEndpoint,
key: usize,
) -> io::Result<AssociatedEndpoint<'_>> {
let handle = endpoint.into_handle();
let result = unsafe { CreateIoCompletionPort(handle.as_raw_handle(), self.raw(), key, 0) };
if result.is_null() {
return Err(io::Error::last_os_error());
}
Ok(AssociatedEndpoint {
port: self,
handle,
key,
})
}
pub fn post(&self, key: usize, bytes_transferred: u32) -> io::Result<()> {
let ok = unsafe {
PostQueuedCompletionStatus(self.raw(), bytes_transferred, key, std::ptr::null())
};
if ok == 0 {
return Err(io::Error::last_os_error());
}
Ok(())
}
#[cfg(test)]
pub(crate) fn post_raw(
&self,
key: usize,
bytes_transferred: u32,
overlapped: *mut OVERLAPPED,
) -> io::Result<()> {
let ok = unsafe {
PostQueuedCompletionStatus(self.raw(), bytes_transferred, key, overlapped.cast_const())
};
if ok == 0 {
return Err(io::Error::last_os_error());
}
Ok(())
}
pub fn get(&self, timeout_ms: u32) -> io::Result<Option<Completion>> {
let mut bytes_transferred: u32 = 0;
let mut key: usize = 0;
let mut overlapped: *mut OVERLAPPED = std::ptr::null_mut();
let ok = unsafe {
GetQueuedCompletionStatus(
self.raw(),
&mut bytes_transferred,
&mut key,
&mut overlapped,
timeout_ms,
)
};
if ok != 0 {
return Ok(Some(Completion {
key,
bytes_transferred,
overlapped,
error: None,
id: self.deregister_dequeued(overlapped),
claimed: Cell::new(false),
}));
}
let error = io::Error::last_os_error();
if overlapped.is_null() {
if error.raw_os_error() == Some(WAIT_TIMEOUT as i32) {
return Ok(None);
}
return Err(error);
}
Ok(Some(Completion {
key,
bytes_transferred,
overlapped,
error: Some(error),
id: self.deregister_dequeued(overlapped),
claimed: Cell::new(false),
}))
}
fn deregister_dequeued(&self, overlapped: *mut OVERLAPPED) -> Option<OperationId> {
let id = self.state.live.remove(overlapped);
if id.is_some() && crate::source_tracking_enabled() {
lock(&self.state.tracked).remove(&(overlapped as usize));
}
id
}
pub(crate) fn raw(&self) -> HANDLE {
self.handle.as_raw_handle()
}
#[cfg(feature = "socket")]
pub(crate) fn live_operations(&self) -> &OperationRegistry {
&self.state.live
}
#[must_use]
pub fn outstanding(&self) -> usize {
self.state.live.len()
}
pub fn run_down(&self) -> io::Result<()> {
while self.outstanding() > 0 {
self.get(RUN_DOWN_POLL_MS)?;
}
Ok(())
}
#[track_caller]
pub(crate) unsafe fn submit_with<P, F>(&self, operation: Operation<P>, issue: F) -> Submitted<P>
where
P: Send,
F: FnOnce(*mut OVERLAPPED) -> io::Result<Issued>,
{
let overlapped = operation.into_overlapped();
let identity = overlapped as usize;
let id = OperationId::mint(overlapped);
let state = &self.state;
state.live.insert(id);
let tracking = crate::source_tracking_enabled();
if tracking {
lock(&state.tracked).insert(
identity,
Track {
location: Location::caller(),
#[cfg(feature = "operation-backtrace")]
backtrace: std::backtrace::Backtrace::capture(),
},
);
}
match issue(overlapped) {
Ok(Issued::Pending) => Submitted::Pending(id),
Ok(Issued::Completed { bytes_transferred }) => {
state.live.remove(overlapped);
if tracking {
lock(&state.tracked).remove(&identity);
}
let mut operation = unsafe { Operation::<P>::from_overlapped(overlapped) };
operation.set_state(OperationState::Completed);
Submitted::Completed {
operation,
bytes_transferred,
}
}
Err(error) => {
state.live.remove(overlapped);
if tracking {
lock(&state.tracked).remove(&identity);
}
let mut operation = unsafe { Operation::<P>::from_overlapped(overlapped) };
operation.set_state(OperationState::Idle);
Submitted::Failed { operation, error }
}
}
}
fn report_outstanding_at_drop(&self, count: usize) {
let tracked = lock(&self.state.tracked);
let mut message = format!(
"windows-overlapped-io-sys: CompletionPort dropped with {count} operation(s) still \
outstanding; call run_down() before dropping to control when this blocks."
);
if tracked.is_empty() {
message.push_str(
" Enable source tracking (WINDOWS_OVERLAPPED_IO_SYS_TRACK=1, or \
set_source_tracking) to identify the submit sites.",
);
} else {
message.push_str(" Sources:");
for track in tracked.values() {
message.push_str("\n - ");
message.push_str(&track.location.to_string());
#[cfg(feature = "operation-backtrace")]
{
message.push_str("\n backtrace:\n");
message.push_str(&track.backtrace.to_string());
}
}
}
eprintln!("{message}");
}
}
impl fmt::Debug for CompletionPort {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("CompletionPort")
.field("outstanding", &self.outstanding())
.finish_non_exhaustive()
}
}
impl Drop for CompletionPort {
fn drop(&mut self) {
let count = self.outstanding();
if count == 0 {
return;
}
self.report_outstanding_at_drop(count);
let _ = self.run_down();
}
}
#[derive(Debug)]
pub struct AssociatedEndpoint<'port> {
port: &'port CompletionPort,
handle: OwnedHandle,
key: usize,
}
impl<'port> AssociatedEndpoint<'port> {
#[must_use]
pub fn handle(&self) -> BorrowedHandle<'_> {
self.handle.as_handle()
}
#[must_use]
pub fn key(&self) -> usize {
self.key
}
#[must_use]
pub fn port(&self) -> &'port CompletionPort {
self.port
}
#[track_caller]
pub unsafe fn submit<P, F>(&self, operation: Operation<P>, issue: F) -> Submitted<P>
where
P: Send,
F: FnOnce(BorrowedHandle<'_>, *mut OVERLAPPED) -> io::Result<Issued>,
{
let handle = self.handle();
unsafe {
self.port
.submit_with(operation, move |overlapped| issue(handle, overlapped))
}
}
pub fn cancel(&self, id: OperationId) -> io::Result<()> {
self.port.state.live.cancel_if_live(id, || {
let ok = unsafe { CancelIoEx(self.raw_handle(), id.as_ptr()) };
if ok == 0 {
return Err(io::Error::last_os_error());
}
Ok(())
})
}
pub fn cancel_all(&self) -> io::Result<()> {
let ok = unsafe { CancelIoEx(self.raw_handle(), std::ptr::null()) };
if ok == 0 {
return Err(io::Error::last_os_error());
}
Ok(())
}
fn raw_handle(&self) -> HANDLE {
self.handle.as_raw_handle()
}
}
#[derive(Debug, Clone, Copy)]
pub enum Issued {
Pending,
Completed {
bytes_transferred: u32,
},
}
#[derive(Debug)]
pub enum Submitted<P> {
Pending(OperationId),
Completed {
operation: Operation<P>,
bytes_transferred: u32,
},
Failed {
operation: Operation<P>,
error: io::Error,
},
}
pub struct Completion {
key: usize,
bytes_transferred: u32,
overlapped: *mut OVERLAPPED,
error: Option<io::Error>,
id: Option<OperationId>,
claimed: Cell<bool>,
}
impl fmt::Debug for Completion {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Completion")
.field("key", &self.key)
.field("bytes_transferred", &self.bytes_transferred)
.field("overlapped", &self.overlapped)
.field("id", &self.id)
.field("error", &self.error)
.finish_non_exhaustive()
}
}
impl Drop for Completion {
fn drop(&mut self) {
if self.claimed.get() || self.overlapped.is_null() {
return;
}
unsafe { crate::operation::reclaim_from_overlapped(self.overlapped) };
}
}
impl Completion {
#[must_use]
pub fn key(&self) -> usize {
self.key
}
#[must_use]
pub fn bytes_transferred(&self) -> u32 {
self.bytes_transferred
}
#[must_use]
pub fn overlapped_ptr(&self) -> *mut OVERLAPPED {
self.overlapped
}
#[must_use]
pub fn id(&self) -> Option<OperationId> {
self.id
}
#[must_use]
pub fn error(&self) -> Option<&io::Error> {
self.error.as_ref()
}
pub unsafe fn claim<P>(&self) -> Operation<P> {
self.claimed.set(true);
let mut operation = unsafe { Operation::<P>::from_overlapped(self.overlapped) };
operation.set_state(OperationState::Completed);
operation
}
}
#[cfg(test)]
mod tests;