use core::{
arch::{asm, naked_asm},
mem::offset_of,
};
use riscv::register::{
sstatus,
stvec::{self, Stvec, TrapMode},
};
pub fn has_hypervisor_extension() -> bool {
with_detect_trap() == 0
}
#[inline]
fn with_detect_trap() -> usize {
let (sie, stvec) = init_detect_trap();
let mut state = DetectState::new(super::registers::read_sscratch());
run_h_extension_probe(&mut state);
restore_detect_trap(sie, stvec);
state.result
}
#[inline]
fn run_h_extension_probe(state: &mut DetectState) {
let saved_sscratch = state.saved_sscratch;
unsafe {
asm!(
"csrw sscratch, {state}",
"csrr {probe_value}, 0x680",
"csrw sscratch, {saved_sscratch}",
state = in(reg) state,
saved_sscratch = in(reg) saved_sscratch,
probe_value = out(reg) _,
options(nostack)
)
}
}
#[inline]
fn init_detect_trap() -> (bool, Stvec) {
let stored_sie = sstatus::read().sie();
unsafe {
sstatus::clear_sie();
}
let stored_stvec = stvec::read();
let trap_addr = on_detect_trap as *const () as usize;
assert_eq!(
trap_addr & 0b11,
0,
"H-extension probe trap vector must be four-byte aligned"
);
let mut stvec = Stvec::from_bits(0);
stvec.set_address(trap_addr);
stvec.set_trap_mode(TrapMode::Direct);
unsafe { stvec::write(stvec) }
(stored_sie, stored_stvec)
}
#[inline]
fn restore_detect_trap(sie: bool, stvec: Stvec) {
unsafe {
asm!("csrw stvec, {}", in(reg) stvec.bits(), options(nomem, nostack));
if sie {
sstatus::set_sie();
};
}
}
#[repr(C)]
struct DetectState {
saved_sscratch: usize,
result: usize,
saved_t0: usize,
saved_t1: usize,
}
impl DetectState {
const fn new(saved_sscratch: usize) -> Self {
Self {
saved_sscratch,
result: 0,
saved_t0: 0,
saved_t1: 0,
}
}
}
#[unsafe(naked)]
unsafe extern "C" fn on_detect_trap() -> ! {
naked_asm!(
".p2align 2",
"csrrw t0, sscratch, t0",
"sd t1, {saved_t1}(t0)",
"csrr t1, sscratch",
"sd t1, {saved_t0}(t0)",
"csrr t1, scause",
"sd t1, {result}(t0)",
"csrr t1, sepc",
"addi t1, t1, 4",
"csrw sepc, t1",
"ld t1, {saved_sscratch}(t0)",
"csrw sscratch, t1",
"ld t1, {saved_t1}(t0)",
"ld t0, {saved_t0}(t0)",
"sret",
saved_sscratch = const offset_of!(DetectState, saved_sscratch),
result = const offset_of!(DetectState, result),
saved_t0 = const offset_of!(DetectState, saved_t0),
saved_t1 = const offset_of!(DetectState, saved_t1),
)
}