use crate::{accept_memory, AcceptError};
pub const LINUX_EFI_UNACCEPTED_MEM_TABLE_GUID: uefi_raw::Guid =
uefi_raw::guid!("d5d1de3c-105c-44f9-9ea9-bcef98120031");
pub const EFI_UNACCEPTED_UNIT_SIZE: u64 = 2 * 1024 * 1024;
#[derive(Copy, Clone, Debug)]
#[repr(C, packed)]
pub struct EfiUnacceptedMemory {
pub version: u32,
pub unit_size: u32,
pub phys_base: u64,
pub size: u64,
}
impl EfiUnacceptedMemory {
pub unsafe fn accept_by_size(&mut self, start: u64, size: u64) -> Result<(), AcceptError> {
let Some(end) = start.checked_add(size) else {
return Err(AcceptError::InvalidAlignment);
};
unsafe { self.accept_range(start, end) }
}
pub unsafe fn accept_range(&mut self, start: u64, end: u64) -> Result<(), AcceptError> {
let Some((range_start, range_end, unit_size)) =
self.clamp_gpa_range_to_bitmap_coverage(start, end)?
else {
return Ok(());
};
let (first_bit, last_bit) = self.addr_to_bit_range(range_start, range_end, unit_size)?;
let phys_base = self.phys_base;
let mut bitmap = BitmapMut::new(unsafe { self.as_bitmap_slice_mut() });
let mut scan = first_bit;
while let Some(run_start) = bitmap.find_next_set(scan, last_bit)? {
let run_end = bitmap
.find_next_zero(run_start, last_bit)?
.unwrap_or(last_bit);
let run_gpa_start = phys_base
.checked_add(
run_start
.checked_mul(unit_size)
.ok_or(AcceptError::InvalidAlignment)?,
)
.ok_or(AcceptError::InvalidAlignment)?;
let run_gpa_end = phys_base
.checked_add(
run_end
.checked_mul(unit_size)
.ok_or(AcceptError::InvalidAlignment)?,
)
.ok_or(AcceptError::InvalidAlignment)?;
unsafe { accept_memory(run_gpa_start, run_gpa_end)? };
let mut clear = run_start;
while clear < run_end {
bitmap.clear_bit(clear)?;
clear += 1;
}
scan = run_end;
}
Ok(())
}
pub fn bitmap_coverage_end(&self) -> Option<u64> {
self.phys_base.checked_add(self.total_coverage_size()?)
}
pub unsafe fn as_bitmap_slice(&self) -> &[u8] {
debug_assert!(self.byte_len().is_ok());
let bitmap_ptr = core::ptr::from_ref(self)
.cast::<u8>()
.wrapping_add(core::mem::size_of::<Self>());
let bitmap_len = self
.byte_len()
.expect("bitmap size must fit usize on this platform");
unsafe { core::slice::from_raw_parts(bitmap_ptr, bitmap_len) }
}
pub unsafe fn as_bitmap_slice_mut(&mut self) -> &mut [u8] {
debug_assert!(self.byte_len().is_ok());
let bitmap_ptr_mut = core::ptr::from_mut(self)
.cast::<u8>()
.wrapping_add(core::mem::size_of::<Self>());
debug_assert!(!bitmap_ptr_mut.is_null());
let bitmap_len = self
.byte_len()
.expect("bitmap size must fit usize on this platform");
unsafe { core::slice::from_raw_parts_mut(bitmap_ptr_mut, bitmap_len) }
}
pub unsafe fn register_range(&mut self, start: u64, end: u64) -> Result<(), AcceptError> {
let table_phys_base = self.phys_base;
let unit_size = self.validated_unit_size()?;
if start >= end {
return Ok(());
}
let unit_mask = unit_size - 1;
if end - start < 2 * unit_size {
return unsafe { Self::try_accept_range(start, end) };
}
let mut current_start = start;
let mut current_end = end;
if current_start & unit_mask != 0 {
let Some(aligned_start) = align_up(current_start, unit_size) else {
return Err(AcceptError::InvalidAlignment);
};
unsafe { Self::try_accept_range(current_start, aligned_start)? };
current_start = aligned_start;
}
if current_end & unit_mask != 0 {
let aligned_end = align_down(current_end, unit_size);
unsafe { Self::try_accept_range(aligned_end, current_end)? };
current_end = aligned_end;
}
let Some(bitmap_coverage) = self.total_coverage_size() else {
return Err(AcceptError::InvalidAlignment);
};
let Some(bitmap_end) = table_phys_base.checked_add(bitmap_coverage) else {
return Err(AcceptError::InvalidAlignment);
};
if current_start < table_phys_base {
let accept_end = current_end.min(table_phys_base);
unsafe { Self::try_accept_range(current_start, accept_end)? };
current_start = accept_end;
}
if current_start >= current_end {
return Ok(());
}
if current_start < bitmap_end {
let bitmap_range_end = current_end.min(bitmap_end);
if current_start < bitmap_range_end {
unsafe {
self.mark_range_as_unaccepted(current_start, bitmap_range_end, unit_size)?
};
}
current_start = bitmap_range_end;
}
if current_start < current_end {
unsafe { Self::try_accept_range(current_start, current_end)? };
}
Ok(())
}
pub fn total_coverage_size(&self) -> Option<u64> {
let unit_size = u64::from(self.unit_size);
self.size.checked_mul(unit_size)?.checked_mul(8)
}
unsafe fn set_unaccepted_bits(&mut self, start: u64, end: u64) -> Result<(), AcceptError> {
let unit_size = self.validated_unit_size()?;
let abs_start = self
.phys_base
.checked_add(start)
.ok_or(AcceptError::InvalidAlignment)?;
let abs_end = self
.phys_base
.checked_add(end)
.ok_or(AcceptError::InvalidAlignment)?;
unsafe { self.mark_range_as_unaccepted(abs_start, abs_end, unit_size) }
}
fn total_bits(&self) -> Result<u64, AcceptError> {
self.size
.checked_mul(8)
.ok_or(AcceptError::InvalidAlignment)
}
fn byte_len(&self) -> Result<usize, AcceptError> {
usize::try_from(self.size).map_err(|_| AcceptError::InvalidAlignment)
}
fn validated_unit_size(&self) -> Result<u64, AcceptError> {
let unit_size = u64::from(self.unit_size);
if unit_size == 0 || !unit_size.is_power_of_two() {
return Err(AcceptError::InvalidAlignment);
}
Ok(unit_size)
}
fn max_phys_addr_exclusive(&self, unit_size: u64) -> Result<u64, AcceptError> {
let total_bits = self.total_bits()?;
let coverage_len = total_bits
.checked_mul(unit_size)
.ok_or(AcceptError::InvalidAlignment)?;
self.phys_base
.checked_add(coverage_len)
.ok_or(AcceptError::InvalidAlignment)
}
fn clamp_gpa_range_to_bitmap_coverage(
&self,
start: u64,
end: u64,
) -> Result<Option<(u64, u64, u64)>, AcceptError> {
if start >= end {
return Ok(None);
}
let unit_size = self.validated_unit_size()?;
let coverage_end = self.max_phys_addr_exclusive(unit_size)?;
let range_start = start.max(self.phys_base);
let range_end = end.min(coverage_end);
if range_start >= range_end {
return Ok(None);
}
Ok(Some((range_start, range_end, unit_size)))
}
fn addr_to_bit_range(
&self,
start: u64,
end: u64,
unit_size: u64,
) -> Result<(u64, u64), AcceptError> {
debug_assert!(start >= self.phys_base);
debug_assert!(start < end);
debug_assert!(unit_size.is_power_of_two());
let rel_start = start - self.phys_base;
let rel_end = end - self.phys_base;
let first_bit = rel_start / unit_size;
let last_bit = rel_end
.checked_add(unit_size - 1)
.ok_or(AcceptError::InvalidAlignment)?
/ unit_size;
Ok((first_bit, last_bit))
}
unsafe fn mark_range_as_unaccepted(
&mut self,
start: u64,
end: u64,
unit_size: u64,
) -> Result<(), AcceptError> {
if start >= end {
return Ok(());
}
debug_assert_eq!(start % unit_size, 0);
debug_assert_eq!(end % unit_size, 0);
let start_bit = (start - self.phys_base) / unit_size;
let end_bit = (end - self.phys_base) / unit_size;
let total_bits = self.total_bits()?;
let clamped_start_bit = start_bit.min(total_bits);
let clamped_end_bit = end_bit.min(total_bits);
if clamped_start_bit >= clamped_end_bit {
return Ok(());
}
let mut bitmap = BitmapMut::new(unsafe { self.as_bitmap_slice_mut() });
for bit in clamped_start_bit..clamped_end_bit {
bitmap.set_bit(bit)?;
}
Ok(())
}
unsafe fn try_accept_range(start: u64, end: u64) -> Result<(), AcceptError> {
if start >= end {
return Ok(());
}
unsafe { accept_memory(start, end) }
}
}
struct BitmapMut<'a> {
bits: &'a mut [u8],
}
impl<'a> BitmapMut<'a> {
fn new(bits: &'a mut [u8]) -> Self {
Self { bits }
}
fn capacity(&self) -> Result<u64, AcceptError> {
let len = u64::try_from(self.bits.len()).map_err(|_| AcceptError::InvalidAlignment)?;
len.checked_mul(8).ok_or(AcceptError::InvalidAlignment)
}
fn get_pos_mask(&self, bit_index: u64) -> Result<(usize, u8), AcceptError> {
if bit_index >= self.capacity()? {
return Err(AcceptError::InvalidAlignment);
}
let byte_index =
usize::try_from(bit_index >> 3).map_err(|_| AcceptError::InvalidAlignment)?;
let mask = 1u8 << (bit_index & 7);
Ok((byte_index, mask))
}
fn is_set(&self, bit_index: u64) -> Result<bool, AcceptError> {
let (byte_index, mask) = self.get_pos_mask(bit_index)?;
Ok((self.bits[byte_index] & mask) != 0)
}
fn set_bit(&mut self, bit_index: u64) -> Result<(), AcceptError> {
let (byte_index, mask) = self.get_pos_mask(bit_index)?;
self.bits[byte_index] |= mask;
Ok(())
}
fn clear_bit(&mut self, bit_index: u64) -> Result<(), AcceptError> {
let (byte_index, mask) = self.get_pos_mask(bit_index)?;
self.bits[byte_index] &= !mask;
Ok(())
}
fn find_next_set(&self, start_bit: u64, end_bit: u64) -> Result<Option<u64>, AcceptError> {
self.find_next_matching(start_bit, end_bit, true)
}
fn find_next_zero(&self, start_bit: u64, end_bit: u64) -> Result<Option<u64>, AcceptError> {
self.find_next_matching(start_bit, end_bit, false)
}
fn find_next_matching(
&self,
start_bit: u64,
end_bit: u64,
target: bool,
) -> Result<Option<u64>, AcceptError> {
let bit_len = self.capacity()?;
if start_bit > end_bit || end_bit > bit_len {
return Err(AcceptError::InvalidAlignment);
}
if start_bit == end_bit {
return Ok(None);
}
let mut scan_bit = start_bit;
while scan_bit < end_bit && (scan_bit & 63) != 0 {
if self.is_set(scan_bit)? == target {
return Ok(Some(scan_bit));
}
scan_bit += 1;
}
while end_bit - scan_bit >= 64 {
let next = scan_bit + 64;
let byte_index =
usize::try_from(scan_bit >> 3).map_err(|_| AcceptError::InvalidAlignment)?;
let word = unsafe {
let ptr = self.bits.as_ptr().add(byte_index).cast::<u64>();
u64::from_le(ptr.read_unaligned())
};
let match_word = if target { word } else { !word };
if match_word != 0 {
let delta = u64::from(match_word.trailing_zeros());
let found = scan_bit + delta;
return Ok(Some(found));
}
scan_bit = next;
}
while scan_bit < end_bit {
if self.is_set(scan_bit)? == target {
return Ok(Some(scan_bit));
}
scan_bit += 1;
}
Ok(None)
}
}
fn align_down(addr: u64, align: u64) -> u64 {
addr & !(align - 1)
}
fn align_up(addr: u64, align: u64) -> Option<u64> {
addr.checked_add(align - 1).map(|v| v & !(align - 1))
}