use std::{
ffi::{CStr, CString},
io,
marker::PhantomData,
};
use memmap2::{MmapMut, MmapOptions};
use rustix::{
fd::{AsFd, AsRawFd, BorrowedFd, OwnedFd},
fs::Mode,
io::Errno,
shm::OFlags,
};
#[derive(Debug)]
pub struct Shm {
fd: OwnedFd,
name: CString,
}
fn to_io_err(e: Errno) -> io::Error {
io::Error::from_raw_os_error(e.raw_os_error())
}
impl Shm {
pub fn open(name: &str, oflags: OFlags, mode: Mode) -> io::Result<Self> {
let cstr = CString::new(name).unwrap();
let fd = rustix::shm::open(&*cstr, oflags, mode).map_err(to_io_err)?;
Ok(Self { fd, name: cstr })
}
pub fn size(&self) -> io::Result<usize> {
rustix::fs::fstat(&self.fd)
.map(|stat| stat.st_size as usize)
.map_err(to_io_err)
}
pub fn set_size(&mut self, size: usize) -> io::Result<()> {
rustix::fs::ftruncate(&self.fd, u64::try_from(size).unwrap_or(u64::MAX)).map_err(to_io_err)
}
pub fn name(&self) -> &str {
let bytes = self.name.as_bytes();
unsafe { std::str::from_utf8_unchecked(bytes) }
}
pub unsafe fn map(&mut self, offset: usize) -> io::Result<BorrowedMap<'_>> {
let size = self.size()?;
if offset >= size {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"The provided offset must not be greater than self.size()",
));
}
let mut opts = MmapOptions::new();
opts.offset(offset as u64);
opts.len(size - offset);
let map = unsafe { opts.map_mut(self.fd.as_raw_fd()) }?;
Ok(BorrowedMap {
map,
_borrowed: PhantomData,
})
}
pub fn as_fd(&self) -> BorrowedFd<'_> {
self.fd.as_fd()
}
pub fn unlink(self) -> io::Result<()> {
rustix::shm::unlink(&*self.name).map_err(to_io_err)
}
pub fn name_ptr(&self) -> &CStr {
&self.name
}
}
#[derive(Debug)]
pub struct UnlinkOnDrop {
pub shm: Shm,
}
impl Drop for UnlinkOnDrop {
fn drop(&mut self) {
_ = rustix::shm::unlink(self.shm.name_ptr())
}
}
#[derive(Debug)]
pub struct BorrowedMap<'shm> {
map: MmapMut,
_borrowed: PhantomData<&'shm ()>,
}
impl BorrowedMap<'_> {
pub fn map(&mut self) -> &mut MmapMut {
&mut self.map
}
pub unsafe fn into_map(self) -> MmapMut {
self.map
}
}
#[cfg(test)]
mod test {
use super::*;
#[test]
fn offset_larger_than_map_fails() {
let mut shm = Shm::open(
"/__psx_shm_oltmp_ahafeufhdmdhkeysmash",
OFlags::RDWR | OFlags::CREATE,
Mode::RUSR | Mode::WUSR,
)
.unwrap();
shm.set_size(20).unwrap();
let err = unsafe { shm.map(21) }.unwrap_err();
assert_eq!(err.kind(), std::io::ErrorKind::InvalidInput);
}
}