use super::drivers::RangeBlockedGuard;
use crate::io_buffers::{IoBuffer, IoVector, IoVectorMut, IoVectorTrait};
use crate::Storage;
use std::ops::Range;
use std::{cmp, io};
use tracing::trace;
pub trait StorageExt: Storage {
#[allow(async_fn_in_trait)] async fn readv(&self, bufv: IoVectorMut<'_>, offset: u64) -> io::Result<()>;
#[allow(async_fn_in_trait)] async fn writev(&self, bufv: IoVector<'_>, offset: u64) -> io::Result<()>;
#[allow(async_fn_in_trait)] async fn read(&self, buf: impl Into<IoVectorMut<'_>>, offset: u64) -> io::Result<()>;
#[allow(async_fn_in_trait)] async fn write(&self, buf: impl Into<IoVector<'_>>, offset: u64) -> io::Result<()>;
#[allow(async_fn_in_trait)] async fn write_zeroes(&self, offset: u64, length: u64) -> io::Result<()>;
#[allow(async_fn_in_trait)] async fn write_allocated_zeroes(&self, offset: u64, length: u64) -> io::Result<()>;
#[allow(async_fn_in_trait)] async fn discard(&self, offset: u64, length: u64) -> io::Result<()>;
#[allow(async_fn_in_trait)] async fn weak_write_blocker(&self, range: Range<u64>) -> RangeBlockedGuard<'_>;
#[allow(async_fn_in_trait)] async fn strong_write_blocker(&self, range: Range<u64>) -> RangeBlockedGuard<'_>;
}
impl<S: Storage> StorageExt for S {
async fn readv(&self, mut bufv: IoVectorMut<'_>, offset: u64) -> io::Result<()> {
if bufv.is_empty() {
return Ok(());
}
let mem_align = self.mem_align();
let req_align = self.req_align();
if is_aligned(&bufv, offset, mem_align, req_align) {
return unsafe { self.pure_readv(bufv, offset) }.await;
}
trace!(
"Unaligned read: 0x{offset:x} + {} (size: {:#x})",
bufv.len(),
self.size().unwrap()
);
let req_align_mask = req_align as u64 - 1;
let len_align_mask = req_align_mask | (mem_align as u64 - 1);
debug_assert!((len_align_mask + 1).is_multiple_of(req_align as u64));
let unpadded_end = offset + bufv.len();
let padded_offset = offset & !req_align_mask;
let padded_end = (unpadded_end + req_align_mask) & !req_align_mask;
let padded_len = (padded_end - padded_offset + len_align_mask) & !(len_align_mask);
let padded_end = padded_offset + padded_len;
let padded_len: usize = (padded_end - padded_offset)
.try_into()
.map_err(|e| io::Error::other(format!("Cannot realign read: {e}")))?;
trace!("Padded read: {padded_offset:#x} + {padded_len}");
let mut bounce_buf = IoBuffer::new(padded_len, mem_align)?;
unsafe { self.pure_readv(bounce_buf.as_mut().into(), padded_offset) }.await?;
let in_buf_ofs = (offset - padded_offset) as usize;
let in_buf_end = (unpadded_end - padded_offset) as usize;
bufv.copy_from_slice(bounce_buf.as_ref_range(in_buf_ofs..in_buf_end).into_slice());
Ok(())
}
async fn writev(&self, bufv: IoVector<'_>, offset: u64) -> io::Result<()> {
if bufv.is_empty() {
return Ok(());
}
let mem_align = self.mem_align();
let req_align = self.req_align();
if is_aligned(&bufv, offset, mem_align, req_align) {
let _sw_guard = self.weak_write_blocker(offset..(offset + bufv.len())).await;
return unsafe { self.pure_writev(bufv, offset) }.await;
}
trace!(
"Unaligned write: {offset:#x} + {} (size: {:#x})",
bufv.len(),
self.size().unwrap()
);
let req_align_mask = req_align - 1;
let len_align_mask = req_align_mask | (mem_align - 1);
let len_align = req_align_mask + 1;
debug_assert!(len_align.is_multiple_of(req_align));
let unpadded_end = offset + bufv.len();
let padded_offset = offset & !(req_align_mask as u64);
let padded_end = (unpadded_end + req_align_mask as u64) & !(req_align_mask as u64);
let padded_len =
(padded_end - padded_offset + len_align_mask as u64) & !(len_align_mask as u64);
let padded_end = padded_offset + padded_len;
let padded_len: usize = (padded_end - padded_offset)
.try_into()
.map_err(|e| io::Error::other(format!("Cannot realign write: {e}")))?;
trace!("Padded write: {padded_offset:#x} + {padded_len}");
let mut bounce_buf = IoBuffer::new(padded_len, mem_align)?;
assert!(padded_len >= len_align && padded_len & len_align_mask == 0);
let _sw_guard = self.strong_write_blocker(padded_offset..padded_end).await;
let in_buf_ofs = (offset - padded_offset) as usize;
let in_buf_end = (unpadded_end - padded_offset) as usize;
let head_len = in_buf_ofs;
let aligned_head_len = (head_len + len_align_mask) & !len_align_mask;
let tail_len = padded_len - in_buf_end;
let aligned_tail_len = (tail_len + len_align_mask) & !len_align_mask;
if aligned_head_len + aligned_tail_len == padded_len {
unsafe { self.pure_readv(bounce_buf.as_mut().into(), padded_offset) }.await?;
} else {
if aligned_head_len > 0 {
let head_bufv = bounce_buf.as_mut_range(0..aligned_head_len).into();
unsafe { self.pure_readv(head_bufv, padded_offset) }.await?;
}
if aligned_tail_len > 0 {
let tail_start = padded_len - aligned_tail_len;
let tail_bufv = bounce_buf.as_mut_range(tail_start..padded_len).into();
unsafe { self.pure_readv(tail_bufv, padded_offset + tail_start as u64) }.await?;
}
}
bufv.copy_into_slice(bounce_buf.as_mut_range(in_buf_ofs..in_buf_end).into_slice());
unsafe { self.pure_writev(bounce_buf.as_ref().into(), padded_offset) }.await
}
async fn read(&self, buf: impl Into<IoVectorMut<'_>>, offset: u64) -> io::Result<()> {
self.readv(buf.into(), offset).await
}
async fn write(&self, buf: impl Into<IoVector<'_>>, offset: u64) -> io::Result<()> {
self.writev(buf.into(), offset).await
}
async fn write_zeroes(&self, offset: u64, length: u64) -> io::Result<()> {
write_efficient_zeroes(self, offset, length, false).await
}
async fn write_allocated_zeroes(&self, offset: u64, length: u64) -> io::Result<()> {
write_efficient_zeroes(self, offset, length, true).await
}
async fn discard(&self, offset: u64, length: u64) -> io::Result<()> {
let discard_align = self.discard_align();
debug_assert!(discard_align.is_power_of_two());
let align_mask = discard_align as u64 - 1;
let unaligned_end = offset
.checked_add(length)
.ok_or_else(|| io::Error::other("Discard wrap-around"))?;
let aligned_offset = (offset + align_mask) & !align_mask;
let aligned_end = unaligned_end & !align_mask;
if aligned_end > aligned_offset {
let _sw_guard = self.weak_write_blocker(aligned_offset..aligned_end).await;
let aligned_len = aligned_end - aligned_offset;
if let Err(err) = unsafe { self.pure_discard(aligned_offset, aligned_len) }.await {
if err.kind() != io::ErrorKind::Unsupported {
return Err(err);
}
}
}
Ok(())
}
async fn weak_write_blocker(&self, range: Range<u64>) -> RangeBlockedGuard<'_> {
self.get_storage_helper().weak_write_blocker(range).await
}
async fn strong_write_blocker(&self, range: Range<u64>) -> RangeBlockedGuard<'_> {
self.get_storage_helper().strong_write_blocker(range).await
}
}
fn is_aligned<V: IoVectorTrait>(bufv: &V, offset: u64, mem_align: usize, req_align: usize) -> bool {
debug_assert!(mem_align.is_power_of_two() && req_align.is_power_of_two());
let req_align_mask = req_align as u64 - 1;
if offset & req_align_mask != 0 {
false
} else if bufv.len() & req_align_mask == 0 {
bufv.is_aligned(mem_align, req_align)
} else {
false
}
}
pub(crate) async fn write_full_zeroes<S: StorageExt>(
storage: S,
mut offset: u64,
mut length: u64,
) -> io::Result<()> {
let buflen = cmp::min(length, 1048576) as usize;
let mut buf = IoBuffer::new(buflen, storage.mem_align())?;
buf.as_mut().into_slice().fill(0);
let req_align = storage.req_align();
let req_align_mask = (req_align - 1) as u64;
while length > 0 {
let mut chunk_length = cmp::min(length, 1048576) as usize;
if offset & req_align_mask != 0 {
chunk_length = cmp::min(chunk_length, req_align - (offset & req_align_mask) as usize);
}
storage
.write(buf.as_ref_range(0..chunk_length), offset)
.await?;
offset += chunk_length as u64;
length -= chunk_length as u64;
}
Ok(())
}
pub(crate) async fn write_efficient_zeroes<S: StorageExt>(
storage: S,
offset: u64,
length: u64,
allocate: bool,
) -> io::Result<()> {
let zero_align = storage.zero_align();
debug_assert!(zero_align.is_power_of_two());
let align_mask = zero_align as u64 - 1;
let unaligned_end = offset
.checked_add(length)
.ok_or_else(|| io::Error::other("Zero-write wrap-around"))?;
let aligned_offset = (offset + align_mask) & !align_mask;
let aligned_end = unaligned_end & !align_mask;
if aligned_end > aligned_offset {
let result = {
let _sw_guard = storage
.weak_write_blocker(aligned_offset..aligned_end)
.await;
if allocate {
unsafe {
storage
.pure_write_allocated_zeroes(aligned_offset, aligned_end - aligned_offset)
}
.await
} else {
unsafe { storage.pure_write_zeroes(aligned_offset, aligned_end - aligned_offset) }
.await
}
};
if let Err(err) = result {
return if err.kind() == io::ErrorKind::Unsupported {
write_full_zeroes(storage, offset, length).await
} else {
Err(err)
};
}
}
let zero_buf = if aligned_offset > offset || aligned_end < unaligned_end {
let mut buf = IoBuffer::new(
cmp::max(aligned_offset - offset, unaligned_end - aligned_end) as usize,
storage.mem_align(),
)?;
buf.as_mut().into_slice().fill(0);
Some(buf)
} else {
None
};
if aligned_offset > offset {
let buf = zero_buf
.as_ref()
.unwrap()
.as_ref_range(0..((aligned_offset - offset) as usize));
storage.write(buf, offset).await?;
}
if aligned_end < unaligned_end {
let buf = zero_buf
.as_ref()
.unwrap()
.as_ref_range(0..((unaligned_end - aligned_end) as usize));
storage.write(buf, aligned_end).await?;
}
Ok(())
}