use core::mem::size_of;
#[cfg(not(feature = "tls"))]
use cpu_local::CURRENT_THREAD_CPU_BASE_OFFSET;
#[cfg(feature = "fp-simd")]
use riscv::register::sstatus;
use riscv::{
interrupt::{
Trap,
supervisor::{Exception as E, Interrupt as I},
},
register::{scause, stval},
};
#[cfg(feature = "tls")]
use super::local_state::{CPU_ENTRY_SCRATCH0_OFFSET, CPU_ENTRY_SCRATCH1_OFFSET};
#[cfg(not(feature = "tls"))]
use super::local_state::{THREAD_SCRATCH0_OFFSET, THREAD_SCRATCH1_OFFSET};
use super::{
TrapFrame,
local_state::{CPU_KERNEL_STACK_POINTER_OFFSET, CPU_USER_TRAP_FRAME_OFFSET},
};
use crate::{TrapOrigin, trap::PageFaultFlags};
#[repr(transparent)]
struct RawTrapFrame(TrapFrame);
const _: () = {
assert!(size_of::<RawTrapFrame>() == size_of::<TrapFrame>());
assert!(core::mem::align_of::<RawTrapFrame>() == core::mem::align_of::<TrapFrame>());
};
pub struct KernelTrapFrame<'a> {
raw: &'a mut RawTrapFrame,
_not_send: core::marker::PhantomData<*mut ()>,
}
impl<'a> KernelTrapFrame<'a> {
pub const fn origin(&self) -> TrapOrigin {
TrapOrigin::Kernel
}
pub const fn snapshot(&self) -> TrapFrame {
self.raw.0
}
pub fn apply_registers(&mut self, updated: &TrapFrame) {
let spp = self.raw.0.sstatus.spp();
let kernel_gp = self.raw.0.regs.gp;
let kernel_tp = self.raw.0.regs.tp;
self.raw.0 = *updated;
self.raw.0.sstatus.set_spp(spp);
self.raw.0.regs.gp = kernel_gp;
self.raw.0.regs.tp = kernel_tp;
}
pub const fn ip(&self) -> usize {
self.raw.0.ip()
}
pub const fn set_ip(&mut self, ip: usize) {
self.raw.0.set_ip(ip);
}
unsafe fn from_raw(raw: &'a mut RawTrapFrame) -> Self {
debug_assert_eq!(raw.0.origin(), TrapOrigin::Kernel);
Self {
raw,
_not_send: core::marker::PhantomData,
}
}
}
impl core::fmt::Debug for KernelTrapFrame<'_> {
fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
self.snapshot().fmt(formatter)
}
}
#[cfg(not(feature = "tls"))]
core::arch::global_asm!(
include_asm_macros!(),
include_str!("trap.S"),
trapframe_size = const size_of::<RawTrapFrame>(),
kernel_stack_pointer_index = const CPU_KERNEL_STACK_POINTER_OFFSET / size_of::<usize>(),
user_trap_frame_index = const CPU_USER_TRAP_FRAME_OFFSET / size_of::<usize>(),
thread_cpu_base_index = const CURRENT_THREAD_CPU_BASE_OFFSET / size_of::<usize>(),
thread_scratch0_index = const THREAD_SCRATCH0_OFFSET / size_of::<usize>(),
thread_scratch1_index = const THREAD_SCRATCH1_OFFSET / size_of::<usize>(),
);
#[cfg(feature = "tls")]
core::arch::global_asm!(
include_asm_macros!(),
include_str!("trap_tls.S"),
trapframe_size = const size_of::<RawTrapFrame>(),
kernel_stack_pointer_index = const CPU_KERNEL_STACK_POINTER_OFFSET / size_of::<usize>(),
user_trap_frame_index = const CPU_USER_TRAP_FRAME_OFFSET / size_of::<usize>(),
entry_scratch0_index = const CPU_ENTRY_SCRATCH0_OFFSET / size_of::<usize>(),
entry_scratch1_index = const CPU_ENTRY_SCRATCH1_OFFSET / size_of::<usize>(),
);
fn handle_breakpoint(tf: &mut KernelTrapFrame<'_>) {
debug!("Exception(Breakpoint) @ {:#x} ", tf.raw.0.sepc);
if crate::trap::breakpoint_handler(tf) {
return;
}
tf.set_ip(tf.ip() + 2);
}
fn handle_page_fault(tf: &mut KernelTrapFrame<'_>, access_flags: PageFaultFlags) {
let vaddr = va!(stval::read());
if crate::trap::call_page_fault_handler_with_parent_irqs(
vaddr,
access_flags,
tf.raw.0.sstatus.spie(),
) {
return;
}
#[cfg(feature = "exception-table")]
if tf.raw.0.fixup_exception() {
return;
}
let snapshot = tf.snapshot();
let bt = snapshot.backtrace();
panic!(
"Unhandled Supervisor Page Fault @ {:#x}, fault_vaddr={:#x} ({:?}):\n{:#x?}\n{}",
tf.raw.0.sepc,
vaddr,
access_flags,
snapshot,
bt.kind("trap")
);
}
#[unsafe(no_mangle)]
unsafe extern "C" fn riscv_trap_handler(raw_tf: *mut RawTrapFrame) {
let raw = unsafe { &mut *raw_tf };
let mut tf = unsafe { KernelTrapFrame::from_raw(raw) };
handle_trap(&mut tf);
}
fn handle_trap(tf: &mut KernelTrapFrame<'_>) {
let scause = scause::read();
if let Ok(cause) = scause.cause().try_into::<I, E>() {
match cause {
Trap::Exception(E::LoadPageFault) => handle_page_fault(tf, PageFaultFlags::READ),
Trap::Exception(E::StorePageFault) => handle_page_fault(tf, PageFaultFlags::WRITE),
Trap::Exception(E::InstructionPageFault) => {
handle_page_fault(tf, PageFaultFlags::EXECUTE)
}
Trap::Exception(E::Breakpoint) => handle_breakpoint(tf),
Trap::Interrupt(_) => {
crate::trap::dispatch_irq(scause.bits());
}
_ => {
let snapshot = tf.snapshot();
let bt = snapshot.backtrace();
panic!(
"Unhandled trap {:?} @ {:#x}, stval={:#x}:\n{:#x?}\n{}",
cause,
tf.raw.0.sepc,
stval::read(),
snapshot,
bt.kind("trap")
);
}
}
} else {
let snapshot = tf.snapshot();
let bt = snapshot.backtrace();
panic!(
"Unknown trap {:#x?} @ {:#x}:\n{:#x?}\n{}",
scause.cause(),
tf.raw.0.sepc,
snapshot,
bt.kind("trap")
);
}
#[cfg(feature = "fp-simd")]
tf.raw.0.sstatus.set_fs(sstatus::read().fs());
}