use super::wait_queue::WaitQueue;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::task::{Context, Poll};
#[repr(align(64))]
pub struct Semaphore {
permits: AtomicUsize,
wait: WaitQueue,
}
impl Semaphore {
#[must_use]
pub const fn new(permits: usize) -> Self {
Self {
permits: AtomicUsize::new(permits),
wait: WaitQueue::new(),
}
}
#[must_use]
#[inline(always)]
pub fn available_permits(&self) -> usize {
self.permits.load(Ordering::Relaxed)
}
#[inline(always)]
pub fn add_permits(&self, n: usize) {
self.permits.fetch_add(n, Ordering::Release);
self.wait.wake_all();
}
#[inline(always)]
pub async fn acquire(&self) -> SemaphorePermit<'_> {
std::future::poll_fn(|cx| self.poll_acquire(cx)).await
}
#[inline(always)]
pub fn try_acquire(&self) -> Result<SemaphorePermit<'_>, TryAcquireError> {
self.try_acquire_one()
.then(|| SemaphorePermit { sem: self })
.ok_or(TryAcquireError::NoPermits)
}
#[inline]
fn try_acquire_one(&self) -> bool {
let mut current = self.permits.load(Ordering::Relaxed);
loop {
if current == 0 {
return false;
}
match self.permits.compare_exchange_weak(
current,
current - 1,
Ordering::Acquire,
Ordering::Relaxed,
) {
Ok(_) => return true,
Err(observed) => current = observed,
}
}
}
#[inline]
fn poll_acquire(&self, cx: &Context<'_>) -> Poll<SemaphorePermit<'_>> {
if !self.wait.has_waiters() && self.try_acquire_one() {
return Poll::Ready(SemaphorePermit { sem: self });
}
let token = self.wait.register(cx.waker());
if self.try_acquire_one() {
self.wait.cancel(token);
return Poll::Ready(SemaphorePermit { sem: self });
}
Poll::Pending
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(align(64))]
pub enum TryAcquireError {
NoPermits,
}
impl std::fmt::Display for TryAcquireError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("no permits available")
}
}
impl std::error::Error for TryAcquireError {}
#[repr(align(64))]
pub struct SemaphorePermit<'a> {
sem: &'a Semaphore,
}
impl Drop for SemaphorePermit<'_> {
#[inline(always)]
fn drop(&mut self) {
self.sem.permits.fetch_add(1, Ordering::Release);
self.sem.wait.wake_one();
}
}