#![allow(unsafe_code)]
use std::os::unix::io::AsRawFd;
use std::sync::Arc;
use io_uring::opcode;
use io_uring::types::Fd;
use crate::Result;
use crate::errors::PagedbError;
use crate::vfs::blocking::offload;
use crate::vfs::iouring::ring::{SharedRing, SubmitError};
use crate::vfs::traits::{
VfsFile, checked_iouring_positioned_offset, checked_read_count, checked_signed_file_len,
write_all_at,
};
use crate::vfs::types::{ReadReq, WriteReq};
pub struct IouringFile {
file: Arc<std::fs::File>,
writable: bool,
ring: Arc<SharedRing>,
}
impl IouringFile {
pub(crate) fn new(file: std::fs::File, writable: bool, ring: Arc<SharedRing>) -> Self {
Self {
file: Arc::new(file),
writable,
ring,
}
}
fn shared(&self) -> (Arc<std::fs::File>, Arc<SharedRing>) {
(Arc::clone(&self.file), Arc::clone(&self.ring))
}
fn check_write_range(offset: u64, len: usize) -> Result<()> {
let len = u64::try_from(len).map_err(|_| {
PagedbError::Io(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"buffer length does not fit in u64",
))
})?;
offset.checked_add(len).ok_or_else(|| {
PagedbError::Io(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"write offset overflow",
))
})?;
Ok(())
}
fn entry_len(len: usize) -> Result<u32> {
u32::try_from(len)
.map_err(|_| PagedbError::Io(std::io::Error::other("buffer too large for u32")))
}
fn run_cycle<B>(
ring: &SharedRing,
entries: &[io_uring::squeue::Entry],
buffers: B,
) -> Result<(Vec<i32>, B)> {
match unsafe { ring.submit_and_collect(entries) } {
Ok(results) => Ok((results, buffers)),
Err(SubmitError::Settled(error)) => Err(PagedbError::Io(error)),
Err(abandoned @ SubmitError::Abandoned(_)) => {
std::mem::forget(buffers);
Err(PagedbError::Io(abandoned.into_io()))
}
}
}
}
impl VfsFile for IouringFile {
async fn read_at(&self, offset: u64, buf: &mut [u8]) -> Result<usize> {
if buf.is_empty() {
return Ok(0);
}
let len = buf.len();
checked_iouring_positioned_offset(offset, len)?;
Self::entry_len(len)?;
let (file, ring) = self.shared();
let (scratch, read) = offload(move || {
let fd = Fd(file.as_raw_fd());
let entry_len = IouringFile::entry_len(len)?;
let mut scratch = vec![0u8; len];
let entry = opcode::Read::new(fd, scratch.as_mut_ptr(), entry_len)
.offset(offset)
.build()
.user_data(0);
let (results, scratch) = IouringFile::run_cycle(&ring, &[entry], scratch)?;
let n = single_result(&results)?;
#[allow(clippy::cast_sign_loss)]
let read = checked_read_count(n as usize, len)?;
Ok((scratch, read))
})
.await?;
buf[..read].copy_from_slice(&scratch[..read]);
Ok(read)
}
async fn read_at_vectored(&self, reqs: &mut [ReadReq<'_>]) -> Result<()> {
if reqs.is_empty() {
return Ok(());
}
let mut plan: Vec<(u64, usize)> = Vec::with_capacity(reqs.len());
for req in reqs.iter() {
checked_iouring_positioned_offset(req.offset, req.buf.len())?;
Self::entry_len(req.buf.len())?;
plan.push((req.offset, req.buf.len()));
}
let (file, ring) = self.shared();
let completed = offload(move || {
let fd = Fd(file.as_raw_fd());
let mut buffers: Vec<Vec<u8>> = plan.iter().map(|&(_, len)| vec![0u8; len]).collect();
let mut entries: Vec<io_uring::squeue::Entry> = Vec::with_capacity(plan.len());
for (i, ((offset, len), scratch)) in plan.iter().zip(buffers.iter_mut()).enumerate() {
let entry_len = IouringFile::entry_len(*len)?;
entries.push(
opcode::Read::new(fd, scratch.as_mut_ptr(), entry_len)
.offset(*offset)
.build()
.user_data(i as u64),
);
}
let (results, buffers) = IouringFile::run_cycle(&ring, &entries, buffers)?;
drop(entries);
let mut out: Vec<(Vec<u8>, usize)> = Vec::with_capacity(plan.len());
for ((_, len), (scratch, res)) in plan.iter().zip(buffers.into_iter().zip(results)) {
if res < 0 {
return Err(PagedbError::Io(std::io::Error::from_raw_os_error(-res)));
}
#[allow(clippy::cast_sign_loss)]
let nread = checked_read_count(res as usize, *len)?;
out.push((scratch, nread));
}
Ok(out)
})
.await?;
for (req, (scratch, nread)) in reqs.iter_mut().zip(completed) {
req.buf[..nread].copy_from_slice(&scratch[..nread]);
for b in &mut req.buf[nread..] {
*b = 0;
}
}
Ok(())
}
async fn write_at(&mut self, offset: u64, buf: &[u8]) -> Result<usize> {
if !self.writable {
return Err(PagedbError::ReadOnly);
}
if buf.is_empty() {
return Ok(0);
}
Self::check_write_range(offset, buf.len())?;
Self::entry_len(buf.len())?;
let (file, ring) = self.shared();
let data = buf.to_vec();
offload(move || {
let fd = Fd(file.as_raw_fd());
let entry_len = IouringFile::entry_len(data.len())?;
let entry = opcode::Write::new(fd, data.as_ptr(), entry_len)
.offset(offset)
.build()
.user_data(0);
let (results, data) = IouringFile::run_cycle(&ring, &[entry], data)?;
let n = single_result(&results)?;
drop(data);
#[allow(clippy::cast_sign_loss)]
let written = n as usize;
Ok(written)
})
.await
}
async fn write_at_vectored(&mut self, reqs: &[WriteReq<'_>]) -> Result<()> {
if !self.writable {
return Err(PagedbError::ReadOnly);
}
if reqs.is_empty() {
return Ok(());
}
for req in reqs {
Self::check_write_range(req.offset, req.buf.len())?;
}
let mut plan: Vec<(usize, u64, Vec<u8>)> = Vec::with_capacity(reqs.len());
for (i, req) in reqs.iter().enumerate() {
if req.buf.is_empty() {
continue;
}
Self::entry_len(req.buf.len())?;
plan.push((i, req.offset, req.buf.to_vec()));
}
if plan.is_empty() {
return Ok(());
}
let (file, ring) = self.shared();
let short_writes = offload(move || {
let fd = Fd(file.as_raw_fd());
let mut entries: Vec<io_uring::squeue::Entry> = Vec::with_capacity(plan.len());
for (slot, (_, offset, data)) in plan.iter().enumerate() {
let entry_len = IouringFile::entry_len(data.len())?;
let user_data = u64::try_from(slot).map_err(|_| {
PagedbError::Io(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"too many vectored write requests",
))
})?;
entries.push(
opcode::Write::new(fd, data.as_ptr(), entry_len)
.offset(*offset)
.build()
.user_data(user_data),
);
}
let (results, plan) = IouringFile::run_cycle(&ring, &entries, plan)?;
drop(entries);
let mut short_writes = Vec::new();
for (slot, &res) in results.iter().enumerate() {
if res < 0 {
return Err(PagedbError::Io(std::io::Error::from_raw_os_error(-res)));
}
let written = usize::try_from(res)
.map_err(|_| PagedbError::Io(std::io::Error::other("negative write result")))?;
let (request_index, _, data) = &plan[slot];
if written > data.len() {
return Err(PagedbError::Io(std::io::Error::other(
"io_uring write overreported bytes",
)));
}
if written == 0 {
return Err(PagedbError::Io(std::io::Error::from(
std::io::ErrorKind::WriteZero,
)));
}
if written < data.len() {
short_writes.push((*request_index, written));
}
}
Ok(short_writes)
})
.await?;
for (request_index, written) in short_writes {
let request = &reqs[request_index];
let written_u64 = u64::try_from(written).map_err(|_| {
PagedbError::Io(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"write count does not fit in u64",
))
})?;
let offset = request.offset.checked_add(written_u64).ok_or_else(|| {
PagedbError::Io(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"write offset overflow",
))
})?;
write_all_at(self, offset, &request.buf[written..]).await?;
}
Ok(())
}
async fn sync(&mut self) -> Result<()> {
let (file, ring) = self.shared();
offload(move || {
let fd = Fd(file.as_raw_fd());
let entry = opcode::Fsync::new(fd).build().user_data(0);
let (results, ()) = IouringFile::run_cycle(&ring, &[entry], ())?;
single_result(&results)?;
Ok(())
})
.await
}
async fn truncate(&mut self, len: u64) -> Result<()> {
if !self.writable {
return Err(PagedbError::ReadOnly);
}
let len = checked_signed_file_len(len, "ftruncate")?;
let file = Arc::clone(&self.file);
offload(move || {
let rc = unsafe { libc::ftruncate(file.as_raw_fd(), len) };
if rc != 0 {
return Err(PagedbError::Io(std::io::Error::last_os_error()));
}
Ok(())
})
.await
}
async fn len(&self) -> Result<u64> {
let file = Arc::clone(&self.file);
offload(move || Ok(file.metadata().map_err(PagedbError::Io)?.len())).await
}
async fn is_empty(&self) -> Result<bool> {
Ok(self.len().await? == 0)
}
fn supports_direct_io(&self) -> bool {
true
}
}
fn single_result(results: &[i32]) -> Result<i32> {
let result = *results
.first()
.ok_or_else(|| PagedbError::Io(std::io::Error::other("io_uring: no CQE for entry")))?;
if result < 0 {
return Err(PagedbError::Io(std::io::Error::from_raw_os_error(-result)));
}
Ok(result)
}