use core::ptr::NonNull;
use core::sync::atomic::{AtomicUsize, Ordering};
use crate::error::PagerError;
use crate::protection::Protection;
use crate::sys;
#[must_use]
pub fn page_size() -> usize {
static CACHE: AtomicUsize = AtomicUsize::new(0);
let cached = CACHE.load(Ordering::Relaxed);
if cached != 0 {
return cached;
}
let queried = sys::page_size();
CACHE.store(queried, Ordering::Relaxed);
queried
}
fn pages_for(len: usize) -> Result<usize, PagerError> {
if len == 0 {
return Err(PagerError::ZeroSize);
}
let page = page_size();
let bumped = len.checked_add(page - 1).ok_or(PagerError::SizeOverflow)?;
Ok(bumped & !(page - 1))
}
pub struct Region {
base: NonNull<u8>,
data: NonNull<u8>,
len: usize,
total: usize,
prot: Protection,
guarded: bool,
}
unsafe impl Send for Region {}
unsafe impl Sync for Region {}
impl Region {
pub fn new(len: usize) -> Result<Self, PagerError> {
let data_len = pages_for(len)?;
let base = sys::map(data_len).map_err(PagerError::Map)?;
Ok(Region {
base,
data: base,
len: data_len,
total: data_len,
prot: Protection::ReadWrite,
guarded: false,
})
}
pub fn with_guard_pages(len: usize) -> Result<Self, PagerError> {
let page = page_size();
let data_len = pages_for(len)?;
let two_guards = page.checked_mul(2).ok_or(PagerError::SizeOverflow)?;
let total = data_len
.checked_add(two_guards)
.ok_or(PagerError::SizeOverflow)?;
let base = sys::map_guarded(total, page, data_len).map_err(PagerError::Map)?;
let data = unsafe { NonNull::new_unchecked(base.as_ptr().add(page)) };
Ok(Region {
base,
data,
len: data_len,
total,
prot: Protection::ReadWrite,
guarded: true,
})
}
#[must_use]
pub fn len(&self) -> usize {
self.len
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.len == 0
}
#[must_use]
pub fn protection(&self) -> Protection {
self.prot
}
#[must_use]
pub fn has_guard_pages(&self) -> bool {
self.guarded
}
#[must_use]
pub fn as_ptr(&self) -> *const u8 {
self.data.as_ptr()
}
#[must_use]
pub fn as_mut_ptr(&mut self) -> *mut u8 {
self.data.as_ptr()
}
#[must_use]
pub fn as_slice(&self) -> Option<&[u8]> {
if !self.prot.is_readable() {
return None;
}
Some(unsafe { core::slice::from_raw_parts(self.data.as_ptr(), self.len) })
}
#[must_use]
pub fn as_mut_slice(&mut self) -> Option<&mut [u8]> {
if !self.prot.is_writable() {
return None;
}
Some(unsafe { core::slice::from_raw_parts_mut(self.data.as_ptr(), self.len) })
}
pub fn write(&mut self, offset: usize, bytes: &[u8]) -> Result<(), PagerError> {
if !self.prot.is_writable() {
return Err(PagerError::NotWritable);
}
let out_of_bounds = || PagerError::OutOfBounds {
offset,
len: bytes.len(),
region_len: self.len,
};
let end = offset.checked_add(bytes.len()).ok_or_else(out_of_bounds)?;
if end > self.len {
return Err(out_of_bounds());
}
if bytes.is_empty() {
return Ok(());
}
unsafe {
core::ptr::copy_nonoverlapping(
bytes.as_ptr(),
self.data.as_ptr().add(offset),
bytes.len(),
);
}
Ok(())
}
pub fn protect(&mut self, protection: Protection) -> Result<(), PagerError> {
if protection == self.prot {
return Ok(());
}
unsafe { sys::protect(self.data.as_ptr(), self.len, protection) }
.map_err(PagerError::Protect)?;
self.prot = protection;
Ok(())
}
}
impl core::fmt::Debug for Region {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("Region")
.field("addr", &self.data.as_ptr())
.field("len", &self.len)
.field("protection", &self.prot)
.field("guard_pages", &self.guarded)
.finish()
}
}
impl Drop for Region {
fn drop(&mut self) {
let _ = unsafe { sys::unmap(self.base.as_ptr(), self.total) };
}
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::expect_used,
clippy::panic,
reason = "tests assert on specific outcomes; a wrong outcome should fail the test loudly"
)]
mod tests {
use super::{Region, page_size};
use crate::{PagerError, Protection};
#[test]
fn test_page_size_is_a_sane_power_of_two() {
let page = page_size();
assert!(page.is_power_of_two());
assert!(page >= 4096);
assert_eq!(page, page_size());
}
#[test]
fn test_new_rounds_up_to_a_whole_page() {
let region = Region::new(1).unwrap();
assert_eq!(region.len(), page_size());
assert!(!region.is_empty());
assert_eq!(region.protection(), Protection::ReadWrite);
assert!(!region.has_guard_pages());
}
#[test]
fn test_new_zero_is_rejected() {
assert!(matches!(Region::new(0), Err(PagerError::ZeroSize)));
}
#[test]
fn test_new_overflow_is_rejected() {
assert!(matches!(
Region::new(usize::MAX),
Err(PagerError::SizeOverflow)
));
}
#[test]
fn test_write_then_read_round_trips() {
let mut region = Region::new(64).unwrap();
region.write(8, &[0xDE, 0xAD, 0xBE, 0xEF]).unwrap();
assert_eq!(
®ion.as_slice().unwrap()[8..12],
&[0xDE, 0xAD, 0xBE, 0xEF]
);
}
#[test]
fn test_write_out_of_bounds_is_rejected() {
let mut region = Region::new(16).unwrap();
let len = region.len();
assert!(matches!(
region.write(len, &[1]),
Err(PagerError::OutOfBounds { .. })
));
assert!(matches!(
region.write(usize::MAX, &[1, 2]),
Err(PagerError::OutOfBounds { .. })
));
}
#[test]
fn test_write_empty_slice_is_ok_even_at_the_end() {
let mut region = Region::new(16).unwrap();
let len = region.len();
assert_eq!(region.write(len, &[]), Ok(()));
}
#[test]
fn test_protect_changes_access_and_gates_the_slices() {
let mut region = Region::new(32).unwrap();
assert!(region.as_mut_slice().is_some());
region.protect(Protection::ReadExecute).unwrap();
assert_eq!(region.protection(), Protection::ReadExecute);
assert!(region.as_mut_slice().is_none());
assert!(region.as_slice().is_some());
assert_eq!(region.write(0, &[1]), Err(PagerError::NotWritable));
region.protect(Protection::None).unwrap();
assert!(region.as_slice().is_none());
}
#[test]
fn test_protect_to_same_protection_is_a_noop() {
let mut region = Region::new(8).unwrap();
assert_eq!(region.protect(Protection::ReadWrite), Ok(()));
assert_eq!(region.protection(), Protection::ReadWrite);
}
#[test]
fn test_guarded_region_is_usable_like_any_other() {
let mut region = Region::with_guard_pages(100).unwrap();
assert!(region.has_guard_pages());
assert!(region.len() >= 100);
region.write(0, &[7; 16]).unwrap();
assert_eq!(®ion.as_slice().unwrap()[..16], &[7; 16]);
}
#[test]
fn test_many_regions_coexist() {
let regions: Vec<Region> = (1..=32).map(|n| Region::new(n * 64).unwrap()).collect();
for (i, region) in regions.iter().enumerate() {
assert!(region.len() >= (i + 1) * 64);
}
}
}