use core::{
marker::PhantomData,
mem::ManuallyDrop,
ptr::NonNull,
sync::atomic::{AtomicU32, Ordering},
};
use crate::{CpuLocalError, CpuPin};
const PREEMPT_NO_PENDING: u32 = 1 << 31;
const PREEMPT_DEPTH_MASK: u32 = !PREEMPT_NO_PENDING;
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct PreemptionSnapshot {
depth: u32,
pending: bool,
}
impl PreemptionSnapshot {
pub const fn depth(self) -> u32 {
self.depth
}
pub const fn is_pending(self) -> bool {
self.pending
}
}
#[must_use = "every entered preemption token must be finished exactly once"]
pub struct PreemptionToken {
owner: NonNull<PreemptionState>,
_not_send_or_sync: PhantomData<*mut ()>,
}
#[must_use = "pending preemption must be released by the external safe-point owner"]
pub struct PendingPreemption {
owner: NonNull<PreemptionState>,
_not_send_or_sync: PhantomData<*mut ()>,
}
#[must_use]
pub enum PreemptionExit {
Nested,
Enabled,
Pending(PendingPreemption),
}
#[repr(transparent)]
pub(crate) struct PreemptionState(AtomicU32);
impl PreemptionState {
pub(crate) const fn new() -> Self {
Self(AtomicU32::new(PREEMPT_NO_PENDING))
}
pub(crate) const fn bootstrap_disabled() -> Self {
Self(AtomicU32::new(PREEMPT_NO_PENDING | 1))
}
fn snapshot(&self) -> PreemptionSnapshot {
let state = self.0.load(Ordering::Relaxed);
PreemptionSnapshot {
depth: state & PREEMPT_DEPTH_MASK,
pending: state & PREEMPT_NO_PENDING == 0,
}
}
fn set_pending(&self) {
self.0.fetch_and(PREEMPT_DEPTH_MASK, Ordering::Relaxed);
}
fn clear_pending(&self) {
self.0.fetch_or(PREEMPT_NO_PENDING, Ordering::Relaxed);
}
#[cfg(any(not(target_arch = "x86_64"), feature = "host-test"))]
fn enter(&self) {
let previous = self.0.fetch_add(1, Ordering::Relaxed);
assert_ne!(
previous & PREEMPT_DEPTH_MASK,
PREEMPT_DEPTH_MASK,
"preemption nesting overflow"
);
}
fn finish(&self) -> PreemptionExit {
loop {
let state = self.0.load(Ordering::Relaxed);
let depth = state & PREEMPT_DEPTH_MASK;
assert!(depth > 0, "unbalanced preemption exit");
if depth == 1 && state & PREEMPT_NO_PENDING == 0 {
return PreemptionExit::Pending(PendingPreemption::new(self));
}
let next = state - 1;
if self
.0
.compare_exchange_weak(state, next, Ordering::Relaxed, Ordering::Relaxed)
.is_ok()
{
return if depth == 1 {
PreemptionExit::Enabled
} else {
PreemptionExit::Nested
};
}
}
}
fn release_pending(&self) {
assert_eq!(
self.0
.compare_exchange(1, 0, Ordering::Relaxed, Ordering::Relaxed),
Ok(1),
"pending preemption no longer owns the final depth"
);
}
fn release_bootstrap(&self) {
assert_eq!(
self.0.compare_exchange(
PREEMPT_NO_PENDING | 1,
PREEMPT_NO_PENDING,
Ordering::Relaxed,
Ordering::Relaxed,
),
Ok(PREEMPT_NO_PENDING | 1),
"bootstrap preemption depth must be released exactly once"
);
}
#[cfg(any(all(target_arch = "x86_64", not(feature = "host-test")), test))]
fn release_initial_switch(&self) -> bool {
loop {
let state = self.0.load(Ordering::Relaxed);
if state == PREEMPT_NO_PENDING {
return false;
}
if state != (PREEMPT_NO_PENDING | 1) && state != 1 {
panic!("initial context switch found invalid preemption state {state:#x}");
}
if self
.0
.compare_exchange_weak(
state,
PREEMPT_NO_PENDING,
Ordering::Relaxed,
Ordering::Relaxed,
)
.is_ok()
{
return true;
}
}
}
}
impl PreemptionToken {
fn new(owner: &PreemptionState) -> Self {
Self {
owner: NonNull::from(owner),
_not_send_or_sync: PhantomData,
}
}
#[doc(hidden)]
pub fn into_raw(self) -> usize {
let token = ManuallyDrop::new(self);
token.owner.as_ptr() as usize
}
#[doc(hidden)]
pub unsafe fn from_raw(raw: usize) -> Option<Self> {
if !raw.is_multiple_of(core::mem::align_of::<PreemptionState>()) {
return None;
}
NonNull::new(raw as *mut PreemptionState).map(|owner| Self {
owner,
_not_send_or_sync: PhantomData,
})
}
fn state(&self) -> &PreemptionState {
unsafe { self.owner.as_ref() }
}
#[cfg(any(all(target_arch = "x86_64", not(feature = "host-test")), test))]
fn handoff_after_context_switch(self, resumed_owner: &PreemptionState) -> Self {
if self.owner == NonNull::from(resumed_owner) {
self
} else {
Self::new(resumed_owner)
}
}
}
impl PendingPreemption {
fn new(owner: &PreemptionState) -> Self {
Self {
owner: NonNull::from(owner),
_not_send_or_sync: PhantomData,
}
}
pub fn release(self) {
unsafe { self.owner.as_ref() }.release_pending();
}
}
#[inline(always)]
pub fn enter_preemption() -> PreemptionToken {
#[cfg(all(target_arch = "x86_64", not(feature = "host-test")))]
{
unsafe { crate::register::enter_x86_preemption() };
let owner = crate::register::current_area()
.unwrap_or_else(|_| crate::register::fatal_register_invariant())
.runtime_anchor()
.preemption_state();
PreemptionToken::new(owner)
}
#[cfg(any(not(target_arch = "x86_64"), feature = "host-test"))]
{
let current = unsafe { crate::current_context_unpinned() }
.unwrap_or_else(|_| crate::register::fatal_register_invariant());
let owner = unsafe { current.as_ref() }.preemption_state();
owner.enter();
PreemptionToken::new(owner)
}
}
pub fn preemption_snapshot(pin: &CpuPin<'_>) -> Result<PreemptionSnapshot, CpuLocalError> {
Ok(selected_state(pin)?.snapshot())
}
pub fn set_preemption_pending(pin: &CpuPin<'_>) -> Result<(), CpuLocalError> {
selected_state(pin)?.set_pending();
Ok(())
}
pub fn clear_preemption_pending(pin: &CpuPin<'_>) -> Result<(), CpuLocalError> {
selected_state(pin)?.clear_pending();
Ok(())
}
pub fn finish_preemption(token: PreemptionToken) -> PreemptionExit {
token.state().finish()
}
#[doc(hidden)]
pub fn handoff_preemption_after_context_switch(
pin: &CpuPin<'_>,
token: PreemptionToken,
) -> Result<PreemptionToken, CpuLocalError> {
let owner = selected_state(pin)?;
#[cfg(all(target_arch = "x86_64", not(feature = "host-test")))]
{
Ok(token.handoff_after_context_switch(owner))
}
#[cfg(any(not(target_arch = "x86_64"), feature = "host-test"))]
{
assert_eq!(
token.owner,
NonNull::from(owner),
"context-owned preemption token changed owner across a context switch"
);
Ok(token)
}
}
#[doc(hidden)]
pub fn release_bootstrap_preemption(pin: &CpuPin<'_>) -> Result<(), CpuLocalError> {
selected_state(pin)?.release_bootstrap();
Ok(())
}
#[doc(hidden)]
pub fn release_initial_context_preemption(pin: &CpuPin<'_>) -> Result<bool, CpuLocalError> {
#[cfg(all(target_arch = "x86_64", not(feature = "host-test")))]
{
Ok(selected_state(pin)?.release_initial_switch())
}
#[cfg(any(not(target_arch = "x86_64"), feature = "host-test"))]
{
let snapshot = selected_state(pin)?.snapshot();
assert_eq!(
snapshot.depth(),
0,
"new context-owned preemption state must start enabled"
);
Ok(false)
}
}
fn selected_state(pin: &CpuPin<'_>) -> Result<&'static PreemptionState, CpuLocalError> {
#[cfg(all(target_arch = "x86_64", not(feature = "host-test")))]
{
Ok(pin.area().runtime_anchor().preemption_state())
}
#[cfg(any(not(target_arch = "x86_64"), feature = "host-test"))]
{
let current = crate::current_context(pin)?;
Ok(unsafe { current.as_ref() }.preemption_state())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn enter_on(state: &PreemptionState) -> PreemptionToken {
state.enter();
PreemptionToken::new(state)
}
#[test]
fn nested_and_final_exits_are_linear() {
let state = PreemptionState::new();
let outer = enter_on(&state);
let inner = enter_on(&state);
assert!(matches!(finish_preemption(inner), PreemptionExit::Nested));
assert_eq!(state.snapshot().depth(), 1);
assert!(matches!(finish_preemption(outer), PreemptionExit::Enabled));
assert_eq!(state.snapshot().depth(), 0);
}
#[test]
fn pending_exit_reserves_depth_until_release() {
let state = PreemptionState::new();
let token = enter_on(&state);
state.set_pending();
let PreemptionExit::Pending(pending) = finish_preemption(token) else {
panic!("final pending exit must retain its depth");
};
assert_eq!(state.snapshot().depth(), 1);
pending.release();
assert_eq!(state.snapshot().depth(), 0);
assert!(state.snapshot().is_pending());
}
#[test]
fn bootstrap_starts_disabled() {
let state = PreemptionState::bootstrap_disabled();
assert_eq!(state.snapshot().depth(), 1);
assert!(!state.snapshot().is_pending());
}
#[test]
fn initial_switch_discards_outgoing_pending_mirror() {
let state = PreemptionState::bootstrap_disabled();
state.set_pending();
assert!(state.release_initial_switch());
assert_eq!(state.snapshot().depth(), 0);
assert!(!state.snapshot().is_pending());
}
#[test]
#[should_panic(expected = "pending preemption no longer owns the final depth")]
fn pending_depth_cannot_be_consumed_twice() {
let state = PreemptionState(AtomicU32::new(1));
let first = PendingPreemption::new(&state);
let duplicate = PendingPreemption::new(&state);
first.release();
duplicate.release();
}
#[test]
#[should_panic(expected = "bootstrap preemption depth must be released exactly once")]
fn bootstrap_depth_cannot_be_released_twice() {
let state = PreemptionState::bootstrap_disabled();
state.release_bootstrap();
state.release_bootstrap();
}
#[test]
fn token_stays_bound_to_the_entry_owner() {
let original_context = PreemptionState::new();
let migrated_context = PreemptionState::new();
let original_cpu = PreemptionState::new();
let migrated_cpu = PreemptionState::new();
let context_token = enter_on(&original_context);
let cpu_token = enter_on(&original_cpu);
migrated_context.enter();
migrated_cpu.enter();
assert!(matches!(
finish_preemption(context_token),
PreemptionExit::Enabled
));
assert!(matches!(
finish_preemption(cpu_token),
PreemptionExit::Enabled
));
assert_eq!(original_context.snapshot().depth(), 0);
assert_eq!(original_cpu.snapshot().depth(), 0);
assert_eq!(migrated_context.snapshot().depth(), 1);
assert_eq!(migrated_cpu.snapshot().depth(), 1);
}
#[test]
fn cpu_owned_switch_token_handoff_follows_the_resumed_cpu() {
let original_cpu = PreemptionState::new();
let resumed_cpu = PreemptionState::new();
let suspended_switch = enter_on(&original_cpu);
assert!(original_cpu.release_initial_switch());
resumed_cpu.enter();
let resumed_switch = suspended_switch.handoff_after_context_switch(&resumed_cpu);
assert!(matches!(
finish_preemption(resumed_switch),
PreemptionExit::Enabled
));
assert_eq!(original_cpu.snapshot().depth(), 0);
assert_eq!(resumed_cpu.snapshot().depth(), 0);
}
#[test]
fn malformed_raw_owner_is_rejected() {
assert!(unsafe { PreemptionToken::from_raw(1) }.is_none());
}
}