use std::ffi::c_void;
use std::io;
use std::os::windows::io::AsRawHandle;
use windows_sys::Win32::Foundation::ERROR_IO_PENDING;
use windows_sys::Win32::System::IO::DeviceIoControl;
use crate::operation::sync_bytes_ptr_from_overlapped;
use crate::{
AssociatedEndpoint, BlockingEndpoint, Completion, IoBuf, IoBufMut, Issued, Operation,
OperationId, Started, Submitted,
};
impl BlockingEndpoint {
pub unsafe fn ioctl(
&mut self,
code: u32,
input: &[u8],
output: &mut [u8],
) -> io::Result<usize> {
let in_len = checked_len(input.len(), "input")?;
let out_len = checked_len(output.len(), "output")?;
let in_ptr = in_ptr(input.as_ptr(), in_len);
let out_ptr = out_ptr(output.as_mut_ptr(), out_len);
let mut operation = Operation::new(());
unsafe {
self.run(&mut operation, |handle, overlapped| {
let ok = DeviceIoControl(
handle.as_raw_handle(),
code,
in_ptr,
in_len,
out_ptr,
out_len,
std::ptr::null_mut(),
overlapped,
);
classify(ok)
})
}
}
}
fn in_ptr(ptr: *const u8, len: u32) -> *const c_void {
if len == 0 {
std::ptr::null()
} else {
ptr.cast()
}
}
fn out_ptr(ptr: *mut u8, len: u32) -> *mut c_void {
if len == 0 {
std::ptr::null_mut()
} else {
ptr.cast()
}
}
fn classify(ok: i32) -> io::Result<()> {
if ok != 0 {
return Ok(());
}
let error = io::Error::last_os_error();
if error.raw_os_error() == Some(ERROR_IO_PENDING as i32) {
Ok(())
} else {
Err(error)
}
}
fn checked_len(len: usize, which: &str) -> io::Result<u32> {
u32::try_from(len).map_err(|_| {
io::Error::new(
io::ErrorKind::InvalidInput,
format!(
"a DeviceIoControl {which} buffer is limited to u32::MAX bytes; {len} does not fit"
),
)
})
}
struct DeviceIoPayload<I, O> {
#[allow(dead_code)]
input: I,
output: O,
}
impl AssociatedEndpoint<'_> {
#[track_caller]
pub unsafe fn ioctl<I: IoBuf, O: IoBufMut>(
&self,
code: u32,
input: I,
mut output: O,
) -> io::Result<Started<DeviceIoControlIo<I, O>, O>> {
let in_len = checked_len(input.bytes_len(), "input")?;
let out_len = checked_len(output.bytes_len(), "output")?;
let skip = self.notification_modes().skip_completion_port_on_success;
let in_ptr = in_ptr(input.stable_ptr(), in_len);
let out_ptr = out_ptr(output.stable_mut_ptr(), out_len);
let operation = Operation::new(DeviceIoPayload { input, output });
let submitted = unsafe {
self.submit(operation, |handle, overlapped| {
let bytes = sync_bytes_ptr_from_overlapped(overlapped);
let ok = DeviceIoControl(
handle.as_raw_handle(),
code,
in_ptr,
in_len,
out_ptr,
out_len,
bytes,
overlapped,
);
classify_issued(ok, skip, bytes)
})
};
finish_device(submitted)
}
}
unsafe fn classify_issued(
ok: i32,
skip_on_success: bool,
sync_bytes: *mut u32,
) -> io::Result<Issued> {
if ok != 0 {
if skip_on_success {
let bytes_transferred = unsafe { *sync_bytes };
return Ok(Issued::Completed { bytes_transferred });
}
return Ok(Issued::Pending);
}
let error = io::Error::last_os_error();
if error.raw_os_error() == Some(ERROR_IO_PENDING as i32) {
Ok(Issued::Pending)
} else {
Err(error)
}
}
fn finish_device<I: IoBuf, O: IoBufMut>(
submitted: Submitted<DeviceIoPayload<I, O>>,
) -> io::Result<Started<DeviceIoControlIo<I, O>, O>> {
match submitted {
Submitted::Pending(id) => Ok(Started::Pending(DeviceIoControlIo {
id,
buffers: std::marker::PhantomData,
})),
Submitted::Completed {
operation,
bytes_transferred,
} => Ok(Started::Completed {
payload: operation.into_payload().output,
bytes_transferred: bytes_transferred as usize,
}),
Submitted::Failed { error, .. } => Err(error),
}
}
#[derive(Debug)]
pub struct DeviceIoControlIo<I, O> {
id: OperationId,
buffers: std::marker::PhantomData<fn() -> (I, O)>,
}
impl<I: IoBuf, O: IoBufMut> DeviceIoControlIo<I, O> {
#[must_use]
pub fn id(&self) -> OperationId {
self.id
}
pub fn claim(self, completion: &Completion) -> Result<(O, io::Result<usize>), Self> {
if completion.id() != Some(self.id) {
return Err(self);
}
let operation = unsafe { completion.claim::<DeviceIoPayload<I, O>>() };
let output = operation.into_payload().output;
let result = match completion.error() {
Some(error) => Err(io::Error::from_raw_os_error(
error.raw_os_error().unwrap_or_default(),
)),
None => Ok(completion.bytes_transferred() as usize),
};
Ok((output, result))
}
}
#[cfg(test)]
mod tests;