use std::marker::PhantomData;
use ax_memory_addr::{PAGE_SIZE_4K, PhysAddr, VirtAddr};
use ax_memory_set::{MappingError, MappingResult};
use axaddrspace::{AddrSpaceResult, NestedPageTableOps, PageSize};
use axvm_types::{GuestPhysAddr, MappingFlags};
use page_table_generic as ptg;
use crate::{AxVmError, AxVmResult, ax_err, host::PagingHandler};
struct GenericFrameAllocator<H>(PhantomData<fn() -> H>);
impl<H> GenericFrameAllocator<H> {
const fn new() -> Self {
Self(PhantomData)
}
}
impl<H> Clone for GenericFrameAllocator<H> {
fn clone(&self) -> Self {
*self
}
}
impl<H> Copy for GenericFrameAllocator<H> {}
impl<H: PagingHandler + 'static> ptg::FrameAllocator for GenericFrameAllocator<H> {
fn alloc_frame(&self) -> Option<ptg::PhysAddr> {
H::alloc_frame()
}
fn dealloc_frame(&self, frame: ptg::PhysAddr) {
H::dealloc_frame(frame);
}
fn phys_to_virt(&self, paddr: ptg::PhysAddr) -> *mut u8 {
H::phys_to_virt(paddr).as_mut_ptr()
}
fn alloc_frames(&self, frames: usize, align: usize) -> Option<ptg::PhysAddr> {
H::alloc_frames(frames, align)
}
fn dealloc_frames(&self, start: ptg::PhysAddr, frames: usize, _frame_size: usize) {
H::dealloc_frames(start, frames);
}
}
pub(crate) struct GenericNestedPageTable<M, H>
where
M: ptg::TableMeta,
M::P: ptg::PageTableEntry<PteConfig = MappingFlags>,
H: PagingHandler + 'static,
{
inner: ptg::PageTable<M, GenericFrameAllocator<H>>,
}
impl<M, H> GenericNestedPageTable<M, H>
where
M: ptg::TableMeta,
M::P: ptg::PageTableEntry<PteConfig = MappingFlags>,
H: PagingHandler + 'static,
{
pub(crate) fn try_new() -> ptg::PagingResult<Self> {
Ok(Self {
inner: ptg::PageTable::new(GenericFrameAllocator::new())?,
})
}
pub(crate) fn root_paddr(&self) -> PhysAddr {
self.inner.root_paddr()
}
pub(crate) fn map(
&mut self,
vaddr: GuestPhysAddr,
paddr: PhysAddr,
size: PageSize,
flags: MappingFlags,
) -> ptg::PagingResult {
self.inner.map_page(
ptg::VirtAddr::from_usize(vaddr.as_usize()),
paddr,
size.into(),
flags,
)
}
pub(crate) fn map_region(
&mut self,
vaddr: GuestPhysAddr,
get_paddr: impl Fn(GuestPhysAddr) -> PhysAddr,
size: usize,
flags: MappingFlags,
allow_huge: bool,
) -> ptg::PagingResult {
self.inner.map_region(
ptg::VirtAddr::from_usize(vaddr.as_usize()),
|current| get_paddr(GuestPhysAddr::from(current.as_usize())),
size,
flags,
allow_huge,
)
}
pub(crate) fn unmap(
&mut self,
vaddr: GuestPhysAddr,
) -> ptg::PagingResult<(PhysAddr, MappingFlags, PageSize)> {
let (paddr, flags, page_size) = self
.inner
.unmap_page(ptg::VirtAddr::from_usize(vaddr.as_usize()))?;
Ok((paddr, flags, nested_page_size(page_size)?))
}
pub(crate) fn unmap_region(&mut self, start: GuestPhysAddr, size: usize) -> ptg::PagingResult {
self.inner
.unmap(ptg::VirtAddr::from_usize(start.as_usize()), size)
}
pub(crate) fn remap(
&mut self,
start: GuestPhysAddr,
paddr: PhysAddr,
flags: MappingFlags,
) -> ptg::PagingResult {
let start = GuestPhysAddr::from(start.as_usize() & !(PAGE_SIZE_4K - 1));
let _ = self.unmap(start);
self.map(start, paddr, PageSize::Size4K, flags)
}
pub(crate) fn protect_region(
&mut self,
start: GuestPhysAddr,
size: usize,
new_flags: MappingFlags,
) -> ptg::PagingResult {
let end = start
.as_usize()
.checked_add(size)
.ok_or_else(|| ptg::PagingError::address_overflow("protect_region"))?;
let mut current = start;
while current.as_usize() < end {
let page_size = self
.inner
.protect_page(ptg::VirtAddr::from_usize(current.as_usize()), new_flags)?;
current += page_size;
}
Ok(())
}
pub(crate) fn query(
&self,
vaddr: GuestPhysAddr,
) -> ptg::PagingResult<(PhysAddr, MappingFlags, PageSize)> {
let (paddr, flags, page_size) = self
.inner
.query(ptg::VirtAddr::from_usize(vaddr.as_usize()))?;
Ok((paddr, flags, nested_page_size(page_size)?))
}
}
pub(crate) enum LeveledPageTable<M3, M4, H, const SUPPORT_L3: bool>
where
M3: ptg::TableMeta,
M4: ptg::TableMeta,
M3::P: ptg::PageTableEntry<PteConfig = MappingFlags>,
M4::P: ptg::PageTableEntry<PteConfig = MappingFlags>,
H: PagingHandler + 'static,
{
L3(GenericNestedPageTable<M3, H>),
L4(GenericNestedPageTable<M4, H>),
}
impl<M3, M4, H, const SUPPORT_L3: bool> LeveledPageTable<M3, M4, H, SUPPORT_L3>
where
M3: ptg::TableMeta,
M4: ptg::TableMeta,
M3::P: ptg::PageTableEntry<PteConfig = MappingFlags>,
M4::P: ptg::PageTableEntry<PteConfig = MappingFlags>,
H: PagingHandler + 'static,
{
pub(crate) fn new(level: usize) -> AxVmResult<Self> {
match level {
3 => {
if !SUPPORT_L3 {
return ax_err!(InvalidInput, "L3 not supported on this architecture");
}
let table = GenericNestedPageTable::try_new().map_err(map_new_error)?;
Ok(Self::L3(table))
}
4 => {
let table = GenericNestedPageTable::try_new().map_err(map_new_error)?;
Ok(Self::L4(table))
}
_ => ax_err!(InvalidInput, "Invalid page table level"),
}
}
pub(crate) fn root_paddr(&self) -> PhysAddr {
match self {
Self::L3(pt) => pt.root_paddr(),
Self::L4(pt) => pt.root_paddr(),
}
}
pub(crate) const fn levels(&self) -> usize {
match self {
Self::L3(_) => 3,
Self::L4(_) => 4,
}
}
pub(crate) fn map(
&mut self,
vaddr: GuestPhysAddr,
paddr: PhysAddr,
size: PageSize,
flags: MappingFlags,
) -> MappingResult {
match self {
Self::L3(pt) => pt.map(vaddr, paddr, size, flags),
Self::L4(pt) => pt.map(vaddr, paddr, size, flags),
}
.map_err(map_error)
}
pub(crate) fn unmap(
&mut self,
vaddr: GuestPhysAddr,
) -> MappingResult<(PhysAddr, MappingFlags, PageSize)> {
match self {
Self::L3(pt) => pt.unmap(vaddr),
Self::L4(pt) => pt.unmap(vaddr),
}
.map_err(map_error)
}
pub(crate) fn map_region(
&mut self,
vaddr: GuestPhysAddr,
get_paddr: impl Fn(GuestPhysAddr) -> PhysAddr,
size: usize,
flags: MappingFlags,
allow_huge: bool,
) -> MappingResult {
match self {
Self::L3(pt) => pt.map_region(vaddr, &get_paddr, size, flags, allow_huge),
Self::L4(pt) => pt.map_region(vaddr, &get_paddr, size, flags, allow_huge),
}
.map_err(map_error)
}
pub(crate) fn unmap_region(&mut self, start: GuestPhysAddr, size: usize) -> MappingResult {
match self {
Self::L3(pt) => pt.unmap_region(start, size),
Self::L4(pt) => pt.unmap_region(start, size),
}
.map_err(map_error)
}
pub(crate) fn remap(
&mut self,
start: GuestPhysAddr,
paddr: PhysAddr,
flags: MappingFlags,
) -> bool {
match self {
Self::L3(pt) => pt.remap(start, paddr, flags),
Self::L4(pt) => pt.remap(start, paddr, flags),
}
.is_ok()
}
pub(crate) fn protect_region(
&mut self,
start: GuestPhysAddr,
size: usize,
new_flags: MappingFlags,
) -> bool {
match self {
Self::L3(pt) => pt.protect_region(start, size, new_flags),
Self::L4(pt) => pt.protect_region(start, size, new_flags),
}
.is_ok()
}
pub(crate) fn query(
&self,
vaddr: GuestPhysAddr,
) -> MappingResult<(PhysAddr, MappingFlags, PageSize)> {
match self {
Self::L3(pt) => pt.query(vaddr),
Self::L4(pt) => pt.query(vaddr),
}
.map_err(map_error)
}
}
impl<M3, M4, H, const SUPPORT_L3: bool> NestedPageTableOps
for LeveledPageTable<M3, M4, H, SUPPORT_L3>
where
M3: ptg::TableMeta,
M4: ptg::TableMeta,
M3::P: ptg::PageTableEntry<PteConfig = MappingFlags>,
M4::P: ptg::PageTableEntry<PteConfig = MappingFlags>,
H: PagingHandler + 'static,
{
fn root_paddr(&self) -> PhysAddr {
LeveledPageTable::root_paddr(self)
}
fn levels(&self) -> usize {
LeveledPageTable::levels(self)
}
fn alloc_frame(&self) -> Option<PhysAddr> {
H::alloc_frame()
}
fn dealloc_frame(&self, paddr: PhysAddr) {
H::dealloc_frame(paddr);
}
fn phys_to_virt(&self, paddr: PhysAddr) -> VirtAddr {
H::phys_to_virt(paddr)
}
fn map(
&mut self,
vaddr: GuestPhysAddr,
paddr: PhysAddr,
size: PageSize,
flags: MappingFlags,
) -> AddrSpaceResult {
Ok(LeveledPageTable::map(self, vaddr, paddr, size, flags)?)
}
fn unmap(
&mut self,
vaddr: GuestPhysAddr,
) -> AddrSpaceResult<(PhysAddr, MappingFlags, PageSize)> {
Ok(LeveledPageTable::unmap(self, vaddr)?)
}
fn map_region(
&mut self,
vaddr: GuestPhysAddr,
get_paddr: impl Fn(GuestPhysAddr) -> PhysAddr,
size: usize,
flags: MappingFlags,
allow_huge: bool,
) -> AddrSpaceResult {
Ok(LeveledPageTable::map_region(
self, vaddr, get_paddr, size, flags, allow_huge,
)?)
}
fn unmap_region(&mut self, start: GuestPhysAddr, size: usize) -> AddrSpaceResult {
Ok(LeveledPageTable::unmap_region(self, start, size)?)
}
fn remap(&mut self, start: GuestPhysAddr, paddr: PhysAddr, flags: MappingFlags) -> bool {
LeveledPageTable::remap(self, start, paddr, flags)
}
fn protect_region(
&mut self,
start: GuestPhysAddr,
size: usize,
new_flags: MappingFlags,
) -> bool {
LeveledPageTable::protect_region(self, start, size, new_flags)
}
fn query(&self, vaddr: GuestPhysAddr) -> AddrSpaceResult<(PhysAddr, MappingFlags, PageSize)> {
Ok(LeveledPageTable::query(self, vaddr)?)
}
}
fn nested_page_size(size: usize) -> ptg::PagingResult<PageSize> {
match size {
0x1000 => Ok(PageSize::Size4K),
0x10_0000 => Ok(PageSize::Size1M),
0x20_0000 => Ok(PageSize::Size2M),
0x4000_0000 => Ok(PageSize::Size1G),
_ => Err(ptg::PagingError::invalid_size(
"Nested page-table level has an unsupported page size",
)),
}
}
pub(crate) fn map_new_error(err: ptg::PagingError) -> AxVmError {
match err {
ptg::PagingError::NoMemory => AxVmError::OutOfMemory {
operation: "allocate nested page table",
},
_ => AxVmError::memory("create nested page table", err),
}
}
pub(crate) fn map_error(err: ptg::PagingError) -> MappingError {
match err {
ptg::PagingError::MappingConflict { .. } => MappingError::AlreadyExists,
ptg::PagingError::AlignmentError { .. }
| ptg::PagingError::AddressOverflow { .. }
| ptg::PagingError::InvalidSize { .. }
| ptg::PagingError::InvalidRange { .. } => MappingError::InvalidParam,
ptg::PagingError::NoMemory
| ptg::PagingError::HierarchyError { .. }
| ptg::PagingError::NotMapped => MappingError::BadState,
}
}