1#![cfg_attr(not(feature = "sync"), allow(unreachable_pub, dead_code))]
2use crate::loom::cell::UnsafeCell;
19use crate::loom::sync::atomic::AtomicUsize;
20use crate::loom::sync::{Mutex, MutexGuard};
21use crate::util::linked_list::{self, LinkedList};
22#[cfg(all(tokio_unstable, feature = "tracing"))]
23use crate::util::trace;
24use crate::util::WakeList;
25
26use std::future::Future;
27use std::marker::PhantomPinned;
28use std::pin::Pin;
29use std::ptr::NonNull;
30use std::sync::atomic::Ordering::*;
31use std::task::{ready, Context, Poll, Waker};
32use std::{cmp, fmt};
33
34pub(crate) struct Semaphore {
36 waiters: Mutex<Waitlist>,
37 permits: AtomicUsize,
39 #[cfg(all(tokio_unstable, feature = "tracing"))]
40 resource_span: tracing::Span,
41}
42
43struct Waitlist {
44 queue: LinkedList<Waiter>,
45 closed: bool,
46}
47
48#[derive(Debug, PartialEq, Eq)]
52pub enum TryAcquireError {
53 Closed,
57
58 NoPermits,
60}
61#[derive(Debug)]
69pub struct AcquireError(());
70
71pub(crate) struct Acquire<'a> {
72 node: Waiter,
73 semaphore: &'a Semaphore,
74 num_permits: usize,
75 queued: bool,
83}
84
85struct Waiter {
87 state: AtomicUsize,
92
93 waker: UnsafeCell<Option<Waker>>,
99
100 pointers: linked_list::Pointers<Waiter>,
113
114 #[cfg(all(tokio_unstable, feature = "tracing"))]
115 ctx: trace::AsyncOpTracingCtx,
116
117 _p: PhantomPinned,
119}
120
121generate_addr_of_methods! {
122 impl<> Waiter {
123 unsafe fn addr_of_pointers(self: NonNull<Self>) -> NonNull<linked_list::Pointers<Waiter>> {
124 &self.pointers
125 }
126 }
127}
128
129impl Semaphore {
130 pub(crate) const MAX_PERMITS: usize = usize::MAX >> 3;
138 const CLOSED: usize = 1;
139 const PERMIT_SHIFT: usize = 1;
143
144 pub(crate) fn new(permits: usize) -> Self {
148 assert!(
149 permits <= Self::MAX_PERMITS,
150 "a semaphore may not have more than MAX_PERMITS permits ({})",
151 Self::MAX_PERMITS
152 );
153
154 #[cfg(all(tokio_unstable, feature = "tracing"))]
155 let resource_span = {
156 let resource_span = tracing::trace_span!(
157 parent: None,
158 "runtime.resource",
159 concrete_type = "Semaphore",
160 kind = "Sync",
161 is_internal = true
162 );
163
164 resource_span.in_scope(|| {
165 tracing::trace!(
166 target: "runtime::resource::state_update",
167 permits = permits,
168 permits.op = "override",
169 )
170 });
171 resource_span
172 };
173
174 Self {
175 permits: AtomicUsize::new(permits << Self::PERMIT_SHIFT),
176 waiters: Mutex::new(Waitlist {
177 queue: LinkedList::new(),
178 closed: false,
179 }),
180 #[cfg(all(tokio_unstable, feature = "tracing"))]
181 resource_span,
182 }
183 }
184
185 #[cfg(not(all(loom, test)))]
189 pub(crate) const fn const_new(permits: usize) -> Self {
190 assert!(permits <= Self::MAX_PERMITS);
191
192 Self {
193 permits: AtomicUsize::new(permits << Self::PERMIT_SHIFT),
194 waiters: Mutex::const_new(Waitlist {
195 queue: LinkedList::new(),
196 closed: false,
197 }),
198 #[cfg(all(tokio_unstable, feature = "tracing"))]
199 resource_span: tracing::Span::none(),
200 }
201 }
202
203 pub(crate) fn new_closed() -> Self {
205 Self {
206 permits: AtomicUsize::new(Self::CLOSED),
207 waiters: Mutex::new(Waitlist {
208 queue: LinkedList::new(),
209 closed: true,
210 }),
211 #[cfg(all(tokio_unstable, feature = "tracing"))]
212 resource_span: tracing::Span::none(),
213 }
214 }
215
216 #[cfg(not(all(loom, test)))]
218 pub(crate) const fn const_new_closed() -> Self {
219 Self {
220 permits: AtomicUsize::new(Self::CLOSED),
221 waiters: Mutex::const_new(Waitlist {
222 queue: LinkedList::new(),
223 closed: true,
224 }),
225 #[cfg(all(tokio_unstable, feature = "tracing"))]
226 resource_span: tracing::Span::none(),
227 }
228 }
229
230 pub(crate) fn available_permits(&self) -> usize {
232 self.permits.load(Acquire) >> Self::PERMIT_SHIFT
233 }
234
235 pub(crate) fn release(&self, added: usize) {
239 if added == 0 {
240 return;
241 }
242
243 self.add_permits_locked(added, self.waiters.lock());
245 }
246
247 pub(crate) fn close(&self) {
250 let mut waiters = self.waiters.lock();
251 self.permits.fetch_or(Self::CLOSED, Release);
259 waiters.closed = true;
260 while let Some(mut waiter) = waiters.queue.pop_back() {
261 let waker = unsafe { waiter.as_mut().waker.with_mut(|waker| (*waker).take()) };
262 if let Some(waker) = waker {
263 waker.wake();
264 }
265 }
266 }
267
268 pub(crate) fn is_closed(&self) -> bool {
270 self.permits.load(Acquire) & Self::CLOSED == Self::CLOSED
271 }
272
273 pub(crate) fn try_acquire(&self, num_permits: usize) -> Result<(), TryAcquireError> {
274 assert!(
275 num_permits <= Self::MAX_PERMITS,
276 "a semaphore may not have more than MAX_PERMITS permits ({})",
277 Self::MAX_PERMITS
278 );
279 let num_permits = num_permits << Self::PERMIT_SHIFT;
280 let mut curr = self.permits.load(Acquire);
281 loop {
282 if curr & Self::CLOSED == Self::CLOSED {
284 return Err(TryAcquireError::Closed);
285 }
286
287 if curr < num_permits {
289 return Err(TryAcquireError::NoPermits);
290 }
291
292 let next = curr - num_permits;
293
294 match self.permits.compare_exchange(curr, next, AcqRel, Acquire) {
295 Ok(_) => {
296 return Ok(());
298 }
299 Err(actual) => curr = actual,
300 }
301 }
302 }
303
304 pub(crate) fn acquire(&self, num_permits: usize) -> Acquire<'_> {
305 Acquire::new(self, num_permits)
306 }
307
308 fn add_permits_locked(&self, mut rem: usize, waiters: MutexGuard<'_, Waitlist>) {
314 let mut wakers = WakeList::new();
315 let mut lock = Some(waiters);
316 let mut is_empty = false;
317 while rem > 0 {
318 let mut waiters = lock.take().unwrap_or_else(|| self.waiters.lock());
319 'inner: while wakers.can_push() {
320 let _assigned = match waiters.queue.last() {
322 Some(waiter) => {
323 let (should_remove, assigned) = waiter.assign_permits(&mut rem);
324 if !should_remove {
325 #[cfg(all(tokio_unstable, feature = "tracing"))]
326 waiter.trace_assigned(assigned);
327 break 'inner;
328 }
329 assigned
330 }
331 None => {
332 is_empty = true;
333 break 'inner;
336 }
337 };
338 let mut waiter = waiters.queue.pop_back().unwrap();
339 if let Some(waker) =
340 unsafe { waiter.as_mut().waker.with_mut(|waker| (*waker).take()) }
341 {
342 wakers.push(waker);
343 }
344 #[cfg(all(tokio_unstable, feature = "tracing"))]
346 unsafe { waiter.as_ref() }.trace_assigned(_assigned);
347 }
348
349 if rem > 0 && is_empty {
350 let permits = rem;
351 assert!(
352 permits <= Self::MAX_PERMITS,
353 "cannot add more than MAX_PERMITS permits ({})",
354 Self::MAX_PERMITS
355 );
356 let prev = self.permits.fetch_add(rem << Self::PERMIT_SHIFT, Release);
357 let prev = prev >> Self::PERMIT_SHIFT;
358 assert!(
359 prev + permits <= Self::MAX_PERMITS,
360 "number of added permits ({}) would overflow MAX_PERMITS ({})",
361 rem,
362 Self::MAX_PERMITS
363 );
364
365 #[cfg(all(tokio_unstable, feature = "tracing"))]
367 self.resource_span.in_scope(|| {
368 tracing::trace!(
369 target: "runtime::resource::state_update",
370 permits = rem,
371 permits.op = "add",
372 )
373 });
374
375 rem = 0;
376 }
377
378 drop(waiters); wakers.wake_all();
381 }
382
383 assert_eq!(rem, 0);
384 }
385
386 pub(crate) fn forget_permits(&self, n: usize) -> usize {
391 if n == 0 {
392 return 0;
393 }
394
395 let mut curr_bits = self.permits.load(Acquire);
396 loop {
397 let curr = curr_bits >> Self::PERMIT_SHIFT;
398 let new = curr.saturating_sub(n);
399 match self.permits.compare_exchange_weak(
400 curr_bits,
401 (new << Self::PERMIT_SHIFT) | (curr_bits & Self::CLOSED),
402 AcqRel,
403 Acquire,
404 ) {
405 Ok(_) => return std::cmp::min(curr, n),
406 Err(actual) => curr_bits = actual,
407 };
408 }
409 }
410
411 fn poll_acquire(
412 &self,
413 cx: &mut Context<'_>,
414 num_permits: usize,
415 node: Pin<&mut Waiter>,
416 queued: &mut bool,
417 ) -> Poll<Result<(), AcquireError>> {
418 let mut acquired = 0;
419
420 let needed = if *queued {
421 node.state.load(Acquire) << Self::PERMIT_SHIFT
422 } else {
423 num_permits << Self::PERMIT_SHIFT
424 };
425
426 let mut lock = None;
427 let mut curr = self.permits.load(Acquire);
430 let mut waiters = loop {
431 if curr & Self::CLOSED > 0 {
433 return Poll::Ready(Err(AcquireError::closed()));
434 }
435
436 let mut remaining = 0;
437 let total = curr
438 .checked_add(acquired)
439 .expect("number of permits must not overflow");
440 let (next, acq) = if total >= needed {
441 let next = curr - (needed - acquired);
442 (next, needed >> Self::PERMIT_SHIFT)
443 } else {
444 remaining = (needed - acquired) - curr;
445 (0, curr >> Self::PERMIT_SHIFT)
446 };
447
448 if remaining > 0 && lock.is_none() {
449 lock = Some(self.waiters.lock());
457 }
458
459 match self.permits.compare_exchange(curr, next, AcqRel, Acquire) {
460 Ok(_) => {
461 acquired += acq;
462 if remaining == 0 {
463 if !*queued {
464 node.state.store(0, Release);
470 *queued = true;
471
472 #[cfg(all(tokio_unstable, feature = "tracing"))]
473 self.resource_span.in_scope(|| {
474 tracing::trace!(
475 target: "runtime::resource::state_update",
476 permits = acquired,
477 permits.op = "sub",
478 );
479 tracing::trace!(
480 target: "runtime::resource::async_op::state_update",
481 permits_obtained = acquired,
482 permits.op = "add",
483 )
484 });
485
486 return Poll::Ready(Ok(()));
487 } else if lock.is_none() {
488 break self.waiters.lock();
489 }
490 }
491 break lock.expect("lock must be acquired before waiting");
492 }
493 Err(actual) => curr = actual,
494 }
495 };
496
497 if waiters.closed {
498 return Poll::Ready(Err(AcquireError::closed()));
499 }
500
501 let was_queued = *queued;
504 *queued = true;
505
506 #[cfg(all(tokio_unstable, feature = "tracing"))]
507 let sub_permits = acquired;
508 let (should_remove, _assigned) = node.assign_permits(&mut acquired);
509
510 #[cfg(all(tokio_unstable, feature = "tracing"))]
514 self.resource_span.in_scope(|| {
515 tracing::trace!(
516 target: "runtime::resource::state_update",
517 permits = sub_permits,
518 permits.op = "sub",
519 )
520 });
521 #[cfg(all(tokio_unstable, feature = "tracing"))]
522 node.trace_assigned(_assigned);
523
524 if should_remove {
525 self.add_permits_locked(acquired, waiters);
526 return Poll::Ready(Ok(()));
527 }
528
529 assert_eq!(acquired, 0);
530 let mut old_waker = None;
531
532 node.waker.with_mut(|waker| {
534 let waker = unsafe { &mut *waker };
536 if waker
538 .as_ref()
539 .map_or(true, |waker| !waker.will_wake(cx.waker()))
540 {
541 old_waker = waker.replace(cx.waker().clone());
542 }
543 });
544
545 if !was_queued {
547 let node = unsafe {
548 let node = Pin::into_inner_unchecked(node) as *mut _;
549 NonNull::new_unchecked(node)
550 };
551
552 waiters.queue.push_front(node);
553 }
554 drop(waiters);
555 drop(old_waker);
556
557 Poll::Pending
558 }
559}
560
561impl fmt::Debug for Semaphore {
562 fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
563 fmt.debug_struct("Semaphore")
564 .field("permits", &self.available_permits())
565 .finish()
566 }
567}
568
569impl Waiter {
570 fn new(
571 num_permits: usize,
572 #[cfg(all(tokio_unstable, feature = "tracing"))] ctx: trace::AsyncOpTracingCtx,
573 ) -> Self {
574 Waiter {
575 waker: UnsafeCell::new(None),
576 state: AtomicUsize::new(num_permits),
577 pointers: linked_list::Pointers::new(),
578 #[cfg(all(tokio_unstable, feature = "tracing"))]
579 ctx,
580 _p: PhantomPinned,
581 }
582 }
583
584 fn assign_permits(&self, n: &mut usize) -> (bool, usize) {
589 let mut curr = self.state.load(Acquire);
590 loop {
591 let assign = cmp::min(curr, *n);
592 let next = curr - assign;
593 match self.state.compare_exchange(curr, next, AcqRel, Acquire) {
594 Ok(_) => {
595 *n -= assign;
596 return (next == 0, assign);
597 }
598 Err(actual) => curr = actual,
599 }
600 }
601 }
602
603 #[cfg(all(tokio_unstable, feature = "tracing"))]
606 fn trace_assigned(&self, assigned: usize) {
607 self.ctx.async_op_span.in_scope(|| {
608 tracing::trace!(
609 target: "runtime::resource::async_op::state_update",
610 permits_obtained = assigned,
611 permits.op = "add",
612 );
613 });
614 }
615}
616
617impl Future for Acquire<'_> {
618 type Output = Result<(), AcquireError>;
619
620 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
621 ready!(crate::trace::trace_leaf(cx));
622
623 #[cfg(all(tokio_unstable, feature = "tracing"))]
624 let _resource_span = self.node.ctx.resource_span.clone().entered();
625 #[cfg(all(tokio_unstable, feature = "tracing"))]
626 let _async_op_span = self.node.ctx.async_op_span.clone().entered();
627 #[cfg(all(tokio_unstable, feature = "tracing"))]
628 let _async_op_poll_span = self.node.ctx.async_op_poll_span.clone().entered();
629
630 let (node, semaphore, needed, queued) = self.project();
631
632 #[cfg(all(tokio_unstable, feature = "tracing"))]
634 let coop = ready!(trace_poll_op!(
635 "poll_acquire",
636 crate::task::coop::poll_proceed(cx),
637 ));
638
639 #[cfg(not(all(tokio_unstable, feature = "tracing")))]
640 let coop = ready!(crate::task::coop::poll_proceed(cx));
641
642 let result = match semaphore.poll_acquire(cx, needed, node, queued) {
643 Poll::Pending => Poll::Pending,
644 Poll::Ready(r) => {
645 coop.made_progress();
646 r?;
647 Poll::Ready(Ok(()))
648 }
649 };
650
651 #[cfg(all(tokio_unstable, feature = "tracing"))]
652 let result = trace_poll_op!("poll_acquire", result);
653
654 if result.is_ready() {
659 *queued = false;
660 }
661
662 result
663 }
664}
665
666impl<'a> Acquire<'a> {
667 fn new(semaphore: &'a Semaphore, num_permits: usize) -> Self {
668 assert!(
669 num_permits <= Semaphore::MAX_PERMITS,
670 "a semaphore may not have more than MAX_PERMITS permits ({})",
671 Semaphore::MAX_PERMITS
672 );
673
674 #[cfg(any(not(tokio_unstable), not(feature = "tracing")))]
675 return Self {
676 node: Waiter::new(num_permits),
677 semaphore,
678 num_permits,
679 queued: false,
680 };
681
682 #[cfg(all(tokio_unstable, feature = "tracing"))]
683 return semaphore.resource_span.in_scope(|| {
684 let async_op_span =
685 tracing::trace_span!("runtime.resource.async_op", source = "Acquire::new");
686 let async_op_poll_span = async_op_span.in_scope(|| {
687 tracing::trace!(
688 target: "runtime::resource::async_op::state_update",
689 permits_requested = num_permits,
690 permits.op = "override",
691 );
692
693 tracing::trace!(
694 target: "runtime::resource::async_op::state_update",
695 permits_obtained = 0usize,
696 permits.op = "override",
697 );
698
699 tracing::trace_span!("runtime.resource.async_op.poll")
700 });
701
702 let ctx = trace::AsyncOpTracingCtx {
703 async_op_span,
704 async_op_poll_span,
705 resource_span: semaphore.resource_span.clone(),
706 };
707
708 Self {
709 node: Waiter::new(num_permits, ctx),
710 semaphore,
711 num_permits,
712 queued: false,
713 }
714 });
715 }
716
717 fn project(self: Pin<&mut Self>) -> (Pin<&mut Waiter>, &Semaphore, usize, &mut bool) {
718 fn is_unpin<T: Unpin>() {}
719 unsafe {
720 is_unpin::<&Semaphore>();
723 is_unpin::<&mut bool>();
724 is_unpin::<usize>();
725
726 let this = self.get_unchecked_mut();
727 (
728 Pin::new_unchecked(&mut this.node),
729 this.semaphore,
730 this.num_permits,
731 &mut this.queued,
732 )
733 }
734 }
735}
736
737impl Drop for Acquire<'_> {
738 fn drop(&mut self) {
739 if !self.queued {
742 return;
743 }
744
745 let mut waiters = self.semaphore.waiters.lock();
749
750 let node = NonNull::from(&mut self.node);
752 unsafe { waiters.queue.remove(node) };
754
755 let acquired_permits = self.num_permits - self.node.state.load(Acquire);
756 if acquired_permits > 0 {
757 self.semaphore.add_permits_locked(acquired_permits, waiters);
758 }
759 }
760}
761
762unsafe impl Sync for Acquire<'_> {}
768
769impl AcquireError {
772 fn closed() -> AcquireError {
773 AcquireError(())
774 }
775}
776
777impl fmt::Display for AcquireError {
778 fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
779 write!(fmt, "semaphore closed")
780 }
781}
782
783impl std::error::Error for AcquireError {}
784
785impl TryAcquireError {
788 #[allow(dead_code)] pub(crate) fn is_closed(&self) -> bool {
791 matches!(self, TryAcquireError::Closed)
792 }
793
794 #[allow(dead_code)] pub(crate) fn is_no_permits(&self) -> bool {
798 matches!(self, TryAcquireError::NoPermits)
799 }
800}
801
802impl fmt::Display for TryAcquireError {
803 fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result {
804 match self {
805 TryAcquireError::Closed => write!(fmt, "semaphore closed"),
806 TryAcquireError::NoPermits => write!(fmt, "no permits available"),
807 }
808 }
809}
810
811impl std::error::Error for TryAcquireError {}
812
813unsafe impl linked_list::Link for Waiter {
817 type Handle = NonNull<Waiter>;
818 type Target = Waiter;
819
820 fn as_raw(handle: &Self::Handle) -> NonNull<Waiter> {
821 *handle
822 }
823
824 unsafe fn from_raw(ptr: NonNull<Waiter>) -> NonNull<Waiter> {
825 ptr
826 }
827
828 unsafe fn pointers(target: NonNull<Waiter>) -> NonNull<linked_list::Pointers<Waiter>> {
829 unsafe { Waiter::addr_of_pointers(target) }
830 }
831}