use rustc_hash::FxHashMap;
use crate::tlb::TranslationCache;
pub const PAGE_SIZE: u64 = 0x1000;
const PAGE_MASK: u64 = PAGE_SIZE - 1;
pub type Perm = u8;
pub mod perm {
use super::Perm;
pub const NONE: Perm = 0;
pub const INIT: Perm = 1 << 0;
pub const READ: Perm = 1 << 1;
pub const WRITE: Perm = 1 << 2;
pub const EXEC: Perm = 1 << 3;
pub const MAP: Perm = 1 << 4;
pub const READ_WATCH: Perm = 1 << 5;
pub const WRITE_WATCH: Perm = 1 << 6;
pub const READ_WRITE: Perm = READ | WRITE;
pub const RW_INIT: Perm = MAP | READ | WRITE | INIT;
pub const RX_INIT: Perm = MAP | READ | EXEC | INIT;
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum FaultKind {
ReadUnmapped,
ReadPerm,
ReadUninit,
WriteUnmapped,
WritePerm,
ExecUnmapped,
ExecViolation,
ReadWatch,
WriteWatch,
AddressOverflow,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct MemFault {
pub kind: FaultKind,
pub addr: u64,
}
impl MemFault {
fn new(kind: FaultKind, addr: u64) -> Self {
Self { kind, addr }
}
pub fn is_write(&self) -> bool {
matches!(
self.kind,
FaultKind::WriteUnmapped | FaultKind::WritePerm | FaultKind::WriteWatch
)
}
}
impl std::fmt::Display for MemFault {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let what = match self.kind {
FaultKind::ReadUnmapped => "read of unmapped memory",
FaultKind::ReadPerm => "read of unreadable memory",
FaultKind::ReadUninit => "read of uninitialized memory",
FaultKind::WriteUnmapped => "write to unmapped memory",
FaultKind::WritePerm => "write to unwritable memory",
FaultKind::ExecUnmapped => "execution of unmapped memory",
FaultKind::ExecViolation => "execution of non-executable memory",
FaultKind::ReadWatch => "watched read",
FaultKind::WriteWatch => "watched write",
FaultKind::AddressOverflow => "address space overflow",
};
write!(f, "{what} at {:#x}", self.addr)
}
}
impl std::error::Error for MemFault {}
#[repr(C)]
#[derive(Clone)]
pub struct PageData {
pub data: [u8; PAGE_SIZE as usize],
pub perm: [Perm; PAGE_SIZE as usize],
}
pub const PAGE_PERM_OFFSET: usize = PAGE_SIZE as usize;
#[derive(Clone)]
struct Page {
inner: Box<PageData>,
}
impl Page {
fn unmapped() -> Self {
let inner = unsafe {
let layout = std::alloc::Layout::new::<PageData>();
let ptr = std::alloc::alloc_zeroed(layout).cast::<PageData>();
if ptr.is_null() {
std::alloc::handle_alloc_error(layout);
}
Box::from_raw(ptr)
};
Self { inner }
}
fn host_ptr(&mut self) -> *mut u8 {
std::ptr::from_mut(&mut *self.inner).cast()
}
}
impl std::ops::Deref for Page {
type Target = PageData;
fn deref(&self) -> &PageData {
&self.inner
}
}
impl std::ops::DerefMut for Page {
fn deref_mut(&mut self) -> &mut PageData {
&mut self.inner
}
}
#[derive(Default)]
pub struct Mmu {
pages: FxHashMap<u64, Page>,
check_uninit: bool,
watchpoints_armed: bool,
tlb: TranslationCache,
}
impl Clone for Mmu {
fn clone(&self) -> Self {
Self {
pages: self.pages.clone(),
check_uninit: self.check_uninit,
watchpoints_armed: self.watchpoints_armed,
tlb: TranslationCache::default(),
}
}
}
fn page_chunks(addr: u64, len: u64) -> Result<impl Iterator<Item = (u64, usize, usize)>, MemFault> {
if len != 0 && addr.checked_add(len - 1).is_none() {
return Err(MemFault::new(FaultKind::AddressOverflow, addr));
}
let mut remaining = len;
let mut cursor = addr;
Ok(std::iter::from_fn(move || {
if remaining == 0 {
return None;
}
let offset = cursor & PAGE_MASK;
let take = (PAGE_SIZE - offset).min(remaining);
let chunk = (cursor >> 12, offset as usize, take as usize);
cursor = cursor.wrapping_add(take);
remaining -= take;
Some(chunk)
}))
}
impl Mmu {
pub fn new() -> Self {
Self::default()
}
pub fn resident_pages(&self) -> usize {
self.pages.len()
}
pub fn check_uninit(&self) -> bool {
self.check_uninit
}
pub fn set_check_uninit(&mut self, enabled: bool) {
self.check_uninit = enabled;
self.tlb.flush();
}
pub fn watchpoints_armed(&self) -> bool {
self.watchpoints_armed
}
pub fn set_watchpoints_armed(&mut self, armed: bool) {
self.watchpoints_armed = armed;
self.tlb.flush();
}
fn caching_allowed(&self) -> bool {
!self.check_uninit && !self.watchpoints_armed
}
pub fn tlb_ptr(&mut self) -> *mut u8 {
std::ptr::from_mut(&mut self.tlb).cast()
}
pub fn flush_tlb(&mut self) {
self.tlb.flush();
}
pub fn cache_translation(&mut self, addr: u64) -> bool {
if !self.caching_allowed() {
return false;
}
let Some(page) = self.pages.get_mut(&(addr >> 12)) else {
return false;
};
let host = page.host_ptr();
self.tlb.insert(addr, host);
true
}
pub fn map(&mut self, addr: u64, len: u64, permissions: Perm) -> Result<(), MemFault> {
self.tlb.flush();
for (index, offset, take) in page_chunks(addr, len)? {
let page = self.pages.entry(index).or_insert_with(Page::unmapped);
page.data[offset..offset + take].fill(0);
page.perm[offset..offset + take].fill(permissions | perm::MAP);
}
Ok(())
}
pub fn unmap(&mut self, addr: u64, len: u64) -> Result<(), MemFault> {
self.tlb.flush();
for (index, offset, take) in page_chunks(addr, len)? {
let Some(page) = self.pages.get_mut(&index) else {
continue;
};
page.data[offset..offset + take].fill(0);
page.perm[offset..offset + take].fill(perm::NONE);
if page.perm.iter().all(|&p| p == perm::NONE) {
self.pages.remove(&index);
}
}
Ok(())
}
pub fn protect(&mut self, addr: u64, len: u64, permissions: Perm) -> Result<(), MemFault> {
self.tlb.flush();
for (index, offset, take) in page_chunks(addr, len)? {
let base = index << 12;
match self.pages.get(&index) {
Some(page) => {
for byte in offset..offset + take {
if page.perm[byte] & perm::MAP == 0 {
let at = base + byte as u64;
return Err(MemFault::new(FaultKind::WriteUnmapped, at));
}
}
}
None => {
return Err(MemFault::new(
FaultKind::WriteUnmapped,
base + offset as u64,
));
}
}
}
for (index, offset, take) in page_chunks(addr, len)? {
let page = self.pages.get_mut(&index).expect("checked above");
for byte in offset..offset + take {
let init = page.perm[byte] & perm::INIT;
page.perm[byte] = permissions | perm::MAP | init;
}
}
Ok(())
}
pub fn permissions(&self, addr: u64) -> Perm {
self.pages
.get(&(addr >> 12))
.map_or(perm::NONE, |page| page.perm[(addr & PAGE_MASK) as usize])
}
pub fn read(&self, addr: u64, out: &mut [u8]) -> Result<(), MemFault> {
self.read_with(addr, out, perm::READ)
}
pub fn read_code(&self, addr: u64, out: &mut [u8]) -> Result<(), MemFault> {
self.read_with(addr, out, perm::EXEC)
}
fn read_with(&self, addr: u64, out: &mut [u8], required: Perm) -> Result<(), MemFault> {
let executing = required & perm::EXEC != 0;
let mut written = 0;
for (index, offset, take) in page_chunks(addr, out.len() as u64)? {
let base = index << 12;
let Some(page) = self.pages.get(&index) else {
let kind = if executing {
FaultKind::ExecUnmapped
} else {
FaultKind::ReadUnmapped
};
return Err(MemFault::new(kind, base + offset as u64));
};
for byte in offset..offset + take {
let at = base + byte as u64;
let held = page.perm[byte];
if held & perm::MAP == 0 {
let kind = if executing {
FaultKind::ExecUnmapped
} else {
FaultKind::ReadUnmapped
};
return Err(MemFault::new(kind, at));
}
if held & required == 0 {
let kind = if executing {
FaultKind::ExecViolation
} else {
FaultKind::ReadPerm
};
return Err(MemFault::new(kind, at));
}
if self.check_uninit && held & perm::INIT == 0 {
return Err(MemFault::new(FaultKind::ReadUninit, at));
}
if self.watchpoints_armed && held & perm::READ_WATCH != 0 {
return Err(MemFault::new(FaultKind::ReadWatch, at));
}
}
out[written..written + take].copy_from_slice(&page.data[offset..offset + take]);
written += take;
}
Ok(())
}
pub fn write(&mut self, addr: u64, bytes: &[u8]) -> Result<(), MemFault> {
let mut read = 0;
for (index, offset, take) in page_chunks(addr, bytes.len() as u64)? {
let base = index << 12;
let Some(page) = self.pages.get(&index) else {
return Err(MemFault::new(
FaultKind::WriteUnmapped,
base + offset as u64,
));
};
for byte in offset..offset + take {
let at = base + byte as u64;
let held = page.perm[byte];
if held & perm::MAP == 0 {
return Err(MemFault::new(FaultKind::WriteUnmapped, at));
}
if held & perm::WRITE == 0 {
return Err(MemFault::new(FaultKind::WritePerm, at));
}
if self.watchpoints_armed && held & perm::WRITE_WATCH != 0 {
return Err(MemFault::new(FaultKind::WriteWatch, at));
}
}
read += take;
}
debug_assert_eq!(read, bytes.len());
let mut written = 0;
for (index, offset, take) in page_chunks(addr, bytes.len() as u64)? {
let page = self.pages.get_mut(&index).expect("checked above");
page.data[offset..offset + take].copy_from_slice(&bytes[written..written + take]);
for byte in offset..offset + take {
page.perm[byte] |= perm::INIT;
}
written += take;
}
Ok(())
}
pub fn write_unchecked(&mut self, addr: u64, bytes: &[u8], permissions: Perm) {
self.tlb.flush();
let mut written = 0;
let chunks = page_chunks(addr, bytes.len() as u64)
.expect("write_unchecked range must fit the address space");
for (index, offset, take) in chunks {
let page = self.pages.entry(index).or_insert_with(Page::unmapped);
page.data[offset..offset + take].copy_from_slice(&bytes[written..written + take]);
page.perm[offset..offset + take].fill(permissions | perm::MAP | perm::INIT);
written += take;
}
}
pub fn snapshot(&self) -> MmuSnapshot {
MmuSnapshot {
pages: self.pages.clone(),
}
}
pub fn restore(&mut self, snapshot: &MmuSnapshot) {
self.pages.clone_from(&snapshot.pages);
self.tlb.flush();
}
}
#[derive(Clone)]
pub struct MmuSnapshot {
pages: FxHashMap<u64, Page>,
}
#[cfg(test)]
mod tests {
use super::*;
fn mapped() -> Mmu {
let mut mmu = Mmu::new();
mmu.map(0x1000, 0x2000, perm::RW_INIT).unwrap();
mmu
}
#[test]
fn the_permission_array_follows_the_data_array() {
let page = Page::unmapped();
let data = std::ptr::from_ref(&page.data) as usize;
let perm = std::ptr::from_ref(&page.perm) as usize;
assert_eq!(perm - data, PAGE_PERM_OFFSET);
assert_eq!(std::mem::size_of::<PageData>(), 2 * PAGE_SIZE as usize);
}
#[test]
fn a_translation_is_cached_only_for_a_resident_page() {
let mut mmu = mapped();
assert!(!mmu.cache_translation(0x9000), "unmapped page");
assert!(mmu.cache_translation(0x1abc));
mmu.write(0x1abc, &[0x5a]).unwrap();
let host = mmu.tlb.lookup(0x1abc).expect("just cached");
assert_eq!(unsafe { *host }, 0x5a);
assert_eq!(unsafe { *host.add(PAGE_PERM_OFFSET) }, perm::RW_INIT);
}
#[test]
fn a_dynamic_check_empties_the_cache_and_keeps_it_empty() {
let mut mmu = mapped();
assert!(mmu.cache_translation(0x1000));
mmu.set_check_uninit(true);
assert!(mmu.tlb.lookup(0x1000).is_none());
assert!(!mmu.cache_translation(0x1000));
mmu.set_check_uninit(false);
assert!(mmu.cache_translation(0x1000));
}
#[test]
fn unmapping_a_page_drops_its_cached_translation() {
let mut mmu = mapped();
assert!(mmu.cache_translation(0x1000));
mmu.unmap(0x1000, PAGE_SIZE).unwrap();
assert!(mmu.tlb.lookup(0x1000).is_none());
}
#[test]
fn a_cloned_mmu_starts_with_an_empty_cache() {
let mut mmu = mapped();
assert!(mmu.cache_translation(0x1000));
let clone = mmu.clone();
assert!(clone.tlb.lookup(0x1000).is_none());
}
#[test]
fn maps_reads_and_writes() {
let mut mmu = mapped();
mmu.write(0x1004, &[1, 2, 3, 4]).unwrap();
let mut out = [0; 4];
mmu.read(0x1004, &mut out).unwrap();
assert_eq!(out, [1, 2, 3, 4]);
}
#[test]
fn unmapped_access_faults_at_the_offending_byte() {
let mmu = mapped();
let mut out = [0; 4];
assert_eq!(
mmu.read(0x2ffe, &mut out).unwrap_err(),
MemFault::new(FaultKind::ReadUnmapped, 0x3000)
);
}
#[test]
fn access_spanning_pages_is_contiguous() {
let mut mmu = mapped();
let bytes: Vec<u8> = (0..16).collect();
mmu.write(0x1ff8, &bytes).unwrap();
let mut out = [0; 16];
mmu.read(0x1ff8, &mut out).unwrap();
assert_eq!(out.to_vec(), bytes);
}
#[test]
fn write_to_read_only_memory_faults_and_changes_nothing() {
let mut mmu = Mmu::new();
mmu.map(0x1000, PAGE_SIZE, perm::RX_INIT).unwrap();
assert_eq!(
mmu.write(0x1000, &[0xff]).unwrap_err(),
MemFault::new(FaultKind::WritePerm, 0x1000)
);
let mut out = [0xaa];
mmu.read(0x1000, &mut out).unwrap();
assert_eq!(out, [0]);
}
#[test]
fn partially_refused_write_is_not_applied() {
let mut mmu = Mmu::new();
mmu.map(0x1000, PAGE_SIZE, perm::RW_INIT).unwrap();
mmu.map(0x2000, PAGE_SIZE, perm::RX_INIT).unwrap();
assert_eq!(
mmu.write(0x1ffc, &[0xff; 8]).unwrap_err(),
MemFault::new(FaultKind::WritePerm, 0x2000)
);
let mut out = [0xaa; 4];
mmu.read(0x1ffc, &mut out).unwrap();
assert_eq!(out, [0; 4]);
}
#[test]
fn fetching_from_non_executable_memory_faults() {
let mut mmu = mapped();
let mut out = [0; 4];
assert_eq!(
mmu.read_code(0x1000, &mut out).unwrap_err(),
MemFault::new(FaultKind::ExecViolation, 0x1000)
);
mmu.protect(0x1000, PAGE_SIZE, perm::READ | perm::EXEC)
.unwrap();
mmu.read_code(0x1000, &mut out).unwrap();
}
#[test]
fn protect_preserves_contents_and_initializedness() {
let mut mmu = mapped();
mmu.write(0x1000, &[7; 4]).unwrap();
mmu.set_check_uninit(true);
mmu.protect(0x1000, PAGE_SIZE, perm::READ).unwrap();
let mut out = [0; 4];
mmu.read(0x1000, &mut out).unwrap();
assert_eq!(out, [7; 4]);
}
#[test]
fn protect_of_unmapped_memory_is_refused_entirely() {
let mut mmu = mapped();
assert!(mmu.protect(0x2000, 0x2000, perm::READ).is_err());
assert_eq!(mmu.permissions(0x2000), perm::RW_INIT);
}
#[test]
fn uninitialized_reads_fault_only_when_checked() {
let mut mmu = Mmu::new();
mmu.map(0x1000, PAGE_SIZE, perm::MAP | perm::READ_WRITE)
.unwrap();
let mut out = [0; 1];
mmu.read(0x1000, &mut out).unwrap();
mmu.set_check_uninit(true);
assert_eq!(
mmu.read(0x1000, &mut out).unwrap_err(),
MemFault::new(FaultKind::ReadUninit, 0x1000)
);
mmu.write(0x1000, &[1]).unwrap();
mmu.read(0x1000, &mut out).unwrap();
}
#[test]
fn watchpoints_fire_only_when_armed() {
let mut mmu = Mmu::new();
mmu.map(0x1000, PAGE_SIZE, perm::RW_INIT | perm::WRITE_WATCH)
.unwrap();
mmu.write(0x1000, &[1]).unwrap();
mmu.set_watchpoints_armed(true);
assert_eq!(
mmu.write(0x1000, &[2]).unwrap_err(),
MemFault::new(FaultKind::WriteWatch, 0x1000)
);
let mut out = [0; 1];
mmu.read(0x1000, &mut out).unwrap();
assert_eq!(out, [1]);
}
#[test]
fn unmap_releases_pages_and_faults_afterwards() {
let mut mmu = mapped();
assert_eq!(mmu.resident_pages(), 2);
mmu.unmap(0x1000, 0x2000).unwrap();
assert_eq!(mmu.resident_pages(), 0);
let mut out = [0; 1];
assert_eq!(
mmu.read(0x1000, &mut out).unwrap_err(),
MemFault::new(FaultKind::ReadUnmapped, 0x1000)
);
}
#[test]
fn partial_unmap_keeps_the_rest_of_the_page() {
let mut mmu = mapped();
mmu.write(0x1000, &[9; 8]).unwrap();
mmu.unmap(0x1000, 4).unwrap();
assert_eq!(mmu.resident_pages(), 2);
let mut out = [0; 4];
mmu.read(0x1004, &mut out).unwrap();
assert_eq!(out, [9; 4]);
}
#[test]
fn snapshot_and_restore_round_trips_contents_and_mappings() {
let mut mmu = mapped();
mmu.write(0x1000, &[1, 2, 3, 4]).unwrap();
let snapshot = mmu.snapshot();
mmu.write(0x1000, &[9, 9, 9, 9]).unwrap();
mmu.map(0x8000, PAGE_SIZE, perm::RW_INIT).unwrap();
mmu.unmap(0x2000, PAGE_SIZE).unwrap();
mmu.restore(&snapshot);
let mut out = [0; 4];
mmu.read(0x1000, &mut out).unwrap();
assert_eq!(out, [1, 2, 3, 4]);
assert_eq!(mmu.permissions(0x8000), perm::NONE);
mmu.read(0x2000, &mut out).unwrap();
}
#[test]
fn access_wrapping_the_address_space_faults() {
let mmu = mapped();
let mut out = [0; 8];
assert_eq!(
mmu.read(u64::MAX - 2, &mut out).unwrap_err(),
MemFault::new(FaultKind::AddressOverflow, u64::MAX - 2)
);
}
#[test]
fn write_unchecked_maps_and_ignores_permissions() {
let mut mmu = Mmu::new();
mmu.write_unchecked(0x1000, &[1, 2, 3, 4], perm::READ | perm::EXEC);
let mut out = [0; 4];
mmu.read(0x1000, &mut out).unwrap();
assert_eq!(out, [1, 2, 3, 4]);
assert_eq!(
mmu.write(0x1000, &[0]).unwrap_err(),
MemFault::new(FaultKind::WritePerm, 0x1000)
);
}
}