use super::layout;
use super::LockError;
use core::alloc::Layout;
use core::ptr::NonNull;
use core::sync::atomic::{AtomicI32, AtomicU32, AtomicUsize, Ordering};
use std::io;
use std::sync::Once;
pub(super) struct Region {
base: NonNull<u8>,
interior_len: usize,
generation: u32,
}
static FORK_GENERATION: AtomicU32 = AtomicU32::new(0);
extern "C" fn note_fork_in_child() {
let _ = FORK_GENERATION.fetch_add(1, Ordering::AcqRel);
}
static ATFORK_RC: AtomicI32 = AtomicI32::new(0);
fn fork_generation() -> u32 {
static REGISTERED: Once = Once::new();
REGISTERED.call_once(|| {
#[cfg(not(miri))]
unsafe {
let rc = libc::pthread_atfork(None, None, Some(note_fork_in_child));
ATFORK_RC.store(rc, Ordering::Release);
}
});
FORK_GENERATION.load(Ordering::Acquire)
}
#[cfg_attr(test, mutants::skip)]
fn fork_tracking() -> Option<LockError> {
let _ = fork_generation();
untracked(ATFORK_RC.load(Ordering::Acquire))
}
fn untracked(rc: i32) -> Option<LockError> {
(rc != 0).then(|| LockError::Untracked {
source: io::Error::from_raw_os_error(rc),
})
}
fn page_size() -> usize {
static PAGE: AtomicUsize = AtomicUsize::new(0);
match PAGE.load(Ordering::Relaxed) {
0 => {
let raw = unsafe { libc::sysconf(libc::_SC_PAGESIZE) };
let page = usize::try_from(raw)
.ok()
.filter(|p| p.is_power_of_two())
.unwrap_or(4096);
PAGE.store(page, Ordering::Relaxed);
page
}
page => page,
}
}
#[cfg_attr(test, mutants::skip)]
mod flags {
pub(super) const INTERIOR_PROT: libc::c_int = libc::PROT_READ | libc::PROT_WRITE;
pub(super) const MAP_FLAGS: libc::c_int = libc::MAP_PRIVATE | libc::MAP_ANONYMOUS;
pub(super) const NO_FILE: libc::c_int = -1;
}
use flags::{INTERIOR_PROT, MAP_FLAGS, NO_FILE};
pub(super) fn memlock_limit() -> Option<u64> {
#[cfg(not(memlock_limit))]
{
None
}
#[cfg(memlock_limit)]
{
let mut lim = libc::rlimit {
rlim_cur: 0,
rlim_max: 0,
};
let rc = unsafe { libc::getrlimit(libc::RLIMIT_MEMLOCK, &mut lim) };
if rc != 0 {
return None;
}
if lim.rlim_cur == libc::RLIM_INFINITY {
return None;
}
#[allow(clippy::useless_conversion)]
u64::try_from(lim.rlim_cur).ok()
}
}
impl Region {
pub(super) fn allocate(layout: Layout) -> Result<(Self, Option<LockError>), LockError> {
let page = page_size();
let align = layout.align();
if align > page {
return Err(LockError::Alignment { align, page });
}
let interior_len = layout::interior_len(layout.size(), page);
let total = layout::total_len(interior_len, page);
let base = unsafe {
libc::mmap(
core::ptr::null_mut(),
total,
INTERIOR_PROT,
MAP_FLAGS,
NO_FILE,
0,
)
};
if base == libc::MAP_FAILED {
return Err(LockError::Map {
bytes: total,
source: io::Error::last_os_error(),
});
}
let Some(base) = NonNull::new(base.cast::<u8>()) else {
unsafe {
let _ = libc::munmap(base, total);
}
return Err(LockError::Map {
bytes: total,
source: io::Error::other("mmap returned a mapping at address zero"),
});
};
let region = Self {
base,
interior_len,
generation: fork_generation(),
};
let lock = region.protect(page)?;
Ok((region, lock))
}
fn protect(&self, page: usize) -> Result<Option<LockError>, LockError> {
self.guard(self.base.as_ptr(), page)?;
self.guard(
unsafe { self.base.as_ptr().add(page + self.interior_len) },
page,
)?;
Ok(self.lock_and_exclude())
}
fn lock_and_exclude(&self) -> Option<LockError> {
let lock = self.lock();
let tracking = fork_tracking();
let dump = self.exclude_from_dumps();
lock.or(tracking).or(dump)
}
fn guard(&self, ptr: *mut u8, page: usize) -> Result<(), LockError> {
#[cfg(miri)]
{
let _ = (ptr, page);
Ok(())
}
#[cfg(not(miri))]
{
let rc = unsafe { libc::mprotect(ptr.cast(), page, libc::PROT_NONE) };
if rc == 0 {
Ok(())
} else {
Err(LockError::Guard {
source: io::Error::last_os_error(),
})
}
}
}
fn lock(&self) -> Option<LockError> {
#[cfg(miri)]
{
Some(LockError::Unavailable)
}
#[cfg(not(miri))]
{
let interior_len = self.interior_len;
let rc = unsafe { libc::mlock(self.interior().as_ptr().cast(), interior_len) };
(rc != 0).then(|| {
let source = io::Error::last_os_error();
LockError::Refused {
bytes: interior_len,
limit: memlock_limit(),
source,
}
})
}
}
fn exclude_from_dumps(&self) -> Option<LockError> {
#[cfg(not(all(any(target_os = "linux", target_os = "android"), not(miri))))]
{
None
}
#[cfg(all(any(target_os = "linux", target_os = "android"), not(miri)))]
{
let interior = self.interior();
let rc = unsafe {
libc::madvise(
interior.as_ptr().cast(),
self.interior_len,
libc::MADV_DONTDUMP,
)
};
(rc != 0).then(|| LockError::Dump {
source: io::Error::last_os_error(),
})
}
}
pub(super) fn same_process(&self) -> bool {
self.generation == fork_generation()
}
pub(super) fn relock(&mut self) -> Option<LockError> {
self.generation = fork_generation();
self.lock_and_exclude()
}
fn interior(&self) -> NonNull<u8> {
unsafe { NonNull::new_unchecked(self.base.as_ptr().add(page_size())) }
}
fn total(&self) -> usize {
layout::total_len(self.interior_len, page_size())
}
#[cfg(test)]
pub(super) fn ptr(&self) -> NonNull<u8> {
self.interior()
}
pub(super) fn value_ptr(&self, layout: Layout) -> NonNull<u8> {
let offset = layout::value_offset(self.interior_len, layout.size(), layout.align());
unsafe { NonNull::new_unchecked(self.interior().as_ptr().add(offset)) }
}
pub(super) fn wipe(&mut self) {
unsafe { super::wipe_raw(self.interior().as_ptr(), self.interior_len) };
}
}
impl Drop for Region {
fn drop(&mut self) {
self.wipe();
unsafe {
let _ = libc::munmap(self.base.as_ptr().cast(), self.total());
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn layout(size: usize, align: usize) -> Layout {
Layout::from_size_align(size, align).unwrap()
}
#[test]
fn the_total_is_the_interior_plus_a_guard_page_each_side() {
let page = page_size();
let (region, _) = Region::allocate(layout(1, 1)).unwrap();
assert_eq!(region.total(), 3 * page);
let (region, _) = Region::allocate(layout(page + 1, 1)).unwrap();
assert_eq!(region.total(), 4 * page);
}
#[test]
fn a_failed_atfork_registration_is_a_degradation_naming_its_errno() {
assert!(untracked(0).is_none());
match untracked(libc::ENOMEM) {
Some(LockError::Untracked { source }) => {
assert_eq!(source.raw_os_error(), Some(libc::ENOMEM));
}
other => panic!("{other:?}"),
}
}
#[test]
fn a_region_belongs_to_the_fork_generation_it_was_locked_in() {
let (mut region, _) = Region::allocate(layout(1, 1)).unwrap();
assert!(region.same_process());
region.generation = region.generation.wrapping_sub(1);
assert!(!region.same_process());
let _ = region.relock();
assert!(region.same_process(), "relock takes ownership here");
}
#[test]
fn an_alignment_of_one_page_is_the_most_the_region_offers() {
let page = page_size();
assert!(Region::allocate(layout(1, page)).is_ok());
assert!(matches!(
Region::allocate(layout(1, page * 2)),
Err(LockError::Alignment { align, page: p }) if align == page * 2 && p == page
));
}
#[test]
fn the_value_sits_against_the_trailing_guard() {
let page = page_size();
let (region, _) = Region::allocate(layout(32, 1)).unwrap();
let end = region.ptr().as_ptr() as usize + page;
assert_eq!(region.value_ptr(layout(32, 1)).as_ptr() as usize + 32, end);
let (region, _) = Region::allocate(layout(24, 16)).unwrap();
let end = region.ptr().as_ptr() as usize + page;
let value = region.value_ptr(layout(24, 16)).as_ptr() as usize;
assert_eq!(value % 16, 0);
assert!(
end - (value + 24) < 16,
"{} bytes of slack",
end - (value + 24)
);
let (region, _) = Region::allocate(layout(page, 1)).unwrap();
assert_eq!(region.value_ptr(layout(page, 1)), region.ptr());
}
#[test]
fn a_wipe_zeroes_the_whole_interior() {
let page = page_size();
let (mut region, _) = Region::allocate(layout(1, 1)).unwrap();
unsafe { core::slice::from_raw_parts_mut(region.ptr().as_ptr(), page).fill(0xEE) };
region.wipe();
let bytes = unsafe { core::slice::from_raw_parts(region.ptr().as_ptr(), page) };
assert!(bytes.iter().all(|&b| b == 0));
}
#[cfg(all(memlock_limit, not(miri)))]
#[test]
fn the_memlock_limit_is_the_soft_limit_when_finite() {
let mut lim = libc::rlimit {
rlim_cur: 0,
rlim_max: 0,
};
let rc = unsafe { libc::getrlimit(libc::RLIMIT_MEMLOCK, &mut lim) };
assert_eq!(rc, 0);
let expected = if lim.rlim_cur == libc::RLIM_INFINITY {
None
} else {
#[allow(clippy::useless_conversion)]
u64::try_from(lim.rlim_cur).ok()
};
assert_eq!(memlock_limit(), expected);
}
}