#![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";
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ShmValidation {
Missing,
Valid { actual_size: usize },
}
pub struct ShmRegion {
name: CString,
lock_path: PathBuf,
process_lock: bool,
ptr: NonNull<u8>,
len: usize,
created: bool,
}
impl ShmRegion {
pub fn validate_existing(name: &str, minimum_size: usize) -> io::Result<ShmValidation> {
let cname = CString::new(name)
.map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "shm name has nul byte"))?;
let raw_fd = loop {
let fd = unsafe { libc::shm_open(cname.as_ptr(), libc::O_RDWR, 0o600) };
if fd >= 0 {
break fd;
}
let error = io::Error::last_os_error();
if error.kind() == io::ErrorKind::Interrupted {
continue;
}
if error.raw_os_error() == Some(libc::ENOENT) {
return Ok(ShmValidation::Missing);
}
return Err(error);
};
let fd = unsafe { OwnedFd::from_raw_fd(raw_fd) };
let actual_size = shm_object_size(&fd, name)?;
validate_minimum_size(name, actual_size, minimum_size)?;
Ok(ShmValidation::Valid { actual_size })
}
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 actual_size = match shm_object_size(&fd, name) {
Ok(actual_size) => actual_size,
Err(error) => {
if created {
let _ = unsafe { libc::shm_unlink(cname.as_ptr()) };
}
return Err(error);
}
};
if let Err(error) = validate_minimum_size(name, actual_size, size) {
if created {
let _ = unsafe { libc::shm_unlink(cname.as_ptr()) };
}
return Err(error);
}
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(())
}
}
fn shm_object_size(fd: &OwnedFd, name: &str) -> io::Result<usize> {
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 {
return Err(io::Error::last_os_error());
}
let actual_size = unsafe { stat.assume_init() }.st_size;
usize::try_from(actual_size).map_err(|_| {
io::Error::new(
io::ErrorKind::InvalidData,
format!("SHM segment {name} reported invalid size {actual_size}"),
)
})
}
fn validate_minimum_size(name: &str, actual_size: usize, minimum_size: usize) -> io::Result<()> {
if actual_size < minimum_size {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!(
"SHM segment {name} size {actual_size} is smaller than requested mapping {minimum_size}"
),
));
}
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> {
use std::os::unix::fs::MetadataExt;
if let Some(dir) = lock_path.parent() {
ensure_lock_dir(dir)?;
}
let file = OpenOptions::new()
.read(true)
.write(true)
.create(true)
.mode(0o600)
.custom_flags(libc::O_CLOEXEC | libc::O_NOFOLLOW)
.open(lock_path)?;
let uid = unsafe { libc::geteuid() };
let owner = file.metadata()?.uid();
if owner != uid {
return Err(io::Error::new(
io::ErrorKind::PermissionDenied,
format!(
"{} is owned by uid {owner} rather than {uid}",
lock_path.display()
),
));
}
Ok(file.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 {
lock_dir().join(format!("{}.lock", shm_name.trim_start_matches('/')))
}
fn lock_dir() -> PathBuf {
let uid = unsafe { libc::geteuid() };
let base = std::env::var_os("XDG_RUNTIME_DIR")
.map(PathBuf::from)
.filter(|dir| dir.is_absolute())
.unwrap_or_else(|| PathBuf::from("/tmp"));
base.join(format!("{SHM_NAMESPACE}-{uid}"))
}
fn ensure_lock_dir(dir: &Path) -> io::Result<()> {
use std::os::unix::fs::DirBuilderExt;
let uid = unsafe { libc::geteuid() };
match std::fs::DirBuilder::new().mode(0o700).create(dir) {
Ok(()) => {}
Err(error) if error.kind() == io::ErrorKind::AlreadyExists => {}
Err(error) => return Err(error),
}
ensure_private_dir(dir, uid)
}
fn ensure_private_dir(dir: &Path, uid: u32) -> io::Result<()> {
use std::os::unix::fs::MetadataExt;
use std::os::unix::fs::PermissionsExt;
let metadata = std::fs::symlink_metadata(dir)?;
if !metadata.is_dir() {
return Err(io::Error::new(
io::ErrorKind::PermissionDenied,
format!("{} is not a directory", dir.display()),
));
}
if metadata.uid() != uid {
return Err(io::Error::new(
io::ErrorKind::PermissionDenied,
format!(
"{} is owned by uid {} rather than {uid}; refusing to lock in a directory \
another user controls",
dir.display(),
metadata.uid()
),
));
}
if metadata.permissions().mode() & 0o077 != 0 {
return Err(io::Error::new(
io::ErrorKind::PermissionDenied,
format!(
"{} is mode {:o}; refusing to lock in a directory others can write to",
dir.display(),
metadata.permissions().mode() & 0o777
),
));
}
Ok(())
}