use core::cell::Cell;
use core::mem;
use core::mem::ManuallyDrop;
use crate::wait_list::WaitList;
#[derive(Debug)]
pub struct Semaphore {
waiters: WaitList<usize, WakeUp>,
permits: Cell<usize>,
total_permits: Cell<usize>,
}
#[derive(Debug, Clone, Copy)]
enum Fairness {
Fair,
Unfair,
}
enum WakeUp {
Unfair(UnfairWoken),
Fair(FairGrant),
}
struct UnfairWoken {
semaphore: *const Semaphore,
}
impl Drop for UnfairWoken {
fn drop(&mut self) {
let semaphore = unsafe { &*self.semaphore };
semaphore.release_permits(0, Fairness::Unfair);
}
}
struct FairGrant {
semaphore: *const Semaphore,
permits: usize,
}
impl Drop for FairGrant {
fn drop(&mut self) {
let semaphore = unsafe { &*self.semaphore };
semaphore.release_permits(self.permits, Fairness::Unfair);
}
}
impl Semaphore {
#[must_use]
pub const fn new(permits: usize) -> Self {
Self {
waiters: WaitList::new(),
permits: Cell::new(permits),
total_permits: Cell::new(permits),
}
}
#[must_use]
pub fn available_permits(&self) -> usize {
self.permits.get()
}
#[must_use]
pub fn total_permits(&self) -> usize {
self.total_permits.get()
}
pub fn add_permits(&self, new_permits: usize) {
self.total_permits.set(
self.total_permits
.get()
.checked_add(new_permits)
.expect("number of permits overflowed"),
);
self.release_permits(new_permits, Fairness::Unfair);
}
pub fn add_permits_fair(&self, new_permits: usize) {
self.total_permits
.set(self.total_permits.get().checked_add(new_permits).unwrap());
self.release_permits(new_permits, Fairness::Fair);
}
pub fn try_acquire(&self, to_acquire: usize) -> Option<Permit<'_>> {
if !self.waiters.borrow().is_empty() {
return None;
}
self.try_acquire_unfair(to_acquire)
}
pub fn try_acquire_unfair(&self, to_acquire: usize) -> Option<Permit<'_>> {
let new_permits = self.permits.get().checked_sub(to_acquire)?;
self.permits.set(new_permits);
Some(Permit {
semaphore: self,
permits: to_acquire,
})
}
pub async fn acquire(&self, to_acquire: usize) -> Permit<'_> {
loop {
if let Some(guard) = self.try_acquire(to_acquire) {
break guard;
}
match self.waiters.wait(to_acquire).await {
WakeUp::Unfair(token) => {
if let Some(guard) = self.try_acquire_unfair(to_acquire) {
mem::forget(token);
break guard;
}
}
WakeUp::Fair(grant) => {
mem::forget(grant);
return Permit {
semaphore: self,
permits: to_acquire,
};
}
}
}
}
pub async fn acquire_unfair(&self, to_acquire: usize) -> Permit<'_> {
loop {
if let Some(guard) = self.try_acquire_unfair(to_acquire) {
break guard;
}
match self.waiters.wait(to_acquire).await {
WakeUp::Unfair(token) => {
if let Some(guard) = self.try_acquire_unfair(to_acquire) {
mem::forget(token);
break guard;
}
}
WakeUp::Fair(grant) => {
mem::forget(grant);
return Permit {
semaphore: self,
permits: to_acquire,
};
}
}
}
}
fn release_permits(&self, permits: usize, fairness: Fairness) {
let mut permits = self.permits.get() + permits;
self.permits.set(permits);
let mut waiters = self.waiters.borrow();
while let Some(&wanted_permits) = waiters.head_input() {
permits = match permits.checked_sub(wanted_permits) {
Some(new_permits) => new_permits,
None => break,
};
let wake = match fairness {
Fairness::Fair => {
self.permits.set(permits);
WakeUp::Fair(FairGrant {
semaphore: self,
permits: wanted_permits,
})
}
Fairness::Unfair => WakeUp::Unfair(UnfairWoken { semaphore: self }),
};
if waiters.wake_one(wake).is_err() {
unreachable!();
}
}
}
}
#[derive(Debug)]
pub struct Permit<'semaphore> {
semaphore: &'semaphore Semaphore,
permits: usize,
}
impl<'semaphore> Permit<'semaphore> {
#[must_use]
pub fn semaphore(&self) -> &'semaphore Semaphore {
self.semaphore
}
#[must_use]
pub fn permits(&self) -> usize {
self.permits
}
pub fn leak(self) {
let this = ManuallyDrop::new(self);
let reduced_permits = this.semaphore.total_permits.get() - this.permits;
this.semaphore.total_permits.set(reduced_permits);
}
pub fn release_fair(self) {
let this = ManuallyDrop::new(self);
this.semaphore.release_permits(this.permits, Fairness::Fair);
}
}
impl Drop for Permit<'_> {
fn drop(&mut self) {
self.semaphore()
.release_permits(self.permits(), Fairness::Unfair);
}
}
#[cfg(test)]
mod tests {
use core::future::Future;
use alloc::boxed::Box;
use super::Semaphore;
use crate::utils::noop_cx;
#[test]
fn fair_grant_to_cancelled_waiter_is_not_leaked() {
let cx = &mut noop_cx();
let sem = Semaphore::new(1);
let initial = sem.try_acquire(1).unwrap();
assert_eq!(sem.available_permits(), 0);
let mut f1 = Box::pin(sem.acquire(1));
assert!(f1.as_mut().poll(cx).is_pending());
initial.release_fair();
assert_eq!(sem.available_permits(), 0);
drop(f1);
assert_eq!(sem.available_permits(), 1);
assert!(sem.try_acquire(1).is_some());
}
#[test]
fn permits_from_cancelled_fair_waiter_wake_next_waiter() {
let cx = &mut noop_cx();
let sem = Semaphore::new(1);
let initial = sem.try_acquire(1).unwrap();
let mut f1 = Box::pin(sem.acquire(1));
let mut f2 = Box::pin(sem.acquire(1));
assert!(f1.as_mut().poll(cx).is_pending());
assert!(f2.as_mut().poll(cx).is_pending());
initial.release_fair();
assert_eq!(sem.available_permits(), 0);
drop(f1);
assert!(f2.as_mut().poll(cx).is_ready());
}
#[test]
fn release_wakes_front_waiter_with_another_queued() {
let cx = &mut noop_cx();
let sem = Semaphore::new(1);
let initial = sem.try_acquire(1).unwrap();
let mut f1 = Box::pin(sem.acquire(1));
let mut f2 = Box::pin(sem.acquire(1));
assert!(f1.as_mut().poll(cx).is_pending());
assert!(f2.as_mut().poll(cx).is_pending());
drop(initial);
assert!(f1.as_mut().poll(cx).is_ready());
}
#[test]
fn cancelled_woken_waiter_does_not_strand_permit() {
let cx = &mut noop_cx();
let sem = Semaphore::new(1);
let initial = sem.try_acquire(1).unwrap();
let mut w1 = Box::pin(sem.acquire(1));
let mut w2 = Box::pin(sem.acquire(1));
assert!(w1.as_mut().poll(cx).is_pending());
assert!(w2.as_mut().poll(cx).is_pending());
drop(initial);
drop(w1);
assert!(w2.as_mut().poll(cx).is_ready());
}
#[test]
fn cancelled_woken_waiter_rewakes_across_queue() {
let cx = &mut noop_cx();
let sem = Semaphore::new(2);
let initial = sem.try_acquire(2).unwrap();
let mut w1 = Box::pin(sem.acquire(1));
let mut w2 = Box::pin(sem.acquire(1));
let mut w3 = Box::pin(sem.acquire(1));
assert!(w1.as_mut().poll(cx).is_pending());
assert!(w2.as_mut().poll(cx).is_pending());
assert!(w3.as_mut().poll(cx).is_pending());
drop(initial);
drop(w1);
assert!(w2.as_mut().poll(cx).is_ready());
assert!(w3.as_mut().poll(cx).is_ready());
}
#[test]
fn parked_acquire_cancelled_then_release() {
let cx = &mut noop_cx();
let sem = Semaphore::new(1);
let held = sem.try_acquire(1).unwrap();
let mut a = Box::pin(sem.acquire(1));
let mut b = Box::pin(sem.acquire(1));
assert!(a.as_mut().poll(cx).is_pending());
assert!(b.as_mut().poll(cx).is_pending());
drop(a);
drop(held);
assert!(b.as_mut().poll(cx).is_ready());
}
}