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;
use windows_sys::Win32::Storage::FileSystem::{
FILE_SEGMENT_ELEMENT, ReadFile, ReadFileScatter, WriteFile, WriteFileGather,
};
use crate::operation::payload_ptr_from_overlapped;
use crate::{
AssociatedEndpoint, BlockingEndpoint, Completion, Issued, Operation, OperationId, Submitted,
};
impl BlockingEndpoint {
pub fn read(&mut self, len: usize, offset: u64) -> io::Result<(Vec<u8>, usize)> {
let buf_len = checked_len(len, "read buffer")?;
let mut buffer = vec![0_u8; len];
let buf_ptr = buffer.as_mut_ptr();
let mut operation = Operation::new(());
operation.set_offset(offset);
let read = 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)
})
}?;
buffer.truncate(read);
Ok((buffer, read))
}
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 scatter_gather_len(pages: usize) -> io::Result<u32> {
if pages == 0 {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"a scatter/gather request must name at least one page",
));
}
let bytes = pages.checked_mul(PAGE_SIZE).ok_or_else(|| {
io::Error::new(
io::ErrorKind::InvalidInput,
format!("{pages} pages of {PAGE_SIZE} bytes overflows a byte count"),
)
})?;
checked_len(bytes, "scatter/gather buffer set")
}
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(&self, len: usize, offset: u64) -> io::Result<FileIo> {
let buf_len = checked_len(len, "read buffer")?;
let mut operation = Operation::new(vec![0_u8; len]);
operation.set_offset(offset);
let submitted = unsafe {
self.submit(operation, |handle, overlapped| {
let payload = payload_ptr_from_overlapped::<Vec<u8>>(overlapped);
let ok = ReadFile(
handle.as_raw_handle(),
(*payload).as_mut_ptr(),
buf_len,
std::ptr::null_mut(),
overlapped,
);
classify_issued(ok)
})
};
finish(submitted)
}
#[track_caller]
pub fn write(&self, data: Vec<u8>, offset: u64) -> io::Result<FileIo> {
let data_len = checked_len(data.len(), "write buffer")?;
let mut operation = Operation::new(data);
operation.set_offset(offset);
let submitted = unsafe {
self.submit(operation, |handle, overlapped| {
let payload = payload_ptr_from_overlapped::<Vec<u8>>(overlapped);
let ok = WriteFile(
handle.as_raw_handle(),
(*payload).as_ptr(),
data_len,
std::ptr::null_mut(),
overlapped,
);
classify_issued(ok)
})
};
finish(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(submitted: Submitted<Vec<u8>>) -> io::Result<FileIo> {
match submitted {
Submitted::Pending(id) => Ok(FileIo { id }),
Submitted::Completed { .. } => Err(io::Error::other(
"file 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 FileIo {
id: OperationId,
}
impl FileIo {
#[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::<Vec<u8>>() };
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) };
}
}
impl BlockingEndpoint {
pub fn read_scatter(&mut self, pages: usize, offset: u64) -> io::Result<(PageBuffers, usize)> {
let total = scatter_gather_len(pages)?;
let buffers = PageBuffers::new(pages);
let segments = buffers.segment_array();
let seg_ptr = segments.as_ptr();
let mut operation = Operation::new(());
operation.set_offset(offset);
let read = unsafe {
self.run(&mut operation, |handle, overlapped| {
let ok = ReadFileScatter(
handle.as_raw_handle(),
seg_ptr,
total,
std::ptr::null(),
overlapped,
);
classify(ok)
})
}?;
Ok((buffers, read))
}
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, pages: usize, offset: u64) -> io::Result<ScatterGatherIo> {
let total = scatter_gather_len(pages)?;
let buffers = PageBuffers::new(pages);
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 ok = ReadFileScatter(
handle.as_raw_handle(),
(*payload).segments.as_ptr(),
total,
std::ptr::null(),
overlapped,
);
classify_issued(ok)
})
};
finish_scatter(submitted)
}
#[track_caller]
pub fn write_gather(&self, buffers: PageBuffers, offset: u64) -> io::Result<ScatterGatherIo> {
let total = checked_len(buffers.len(), "scatter/gather buffer set")?;
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 ok = WriteFileGather(
handle.as_raw_handle(),
(*payload).segments.as_ptr(),
total,
std::ptr::null(),
overlapped,
);
classify_issued(ok)
})
};
finish_scatter(submitted)
}
}
fn finish_scatter(submitted: Submitted<ScatterPayload>) -> io::Result<ScatterGatherIo> {
match submitted {
Submitted::Pending(id) => Ok(ScatterGatherIo { id }),
Submitted::Completed { .. } => Err(io::Error::other(
"scatter/gather 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 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;