use std::{
ffi::{c_void, CString},
num::NonZeroUsize,
os::fd::{AsRawFd, FromRawFd, OwnedFd},
ptr::NonNull,
};
use nix::{
errno::Errno,
fcntl::{open, OFlag},
libc::{c_char, munmap, shm_unlink, unlink, S_IRUSR, S_IWUSR},
sys::{
mman::{shm_open, MapFlags, ProtFlags},
stat::Mode,
},
unistd::ftruncate,
};
use crate::page_size::PageSize;
#[derive(Clone)]
pub struct Shmem {
pub id: String,
pub(crate) size: i64,
pub fd: i32,
addr: NonNull<()>,
page_size: PageSize,
}
impl Shmem {
pub fn open_or_create(
id: &str,
uplined_size: i64,
#[cfg(target_os = "linux")] page_size: PageSize,
) -> Result<Shmem, ShmemError> {
#[cfg(not(target_os = "linux"))]
let page_size = PageSize::Standard;
let mode = Mode::from_bits(S_IRUSR | S_IWUSR).unwrap();
let fd = if page_size.is_gigantic() {
let path = format!("/mnt/gigantic/{}", id);
let fd = open::<str>(
&path,
OFlag::O_RDWR | OFlag::O_CREAT,
mode,
)?;
unsafe { OwnedFd::from_raw_fd(fd) }
} else if page_size.is_huge() {
let path = format!("/mnt/hugepages/{}", id);
let fd = open::<str>(
&path,
OFlag::O_RDWR | OFlag::O_CREAT,
mode,
)?;
unsafe { OwnedFd::from_raw_fd(fd) }
} else {
let path = CString::new(id).unwrap();
let fd = shm_open(
path.as_c_str(),
OFlag::O_RDWR | OFlag::O_CREAT,
mode,
)?;
let stat = nix::sys::stat::fstat(fd.as_raw_fd())?;
if stat.st_size != uplined_size {
ftruncate(&fd, uplined_size)?;
}
fd
};
#[cfg_attr(not(target_os = "linux"), allow(unused_mut))]
let mut map_flags = MapFlags::MAP_SHARED;
#[cfg(target_os = "linux")]
if page_size.is_huge() {
map_flags |= MapFlags::MAP_HUGETLB;
}
#[cfg(target_os = "linux")]
if page_size.is_gigantic() {
map_flags |= MapFlags::MAP_HUGETLB;
map_flags |= MapFlags::MAP_HUGE_1GB;
}
let addr = unsafe {
nix::sys::mman::mmap(
None,
NonZeroUsize::new_unchecked(uplined_size as usize),
ProtFlags::PROT_READ | ProtFlags::PROT_WRITE,
map_flags,
&fd,
0,
)
};
let addr = match addr {
Ok(addr) => addr.cast(),
Err(e) => return Err(ShmemError::Errno(e as i32)),
};
Ok(Shmem {
id: id.to_string(),
size: uplined_size,
fd: fd.as_raw_fd(),
addr,
page_size,
})
}
pub fn close(self) -> Result<(), ShmemError> {
println!("Unmapping shared memory: {}", self.id);
let res = unsafe {
munmap(
self.addr.as_ptr() as *mut c_void,
self.size as usize,
)
};
if res != 0 {
let err = std::io::Error::last_os_error();
println!("failed to unmap shared memory from the virtual memory space: {res} -> {err:?}")
}
println!("Closing shared memory: {}", self.id);
if self.page_size.is_gigantic() {
let path = format!("/mnt/gigantic/{}", self.id);
let c_path = CString::new(path).unwrap();
if unsafe { unlink(c_path.as_ptr()) } != 0 {
return Err(ShmemError::UnlinkError);
}
} else if self.page_size.is_huge() {
let path = format!("/mnt/hugepages/{}", self.id);
let c_path = CString::new(path).unwrap();
if unsafe { unlink(c_path.as_ptr()) } != 0 {
return Err(ShmemError::UnlinkError);
}
} else {
let storage_id: *const c_char =
self.id.as_bytes().as_ptr() as *const c_char;
if unsafe { shm_unlink(storage_id) } != 0 {
println!("failed to reclaim shared memory")
}
}
println!("fd closed: {}", self.id);
Ok(())
}
pub fn get_mut_ptr(&self) -> *mut u8 {
self.addr.as_ptr() as *mut u8
}
}
pub fn cleanup_shmem(
id: &str,
size: i64,
#[cfg(target_os = "linux")] page_size: PageSize,
) -> Result<(), ShmemError> {
let shmem = Shmem::open_or_create(
id,
size,
#[cfg(target_os = "linux")]
page_size,
)?;
shmem.close()
}
#[derive(Clone, Copy)]
pub enum ShmemError {
BadFileDescriptor,
AllocationFailedErr,
InvalidPermissions,
UnlinkError,
Errno(i32),
}
impl std::error::Error for ShmemError {}
impl core::fmt::Debug for ShmemError {
fn fmt(&self, f: &mut core::fmt::Formatter) -> core::fmt::Result {
match self {
ShmemError::BadFileDescriptor => {
f.write_str("Bad file descriptor for shared memory object")
}
ShmemError::AllocationFailedErr => {
f.write_str("Failed to allocate shared memory. If using huge/gigantic pages, have you preallocated pages?")
}
ShmemError::InvalidPermissions => {
f.write_str("Invalid permissions for opening/mmaping shared memory object")
}
ShmemError::UnlinkError => {
f.write_str("Failed to unlink huge page")
}
ShmemError::Errno(e) => write!(f, "Other system error: {}", e),
}
}
}
impl core::fmt::Display for ShmemError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{:?}", self)
}
}
impl From<Errno> for ShmemError {
fn from(e: Errno) -> Self {
match e {
Errno::EBADF => ShmemError::BadFileDescriptor,
Errno::ENOMEM => ShmemError::AllocationFailedErr,
Errno::EACCES => ShmemError::InvalidPermissions,
e => ShmemError::Errno(e as i32),
}
}
}