Skip to main content

rivet/sync/
semaphore.rs

1//! Async counting semaphore.
2//!
3//! `acquire()` returns a real [`Future`] — call it with `.await` from an
4//! `async fn` task. The waiting task's identity is read from
5//! [`crate::executor::current_task`], so no manual priority/index
6//! bookkeeping is needed by the caller.
7//!
8//! Safe to `release()` from either task context or an ISR.
9
10use core::future::Future;
11use core::pin::Pin;
12use core::task::{Context, Poll};
13
14use crate::waker;
15
16/// An async counting semaphore.
17///
18/// ```ignore
19/// static SEM: rivet::sync::Semaphore<1> = rivet::sync::Semaphore::new(0);
20///
21/// #[rivet::task(priority = 1)]
22/// async fn waiter() {
23///     loop {
24///         SEM.acquire().await;
25///         // got the semaphore
26///     }
27/// }
28///
29/// // ISR or another task:
30/// fn signaler() {
31///     SEM.release();
32/// }
33/// ```
34pub struct Semaphore<const MAX: u8> {
35    /// Current count. 0 = taken, >0 = available.
36    count: crate::sync::atomic::AtomicU8,
37    /// Per-priority waiter bitmap, mirroring `waker::PRIORITY_QUEUES`
38    /// (plan.md [B9]): bit i of `waiters[p]` = task `(p, i)` waiting.
39    /// Multiple waiters are supported; `release()` wakes the
40    /// highest-priority one.
41    waiters: [crate::sync::atomic::AtomicU32; 32],
42}
43
44impl<const MAX: u8> Semaphore<MAX> {
45    /// Create a new semaphore with the given initial count.
46    #[cfg(not(loom))]
47    pub const fn new(initial: u8) -> Self {
48        Self {
49            count: crate::sync::atomic::AtomicU8::new(initial),
50            // Inline const avoids a named `const` item with interior
51            // mutability.
52            waiters: [const { crate::sync::atomic::AtomicU32::new(0) }; 32],
53        }
54    }
55
56    /// Loom's atomics are not const-constructible; runtime constructor used
57    /// by the loom models.
58    #[cfg(loom)]
59    pub fn new(initial: u8) -> Self {
60        Self {
61            count: crate::sync::atomic::AtomicU8::new(initial),
62            waiters: core::array::from_fn(|_| crate::sync::atomic::AtomicU32::new(0)),
63        }
64    }
65
66    /// Try to acquire without blocking. Returns true if acquired.
67    pub fn try_acquire(&self) -> bool {
68        loop {
69            let c = self.count.load(crate::sync::atomic::Ordering::Acquire);
70            if c == 0 {
71                return false;
72            }
73            if self
74                .count
75                .compare_exchange_weak(
76                    c,
77                    c - 1,
78                    crate::sync::atomic::Ordering::AcqRel,
79                    crate::sync::atomic::Ordering::Acquire,
80                )
81                .is_ok()
82            {
83                return true;
84            }
85        }
86    }
87
88    /// Acquire the semaphore, yielding the task if it's currently taken.
89    ///
90    /// # Panics
91    /// Panics if polled outside of a task context (i.e. not from within
92    /// the executor's poll of a `#[rivet::task]`).
93    pub fn acquire(&self) -> Acquire<'_, MAX> {
94        Acquire {
95            sem: self,
96            registered: None,
97        }
98    }
99
100    /// Release the semaphore. If tasks are waiting, the highest-priority
101    /// one is woken and handed the token directly. Safe to call from ISR
102    /// context.
103    pub fn release(&self) {
104        // Hand the token directly to the highest-priority waiter instead
105        // of incrementing count, so a concurrent try_acquire() by a third
106        // party can't steal it out from under the woken task.
107        //
108        // The waiter bit is claimed with a CAS loop, not a plain load +
109        // clear: two concurrent `release()` calls must not both pick the
110        // same waiter (loom found exactly that race in the initial
111        // check-then-act version).
112        for (prio, queue) in self.waiters.iter().enumerate().rev() {
113            loop {
114                let q = queue.load(crate::sync::atomic::Ordering::Acquire);
115                if q == 0 {
116                    break; // no waiter at this priority — try the next
117                }
118                let bit = q & q.wrapping_neg();
119                match queue.compare_exchange_weak(
120                    q,
121                    q & !bit,
122                    crate::sync::atomic::Ordering::AcqRel,
123                    crate::sync::atomic::Ordering::Acquire,
124                ) {
125                    Ok(_) => {
126                        self.count.store(1, crate::sync::atomic::Ordering::Release);
127                        waker::wake_task(crate::task::TaskId::new(
128                            prio as u8,
129                            bit.trailing_zeros() as u8,
130                        ));
131                        return;
132                    }
133                    Err(_) => continue, // another release took it — retry
134                }
135            }
136        }
137
138        // No waiter: increment (bounded by MAX).
139        let mut c = self.count.load(crate::sync::atomic::Ordering::Acquire);
140        loop {
141            if c >= MAX {
142                return;
143            }
144            match self.count.compare_exchange_weak(
145                c,
146                c + 1,
147                crate::sync::atomic::Ordering::AcqRel,
148                crate::sync::atomic::Ordering::Acquire,
149            ) {
150                Ok(_) => return,
151                Err(actual) => c = actual,
152            }
153        }
154    }
155
156    fn register_waiter(&self, id: crate::task::TaskId) {
157        let mask = 1u32 << id.index();
158        self.waiters[id.priority() as usize].fetch_or(mask, crate::sync::atomic::Ordering::Release);
159    }
160
161    /// Debug snapshot of the waiter bitmaps (test-only).
162    #[cfg(any(loom, feature = "test-support"))]
163    #[doc(hidden)]
164    pub fn debug_waiters(&self) -> [u32; 32] {
165        let mut w = [0u32; 32];
166        for (i, q) in self.waiters.iter().enumerate() {
167            w[i] = q.load(crate::sync::atomic::Ordering::Acquire);
168        }
169        w
170    }
171
172    fn remove_waiter(&self, id: crate::task::TaskId) {
173        let mask = 1u32 << id.index();
174        self.waiters[id.priority() as usize]
175            .fetch_and(!mask, crate::sync::atomic::Ordering::AcqRel);
176    }
177}
178
179/// Future returned by [`Semaphore::acquire`].
180pub struct Acquire<'a, const MAX: u8> {
181    sem: &'a Semaphore<MAX>,
182    /// `Some(id)` while registered as a waiter; cleared on
183    /// completion, and cancelled in [`Drop`] so a dropped acquire never
184    /// leaves a stale registration behind (plan.md §2.5).
185    registered: Option<crate::task::TaskId>,
186}
187
188impl<'a, const MAX: u8> Drop for Acquire<'a, MAX> {
189    fn drop(&mut self) {
190        if let Some(id) = self.registered.take() {
191            self.sem.remove_waiter(id);
192        }
193    }
194}
195
196impl<'a, const MAX: u8> Future for Acquire<'a, MAX> {
197    type Output = ();
198
199    fn poll(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<()> {
200        // SAFETY: `Acquire` holds only a `&Semaphore` and an `Option`;
201        // no `!Unpin` fields, so projecting is sound.
202        let this = unsafe { self.get_unchecked_mut() };
203        if this.sem.try_acquire() {
204            return Poll::Ready(());
205        }
206
207        let id = crate::executor::current_task()
208            .expect("Semaphore::acquire().await polled outside of a task context");
209        this.sem.register_waiter(id);
210        this.registered = Some(id);
211
212        // Re-check: release() may have fired between our first try_acquire()
213        // and register_waiter() above.
214        if this.sem.try_acquire() {
215            if let Some(id) = this.registered.take() {
216                this.sem.remove_waiter(id);
217            }
218            return Poll::Ready(());
219        }
220
221        Poll::Pending
222    }
223}
224
225// Safety: Semaphore uses atomics; safe to share across contexts.
226unsafe impl<const MAX: u8> Sync for Semaphore<MAX> {}
227
228#[cfg(test)]
229mod tests {
230    use super::*;
231
232    #[test]
233    fn semaphore_try_acquire_release() {
234        crate::kernel_test! {
235            let sem: Semaphore<3> = Semaphore::new(1);
236            assert!(sem.try_acquire());
237            assert!(!sem.try_acquire());
238            sem.release();
239            assert!(sem.try_acquire());
240        }
241    }
242
243    #[test]
244    fn semaphore_counting() {
245        crate::kernel_test! {
246            let sem: Semaphore<3> = Semaphore::new(2);
247            assert!(sem.try_acquire());
248            assert!(sem.try_acquire());
249            assert!(!sem.try_acquire());
250            sem.release();
251            assert!(sem.try_acquire());
252            assert!(!sem.try_acquire());
253        }
254    }
255
256    #[test]
257    fn acquire_future_ready_when_available() {
258        crate::kernel_test! {
259            let sem: Semaphore<1> = Semaphore::new(1);
260            let waker = crate::waker::task_waker(crate::task::TaskId::new(0, 0));
261            let mut cx = Context::from_waker(&waker);
262            let mut fut = sem.acquire();
263            // SAFETY: `fut` is a local `Acquire` future; `Unpin`, never
264            // moved while pinned — sound for this single poll.
265            let pinned = unsafe { Pin::new_unchecked(&mut fut) };
266            assert_eq!(pinned.poll(&mut cx), Poll::Ready(()));
267        }
268    }
269
270    #[test]
271    #[should_panic(expected = "outside of a task context")]
272    fn acquire_future_panics_without_task_context() {
273        crate::kernel_test! {
274            let sem: Semaphore<1> = Semaphore::new(0);
275            let waker = crate::waker::task_waker(crate::task::TaskId::new(0, 0));
276            let mut cx = Context::from_waker(&waker);
277            let mut fut = sem.acquire();
278            // SAFETY: `fut` is a local `Acquire` future; `Unpin`, never
279            // moved while pinned — sound for this single poll.
280            let pinned = unsafe { Pin::new_unchecked(&mut fut) };
281            let _ = pinned.poll(&mut cx);
282        }
283    }
284}
285
286#[cfg(test)]
287mod b9_tests {
288    use super::*;
289
290    #[test]
291    fn two_waiters_both_woken() {
292        crate::kernel_test! {
293            let sem: Semaphore<1> = Semaphore::new(0);
294
295            // Register two waiters directly (simulating two Pending polls).
296            sem.register_waiter(crate::task::TaskId::new(1, 0));
297            sem.register_waiter(crate::task::TaskId::new(2, 0));
298            assert_eq!(sem.debug_waiters()[1], 1, "waiter (1,0)");
299            assert_eq!(sem.debug_waiters()[2], 1, "waiter (2,0)");
300
301            sem.release();
302            assert_eq!(crate::waker::next_ready(), Some(crate::task::TaskId::new(2, 0)), "highest priority first");
303
304            sem.release();
305            assert_eq!(crate::waker::next_ready(), Some(crate::task::TaskId::new(1, 0)), "second waiter woken");
306            assert_eq!(crate::waker::next_ready(), None);
307            // Both registrations consumed; no leak.
308            assert_eq!(sem.debug_waiters()[1], 0);
309            assert_eq!(sem.debug_waiters()[2], 0);
310        }
311    }
312
313    #[test]
314    fn remove_waiter_on_drop_clears_registration() {
315        crate::kernel_test! {
316            let sem: Semaphore<1> = Semaphore::new(0);
317            sem.register_waiter(crate::task::TaskId::new(3, 1));
318            assert_eq!(sem.debug_waiters()[3], 1 << 1);
319            sem.remove_waiter(crate::task::TaskId::new(3, 1));
320            assert_eq!(sem.debug_waiters()[3], 0);
321        }
322    }
323}