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::payload_ptr_from_overlapped;
use crate::{
AssociatedEndpoint, BlockingEndpoint, Completion, Issued, Operation, OperationId, Submitted,
};
impl BlockingEndpoint {
pub unsafe fn ioctl(
&mut self,
code: u32,
input: &[u8],
output_len: usize,
) -> io::Result<(Vec<u8>, usize)> {
let in_len = checked_len(input.len(), "input")?;
let out_len = checked_len(output_len, "output")?;
let mut output = vec![0_u8; output_len];
let in_ptr = in_ptr(input);
let out_ptr = out_ptr(&mut output);
let mut operation = Operation::new(());
let returned = 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)
})
}?;
output.truncate(returned);
Ok((output, returned))
}
}
fn in_ptr(input: &[u8]) -> *const c_void {
if input.is_empty() {
std::ptr::null()
} else {
input.as_ptr().cast()
}
}
fn out_ptr(output: &mut [u8]) -> *mut c_void {
if output.is_empty() {
std::ptr::null_mut()
} else {
output.as_mut_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 {
input: Vec<u8>,
output: Vec<u8>,
}
impl AssociatedEndpoint<'_> {
#[track_caller]
pub unsafe fn ioctl(
&self,
code: u32,
input: Vec<u8>,
output_len: usize,
) -> io::Result<DeviceIoControlIo> {
let in_len = checked_len(input.len(), "input")?;
let out_len = checked_len(output_len, "output")?;
let operation = Operation::new(DeviceIoPayload {
input,
output: vec![0_u8; output_len],
});
let submitted = unsafe {
self.submit(operation, |handle, overlapped| {
let payload = payload_ptr_from_overlapped::<DeviceIoPayload>(overlapped);
let in_ptr = in_ptr(&(*payload).input);
let out_ptr = out_ptr(&mut (*payload).output);
let ok = DeviceIoControl(
handle.as_raw_handle(),
code,
in_ptr,
in_len,
out_ptr,
out_len,
std::ptr::null_mut(),
overlapped,
);
classify_issued(ok)
})
};
finish_device(submitted)
}
}
fn classify_issued(ok: i32) -> io::Result<Issued> {
if ok != 0 {
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(submitted: Submitted<DeviceIoPayload>) -> io::Result<DeviceIoControlIo> {
match submitted {
Submitted::Pending(id) => Ok(DeviceIoControlIo { id }),
Submitted::Completed { .. } => Err(io::Error::other(
"device adapter observed a synchronous completion; the endpoint must not be in \
FILE_SKIP_COMPLETION_PORT_ON_SUCCESS mode",
)),
Submitted::Failed { error, .. } => Err(error),
}
}
#[derive(Debug)]
pub struct DeviceIoControlIo {
id: OperationId,
}
impl DeviceIoControlIo {
#[must_use]
pub fn id(&self) -> OperationId {
self.id
}
pub fn claim(self, completion: &Completion) -> Result<(Vec<u8>, io::Result<usize>), Self> {
if completion.id() != Some(self.id) {
return Err(self);
}
let operation = unsafe { completion.claim::<DeviceIoPayload>() };
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;