use std::{
future::Future,
panic::Location,
pin::Pin,
sync::Arc,
task::{Context, Poll, Waker},
};
use parking_lot::Mutex;
use crate::flash::{
diag::PrimKind,
flash_ambient,
ids::{Backend, trace_native_from_ambient},
system,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct AcquireError;
impl std::fmt::Display for AcquireError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("semaphore closed")
}
}
impl std::error::Error for AcquireError {}
struct Inner {
wakers: Vec<Waker>,
permits: usize,
}
pub struct Semaphore {
backend: Backend,
inner: Mutex<Inner>,
}
impl Semaphore {
#[must_use]
#[track_caller]
pub fn new(permits: usize) -> Self {
Self {
inner: Mutex::new(Inner {
permits,
wakers: Vec::new(),
}),
backend: if flash_ambient() {
let cvid = system::next_condvar_id();
system::describe_cvid(cvid, PrimKind::Semaphore, Location::caller());
Backend::Engine(cvid)
} else {
Backend::Native
},
}
}
#[must_use = "the permit is released as soon as it is dropped"]
pub fn acquire_owned(self: Arc<Self>) -> AcquireOwned {
AcquireOwned {
sem: self,
pending: None,
}
}
fn release(&self) {
let mut inner = self.inner.lock();
inner.permits += 1;
let waker = match self.backend {
Backend::Engine(_) => None,
Backend::Native if inner.wakers.is_empty() => None,
Backend::Native => Some(inner.wakers.remove(0)),
};
drop(inner);
match self.backend {
Backend::Engine(cvid) => system::signal_channel(cvid, false),
Backend::Native => {
trace_native_from_ambient("semaphore", "release");
if let Some(waker) = waker {
waker.wake();
}
}
}
}
}
enum Parked {
Engine(system::AsyncHandle),
Real(Waker),
}
pub struct OwnedSemaphorePermit {
sem: Arc<Semaphore>,
}
impl Drop for OwnedSemaphorePermit {
fn drop(&mut self) {
self.sem.release();
}
}
pub struct AcquireOwned {
sem: Arc<Semaphore>,
pending: Option<Parked>,
}
impl Unpin for AcquireOwned {}
impl Future for AcquireOwned {
type Output = Result<OwnedSemaphorePermit, AcquireError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.get_mut();
if let Some(Parked::Engine(handle)) = this.pending.as_ref() {
if handle.granted() {
this.pending = None;
} else {
return Poll::Pending;
}
}
let mut inner = this.sem.inner.lock();
if inner.permits > 0 {
inner.permits -= 1;
return Poll::Ready(Ok(OwnedSemaphorePermit {
sem: Arc::clone(&this.sem),
}));
}
match this.sem.backend {
Backend::Engine(cvid) => {
let (handle, adv) = system::register_channel_async(cvid, cx.waker().clone());
this.pending = Some(Parked::Engine(handle));
drop(inner);
adv.fire();
}
Backend::Native => {
trace_native_from_ambient("semaphore", "acquire_park");
let waker = cx.waker().clone();
inner.wakers.push(waker.clone());
this.pending = Some(Parked::Real(waker));
drop(inner);
}
}
Poll::Pending
}
}
impl Drop for AcquireOwned {
fn drop(&mut self) {
match self.pending.take() {
Some(Parked::Real(waker)) => {
self.sem
.inner
.lock()
.wakers
.retain(|w| !w.will_wake(&waker));
}
Some(Parked::Engine(handle)) => system::cancel_async_wait(&handle),
None => {}
}
}
}
#[cfg(test)]
mod tests {
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};
use kithara_test_utils::kithara;
use super::Semaphore;
use crate::{
flash,
tokio::task::{spawn, yield_now},
};
struct Consts;
impl Consts {
const PERMITS: usize = 3;
const TASKS: usize = 16;
}
#[kithara::test(tokio, multi_thread)]
async fn contention_no_lost_wakeup() {
flash::reset();
let sem = Arc::new(Semaphore::new(Consts::PERMITS));
let done = Arc::new(AtomicUsize::new(0));
let handles: Vec<_> = (0..Consts::TASKS)
.map(|_| {
let sem = Arc::clone(&sem);
let done = Arc::clone(&done);
spawn(async move {
let permit = Arc::clone(&sem).acquire_owned().await.expect("not closed");
yield_now().await;
drop(permit);
done.fetch_add(1, Ordering::SeqCst);
})
})
.collect();
for handle in handles {
handle.await.expect("task joined");
}
assert_eq!(done.load(Ordering::SeqCst), Consts::TASKS);
}
#[kithara::test(tokio, multi_thread)]
async fn release_wakes_parked_acquirer() {
flash::reset();
let sem = Arc::new(Semaphore::new(1));
let held = Arc::clone(&sem)
.acquire_owned()
.await
.expect("first permit");
let sem2 = Arc::clone(&sem);
let waiter = spawn(async move {
let _permit = sem2.acquire_owned().await.expect("second permit");
});
yield_now().await;
drop(held);
waiter.await.expect("waiter joined");
}
}