use std::alloc::{self, Layout};
use std::fmt;
use std::io;
use std::os::windows::io::AsRawHandle;
use std::ptr::NonNull;
use std::slice;
use windows_sys::Win32::Foundation::{ERROR_IO_PENDING, FALSE};
use windows_sys::Win32::Storage::FileSystem::{
FILE_SEGMENT_ELEMENT, ReadFile, ReadFileScatter, WriteFile, WriteFileGather,
};
use windows_sys::Win32::System::IO::{GetOverlappedResult, OVERLAPPED};
use crate::operation::{payload_ptr_from_overlapped, sync_bytes_ptr_from_overlapped};
use crate::{
AssociatedEndpoint, BlockingEndpoint, Completion, IoBuf, IoBufMut, Issued, Operation,
OperationId, Started, Submitted,
};
impl BlockingEndpoint {
pub fn read(&mut self, buffer: &mut [u8], offset: u64) -> io::Result<usize> {
let buf_len = checked_len(buffer.len(), "read buffer")?;
let buf_ptr = buffer.as_mut_ptr();
let mut operation = Operation::new(());
operation.set_offset(offset);
unsafe {
self.run(&mut operation, |handle, overlapped| {
let ok = ReadFile(
handle.as_raw_handle(),
buf_ptr,
buf_len,
std::ptr::null_mut(),
overlapped,
);
classify(ok)
})
}
}
pub fn write(&mut self, data: &[u8], offset: u64) -> io::Result<usize> {
let data_ptr = data.as_ptr();
let data_len = checked_len(data.len(), "write buffer")?;
let mut operation = Operation::new(());
operation.set_offset(offset);
let written = unsafe {
self.run(&mut operation, |handle, overlapped| {
let ok = WriteFile(
handle.as_raw_handle(),
data_ptr,
data_len,
std::ptr::null_mut(),
overlapped,
);
classify(ok)
})
}?;
Ok(written)
}
}
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 {which} is limited to u32::MAX bytes; {len} does not fit"),
)
})
}
impl AssociatedEndpoint<'_> {
#[track_caller]
pub fn read<B: IoBufMut>(
&self,
mut buffer: B,
offset: u64,
) -> io::Result<Started<FileIo<B>, B>> {
let buf_len = checked_len(buffer.bytes_len(), "read buffer")?;
let skip = self.notification_modes().skip_completion_port_on_success;
let buf_ptr = buffer.stable_mut_ptr();
let mut operation = Operation::new(buffer);
operation.set_offset(offset);
let submitted = unsafe {
self.submit(operation, |handle, overlapped| {
let bytes = sync_bytes_ptr_from_overlapped(overlapped);
let ok = ReadFile(handle.as_raw_handle(), buf_ptr, buf_len, bytes, overlapped);
classify_issued(ok, skip, bytes)
})
};
finish(submitted)
}
#[track_caller]
pub fn write<B: IoBuf>(&self, buffer: B, offset: u64) -> io::Result<Started<FileIo<B>, B>> {
let data_len = checked_len(buffer.bytes_len(), "write buffer")?;
let skip = self.notification_modes().skip_completion_port_on_success;
let data_ptr = buffer.stable_ptr();
let mut operation = Operation::new(buffer);
operation.set_offset(offset);
let submitted = unsafe {
self.submit(operation, |handle, overlapped| {
let bytes = sync_bytes_ptr_from_overlapped(overlapped);
let ok = WriteFile(
handle.as_raw_handle(),
data_ptr,
data_len,
bytes,
overlapped,
);
classify_issued(ok, skip, bytes)
})
};
finish(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)
}
}
unsafe fn classify_scatter(
ok: i32,
skip_on_success: bool,
handle: std::os::windows::io::RawHandle,
overlapped: *mut OVERLAPPED,
) -> io::Result<Issued> {
if ok != 0 {
if skip_on_success {
let mut bytes_transferred = 0_u32;
let got =
unsafe { GetOverlappedResult(handle, overlapped, &mut bytes_transferred, FALSE) };
if got == 0 {
return Err(io::Error::last_os_error());
}
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<B: IoBuf>(submitted: Submitted<B>) -> io::Result<Started<FileIo<B>, B>> {
match submitted {
Submitted::Pending(id) => Ok(Started::Pending(FileIo {
id,
buffer: std::marker::PhantomData,
})),
Submitted::Completed {
operation,
bytes_transferred,
} => Ok(Started::Completed {
payload: operation.into_payload(),
bytes_transferred: bytes_transferred as usize,
}),
Submitted::Failed { error, .. } => Err(error),
}
}
#[derive(Debug)]
pub struct FileIo<B> {
id: OperationId,
buffer: std::marker::PhantomData<fn() -> B>,
}
impl<B: IoBuf> FileIo<B> {
#[must_use]
pub fn id(&self) -> OperationId {
self.id
}
pub fn claim(self, completion: &Completion) -> Result<(B, io::Result<usize>), Self> {
if completion.id() != Some(self.id) {
return Err(self);
}
let operation = unsafe { completion.claim::<B>() };
let buffer = operation.into_payload();
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((buffer, result))
}
}
pub const PAGE_SIZE: usize = 4096;
pub const FILE_FLAG_NO_BUFFERING: u32 =
windows_sys::Win32::Storage::FileSystem::FILE_FLAG_NO_BUFFERING;
pub struct PageBuffers {
ptr: NonNull<u8>,
pages: usize,
}
unsafe impl Send for PageBuffers {}
unsafe impl Sync for PageBuffers {}
impl PageBuffers {
#[must_use]
pub fn new(pages: usize) -> Self {
assert!(pages > 0, "PageBuffers requires at least one page");
let size = pages
.checked_mul(PAGE_SIZE)
.expect("page buffer size overflow");
let layout = Layout::from_size_align(size, PAGE_SIZE).expect("valid page layout");
let raw = unsafe { alloc::alloc_zeroed(layout) };
let ptr = NonNull::new(raw).unwrap_or_else(|| alloc::handle_alloc_error(layout));
Self { ptr, pages }
}
#[must_use]
pub fn pages(&self) -> usize {
self.pages
}
#[must_use]
pub fn len(&self) -> usize {
self.pages * PAGE_SIZE
}
#[must_use]
pub fn is_empty(&self) -> bool {
false
}
#[must_use]
pub fn as_bytes(&self) -> &[u8] {
unsafe { slice::from_raw_parts(self.ptr.as_ptr(), self.len()) }
}
#[must_use]
pub fn as_bytes_mut(&mut self) -> &mut [u8] {
unsafe { slice::from_raw_parts_mut(self.ptr.as_ptr(), self.len()) }
}
fn segment_array(&self) -> Vec<FILE_SEGMENT_ELEMENT> {
let mut segments = Vec::with_capacity(self.pages + 1);
for i in 0..self.pages {
let page = unsafe { self.ptr.as_ptr().add(i * PAGE_SIZE) };
segments.push(FILE_SEGMENT_ELEMENT {
Buffer: page.cast(),
});
}
segments.push(FILE_SEGMENT_ELEMENT { Alignment: 0 });
segments
}
}
impl fmt::Debug for PageBuffers {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("PageBuffers")
.field("pages", &self.pages)
.finish_non_exhaustive()
}
}
impl Drop for PageBuffers {
fn drop(&mut self) {
let layout = Layout::from_size_align(self.len(), PAGE_SIZE).expect("valid page layout");
unsafe { alloc::dealloc(self.ptr.as_ptr(), layout) };
}
}
unsafe impl crate::IoBuf for PageBuffers {
fn stable_ptr(&self) -> *const u8 {
self.ptr.as_ptr()
}
fn bytes_len(&self) -> usize {
self.len()
}
}
unsafe impl crate::IoBufMut for PageBuffers {
fn stable_mut_ptr(&mut self) -> *mut u8 {
self.ptr.as_ptr()
}
}
impl BlockingEndpoint {
pub fn read_scatter(&mut self, buffers: &mut PageBuffers, offset: u64) -> io::Result<usize> {
let total = checked_len(buffers.len(), "scatter/gather buffer set")?;
let segments = buffers.segment_array();
let seg_ptr = segments.as_ptr();
let mut operation = Operation::new(());
operation.set_offset(offset);
unsafe {
self.run(&mut operation, |handle, overlapped| {
let ok = ReadFileScatter(
handle.as_raw_handle(),
seg_ptr,
total,
std::ptr::null(),
overlapped,
);
classify(ok)
})
}
}
pub fn write_gather(&mut self, buffers: &PageBuffers, offset: u64) -> io::Result<usize> {
let segments = buffers.segment_array();
let total = checked_len(buffers.len(), "scatter/gather buffer set")?;
let seg_ptr = segments.as_ptr();
let mut operation = Operation::new(());
operation.set_offset(offset);
let written = unsafe {
self.run(&mut operation, |handle, overlapped| {
let ok = WriteFileGather(
handle.as_raw_handle(),
seg_ptr,
total,
std::ptr::null(),
overlapped,
);
classify(ok)
})
}?;
Ok(written)
}
}
struct ScatterPayload {
buffers: PageBuffers,
segments: Vec<FILE_SEGMENT_ELEMENT>,
}
unsafe impl Send for ScatterPayload {}
impl AssociatedEndpoint<'_> {
#[track_caller]
pub fn read_scatter(
&self,
buffers: PageBuffers,
offset: u64,
) -> io::Result<Started<ScatterGatherIo, PageBuffers>> {
let total = checked_len(buffers.len(), "scatter/gather buffer set")?;
let skip = self.notification_modes().skip_completion_port_on_success;
let segments = buffers.segment_array();
let mut operation = Operation::new(ScatterPayload { buffers, segments });
operation.set_offset(offset);
let submitted = unsafe {
self.submit(operation, |handle, overlapped| {
let payload = payload_ptr_from_overlapped::<ScatterPayload>(overlapped);
let raw = handle.as_raw_handle();
let ok = ReadFileScatter(
raw,
(*payload).segments.as_ptr(),
total,
std::ptr::null(),
overlapped,
);
classify_scatter(ok, skip, raw, overlapped)
})
};
finish_scatter(submitted)
}
#[track_caller]
pub fn write_gather(
&self,
buffers: PageBuffers,
offset: u64,
) -> io::Result<Started<ScatterGatherIo, PageBuffers>> {
let total = checked_len(buffers.len(), "scatter/gather buffer set")?;
let skip = self.notification_modes().skip_completion_port_on_success;
let segments = buffers.segment_array();
let mut operation = Operation::new(ScatterPayload { buffers, segments });
operation.set_offset(offset);
let submitted = unsafe {
self.submit(operation, |handle, overlapped| {
let payload = payload_ptr_from_overlapped::<ScatterPayload>(overlapped);
let raw = handle.as_raw_handle();
let ok = WriteFileGather(
raw,
(*payload).segments.as_ptr(),
total,
std::ptr::null(),
overlapped,
);
classify_scatter(ok, skip, raw, overlapped)
})
};
finish_scatter(submitted)
}
}
fn finish_scatter(
submitted: Submitted<ScatterPayload>,
) -> io::Result<Started<ScatterGatherIo, PageBuffers>> {
match submitted {
Submitted::Pending(id) => Ok(Started::Pending(ScatterGatherIo { id })),
Submitted::Completed {
operation,
bytes_transferred,
} => Ok(Started::Completed {
payload: operation.into_payload().buffers,
bytes_transferred: bytes_transferred as usize,
}),
Submitted::Failed { error, .. } => Err(error),
}
}
#[derive(Debug)]
pub struct ScatterGatherIo {
id: OperationId,
}
impl ScatterGatherIo {
#[must_use]
pub fn id(&self) -> OperationId {
self.id
}
pub fn claim(self, completion: &Completion) -> Result<(PageBuffers, io::Result<usize>), Self> {
if completion.id() != Some(self.id) {
return Err(self);
}
let operation = unsafe { completion.claim::<ScatterPayload>() };
let buffers = operation.into_payload().buffers;
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((buffers, result))
}
}
#[cfg(test)]
mod tests;