use core::cell::UnsafeCell;
use core::ops::{Deref, DerefMut};
use super::sched;
use super::tcb::{self, MAX_PTASKS, NO_TASK};
use crate::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum LockError {
Recursive,
Timeout,
TooManyHeldMutexes,
NotInTask,
Poisoned,
}
#[repr(C)]
pub struct PriorityMutex<T> {
locked: AtomicBool,
owner: AtomicUsize,
waiters: [AtomicUsize; MAX_PTASKS],
poisoned: AtomicBool,
data: UnsafeCell<T>,
}
unsafe impl<T: Send> Sync for PriorityMutex<T> {}
impl<T> PriorityMutex<T> {
#[cfg(not(loom))]
pub const fn new(value: T) -> Self {
Self {
locked: AtomicBool::new(false),
owner: AtomicUsize::new(NO_TASK),
waiters: [const { AtomicUsize::new(NO_TASK) }; MAX_PTASKS],
poisoned: AtomicBool::new(false),
data: UnsafeCell::new(value),
}
}
#[cfg(loom)]
pub fn new(value: T) -> Self {
Self {
locked: AtomicBool::new(false),
owner: AtomicUsize::new(NO_TASK),
waiters: core::array::from_fn(|_| AtomicUsize::new(NO_TASK)),
poisoned: AtomicBool::new(false),
data: UnsafeCell::new(value),
}
}
pub fn is_poisoned(&self) -> bool {
self.poisoned.load(Ordering::Acquire)
}
pub fn try_lock(&self) -> Option<PriorityMutexGuard<'_, T>> {
let me = sched::current()?;
self.try_acquire_guarded(me)
}
pub fn lock(&self) -> PriorityMutexGuard<'_, T> {
match self.lock_timeout(None) {
Ok(g) => g,
Err(LockError::Recursive) => panic!(
"PriorityMutex::lock: recursive lock from the same task \
(self-deadlock; use try_lock or restructure)"
),
Err(LockError::TooManyHeldMutexes) => panic!(
"PriorityMutex::lock: task already holds MAX_HELD={} mutexes",
tcb::MAX_HELD
),
Err(LockError::NotInTask) => {
panic!("PriorityMutex::lock() outside preemptive task context")
}
Err(LockError::Poisoned) => panic!(
"PriorityMutex::lock: mutex poisoned by a faulting holder (data may be inconsistent; use lock_timeout/try_lock to recover)"
),
Err(LockError::Timeout) => unreachable!("lock() has no timeout"),
}
}
pub fn lock_timeout(
&self,
timeout: Option<crate::time::Duration>,
) -> Result<PriorityMutexGuard<'_, T>, LockError> {
let me = sched::current().ok_or(LockError::NotInTask)?;
let deadline = timeout.map(|d| crate::port::board::now_us().wrapping_add(d.as_micros()));
loop {
if self.poisoned.load(Ordering::Acquire) {
return Err(LockError::Poisoned);
}
if let Some(g) = self.try_acquire_guarded(me) {
return Ok(g);
}
if self.owner.load(Ordering::Acquire) == me {
return Err(LockError::Recursive);
}
if let Some(d) = deadline {
if crate::port::board::now_us() >= d {
self.remove_waiter(me);
crate::timer::cancel_ptask_deadline(me);
return Err(LockError::Timeout);
}
}
let outcome = crate::critical::enter(|| {
if let Some(g) = self.try_acquire_guarded(me) {
Ok(g)
} else {
self.boost_holder(me);
self.add_waiter(me);
if let Some(d) = deadline {
let _ = crate::timer::register_ptask_deadline(d, me);
}
sched::block_current();
Err(LockError::Timeout) }
});
match outcome {
Ok(g) => return Ok(g),
Err(_) => {
crate::port::arch::request_reschedule();
}
}
}
}
fn try_acquire_guarded(&self, me: usize) -> Option<PriorityMutexGuard<'_, T>> {
if self.poisoned.load(Ordering::Acquire) {
return None;
}
if !self.try_acquire(me) {
return None;
}
match self.push_held(me) {
Ok(()) => {
#[cfg(feature = "trace")]
crate::trace::mutex_lock_acquired(me as u16, self as *const _ as usize as u32);
Some(PriorityMutexGuard { mutex: self })
}
Err(_) => {
self.owner.store(NO_TASK, Ordering::Release);
self.locked.store(false, Ordering::Release);
None
}
}
}
fn try_acquire(&self, me: usize) -> bool {
if self
.locked
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
{
self.owner.store(me, Ordering::Release);
true
} else {
false
}
}
fn push_held(&self, me: usize) -> Result<(), LockError> {
let tcb = tcb::get(me).ok_or(LockError::NotInTask)?;
if tcb.push_held(
self as *const _ as *const (),
Self::highest_waiter_priority_erased,
) {
Ok(())
} else {
Err(LockError::TooManyHeldMutexes)
}
}
fn highest_waiter_priority_erased(ptr: *const ()) -> u8 {
unsafe { (&*(ptr as *const PriorityMutex<T>)).highest_waiter_priority() }
}
fn highest_waiter_priority(&self) -> u8 {
let mut max = 0u8;
for slot in &self.waiters {
let id = slot.load(Ordering::Acquire);
if id != NO_TASK {
if let Some(w) = tcb::get(id) {
let b = w.base_priority.load(Ordering::Acquire);
if b > max {
max = b;
}
}
}
}
max
}
fn boost_holder(&self, me: usize) {
let owner_id = self.owner.load(Ordering::Acquire);
if owner_id != NO_TASK {
if let (Some(me_tcb), Some(owner_tcb)) = (tcb::get(me), tcb::get(owner_id)) {
let my_base = me_tcb.base_priority.load(Ordering::Acquire);
let owner_eff = owner_tcb.effective_priority.load(Ordering::Acquire);
if my_base > owner_eff {
owner_tcb.set_effective_priority(owner_id, my_base);
#[cfg(feature = "trace")]
crate::trace::priority_inherit(
owner_id as u16,
self as *const _ as usize as u32,
);
}
}
}
}
fn add_waiter(&self, id: usize) {
for slot in &self.waiters {
if slot
.compare_exchange(NO_TASK, id, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
{
return;
}
}
}
fn remove_waiter(&self, id: usize) {
for slot in &self.waiters {
let _ = slot.compare_exchange(id, NO_TASK, Ordering::AcqRel, Ordering::Acquire);
}
}
fn wake_all_waiters(&self) {
for slot in &self.waiters {
let id = slot.swap(NO_TASK, Ordering::AcqRel);
if id != NO_TASK {
sched::unblock(id);
crate::timer::cancel_ptask_deadline(id);
}
}
}
}
pub unsafe fn poison_mutex(ptr: *const ()) {
unsafe {
let m = &*(ptr as *const PriorityMutex<()>);
m.poisoned.store(true, Ordering::Release);
m.wake_all_waiters();
}
}
pub struct PriorityMutexGuard<'a, T> {
mutex: &'a PriorityMutex<T>,
}
impl<'a, T> Deref for PriorityMutexGuard<'a, T> {
type Target = T;
fn deref(&self) -> &T {
unsafe { &*self.mutex.data.get() }
}
}
impl<'a, T> DerefMut for PriorityMutexGuard<'a, T> {
fn deref_mut(&mut self) -> &mut T {
unsafe { &mut *self.mutex.data.get() }
}
}
impl<'a, T> Drop for PriorityMutexGuard<'a, T> {
fn drop(&mut self) {
crate::critical::enter(|| {
let owner = self.mutex.owner.swap(NO_TASK, Ordering::AcqRel);
#[cfg(feature = "trace")]
if owner != NO_TASK {
crate::trace::mutex_unlock(
owner as u16,
self.mutex as *const _ as usize as u32,
);
}
if owner != NO_TASK {
if let Some(t) = tcb::get(owner) {
t.remove_held(self.mutex as *const _ as *const ());
let base = t.base_priority.load(Ordering::Acquire);
let mut eff = base;
for slot in &t.held {
let ptr = slot.ptr.load(Ordering::Acquire);
if !ptr.is_null() {
let hwp = unsafe {
let f: fn(*const ()) -> u8 =
core::mem::transmute(slot.hwp.load(Ordering::Acquire));
f(ptr)
};
if hwp > eff {
eff = hwp;
}
}
}
t.set_effective_priority(owner, eff);
}
}
self.mutex.locked.store(false, Ordering::Release);
self.mutex.wake_all_waiters();
});
crate::port::arch::request_reschedule();
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::preempt::tcb as tcbmod;
#[test]
fn lock_unlock_basic() {
crate::kernel_test! {
let a = tcbmod::register(0x1000, 1).unwrap();
sched::set_current(a);
let m: PriorityMutex<u32> = PriorityMutex::new(0);
{
let mut guard = m.lock();
*guard = 42;
}
assert_eq!(*m.lock(), 42);
}
}
#[test]
fn try_lock_behavior() {
crate::kernel_test! {
let a = tcbmod::register(0x1000, 1).unwrap();
sched::set_current(a);
let m: PriorityMutex<u32> = PriorityMutex::new(0);
let g = m.try_lock().expect("free mutex must lock");
assert!(m.try_lock().is_none(), "held mutex must not lock");
drop(g);
assert!(m.try_lock().is_some(), "released mutex must lock");
}
}
#[test]
#[should_panic(expected = "recursive lock")]
fn recursive_lock_panics() {
crate::kernel_test! {
let a = tcbmod::register(0x1000, 1).unwrap();
sched::set_current(a);
let m: PriorityMutex<u32> = PriorityMutex::new(0);
let _g = m.lock();
let _ = m.lock(); }
}
#[test]
fn priority_inheritance_boosts_holder() {
crate::kernel_test! {
let low = tcbmod::register(0x1000, 1).unwrap();
let high = tcbmod::register(0x2000, 5).unwrap();
let m: PriorityMutex<u32> = PriorityMutex::new(0);
sched::set_current(low);
let guard = m.lock();
assert_eq!(
tcbmod::get(low).unwrap().effective_priority.load(Ordering::Acquire),
1
);
sched::set_current(high);
let owner_id = low;
let my_base = tcbmod::get(high).unwrap().base_priority.load(Ordering::Acquire);
let owner_tcb = tcbmod::get(owner_id).unwrap();
if my_base > owner_tcb.effective_priority.load(Ordering::Acquire) {
owner_tcb.effective_priority.store(my_base, Ordering::Release);
}
assert_eq!(
tcbmod::get(low).unwrap().effective_priority.load(Ordering::Acquire),
5
);
drop(guard);
assert_eq!(
tcbmod::get(low).unwrap().effective_priority.load(Ordering::Acquire),
1
);
}
}
#[test]
fn b11_nested_unlock_keeps_boost_from_other_mutex() {
crate::kernel_test! {
static A: PriorityMutex<u32> = PriorityMutex::new(0);
static B: PriorityMutex<u32> = PriorityMutex::new(0);
let holder = tcbmod::register(0x1000, 1).unwrap();
sched::set_current(holder);
let ga = A.lock();
let gb = B.lock();
let waiter = tcbmod::register(0x3000, 8).unwrap();
sched::set_current(waiter);
A.waiters[0].store(waiter, Ordering::Release);
let holder_tcb = tcbmod::get(holder).unwrap();
holder_tcb.effective_priority.store(8, Ordering::Release);
sched::set_current(holder);
drop(gb);
assert_eq!(
holder_tcb.effective_priority.load(Ordering::Acquire),
8,
"[B11] unlocking B must not drop the boost held for A"
);
drop(ga);
assert_eq!(
holder_tcb.effective_priority.load(Ordering::Acquire),
1,
"[B11] after unlocking both, effective priority = base"
);
}
}
#[test]
fn b1_retest_inside_critical_section_catches_release() {
crate::kernel_test! {
let a = tcbmod::register(0x1000, 1).unwrap();
let b = tcbmod::register(0x2000, 5).unwrap();
let m: PriorityMutex<u32> = PriorityMutex::new(0);
sched::set_current(a);
let guard_a = m.lock();
sched::set_current(b);
assert!(!m.try_acquire(b), "A holds the mutex");
sched::set_current(a);
drop(guard_a);
sched::set_current(b);
let guard_b = m.lock();
assert!(guard_b.mutex.locked.load(Ordering::Acquire));
drop(guard_b);
}
}
}