use alloc::sync::Arc;
use core::{
hint::spin_loop,
ptr,
sync::atomic::{AtomicPtr, AtomicU64, Ordering},
};
use crate::thread::ThreadWakeHandle;
const REGISTRATION_PHASE_BITS: u32 = 2;
const REGISTRATION_PHASE_MASK: u64 = (1 << REGISTRATION_PHASE_BITS) - 1;
const REGISTRATION_GENERATION_MAX: u64 = u64::MAX >> REGISTRATION_PHASE_BITS;
const IRQ_NOTIFY_CAS_BUDGET: usize = 8;
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
#[repr(u64)]
enum RegistrationPhase {
Detached = 0,
Attached = 1,
Notifying = 2,
}
const fn registration_state(generation: u64, phase: RegistrationPhase) -> u64 {
(generation << REGISTRATION_PHASE_BITS) | phase as u64
}
const fn registration_generation(state: u64) -> u64 {
state >> REGISTRATION_PHASE_BITS
}
#[repr(align(8))]
struct WaiterSentinel {
_tag: u8,
}
static PENDING_WAITER_SENTINEL: WaiterSentinel = WaiterSentinel { _tag: 1 };
static NOTIFYING_WAITER_SENTINEL: WaiterSentinel = WaiterSentinel { _tag: 2 };
static NOTIFYING_PENDING_WAITER_SENTINEL: WaiterSentinel = WaiterSentinel { _tag: 3 };
fn waiter_sentinel(sentinel: &'static WaiterSentinel) -> *mut IrqWaitNode {
ptr::from_ref(sentinel).cast_mut().cast()
}
fn pending_waiter() -> *mut IrqWaitNode {
waiter_sentinel(&PENDING_WAITER_SENTINEL)
}
fn notifying_waiter() -> *mut IrqWaitNode {
waiter_sentinel(&NOTIFYING_WAITER_SENTINEL)
}
fn notifying_pending_waiter() -> *mut IrqWaitNode {
waiter_sentinel(&NOTIFYING_PENDING_WAITER_SENTINEL)
}
fn is_notification_sentinel(waiter: *mut IrqWaitNode) -> bool {
waiter == notifying_waiter() || waiter == notifying_pending_waiter()
}
fn registration_phase(state: u64) -> RegistrationPhase {
match state & REGISTRATION_PHASE_MASK {
0 => RegistrationPhase::Detached,
1 => RegistrationPhase::Attached,
2 => RegistrationPhase::Notifying,
_ => unreachable!("registration phase exceeds its bit mask"),
}
}
#[derive(Debug)]
enum IrqWaitWake {
Thread(ThreadWakeHandle),
}
impl IrqWaitWake {
fn wake(&self) -> crate::thread::WakeResult {
match self {
Self::Thread(wake) => wake.wake(),
}
}
}
#[derive(Debug)]
struct IrqWaitNode {
wake: IrqWaitWake,
state: AtomicU64,
}
impl IrqWaitNode {
fn new(wake: IrqWaitWake) -> Self {
Self {
wake,
state: AtomicU64::new(registration_state(0, RegistrationPhase::Detached)),
}
}
fn reserve(&self) -> Option<u64> {
let mut state = self.state.load(Ordering::Acquire);
loop {
if registration_phase(state) != RegistrationPhase::Detached {
return None;
}
let generation = registration_generation(state)
.checked_add(1)
.filter(|generation| *generation <= REGISTRATION_GENERATION_MAX)
.expect("IRQ wait registration generation exhausted");
let attached = registration_state(generation, RegistrationPhase::Attached);
match self.state.compare_exchange_weak(
state,
attached,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => return Some(generation),
Err(observed) => state = observed,
}
}
}
fn cancel(&self, generation: u64) {
self.state
.compare_exchange(
registration_state(generation, RegistrationPhase::Attached),
registration_state(generation, RegistrationPhase::Detached),
Ordering::Release,
Ordering::Acquire,
)
.expect("only an attached IRQ wait registration can be cancelled");
}
fn begin_notification(&self) -> u64 {
let state = self.state.load(Ordering::Acquire);
let generation = registration_generation(state);
assert_eq!(
registration_phase(state),
RegistrationPhase::Attached,
"an IRQ wait cell took a registration it no longer owned"
);
self.state
.compare_exchange(
state,
registration_state(generation, RegistrationPhase::Notifying),
Ordering::AcqRel,
Ordering::Acquire,
)
.expect("IRQ wait registration ownership changed after cell removal");
generation
}
fn finish_notification(&self, generation: u64) {
self.state
.compare_exchange(
registration_state(generation, RegistrationPhase::Notifying),
registration_state(generation, RegistrationPhase::Detached),
Ordering::Release,
Ordering::Acquire,
)
.expect("IRQ wait notification generation changed while in flight");
}
fn is_attached(&self, generation: u64) -> bool {
self.state.load(Ordering::Acquire)
== registration_state(generation, RegistrationPhase::Attached)
}
fn is_quiescent(&self, generation: u64) -> bool {
let state = self.state.load(Ordering::Acquire);
registration_generation(state) != generation
|| registration_phase(state) == RegistrationPhase::Detached
}
}
fn publish_cell_owner(node: Arc<IrqWaitNode>) -> *mut IrqWaitNode {
Arc::into_raw(node).cast_mut()
}
unsafe fn take_cell_owner(node: *mut IrqWaitNode) -> Arc<IrqWaitNode> {
unsafe {
Arc::from_raw(node)
}
}
#[derive(Debug)]
pub struct IrqWaitRegistration {
node: Arc<IrqWaitNode>,
}
impl IrqWaitRegistration {
pub fn new(wake: ThreadWakeHandle) -> Self {
Self {
node: Arc::new(IrqWaitNode::new(IrqWaitWake::Thread(wake))),
}
}
}
#[must_use = "an IRQ wait token must enter its drain lifetime before storage is reused"]
pub struct IrqWaitToken<'cell> {
registration: Arc<IrqWaitNode>,
generation: u64,
cell: &'cell IrqWaitCell,
}
impl IrqWaitToken<'_> {
pub const fn generation(&self) -> u64 {
self.generation
}
pub fn is_attached(&self) -> bool {
self.registration.is_attached(self.generation)
}
pub fn detach(self) -> IrqWaitDrain {
let cell = self.cell;
cell.detach(self)
}
fn belongs_to(&self, cell: &IrqWaitCell) -> bool {
ptr::eq(self.cell, cell)
}
}
impl core::fmt::Debug for IrqWaitToken<'_> {
fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
formatter
.debug_struct("IrqWaitToken")
.field("generation", &self.generation)
.field("attached", &self.is_attached())
.finish()
}
}
#[must_use = "an IRQ wait drain must finish before registration storage is reused"]
pub struct IrqWaitDrain {
registration: Arc<IrqWaitNode>,
generation: u64,
}
impl IrqWaitDrain {
pub fn is_quiescent(&self) -> bool {
self.registration.is_quiescent(self.generation)
}
pub fn finish(self) {
while !self.is_quiescent() {
spin_loop();
}
}
}
impl core::fmt::Debug for IrqWaitDrain {
fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
formatter
.debug_struct("IrqWaitDrain")
.field("generation", &self.generation)
.field("quiescent", &self.is_quiescent())
.finish()
}
}
#[derive(Debug)]
pub enum IrqRegisterResult<'cell> {
Registered(IrqWaitToken<'cell>),
ConsumedPending,
NotificationInFlight(IrqWaitToken<'cell>),
Occupied,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum IrqNotifyResult {
Notified,
Pending,
}
#[derive(Debug)]
pub struct IrqWaitCell {
waiter: AtomicPtr<IrqWaitNode>,
}
impl IrqWaitCell {
pub const fn new() -> Self {
Self {
waiter: AtomicPtr::new(ptr::null_mut()),
}
}
pub fn register<'cell>(
&'cell self,
registration: &IrqWaitRegistration,
) -> IrqRegisterResult<'cell> {
let registration = Arc::clone(®istration.node);
let Some(generation) = registration.reserve() else {
return IrqRegisterResult::Occupied;
};
let registration_ptr = publish_cell_owner(Arc::clone(®istration));
let token = IrqWaitToken {
registration,
generation,
cell: self,
};
let pending = pending_waiter();
let mut observed = self.waiter.load(Ordering::Acquire);
loop {
if observed == pending {
match self.waiter.compare_exchange(
pending,
ptr::null_mut(),
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => {
token.registration.cancel(generation);
unsafe { drop(take_cell_owner(registration_ptr)) };
return IrqRegisterResult::ConsumedPending;
}
Err(current) => {
observed = current;
continue;
}
}
}
if !observed.is_null() {
token.registration.cancel(generation);
unsafe { drop(take_cell_owner(registration_ptr)) };
return IrqRegisterResult::Occupied;
}
match self.waiter.compare_exchange(
ptr::null_mut(),
registration_ptr,
Ordering::Release,
Ordering::Acquire,
) {
Ok(_) => break,
Err(current) => observed = current,
}
}
if self.waiter.load(Ordering::Acquire) == registration_ptr {
IrqRegisterResult::Registered(token)
} else {
IrqRegisterResult::NotificationInFlight(token)
}
}
fn detach(&self, token: IrqWaitToken<'_>) -> IrqWaitDrain {
assert!(
token.belongs_to(self),
"an IRQ wait token must be detached by its publishing cell"
);
let registration = token.registration;
let state = registration.state.load(Ordering::Acquire);
if registration_generation(state) == token.generation {
let registration_ptr = Arc::as_ptr(®istration).cast_mut();
if self
.waiter
.compare_exchange(
registration_ptr,
ptr::null_mut(),
Ordering::AcqRel,
Ordering::Acquire,
)
.is_ok()
{
registration.cancel(token.generation);
unsafe { drop(take_cell_owner(registration_ptr)) };
}
}
IrqWaitDrain {
registration,
generation: token.generation,
}
}
pub fn notify(&self) -> IrqNotifyResult {
let _preempt = (!crate::runtime::task_runtime::in_hard_irq())
.then(crate::runtime::lock::PreemptScope::enter);
self.notify_claimed()
}
fn notify_claimed(&self) -> IrqNotifyResult {
let pending = pending_waiter();
let notifying = notifying_waiter();
let notifying_pending = notifying_pending_waiter();
let mut observed = self.waiter.load(Ordering::Acquire);
for _ in 0..IRQ_NOTIFY_CAS_BUDGET {
if observed == pending {
match self.waiter.compare_exchange(
pending,
pending,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => return IrqNotifyResult::Pending,
Err(current) => {
observed = current;
continue;
}
}
}
if observed == notifying_pending {
return IrqNotifyResult::Pending;
}
if observed == notifying {
match self.waiter.compare_exchange(
notifying,
notifying_pending,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => return IrqNotifyResult::Pending,
Err(current) => {
observed = current;
continue;
}
}
}
if observed.is_null() {
match self.waiter.compare_exchange(
ptr::null_mut(),
pending,
Ordering::Release,
Ordering::Acquire,
) {
Ok(_) => return IrqNotifyResult::Pending,
Err(current) => {
observed = current;
continue;
}
}
}
match self.waiter.compare_exchange(
observed,
notifying,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(waiter) => {
let registration = unsafe { take_cell_owner(waiter) };
let (generation, result) = Self::wake_registration(®istration);
self.finish_notification(result);
registration.finish_notification(generation);
return Self::notification_result(result);
}
Err(current) => observed = current,
}
}
let waiter = self.waiter.swap(pending, Ordering::AcqRel);
if waiter.is_null() || waiter == pending || is_notification_sentinel(waiter) {
return IrqNotifyResult::Pending;
}
let registration = unsafe { take_cell_owner(waiter) };
let (generation, result) = Self::wake_registration(®istration);
registration.finish_notification(generation);
Self::notification_result(result)
}
pub fn is_pending(&self) -> bool {
matches!(
self.waiter.load(Ordering::Acquire),
waiter if waiter == pending_waiter() || waiter == notifying_pending_waiter()
)
}
fn wake_registration(registration: &IrqWaitNode) -> (u64, crate::thread::WakeResult) {
let generation = registration.begin_notification();
let result = registration.wake.wake();
(generation, result)
}
fn finish_notification(&self, result: crate::thread::WakeResult) {
let pending = pending_waiter();
let notifying = notifying_waiter();
let notifying_pending = notifying_pending_waiter();
let delivered = matches!(
result,
crate::thread::WakeResult::Notified | crate::thread::WakeResult::AlreadyPending
);
let mut observed = self.waiter.load(Ordering::Acquire);
loop {
let next = if observed == notifying {
if delivered { ptr::null_mut() } else { pending }
} else if observed == notifying_pending {
pending
} else if observed == pending {
return;
} else {
panic!("IRQ wait cell notification ownership changed while wake was in flight");
};
match self.waiter.compare_exchange_weak(
observed,
next,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => return,
Err(current) => observed = current,
}
}
}
const fn notification_result(result: crate::thread::WakeResult) -> IrqNotifyResult {
match result {
crate::thread::WakeResult::Notified | crate::thread::WakeResult::AlreadyPending => {
IrqNotifyResult::Notified
}
crate::thread::WakeResult::Exited | crate::thread::WakeResult::Unavailable => {
IrqNotifyResult::Pending
}
}
}
}
impl Drop for IrqWaitCell {
fn drop(&mut self) {
let waiter = core::mem::replace(self.waiter.get_mut(), ptr::null_mut());
if waiter.is_null() || waiter == pending_waiter() {
return;
}
assert!(
!is_notification_sentinel(waiter),
"exclusive IRQ wait cell teardown found an in-flight notifier",
);
let registration = unsafe { take_cell_owner(waiter) };
let state = registration.state.load(Ordering::Acquire);
assert_eq!(
registration_phase(state),
RegistrationPhase::Attached,
"exclusive IRQ wait cell teardown found an in-flight notifier",
);
registration.cancel(registration_generation(state));
}
}
impl Default for IrqWaitCell {
fn default() -> Self {
Self::new()
}
}
#[cfg(all(test, not(miri)))]
mod loom_tests;