1use core::future::Future;
11use core::pin::Pin;
12use core::task::{Context, Poll};
13
14use crate::waker;
15
16pub struct Semaphore<const MAX: u8> {
35 count: crate::sync::atomic::AtomicU8,
37 waiters: [crate::sync::atomic::AtomicU32; 32],
42}
43
44impl<const MAX: u8> Semaphore<MAX> {
45 #[cfg(not(loom))]
47 pub const fn new(initial: u8) -> Self {
48 Self {
49 count: crate::sync::atomic::AtomicU8::new(initial),
50 waiters: [const { crate::sync::atomic::AtomicU32::new(0) }; 32],
53 }
54 }
55
56 #[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 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 pub fn acquire(&self) -> Acquire<'_, MAX> {
94 Acquire {
95 sem: self,
96 registered: None,
97 }
98 }
99
100 pub fn release(&self) {
104 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; }
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, }
135 }
136 }
137
138 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 #[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
179pub struct Acquire<'a, const MAX: u8> {
181 sem: &'a Semaphore<MAX>,
182 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 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 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
225unsafe 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 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 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 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 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}