use core::{
arch::naked_asm,
mem::{align_of, offset_of, size_of},
ptr::NonNull,
};
use ax_memory_addr::VirtAddr;
use cpu_local::{CurrentThreadHeader, PreparedThreadSwitch};
use riscv::register::sstatus::{self, FS};
use crate::{KernelTlsBase, TaskLocalState};
#[allow(missing_docs)]
#[repr(C)]
#[derive(Debug, Default, Clone, Copy)]
pub struct GeneralRegisters {
pub zero: usize,
pub ra: usize,
pub sp: usize,
pub gp: usize,
pub tp: usize,
pub t0: usize,
pub t1: usize,
pub t2: usize,
pub s0: usize,
pub s1: usize,
pub a0: usize,
pub a1: usize,
pub a2: usize,
pub a3: usize,
pub a4: usize,
pub a5: usize,
pub a6: usize,
pub a7: usize,
pub s2: usize,
pub s3: usize,
pub s4: usize,
pub s5: usize,
pub s6: usize,
pub s7: usize,
pub s8: usize,
pub s9: usize,
pub s10: usize,
pub s11: usize,
pub t3: usize,
pub t4: usize,
pub t5: usize,
pub t6: usize,
}
#[repr(C)]
#[derive(Debug, Clone, Copy)]
pub struct FpState {
pub fp: [u64; 32],
pub fcsr: usize,
pub fs: FS,
}
impl Default for FpState {
fn default() -> Self {
Self {
fs: FS::Initial,
fp: [0; 32],
fcsr: 0,
}
}
}
#[cfg(feature = "fp-simd")]
impl FpState {
#[inline]
pub fn restore(&self) {
unsafe { restore_fp_registers(self) }
}
#[inline]
pub fn save(&mut self) {
unsafe { save_fp_registers(self) }
}
#[inline]
pub fn clear() {
unsafe { clear_fp_registers() }
}
pub fn switch_to(&mut self, next_fp_state: &FpState) {
let current_fs = sstatus::read().fs();
if current_fs == FS::Dirty {
self.save();
self.fs = FS::Clean;
}
if matches!(next_fp_state.fs, FS::Clean | FS::Initial) {
unsafe { sstatus::set_fs(FS::Dirty) };
}
match next_fp_state.fs {
FS::Clean => next_fp_state.restore(),
FS::Initial => FpState::clear(), FS::Off => {} FS::Dirty => unreachable!("FP state of the next task should not be dirty"),
}
unsafe { sstatus::set_fs(next_fp_state.fs) };
}
}
#[repr(C)]
#[derive(Debug, Clone, Copy)]
pub struct TrapFrame {
pub regs: GeneralRegisters,
pub sepc: usize,
pub sstatus: sstatus::Sstatus,
}
impl Default for TrapFrame {
fn default() -> Self {
Self {
regs: GeneralRegisters::default(),
sepc: 0,
sstatus: sstatus::Sstatus::from_bits(0),
}
}
}
impl TrapFrame {
pub fn origin(&self) -> crate::TrapOrigin {
match self.sstatus.spp() {
sstatus::SPP::Supervisor => crate::TrapOrigin::Kernel,
sstatus::SPP::User => crate::TrapOrigin::User,
}
}
pub const fn arg0(&self) -> usize {
self.regs.a0
}
pub const fn set_arg0(&mut self, a0: usize) {
self.regs.a0 = a0;
}
pub const fn arg1(&self) -> usize {
self.regs.a1
}
pub const fn set_arg1(&mut self, a1: usize) {
self.regs.a1 = a1;
}
pub const fn arg2(&self) -> usize {
self.regs.a2
}
pub const fn set_arg2(&mut self, a2: usize) {
self.regs.a2 = a2;
}
pub const fn arg3(&self) -> usize {
self.regs.a3
}
pub const fn set_arg3(&mut self, a3: usize) {
self.regs.a3 = a3;
}
pub const fn arg4(&self) -> usize {
self.regs.a4
}
pub const fn set_arg4(&mut self, a4: usize) {
self.regs.a4 = a4;
}
pub const fn arg5(&self) -> usize {
self.regs.a5
}
pub const fn set_arg5(&mut self, a5: usize) {
self.regs.a5 = a5;
}
pub const fn sysno(&self) -> usize {
self.regs.a7
}
pub const fn set_sysno(&mut self, a7: usize) {
self.regs.a7 = a7;
}
pub const fn ip(&self) -> usize {
self.sepc
}
pub const fn set_ip(&mut self, pc: usize) {
self.sepc = pc;
}
pub const fn sp(&self) -> usize {
self.regs.sp
}
pub const fn set_sp(&mut self, sp: usize) {
self.regs.sp = sp;
}
pub const fn retval(&self) -> usize {
self.regs.a0
}
pub const fn set_retval(&mut self, a0: usize) {
self.regs.a0 = a0;
}
pub const fn set_ra(&mut self, ra: usize) {
self.regs.ra = ra;
}
pub const fn tls(&self) -> usize {
self.regs.tp
}
pub const fn set_tls(&mut self, tls_area: usize) {
self.regs.tp = tls_area;
}
pub fn backtrace(&self) -> axbacktrace::Backtrace {
axbacktrace::Backtrace::capture_trap(self.regs.s0 as _, self.sepc as _, self.regs.ra as _)
}
}
#[allow(missing_docs)]
#[repr(C)]
#[derive(Debug, Default)]
pub struct TaskContext {
pub ra: usize, pub sp: usize,
pub s0: usize, pub s1: usize,
pub s2: usize, pub s3: usize,
pub s4: usize,
pub s5: usize,
pub s6: usize,
pub s7: usize,
pub s8: usize,
pub s9: usize,
pub s10: usize,
pub s11: usize,
task_local: TaskLocalState,
#[cfg(feature = "uspace")]
page_table_root: ax_memory_addr::PhysAddr,
#[cfg(feature = "fp-simd")]
pub fp_state: FpState,
}
const _: () = {
assert!(size_of::<KernelTlsBase>() == size_of::<usize>());
assert!(align_of::<KernelTlsBase>() == align_of::<usize>());
assert!(offset_of!(TaskContext, ra) == 0);
assert!(offset_of!(TaskContext, sp) == offset_of!(TaskContext, ra) + size_of::<usize>());
assert!(offset_of!(TaskContext, task_local) % size_of::<usize>() == 0);
};
impl TaskContext {
pub fn new() -> Self {
Self {
#[cfg(feature = "uspace")]
page_table_root: crate::asm::read_kernel_page_table(),
..Self::default()
}
}
pub fn init(&mut self, entry: usize, kstack_top: VirtAddr, tls_area: KernelTlsBase) {
self.sp = kstack_top.as_usize();
self.ra = entry;
self.task_local.set_kernel_tls(tls_area);
}
pub fn set_current_header(&mut self, header: NonNull<CurrentThreadHeader>) {
self.task_local.set_current_header(header);
}
pub const fn current_header(&self) -> Option<NonNull<CurrentThreadHeader>> {
self.task_local.current_header()
}
#[cfg(feature = "uspace")]
pub fn set_page_table_root(&mut self, page_table_root: ax_memory_addr::PhysAddr) {
self.page_table_root = page_table_root;
}
pub fn prepare_switch_to(&mut self, _next_ctx: &Self) {
#[cfg(feature = "fp-simd")]
{
self.fp_state.switch_to(&_next_ctx.fp_state);
}
#[cfg(feature = "uspace")]
if self.page_table_root != _next_ctx.page_table_root {
unsafe { crate::asm::write_user_page_table(_next_ctx.page_table_root) };
crate::asm::flush_tlb(None);
}
}
#[inline(always)]
pub unsafe fn switch_to_prepared(
&mut self,
next_ctx: &Self,
prepared: PreparedThreadSwitch<'_>,
) {
assert_eq!(
next_ctx.current_header(),
Some(prepared.next_header()),
"prepared switch token must belong to the next task context",
);
unsafe { prepared.commit() };
unsafe { context_switch_raw(self, next_ctx) }
}
}
#[cfg(feature = "fp-simd")]
#[unsafe(naked)]
unsafe extern "C" fn save_fp_registers(fp_state: &mut FpState) {
naked_asm!(
include_fp_asm_macros!(),
"
PUSH_FLOAT_REGS a0
frcsr t0
STR t0, a0, 32
ret"
)
}
#[cfg(feature = "fp-simd")]
#[unsafe(naked)]
unsafe extern "C" fn restore_fp_registers(fp_state: &FpState) {
naked_asm!(
include_fp_asm_macros!(),
"
POP_FLOAT_REGS a0
LDR t0, a0, 32
fscsr x0, t0
ret"
)
}
#[cfg(feature = "fp-simd")]
#[unsafe(naked)]
unsafe extern "C" fn clear_fp_registers() {
naked_asm!(
include_fp_asm_macros!(),
"
CLEAR_FLOAT_REGS
ret"
)
}
#[cfg(feature = "tls")]
#[unsafe(naked)]
unsafe extern "C" fn context_switch_raw(_current_task: &mut TaskContext, _next_task: &TaskContext) {
naked_asm!(
include_asm_macros!(),
"
// save old context (callee-saved registers)
STR ra, a0, {ra_index}
STR sp, a0, {sp_index}
STR s0, a0, {s0_index}
STR s1, a0, {s1_index}
STR s2, a0, {s2_index}
STR s3, a0, {s3_index}
STR s4, a0, {s4_index}
STR s5, a0, {s5_index}
STR s6, a0, {s6_index}
STR s7, a0, {s7_index}
STR s8, a0, {s8_index}
STR s9, a0, {s9_index}
STR s10, a0, {s10_index}
STR s11, a0, {s11_index}
STR tp, a0, {kernel_tls_index}
// restore new context
LDR s11, a1, {s11_index}
LDR s10, a1, {s10_index}
LDR s9, a1, {s9_index}
LDR s8, a1, {s8_index}
LDR s7, a1, {s7_index}
LDR s6, a1, {s6_index}
LDR s5, a1, {s5_index}
LDR s4, a1, {s4_index}
LDR s3, a1, {s3_index}
LDR s2, a1, {s2_index}
LDR s1, a1, {s1_index}
LDR s0, a1, {s0_index}
LDR sp, a1, {sp_index}
LDR tp, a1, {kernel_tls_index}
LDR ra, a1, {ra_index}
ret",
ra_index = const offset_of!(TaskContext, ra) / size_of::<usize>(),
sp_index = const offset_of!(TaskContext, sp) / size_of::<usize>(),
s0_index = const offset_of!(TaskContext, s0) / size_of::<usize>(),
s1_index = const offset_of!(TaskContext, s1) / size_of::<usize>(),
s2_index = const offset_of!(TaskContext, s2) / size_of::<usize>(),
s3_index = const offset_of!(TaskContext, s3) / size_of::<usize>(),
s4_index = const offset_of!(TaskContext, s4) / size_of::<usize>(),
s5_index = const offset_of!(TaskContext, s5) / size_of::<usize>(),
s6_index = const offset_of!(TaskContext, s6) / size_of::<usize>(),
s7_index = const offset_of!(TaskContext, s7) / size_of::<usize>(),
s8_index = const offset_of!(TaskContext, s8) / size_of::<usize>(),
s9_index = const offset_of!(TaskContext, s9) / size_of::<usize>(),
s10_index = const offset_of!(TaskContext, s10) / size_of::<usize>(),
s11_index = const offset_of!(TaskContext, s11) / size_of::<usize>(),
kernel_tls_index = const (offset_of!(TaskContext, task_local)
+ offset_of!(TaskLocalState, kernel_tls)) / size_of::<usize>(),
)
}
#[cfg(not(feature = "tls"))]
#[unsafe(naked)]
unsafe extern "C" fn context_switch_raw(_current_task: &mut TaskContext, _next_task: &TaskContext) {
naked_asm!(
include_asm_macros!(),
"
// Save old context. CPU-owned sscratch and task-owned tp are not
// folded into the generic callee-saved register image.
STR ra, a0, {ra_index}
STR sp, a0, {sp_index}
STR s0, a0, {s0_index}
STR s1, a0, {s1_index}
STR s2, a0, {s2_index}
STR s3, a0, {s3_index}
STR s4, a0, {s4_index}
STR s5, a0, {s5_index}
STR s6, a0, {s6_index}
STR s7, a0, {s7_index}
STR s8, a0, {s8_index}
STR s9, a0, {s9_index}
STR s10, a0, {s10_index}
STR s11, a0, {s11_index}
// Restore all next state, then make tp current immediately before the
// direct return into the next task. gp remains the psABI global pointer.
LDR s11, a1, {s11_index}
LDR s10, a1, {s10_index}
LDR s9, a1, {s9_index}
LDR s8, a1, {s8_index}
LDR s7, a1, {s7_index}
LDR s6, a1, {s6_index}
LDR s5, a1, {s5_index}
LDR s4, a1, {s4_index}
LDR s3, a1, {s3_index}
LDR s2, a1, {s2_index}
LDR s1, a1, {s1_index}
LDR s0, a1, {s0_index}
LDR sp, a1, {sp_index}
LDR tp, a1, {current_header_index}
LDR ra, a1, {ra_index}
ret",
ra_index = const offset_of!(TaskContext, ra) / size_of::<usize>(),
sp_index = const offset_of!(TaskContext, sp) / size_of::<usize>(),
s0_index = const offset_of!(TaskContext, s0) / size_of::<usize>(),
s1_index = const offset_of!(TaskContext, s1) / size_of::<usize>(),
s2_index = const offset_of!(TaskContext, s2) / size_of::<usize>(),
s3_index = const offset_of!(TaskContext, s3) / size_of::<usize>(),
s4_index = const offset_of!(TaskContext, s4) / size_of::<usize>(),
s5_index = const offset_of!(TaskContext, s5) / size_of::<usize>(),
s6_index = const offset_of!(TaskContext, s6) / size_of::<usize>(),
s7_index = const offset_of!(TaskContext, s7) / size_of::<usize>(),
s8_index = const offset_of!(TaskContext, s8) / size_of::<usize>(),
s9_index = const offset_of!(TaskContext, s9) / size_of::<usize>(),
s10_index = const offset_of!(TaskContext, s10) / size_of::<usize>(),
s11_index = const offset_of!(TaskContext, s11) / size_of::<usize>(),
current_header_index = const (offset_of!(TaskContext, task_local)
+ offset_of!(TaskLocalState, current_header)) / size_of::<usize>(),
)
}