#![cfg(unix)]
use std::ffi::CString;
use std::fs::OpenOptions;
use std::io;
use std::os::fd::{AsRawFd, FromRawFd, OwnedFd};
use std::os::unix::fs::OpenOptionsExt;
use std::path::{Path, PathBuf};
use std::ptr::NonNull;
pub const SHM_NAMESPACE: &str = "orbit";
pub struct ShmRegion {
name: CString,
lock_path: PathBuf,
process_lock: bool,
ptr: NonNull<u8>,
len: usize,
created: bool,
}
impl ShmRegion {
pub fn open_or_create(name: &str, size: usize) -> io::Result<Self> {
let (region, initialization_lock) = Self::open_or_create_inner(name, size, false)?;
debug_assert!(initialization_lock.is_none());
Ok(region)
}
pub fn open_or_create_locked(name: &str, size: usize) -> io::Result<(Self, ShmRegionLock)> {
let (region, initialization_lock) = Self::open_or_create_inner(name, size, true)?;
Ok((
region,
initialization_lock.expect("locked SHM open must return its initialization lock"),
))
}
fn open_or_create_inner(
name: &str,
size: usize,
process_lock: bool,
) -> io::Result<(Self, Option<ShmRegionLock>)> {
let cname = CString::new(name)
.map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "shm name has nul byte"))?;
let lock_path = lock_file_path(name);
let initialization_lock = if process_lock {
Some(lock_path_exclusive(&lock_path)?)
} else {
None
};
let (raw_fd, created) = unsafe {
let fd = libc::shm_open(
cname.as_ptr(),
libc::O_RDWR | libc::O_CREAT | libc::O_EXCL,
0o600,
);
if fd >= 0 {
(fd, true)
} else {
let err = io::Error::last_os_error();
if err.raw_os_error() != Some(libc::EEXIST) {
return Err(err);
}
let fd = libc::shm_open(cname.as_ptr(), libc::O_RDWR, 0o600);
if fd < 0 {
return Err(io::Error::last_os_error());
}
(fd, false)
}
};
let fd = unsafe { OwnedFd::from_raw_fd(raw_fd) };
if created {
let rc = unsafe { libc::ftruncate(fd.as_raw_fd(), size as libc::off_t) };
if rc != 0 {
let err = io::Error::last_os_error();
let _ = unsafe { libc::shm_unlink(cname.as_ptr()) };
return Err(err);
}
}
let mut stat = std::mem::MaybeUninit::<libc::stat>::uninit();
let stat_rc = unsafe { libc::fstat(fd.as_raw_fd(), stat.as_mut_ptr()) };
if stat_rc != 0 {
let err = io::Error::last_os_error();
if created {
let _ = unsafe { libc::shm_unlink(cname.as_ptr()) };
}
return Err(err);
}
let actual_size = unsafe { stat.assume_init() }.st_size;
if actual_size < 0 || (actual_size as usize) < size {
if created {
let _ = unsafe { libc::shm_unlink(cname.as_ptr()) };
}
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!(
"SHM segment {name} size {actual_size} is smaller than requested mapping {size}"
),
));
}
let ptr = unsafe {
libc::mmap(
std::ptr::null_mut(),
size,
libc::PROT_READ | libc::PROT_WRITE,
libc::MAP_SHARED,
fd.as_raw_fd(),
0,
)
};
if ptr == libc::MAP_FAILED {
let err = io::Error::last_os_error();
if created {
let _ = unsafe { libc::shm_unlink(cname.as_ptr()) };
}
return Err(err);
}
let ptr = NonNull::new(ptr.cast::<u8>()).expect("mmap returned non-null on success");
Ok((
Self {
name: cname,
lock_path,
process_lock,
ptr,
len: size,
created,
},
initialization_lock,
))
}
pub fn as_ptr(&self) -> *mut u8 {
self.ptr.as_ptr()
}
pub fn len(&self) -> usize {
self.len
}
pub fn is_empty(&self) -> bool {
self.len == 0
}
pub fn created(&self) -> bool {
self.created
}
pub fn lock_exclusive(&self) -> io::Result<ShmRegionLock> {
if !self.process_lock {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"SHM region was opened without a process lock",
));
}
lock_path_exclusive(&self.lock_path)
}
#[cfg(test)]
pub(crate) fn try_lock_exclusive(&self) -> io::Result<ShmRegionLock> {
if !self.process_lock {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"SHM region was opened without a process lock",
));
}
let lock_fd = open_lock_file(&self.lock_path)?;
let rc = unsafe { libc::flock(lock_fd.as_raw_fd(), libc::LOCK_EX | libc::LOCK_NB) };
if rc == 0 {
Ok(ShmRegionLock { lock_fd })
} else {
Err(io::Error::last_os_error())
}
}
pub fn unlink(&self) -> io::Result<()> {
let rc = unsafe { libc::shm_unlink(self.name.as_ptr()) };
let shm_error = if rc != 0 {
let err = io::Error::last_os_error();
if err.raw_os_error() == Some(libc::ENOENT) {
None
} else {
Some(err)
}
} else {
None
};
let lock_error = match std::fs::remove_file(&self.lock_path) {
Ok(()) => None,
Err(error) if error.kind() == io::ErrorKind::NotFound => None,
Err(error) => Some(error),
};
if let Some(error) = shm_error.or(lock_error) {
return Err(error);
}
Ok(())
}
}
pub struct ShmRegionLock {
lock_fd: OwnedFd,
}
impl Drop for ShmRegionLock {
fn drop(&mut self) {
let _ = unsafe { libc::flock(self.lock_fd.as_raw_fd(), libc::LOCK_UN) };
}
}
fn lock_path_exclusive(lock_path: &Path) -> io::Result<ShmRegionLock> {
lock_fd_exclusive(open_lock_file(lock_path)?)
}
fn open_lock_file(lock_path: &Path) -> io::Result<OwnedFd> {
OpenOptions::new()
.read(true)
.write(true)
.create(true)
.mode(0o600)
.custom_flags(libc::O_CLOEXEC | libc::O_NOFOLLOW)
.open(lock_path)
.map(Into::into)
}
fn lock_fd_exclusive(lock_fd: OwnedFd) -> io::Result<ShmRegionLock> {
loop {
let rc = unsafe { libc::flock(lock_fd.as_raw_fd(), libc::LOCK_EX) };
if rc == 0 {
return Ok(ShmRegionLock { lock_fd });
}
let error = io::Error::last_os_error();
if error.kind() != io::ErrorKind::Interrupted {
return Err(error);
}
}
}
impl Drop for ShmRegion {
fn drop(&mut self) {
unsafe {
libc::munmap(self.ptr.as_ptr().cast(), self.len);
}
}
}
unsafe impl Send for ShmRegion {}
unsafe impl Sync for ShmRegion {}
pub fn ring_segment_name(fleet_name: &str, kind: u8) -> String {
let uid = unsafe { libc::geteuid() };
ring_segment_name_for_uid(fleet_name, kind, uid)
}
pub fn ring_segment_name_for_uid(fleet_name: &str, kind: u8, uid: u32) -> String {
format!("/{SHM_NAMESPACE}-{fleet_name}-{kind}-{uid}")
}
fn lock_file_path(shm_name: &str) -> PathBuf {
PathBuf::from("/tmp").join(format!("{}.lock", shm_name.trim_start_matches('/')))
}