Skip to main content

dtact_util/sync/
semaphore.rs

1//! Counting semaphore — the usual building block for concurrency limits
2//! (e.g. capping in-flight connections/requests).
3
4use super::wait_queue::WaitQueue;
5use std::sync::atomic::{AtomicUsize, Ordering};
6use std::task::{Context, Poll};
7
8/// A counting semaphore: `.acquire().await` waits (without blocking the
9/// OS thread) until a permit is available, then holds it until the
10/// returned [`SemaphorePermit`] is dropped.
11#[repr(align(64))]
12pub struct Semaphore {
13    permits: AtomicUsize,
14    wait: WaitQueue,
15}
16
17impl Semaphore {
18    /// Create a semaphore with `permits` available immediately.
19    #[must_use]
20    pub const fn new(permits: usize) -> Self {
21        Self {
22            permits: AtomicUsize::new(permits),
23            wait: WaitQueue::new(),
24        }
25    }
26
27    /// Current number of permits available for immediate acquisition.
28    #[must_use]
29    #[inline(always)]
30    pub fn available_permits(&self) -> usize {
31        self.permits.load(Ordering::Relaxed)
32    }
33
34    /// Add `n` permits, waking waiters as needed.
35    #[inline(always)]
36    pub fn add_permits(&self, n: usize) {
37        self.permits.fetch_add(n, Ordering::Release);
38        self.wait.wake_all();
39    }
40
41    /// Acquire one permit, waiting if none are currently available.
42    #[inline(always)]
43    pub async fn acquire(&self) -> SemaphorePermit<'_> {
44        std::future::poll_fn(|cx| self.poll_acquire(cx)).await
45    }
46
47    /// Acquire one permit if immediately available, without waiting.
48    ///
49    /// # Errors
50    /// Returns [`TryAcquireError::NoPermits`] if none are currently free.
51    #[inline(always)]
52    pub fn try_acquire(&self) -> Result<SemaphorePermit<'_>, TryAcquireError> {
53        // See `Mutex::try_lock`'s comment on why this must be `.then(||
54        // ...)`, not `.then_some(...)` — the permit's `Drop` releases a
55        // permit back, which `then_some`'s eager evaluation would do even
56        // when acquisition failed.
57        self.try_acquire_one()
58            .then(|| SemaphorePermit { sem: self })
59            .ok_or(TryAcquireError::NoPermits)
60    }
61
62    #[inline]
63    fn try_acquire_one(&self) -> bool {
64        let mut current = self.permits.load(Ordering::Relaxed);
65        loop {
66            if current == 0 {
67                return false;
68            }
69            match self.permits.compare_exchange_weak(
70                current,
71                current - 1,
72                Ordering::Acquire,
73                Ordering::Relaxed,
74            ) {
75                Ok(_) => return true,
76                Err(observed) => current = observed,
77            }
78        }
79    }
80
81    #[inline]
82    fn poll_acquire(&self, cx: &Context<'_>) -> Poll<SemaphorePermit<'_>> {
83        // Skip the fast-path acquire if anyone is already waiting for a
84        // permit — same starvation risk as `Mutex::poll_lock` (a fresh
85        // acquirer could otherwise perpetually beat an already-waiting
86        // one to every freed permit); see `WaitQueue::has_waiters`'s doc.
87        // `try_acquire` itself is unaffected — it's an explicit
88        // "don't wait" API and must always attempt regardless of waiters.
89        if !self.wait.has_waiters() && self.try_acquire_one() {
90            return Poll::Ready(SemaphorePermit { sem: self });
91        }
92        // See `Mutex::poll_lock` for why registration comes before the
93        // re-check, not after.
94        let token = self.wait.register(cx.waker());
95        if self.try_acquire_one() {
96            self.wait.cancel(token);
97            return Poll::Ready(SemaphorePermit { sem: self });
98        }
99        Poll::Pending
100    }
101}
102
103/// Error returned by [`Semaphore::try_acquire`].
104#[derive(Debug, Clone, Copy, PartialEq, Eq)]
105#[repr(align(64))]
106pub enum TryAcquireError {
107    /// No permits were immediately available.
108    NoPermits,
109}
110
111impl std::fmt::Display for TryAcquireError {
112    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
113        f.write_str("no permits available")
114    }
115}
116
117impl std::error::Error for TryAcquireError {}
118
119/// RAII permit held on a [`Semaphore`]. Returns the permit (and wakes one
120/// waiter, if any) on drop.
121#[repr(align(64))]
122pub struct SemaphorePermit<'a> {
123    sem: &'a Semaphore,
124}
125
126impl Drop for SemaphorePermit<'_> {
127    #[inline(always)]
128    fn drop(&mut self) {
129        self.sem.permits.fetch_add(1, Ordering::Release);
130        self.sem.wait.wake_one();
131    }
132}