diff --git a/tokio/src/sync/mpsc/chan.rs b/tokio/src/sync/mpsc/chan.rs index 7a15e8b3a..847a0b708 100644 --- a/tokio/src/sync/mpsc/chan.rs +++ b/tokio/src/sync/mpsc/chan.rs @@ -408,7 +408,7 @@ impl Semaphore for (crate::sync::semaphore_ll::Semaphore, usize) { } fn drop_permit(&self, permit: &mut Permit) { - permit.release(&self.0); + permit.release(1, &self.0); } fn add_permit(&self) { @@ -425,17 +425,17 @@ impl Semaphore for (crate::sync::semaphore_ll::Semaphore, usize) { permit: &mut Permit, ) -> Poll> { permit - .poll_acquire(cx, &self.0) + .poll_acquire(cx, 1, &self.0) .map_err(|_| ClosedError::new()) } fn try_acquire(&self, permit: &mut Permit) -> Result<(), TrySendError> { - permit.try_acquire(&self.0)?; + permit.try_acquire(1, &self.0)?; Ok(()) } fn forget(&self, permit: &mut Self::Permit) { - permit.forget() + permit.forget(1); } fn close(&self) { diff --git a/tokio/src/sync/mutex.rs b/tokio/src/sync/mutex.rs index fe891159c..484513576 100644 --- a/tokio/src/sync/mutex.rs +++ b/tokio/src/sync/mutex.rs @@ -156,7 +156,7 @@ impl Mutex { lock: self, permit: semaphore::Permit::new(), }; - poll_fn(|cx| guard.permit.poll_acquire(cx, &self.s)) + poll_fn(|cx| guard.permit.poll_acquire(cx, 1, &self.s)) .await .unwrap_or_else(|_| { // The semaphore was closed. but, we never explicitly close it, and we have a @@ -169,7 +169,7 @@ impl Mutex { /// Try to acquire the lock pub fn try_lock(&self) -> Result, TryLockError> { let mut permit = semaphore::Permit::new(); - match permit.try_acquire(&self.s) { + match permit.try_acquire(1, &self.s) { Ok(_) => Ok(MutexGuard { lock: self, permit }), Err(_) => Err(TryLockError(())), } @@ -178,7 +178,7 @@ impl Mutex { impl<'a, T> Drop for MutexGuard<'a, T> { fn drop(&mut self) { - self.permit.release(&self.lock.s); + self.permit.release(1, &self.lock.s); } } diff --git a/tokio/src/sync/semaphore.rs b/tokio/src/sync/semaphore.rs index 2cfb5d349..13d5cfb27 100644 --- a/tokio/src/sync/semaphore.rs +++ b/tokio/src/sync/semaphore.rs @@ -60,7 +60,7 @@ impl Semaphore { sem: &self, ll_permit: ll::Permit::new(), }; - poll_fn(|cx| permit.ll_permit.poll_acquire(cx, &self.ll_sem)) + poll_fn(|cx| permit.ll_permit.poll_acquire(cx, 1, &self.ll_sem)) .await .unwrap(); permit @@ -69,7 +69,7 @@ impl Semaphore { /// Try to acquire a permit form the semaphore pub fn try_acquire(&self) -> Result, TryAcquireError> { let mut ll_permit = ll::Permit::new(); - match ll_permit.try_acquire(&self.ll_sem) { + match ll_permit.try_acquire(1, &self.ll_sem) { Ok(_) => Ok(SemaphorePermit { sem: self, ll_permit, @@ -84,12 +84,12 @@ impl<'a> SemaphorePermit<'a> { /// This can be used to reduce the amount of permits available from a /// semaphore. pub fn forget(mut self) { - self.ll_permit.forget(); + self.ll_permit.forget(1); } } impl<'a> Drop for SemaphorePermit<'_> { fn drop(&mut self) { - self.ll_permit.release(&self.sem.ll_sem); + self.ll_permit.release(1, &self.sem.ll_sem); } } diff --git a/tokio/src/sync/semaphore_ll.rs b/tokio/src/sync/semaphore_ll.rs index 0ce858386..baed5f0a6 100644 --- a/tokio/src/sync/semaphore_ll.rs +++ b/tokio/src/sync/semaphore_ll.rs @@ -10,17 +10,15 @@ //! section. If no permits are available, then acquiring the semaphore returns //! `Pending`. The task is woken once a permit becomes available. -use crate::loom::{ - cell::CausalCell, - future::AtomicWaker, - sync::atomic::{AtomicPtr, AtomicUsize}, - thread, -}; +use crate::loom::cell::CausalCell; +use crate::loom::future::AtomicWaker; +use crate::loom::sync::atomic::{AtomicPtr, AtomicUsize}; +use crate::loom::thread; +use std::cmp; use std::fmt; use std::ptr::{self, NonNull}; use std::sync::atomic::Ordering::{self, AcqRel, Acquire, Relaxed, Release}; -use std::sync::Arc; use std::task::Poll::{Pending, Ready}; use std::task::{Context, Poll}; use std::usize; @@ -32,13 +30,13 @@ pub(crate) struct Semaphore { state: AtomicUsize, /// waiter queue head pointer. - head: CausalCell>, + head: CausalCell>, /// Coordinates access to the queue head. rx_lock: AtomicUsize, /// Stub waiter node used as part of the MPSC channel algorithm. - stub: Box, + stub: Box, } /// A semaphore permit @@ -54,7 +52,7 @@ pub(crate) struct Semaphore { /// before dropping the permit. #[derive(Debug)] pub(crate) struct Permit { - waiter: Option>, + waiter: Option>, state: PermitState, } @@ -64,29 +62,24 @@ pub(crate) struct AcquireError(()); /// Error returned by `Permit::try_acquire`. #[derive(Debug)] -pub(crate) struct TryAcquireError { - pub(crate) kind: ErrorKind, -} - -#[derive(Debug)] -pub(crate) enum ErrorKind { +pub(crate) enum TryAcquireError { Closed, NoPermits, } /// Node used to notify the semaphore waiter when permit is available. #[derive(Debug)] -struct WaiterNode { +struct Waiter { /// Stores waiter state. /// - /// See `NodeState` for more details. + /// See `WaiterState` for more details. state: AtomicUsize, /// Task to wake when a permit is made available. waker: AtomicWaker, /// Next pointer in the queue of waiting senders. - next: AtomicPtr, + next: AtomicPtr, } /// Semaphore state @@ -103,46 +96,55 @@ struct WaiterNode { struct SemState(usize); /// Permit state -#[derive(Debug, Copy, Clone, Eq, PartialEq)] +#[derive(Debug, Copy, Clone)] enum PermitState { - /// The permit has not been requested. - Idle, - - /// Currently waiting for a permit to be made available and assigned to the + /// Currently waiting for permits to be made available and assigned to the /// waiter. - Waiting, + Waiting(u16), - /// The permit has been acquired. - Acquired, + /// The number of acquired permits + Acquired(u16), } -/// Waiter node state -#[derive(Debug, Copy, Clone, Eq, PartialEq)] -#[repr(usize)] -enum NodeState { - /// Not waiting for a permit and the node is not in the wait queue. - /// - /// This is the initial state. - Idle = 0, +/// State for an individual waker node +#[derive(Debug, Copy, Clone)] +struct WaiterState(usize); - /// Not waiting for a permit but the node is in the wait queue. - /// - /// This happens when the waiter has previously requested a permit, but has - /// since canceled the request. The node cannot be removed by the waiter, so - /// this state informs the receiver to skip the node when it pops it from - /// the wait queue. - Queued = 1, +/// Waiter node is in the semaphore queue +const QUEUED: usize = 0b001; - /// Waiting for a permit and the node is in the wait queue. - QueuedWaiting = 2, +/// Semaphore has been closed, no more permits will be issued. +const CLOSED: usize = 0b10; - /// The waiter has been assigned a permit and the node has been removed from - /// the queue. - Assigned = 3, +/// The permit that owns the `Waiter` dropped. +const DROPPED: usize = 0b100; - /// The semaphore has been closed. No more permits will be issued. - Closed = 4, -} +/// Represents "one requested permit" in the waiter state +const PERMIT_ONE: usize = 0b1000; + +/// Masks the waiter state to only contain bits tracking number of requested +/// permits. +const PERMIT_MASK: usize = usize::MAX - (PERMIT_ONE - 1); + +/// How much to shift a permit count to pack it into the waker state +const PERMIT_SHIFT: u32 = PERMIT_ONE.trailing_zeros(); + +/// Flag differentiating between available permits and waiter pointers. +/// +/// If we assume pointers are properly aligned, then the least significant bit +/// will always be zero. So, we use that bit to track if the value represents a +/// number. +const NUM_FLAG: usize = 0b01; + +/// Signal the semaphore is closed +const CLOSED_FLAG: usize = 0b10; + +/// Maximum number of permits a semaphore can manage +const MAX_PERMITS: usize = usize::MAX >> NUM_SHIFT; + +/// When representing "numbers", the state has to be shifted this much (to get +/// rid of the flag bit). +const NUM_SHIFT: usize = 2; // ===== impl Semaphore ===== @@ -153,8 +155,8 @@ impl Semaphore { /// /// Panics if `permits` is zero. pub(crate) fn new(permits: usize) -> Semaphore { - let stub = Box::new(WaiterNode::new()); - let ptr = NonNull::new(&*stub as *const _ as *mut _).unwrap(); + let stub = Box::new(Waiter::new()); + let ptr = NonNull::from(&*stub); // Allocations are aligned debug_assert!(ptr.as_ptr() as usize & NUM_FLAG == 0); @@ -171,32 +173,63 @@ impl Semaphore { /// Returns the current number of available permits pub(crate) fn available_permits(&self) -> usize { - let curr = SemState::load(&self.state, Acquire); + let curr = SemState(self.state.load(Acquire)); curr.available_permits() } - /// Poll for a permit - fn poll_permit( + /// Try to acquire the requested number of permits, registering the waiter + /// if not enough permits are available. + fn poll_acquire( &self, - mut permit: Option<(&mut Context<'_>, &mut Permit)>, + cx: &mut Context<'_>, + num_permits: u16, + permit: &mut Permit, ) -> Poll> { + self.poll_acquire2(num_permits, || { + let waiter = permit.waiter.get_or_insert_with(|| Box::new(Waiter::new())); + + waiter.waker.register_by_ref(cx.waker()); + + Some(NonNull::from(&**waiter)) + }) + } + + fn try_acquire(&self, num_permits: u16) -> Result<(), TryAcquireError> { + match self.poll_acquire2(num_permits, || None) { + Poll::Ready(res) => res.map_err(to_try_acquire), + Poll::Pending => Err(TryAcquireError::NoPermits), + } + } + + /// Poll for a permit + /// + /// Tries to acquire available permits first. If unable to acquire a + /// sufficient number of permits, the caller's waiter is pushed onto the + /// semaphore's wait queue. + fn poll_acquire2( + &self, + num_permits: u16, + mut get_waiter: F, + ) -> Poll> + where + F: FnMut() -> Option>, + { + let num_permits = num_permits as usize; + // Load the current state - let mut curr = SemState::load(&self.state, Acquire); + let mut curr = SemState(self.state.load(Acquire)); - // Tracks a *mut WaiterNode representing an Arc clone. - // - // This avoids having to bump the ref count unless required. - let mut maybe_strong: Option> = None; + // Saves a ref to the waiter node + let mut maybe_waiter: Option> = None; - macro_rules! undo_strong { + /// Used in branches where we attempt to push the waiter into the wait + /// queue but fail due to permits becoming available or the wait queue + /// transitioning to "closed". In this case, the waiter must be + /// transitioned back to the "idle" state. + macro_rules! revert_to_idle { () => { - if let Some(waiter) = maybe_strong { - // The waiter was cloned, but never got queued. - // Before entering `poll_permit`, the waiter was in the - // `Idle` state. We must transition the node back to the - // idle state. - let waiter = unsafe { Arc::from_raw(waiter.as_ptr()) }; - waiter.revert_to_idle(); + if let Some(waiter) = maybe_waiter { + unsafe { waiter.as_ref() }.revert_to_idle(); } }; } @@ -205,64 +238,82 @@ impl Semaphore { let mut next = curr; if curr.is_closed() { - undo_strong!(); + revert_to_idle!(); return Ready(Err(AcquireError::closed())); } - if !next.acquire_permit(&self.stub) { - debug_assert!(curr.waiter().is_some()); + let acquired = next.acquire_permits(num_permits, &self.stub); - if maybe_strong.is_none() { - if let Some((ref mut cx, ref mut permit)) = permit { - // Get the Sender's waiter node, or initialize one - let waiter = permit - .waiter - .get_or_insert_with(|| Arc::new(WaiterNode::new())); + if !acquired { + // There are not enough available permits to satisfy the + // request. The permit transitions to a waiting state. + debug_assert!(curr.waiter().is_some() || curr.available_permits() < num_permits); - waiter.register(cx); + if let Some(waiter) = maybe_waiter.as_ref() { + // Safety: the caller owns the waiter. + let w = unsafe { waiter.as_ref() }; + w.set_permits_to_acquire(num_permits - curr.available_permits()); + } else { + // Get the waiter for the permit. + if let Some(waiter) = get_waiter() { + // Safety: the caller owns the waiter. + let w = unsafe { waiter.as_ref() }; - if !waiter.to_queued_waiting() { + // If there are any currently available permits, the + // waiter acquires those immediately and waits for the + // remaining permits to become available. + if !w.to_queued(num_permits - curr.available_permits()) { // The node is alrady queued, there is no further work // to do. return Pending; } - maybe_strong = Some(WaiterNode::into_non_null(waiter.clone())); + maybe_waiter = Some(waiter); } else { - // If no `waiter`, then the task is not registered and there - // is no further work to do. + // No waiter, this indicates the caller does not wish to + // "wait", so there is nothing left to do. return Pending; } } - next.set_waiter(maybe_strong.unwrap()); + next.set_waiter(maybe_waiter.unwrap()); } debug_assert_ne!(curr.0, 0); debug_assert_ne!(next.0, 0); - match next.compare_exchange(&self.state, curr, AcqRel, Acquire) { + match self.state.compare_exchange(curr.0, next.0, AcqRel, Acquire) { Ok(_) => { - match curr.waiter() { - Some(prev_waiter) => { - let waiter = maybe_strong.unwrap(); + if acquired { + // Successfully acquire permits **without** queuing the + // waiter node. The waiter node is not currently in the + // queue. + revert_to_idle!(); + return Ready(Ok(())); + } else { + // The node is pushed into the queue, the final step is + // to set the node's "next" pointer to return the wait + // queue into a consistent state. - // Finish pushing - unsafe { - prev_waiter.as_ref().next.store(waiter.as_ptr(), Release); - } + let prev_waiter = + curr.waiter().unwrap_or_else(|| NonNull::from(&*self.stub)); - return Pending; + let waiter = maybe_waiter.unwrap(); + + // Link the nodes. + // + // Safety: the mpsc algorithm guarantees the old tail of + // the queue is not removed from the queue during the + // push process. + unsafe { + prev_waiter.as_ref().store_next(waiter); } - None => { - undo_strong!(); - return Ready(Ok(())); - } + return Pending; } } Err(actual) => { - curr = actual; + curr = SemState(actual); } } } @@ -332,27 +383,13 @@ impl Semaphore { /// This function is called by `add_permits` after the add lock has been /// acquired. fn add_permits_locked2(&self, mut n: usize, closed: bool) { - while n > 0 || closed { - let waiter = match self.pop(n, closed) { - Some(waiter) => waiter, - None => { - return; - } - }; - - if waiter.notify(closed) { - n = n.saturating_sub(1); - } + // If closing the semaphore, we want to drain the entire queue. The + // number of permits being assigned doesn't matter. + if closed { + n = usize::MAX; } - } - /// Pop a waiter - /// - /// `rem` represents the remaining number of times the caller will pop. If - /// there are no more waiters to pop, `rem` is used to set the available - /// permits. - fn pop(&self, rem: usize, closed: bool) -> Option> { - 'outer: loop { + 'outer: while n > 0 { unsafe { let mut head = self.head.with(|head| *head); let mut next_ptr = head.as_ref().next.load(Acquire); @@ -360,12 +397,14 @@ impl Semaphore { let stub = self.stub(); if head == stub { + // The stub node indicates an empty queue. Any remaining + // permits get assigned back to the semaphore. let next = match NonNull::new(next_ptr) { Some(next) => next, None => { // This loop is not part of the standard intrusive mpsc // channel algorithm. This is where we atomically pop - // the last task and add `rem` to the remaining capacity. + // the last task and add `n` to the remaining capacity. // // This modification to the pop algorithm works because, // at this point, we have not done any work (only done @@ -384,27 +423,26 @@ impl Semaphore { loop { if curr.has_waiter(&self.stub) { - // Inconsistent + // A waiter is being added concurrently. + // This is the MPSC queue's "inconsistent" + // state and we must loop and try again. thread::yield_now(); continue 'outer; } - // When closing the semaphore, nodes are popped - // with `rem == 0`. In this case, we are not - // adding permits, but notifying waiters of the - // semaphore's closed state. - if rem == 0 { + // If closing, nothing more to do. + if closed { debug_assert!(curr.is_closed(), "state = {:?}", curr); - return None; + return; } let mut next = curr; - next.release_permits(rem, &self.stub); + next.release_permits(n, &self.stub); - match next.compare_exchange(&self.state, curr, AcqRel, Acquire) { - Ok(_) => return None, + match self.state.compare_exchange(curr.0, next.0, AcqRel, Acquire) { + Ok(_) => return, Err(actual) => { - curr = actual; + curr = SemState(actual); } } } @@ -416,10 +454,19 @@ impl Semaphore { next_ptr = next.as_ref().next.load(Acquire); } + // `head` points to a waiter assign permits to the waiter. If + // all requested permits are satisfied, then we can continue, + // otherwise the node stays in the wait queue. + if !head.as_ref().assign_permits(&mut n, closed) { + assert_eq!(n, 0); + return; + } + if let Some(next) = NonNull::new(next_ptr) { self.head.with_mut(|head| *head = next); - return Some(Arc::from_raw(head.as_ptr())); + self.remove_queued(head, closed); + continue 'outer; } let state = SemState::load(&self.state, Acquire); @@ -440,7 +487,8 @@ impl Semaphore { if let Some(next) = NonNull::new(next_ptr) { self.head.with_mut(|head| *head = next); - return Some(Arc::from_raw(head.as_ptr())); + self.remove_queued(head, closed); + continue 'outer; } // Inconsistent state, loop @@ -449,37 +497,89 @@ impl Semaphore { } } - unsafe fn push_stub(&self, closed: bool) { - let stub = self.stub(); + /// The wait node has had all of its permits assigned and has been removed + /// from the wait queue. + /// + /// Attempt to remove the QUEUED bit from the node. If additional permits + /// are concurrently requested, the node must be pushed back into the wait + /// queued. + fn remove_queued(&self, waiter: NonNull, closed: bool) { + let mut curr = WaiterState(unsafe { waiter.as_ref() }.state.load(Acquire)); + loop { + if curr.is_dropped() { + // The Permit dropped, it is on us to release the memory + let _ = unsafe { Box::from_raw(waiter.as_ptr()) }; + return; + } + + // The node is removed from the queue. We attempt to unset the + // queued bit, but concurrently the waiter has requested more + // permits. When the waiter requested more permits, it saw the + // queued bit set so took no further action. This requires us to + // push the node back into the queue. + if curr.permits_to_acquire() > 0 { + // More permits are requested. The waiter must be re-queued + unsafe { + self.push_waiter(waiter, closed); + } + return; + } + + let mut next = curr; + next.unset_queued(); + + let w = unsafe { waiter.as_ref() }; + + match w.state.compare_exchange(curr.0, next.0, AcqRel, Acquire) { + Ok(_) => return, + Err(actual) => { + curr = WaiterState(actual); + } + } + } + } + + unsafe fn push_stub(&self, closed: bool) { + self.push_waiter(self.stub(), closed); + } + + unsafe fn push_waiter(&self, waiter: NonNull, closed: bool) { // Set the next pointer. This does not require an atomic operation as // this node is not accessible. The write will be flushed with the next // operation - stub.as_ref().next.store(ptr::null_mut(), Relaxed); + waiter.as_ref().next.store(ptr::null_mut(), Relaxed); // Update the tail to point to the new node. We need to see the previous // node in order to update the next pointer as well as release `task` // to any other threads calling `push`. - let prev = SemState::new_ptr(stub, closed).swap(&self.state, AcqRel); + let next = SemState::new_ptr(waiter, closed); + let prev = SemState(self.state.swap(next.0, AcqRel)); debug_assert_eq!(closed, prev.is_closed()); - // The stub is only pushed when there are pending tasks. Because of + // This function is only called when there are pending tasks. Because of // this, the state must *always* be in pointer mode. let prev = prev.waiter().unwrap(); - // We don't want the *existing* pointer to be a stub. - debug_assert_ne!(prev, stub); + // No cycles plz + debug_assert_ne!(prev, waiter); // Release `task` to the consume end. - prev.as_ref().next.store(stub.as_ptr(), Release); + prev.as_ref().next.store(waiter.as_ptr(), Release); } - fn stub(&self) -> NonNull { + fn stub(&self) -> NonNull { unsafe { NonNull::new_unchecked(&*self.stub as *const _ as *mut _) } } } +impl Drop for Semaphore { + fn drop(&mut self) { + self.close(); + } +} + impl fmt::Debug for Semaphore { fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result { fmt.debug_struct("Semaphore") @@ -501,15 +601,20 @@ impl Permit { /// /// The permit begins in the "unacquired" state. pub(crate) fn new() -> Permit { + use PermitState::Acquired; + Permit { waiter: None, - state: PermitState::Idle, + state: Acquired(0), } } /// Returns true if the permit has been acquired pub(crate) fn is_acquired(&self) -> bool { - self.state == PermitState::Acquired + match self.state { + PermitState::Acquired(num) if num > 0 => true, + _ => false, + } } /// Try to acquire the permit. If no permits are available, the current task @@ -517,70 +622,127 @@ impl Permit { pub(crate) fn poll_acquire( &mut self, cx: &mut Context<'_>, + num_permits: u16, semaphore: &Semaphore, ) -> Poll> { + use std::cmp::Ordering::*; + use PermitState::*; + match self.state { - PermitState::Idle => {} - PermitState::Waiting => { + Waiting(requested) => { + // There must be a waiter let waiter = self.waiter.as_ref().unwrap(); - if waiter.acquire(cx)? { - self.state = PermitState::Acquired; - return Ready(Ok(())); - } else { - return Pending; - } - } - PermitState::Acquired => { - return Ready(Ok(())); - } - } + match requested.cmp(&num_permits) { + Less => { + let delta = num_permits - requested; + + // Request additional permits. If the waiter has been + // dequeued, it must be re-queued. + if !waiter.try_inc_permits_to_acquire(delta as usize) { + let waiter = NonNull::from(&**waiter); + + // Ignore the result. The check for + // `permits_to_acquire()` will converge the state as + // needed + let _ = semaphore.poll_acquire2(delta, || Some(waiter))?; + } + + self.state = Waiting(num_permits); + } + Greater => { + let delta = requested - num_permits; + let to_release = waiter.try_dec_permits_to_acquire(delta as usize); + + semaphore.add_permits(to_release); + self.state = Waiting(num_permits); + } + Equal => {} + } + + if waiter.permits_to_acquire()? == 0 { + self.state = Acquired(requested); + return Ready(Ok(())); + } + + waiter.waker.register_by_ref(cx.waker()); + + if waiter.permits_to_acquire()? == 0 { + self.state = Acquired(requested); + return Ready(Ok(())); + } - match semaphore.poll_permit(Some((cx, self)))? { - Ready(()) => { - self.state = PermitState::Acquired; - Ready(Ok(())) - } - Pending => { - self.state = PermitState::Waiting; Pending } + Acquired(acquired) => { + if acquired >= num_permits { + Ready(Ok(())) + } else { + match semaphore.poll_acquire(cx, num_permits - acquired, self)? { + Ready(()) => { + self.state = Acquired(num_permits); + Ready(Ok(())) + } + Pending => { + self.state = Waiting(num_permits); + Pending + } + } + } + } } } /// Try to acquire the permit. - pub(crate) fn try_acquire(&mut self, semaphore: &Semaphore) -> Result<(), TryAcquireError> { + pub(crate) fn try_acquire( + &mut self, + num_permits: u16, + semaphore: &Semaphore, + ) -> Result<(), TryAcquireError> { + use PermitState::*; + match self.state { - PermitState::Idle => {} - PermitState::Waiting => { + Waiting(requested) => { + // There must be a waiter let waiter = self.waiter.as_ref().unwrap(); - if waiter.acquire2().map_err(to_try_acquire)? { - self.state = PermitState::Acquired; - return Ok(()); + if requested > num_permits { + let delta = requested - num_permits; + let to_release = waiter.try_dec_permits_to_acquire(delta as usize); + + semaphore.add_permits(to_release); + self.state = Waiting(num_permits); + } + + let res = waiter.permits_to_acquire().map_err(to_try_acquire)?; + + if res == 0 { + if requested < num_permits { + // Try to acquire the additional permits + semaphore.try_acquire(num_permits - requested)?; + } + + self.state = Acquired(num_permits); + Ok(()) } else { - return Err(TryAcquireError::no_permits()); + Err(TryAcquireError::NoPermits) } } - PermitState::Acquired => { - return Ok(()); - } - } + Acquired(acquired) => { + if acquired < num_permits { + semaphore.try_acquire(num_permits - acquired)?; + self.state = Acquired(num_permits); + } - match semaphore.poll_permit(None).map_err(to_try_acquire)? { - Ready(()) => { - self.state = PermitState::Acquired; Ok(()) } - Pending => Err(TryAcquireError::no_permits()), } } /// Release a permit back to the semaphore - pub(crate) fn release(&mut self, semaphore: &Semaphore) { - if self.forget2() { - semaphore.add_permits(1); - } + pub(crate) fn release(&mut self, n: u16, semaphore: &Semaphore) { + let n = self.forget(n); + semaphore.add_permits(n as usize); } /// Forget the permit **without** releasing it back to the semaphore. @@ -590,22 +752,37 @@ impl Permit { /// /// Repeatedly calling `forget` without associated calls to `add_permit` /// will result in the semaphore losing all permits. - pub(crate) fn forget(&mut self) { - self.forget2(); - } + /// + /// Will forget **at most** the number of acquired permits. This number is + /// returned. + pub(crate) fn forget(&mut self, n: u16) -> u16 { + use PermitState::*; - /// Returns `true` if the permit was acquired - fn forget2(&mut self) -> bool { match self.state { - PermitState::Idle => false, - PermitState::Waiting => { - let ret = self.waiter.as_ref().unwrap().cancel_interest(); - self.state = PermitState::Idle; - ret + Waiting(requested) => { + let n = cmp::min(n, requested); + + // Decrement + let acquired = self + .waiter + .as_ref() + .unwrap() + .try_dec_permits_to_acquire(n as usize) as u16; + + if n == requested { + self.state = Acquired(0); + } else if acquired == requested - n { + self.state = Waiting(acquired); + } else { + self.state = Waiting(requested - n); + } + + acquired } - PermitState::Acquired => { - self.state = PermitState::Idle; - true + Acquired(acquired) => { + let n = cmp::min(n, acquired); + self.state = Acquired(acquired - n); + n } } } @@ -617,6 +794,20 @@ impl Default for Permit { } } +impl Drop for Permit { + fn drop(&mut self) { + if let Some(waiter) = self.waiter.take() { + // Set the dropped flag + let state = WaiterState(waiter.state.fetch_or(DROPPED, AcqRel)); + + if state.is_queued() { + // The waiter is stored in the queue. The semaphore will drop it + std::mem::forget(waiter); + } + } + } +} + // ===== impl AcquireError ==== impl AcquireError { @@ -626,7 +817,7 @@ impl AcquireError { } fn to_try_acquire(_: AcquireError) -> TryAcquireError { - TryAcquireError::closed() + TryAcquireError::Closed } impl fmt::Display for AcquireError { @@ -640,22 +831,10 @@ impl ::std::error::Error for AcquireError {} // ===== impl TryAcquireError ===== impl TryAcquireError { - fn closed() -> TryAcquireError { - TryAcquireError { - kind: ErrorKind::Closed, - } - } - - fn no_permits() -> TryAcquireError { - TryAcquireError { - kind: ErrorKind::NoPermits, - } - } - /// Returns true if the error was caused by a closed semaphore. pub(crate) fn is_closed(&self) -> bool { - match self.kind { - ErrorKind::Closed => true, + match self { + TryAcquireError::Closed => true, _ => false, } } @@ -663,8 +842,8 @@ impl TryAcquireError { /// Returns true if the error was caused by calling `try_acquire` on a /// semaphore with no available permits. pub(crate) fn is_no_permits(&self) -> bool { - match self.kind { - ErrorKind::NoPermits => true, + match self { + TryAcquireError::NoPermits => true, _ => false, } } @@ -672,177 +851,187 @@ impl TryAcquireError { impl fmt::Display for TryAcquireError { fn fmt(&self, fmt: &mut fmt::Formatter<'_>) -> fmt::Result { - let descr = match self.kind { - ErrorKind::Closed => "semaphore closed", - ErrorKind::NoPermits => "no permits available", - }; - write!(fmt, "{}", descr) + match self { + TryAcquireError::Closed => write!(fmt, "{}", "semaphore closed"), + TryAcquireError::NoPermits => write!(fmt, "{}", "no permits available"), + } } } impl ::std::error::Error for TryAcquireError {} -// ===== impl WaiterNode ===== +// ===== impl Waiter ===== -impl WaiterNode { - fn new() -> WaiterNode { - WaiterNode { - state: AtomicUsize::new(NodeState::new().to_usize()), +impl Waiter { + fn new() -> Waiter { + Waiter { + state: AtomicUsize::new(0), waker: AtomicWaker::new(), next: AtomicPtr::new(ptr::null_mut()), } } - fn acquire(&self, cx: &mut Context<'_>) -> Result { - if self.acquire2()? { - return Ok(true); - } + fn permits_to_acquire(&self) -> Result { + let state = WaiterState(self.state.load(Acquire)); - self.waker.register_by_ref(cx.waker()); - - self.acquire2() - } - - fn acquire2(&self) -> Result { - use self::NodeState::*; - - match Idle.compare_exchange(&self.state, Assigned, AcqRel, Acquire) { - Ok(_) => Ok(true), - Err(Closed) => Err(AcquireError::closed()), - Err(_) => Ok(false), + if state.is_closed() { + Err(AcquireError(())) + } else { + Ok(state.permits_to_acquire()) } } - fn register(&self, cx: &mut Context<'_>) { - self.waker.register_by_ref(cx.waker()) - } - - /// Returns `true` if the permit has been acquired - fn cancel_interest(&self) -> bool { - use self::NodeState::*; - - match Queued.compare_exchange(&self.state, QueuedWaiting, AcqRel, Acquire) { - // Successfully removed interest from the queued node. The permit - // has not been assigned to the node. - Ok(_) => false, - // The semaphore has been closed, there is no further action to - // take. - Err(Closed) => false, - // The permit has been assigned. It must be acquired in order to - // be released back to the semaphore. - Err(Assigned) => { - match self.acquire2() { - Ok(true) => true, - // Not a reachable state - Ok(false) => panic!(), - // The semaphore has been closed, no further action to take. - Err(_) => false, - } - } - Err(state) => panic!("unexpected state = {:?}", state), - } - } - - /// Transition the state to `QueuedWaiting`. + /// Only increments the number of permits *if* the waiter is currently + /// queued. /// - /// This step can only happen from `Queued` or from `Idle`. + /// # Returns /// - /// Returns `true` if transitioning into a queued state. - fn to_queued_waiting(&self) -> bool { - use self::NodeState::*; - - let mut curr = NodeState::load(&self.state, Acquire); + /// `true` if the number of permits to acquire has been incremented. `false` + /// otherwise. On `false`, the caller should use `Semaphore::poll_acquire`. + fn try_inc_permits_to_acquire(&self, n: usize) -> bool { + let mut curr = WaiterState(self.state.load(Acquire)); loop { - debug_assert!(curr == Idle || curr == Queued, "actual = {:?}", curr); - let next = QueuedWaiting; + if !curr.is_queued() { + assert_eq!(0, curr.permits_to_acquire()); + return false; + } - match next.compare_exchange(&self.state, curr, AcqRel, Acquire) { + let mut next = curr; + next.set_permits_to_acquire(n + curr.permits_to_acquire()); + + match self.state.compare_exchange(curr.0, next.0, AcqRel, Acquire) { + Ok(_) => return true, + Err(actual) => curr = WaiterState(actual), + } + } + } + + /// Try to decrement the number of permits to acquire. This returns the + /// actual number of permits that were decremented. The delta betweeen `n` + /// and the return has been assigned to the permit and the caller must + /// assign these back to the semaphore. + fn try_dec_permits_to_acquire(&self, n: usize) -> usize { + let mut curr = WaiterState(self.state.load(Acquire)); + + loop { + if !curr.is_queued() { + assert_eq!(0, curr.permits_to_acquire()); + } + + let delta = cmp::min(n, curr.permits_to_acquire()); + let rem = curr.permits_to_acquire() - delta; + + let mut next = curr; + next.set_permits_to_acquire(rem); + + match self.state.compare_exchange(curr.0, next.0, AcqRel, Acquire) { + Ok(_) => return n - delta, + Err(actual) => curr = WaiterState(actual), + } + } + } + + /// Store the number of remaining permits needed to satisfy the waiter and + /// transition to the "QUEUED" state. + /// + /// # Returns + /// + /// `true` if the `QUEUED` bit was set as part of the transition. + fn to_queued(&self, num_permits: usize) -> bool { + let mut curr = WaiterState(self.state.load(Acquire)); + + // The waiter should **not** be waiting for any permits. + debug_assert_eq!(curr.permits_to_acquire(), 0); + + loop { + let mut next = curr; + next.set_permits_to_acquire(num_permits); + next.set_queued(); + + match self.state.compare_exchange(curr.0, next.0, AcqRel, Acquire) { Ok(_) => { if curr.is_queued() { return false; } else { - // Transitioned to queued, reset next pointer + // Make sure the next pointer is null self.next.store(ptr::null_mut(), Relaxed); return true; } } - Err(actual) => { - curr = actual; - } + Err(actual) => curr = WaiterState(actual), } } } - /// Notify the waiter + /// Set the number of permits to acquire. /// - /// Returns `true` if the waiter accepts the notification - fn notify(&self, closed: bool) -> bool { - use self::NodeState::*; + /// This function is only called when the waiter is being inserted into the + /// wait queue. Because of this, there are no concurrent threads that can + /// modify the state and using `store` is safe. + fn set_permits_to_acquire(&self, num_permits: usize) { + debug_assert!(WaiterState(self.state.load(Acquire)).is_queued()); - // Assume QueuedWaiting state - let mut curr = QueuedWaiting; + let mut state = WaiterState(QUEUED); + state.set_permits_to_acquire(num_permits); + + self.state.store(state.0, Release); + } + + /// Assign permits to the waiter. + /// + /// Returns `true` if the waiter should be removed from the queue + fn assign_permits(&self, n: &mut usize, closed: bool) -> bool { + let mut curr = WaiterState(self.state.load(Acquire)); loop { - let next = match curr { - Queued => Idle, - QueuedWaiting => { - if closed { - Closed + let mut next = curr; + + // Number of permits to assign to this waiter + let assign = cmp::min(curr.permits_to_acquire(), *n); + + // Assign the permits + next.set_permits_to_acquire(curr.permits_to_acquire() - assign); + + if closed { + next.set_closed(); + } + + match self.state.compare_exchange(curr.0, next.0, AcqRel, Acquire) { + Ok(_) => { + // Update `n` + *n -= assign; + + if next.permits_to_acquire() == 0 { + if curr.permits_to_acquire() > 0 { + self.waker.wake(); + } + + return true; } else { - Assigned + return false; } } - actual => panic!("actual = {:?}", actual), - }; - - match next.compare_exchange(&self.state, curr, AcqRel, Acquire) { - Ok(_) => match curr { - QueuedWaiting => { - self.waker.wake(); - return true; - } - _ => return false, - }, - Err(actual) => curr = actual, + Err(actual) => curr = WaiterState(actual), } } } fn revert_to_idle(&self) { - use self::NodeState::Idle; - - // There are no other handles to the node - NodeState::store(&self.state, Idle, Relaxed); + // An idle node is not waiting on any permits + self.state.store(0, Relaxed); } - #[allow(clippy::wrong_self_convention)] // https://github.com/rust-lang/rust-clippy/issues/4293 - fn into_non_null(self: Arc) -> NonNull { - let ptr = Arc::into_raw(self); - unsafe { NonNull::new_unchecked(ptr as *mut _) } + fn store_next(&self, next: NonNull) { + self.next.store(next.as_ptr(), Release); } } -// ===== impl State ===== - -/// Flag differentiating between available permits and waiter pointers. -/// -/// If we assume pointers are properly aligned, then the least significant bit -/// will always be zero. So, we use that bit to track if the value represents a -/// number. -const NUM_FLAG: usize = 0b01; - -const CLOSED_FLAG: usize = 0b10; - -const MAX_PERMITS: usize = usize::MAX >> NUM_SHIFT; - -/// When representing "numbers", the state has to be shifted this much (to get -/// rid of the flag bit). -const NUM_SHIFT: usize = 2; +// ===== impl SemState ===== impl SemState { /// Returns a new default `State` value. - fn new(permits: usize, stub: &WaiterNode) -> SemState { + fn new(permits: usize, stub: &Waiter) -> SemState { assert!(permits <= MAX_PERMITS); if permits > 0 { @@ -853,7 +1042,7 @@ impl SemState { } /// Returns a `State` tracking `ptr` as the tail of the queue. - fn new_ptr(tail: NonNull, closed: bool) -> SemState { + fn new_ptr(tail: NonNull, closed: bool) -> SemState { let mut val = tail.as_ptr() as usize; if closed { @@ -877,25 +1066,27 @@ impl SemState { self.0 & NUM_FLAG == NUM_FLAG } - fn has_waiter(self, stub: &WaiterNode) -> bool { + fn has_waiter(self, stub: &Waiter) -> bool { !self.has_available_permits() && !self.is_stub(stub) } - /// Try to acquire a permit + /// Try to atomically acquire specified number of permits. /// /// # Return /// - /// Returns `true` if the permit was acquired, `false` otherwise. If `false` - /// is returned, it can be assumed that `State` represents the head pointer - /// in the mpsc channel. - fn acquire_permit(&mut self, stub: &WaiterNode) -> bool { - if !self.has_available_permits() { + /// Returns `true` if the specified number of permits were acquired, `false` + /// otherwise. Returning false does not mean that there are no more + /// available permits. + fn acquire_permits(&mut self, num: usize, stub: &Waiter) -> bool { + debug_assert!(num > 0); + + if self.available_permits() < num { return false; } debug_assert!(self.waiter().is_none()); - self.0 -= 1 << NUM_SHIFT; + self.0 -= num << NUM_SHIFT; if self.0 == NUM_FLAG { // Set the state to the stub pointer. @@ -908,7 +1099,7 @@ impl SemState { /// Release permits /// /// Returns `true` if the permits were accepted. - fn release_permits(&mut self, permits: usize, stub: &WaiterNode) { + fn release_permits(&mut self, permits: usize, stub: &Waiter) { debug_assert!(permits > 0); if self.is_stub(stub) { @@ -926,7 +1117,7 @@ impl SemState { } /// Returns the waiter, if one is set. - fn waiter(self) -> Option> { + fn waiter(self) -> Option> { if self.is_waiter() { let waiter = NonNull::new(self.as_ptr()).expect("null pointer stored"); @@ -937,22 +1128,21 @@ impl SemState { } /// Assumes `self` represents a pointer - fn as_ptr(self) -> *mut WaiterNode { - (self.0 & !CLOSED_FLAG) as *mut WaiterNode + fn as_ptr(self) -> *mut Waiter { + (self.0 & !CLOSED_FLAG) as *mut Waiter } /// Set to a pointer to a waiter. /// /// This can only be done from the full state. - fn set_waiter(&mut self, waiter: NonNull) { + fn set_waiter(&mut self, waiter: NonNull) { let waiter = waiter.as_ptr() as usize; - debug_assert!(waiter & NUM_FLAG == 0); debug_assert!(!self.is_closed()); self.0 = waiter; } - fn is_stub(self, stub: &WaiterNode) -> bool { + fn is_stub(self, stub: &Waiter) -> bool { self.as_ptr() as usize == stub as *const _ as usize } @@ -962,28 +1152,6 @@ impl SemState { SemState(value) } - /// Swap the values - fn swap(self, cell: &AtomicUsize, ordering: Ordering) -> SemState { - let prev = SemState(cell.swap(self.to_usize(), ordering)); - debug_assert_eq!(prev.is_closed(), self.is_closed()); - prev - } - - /// Compare and exchange the current value into the provided cell - fn compare_exchange( - self, - cell: &AtomicUsize, - prev: SemState, - success: Ordering, - failure: Ordering, - ) -> Result { - debug_assert_eq!(prev.is_closed(), self.is_closed()); - - let res = cell.compare_exchange(prev.to_usize(), self.to_usize(), success, failure); - - res.map(SemState).map_err(SemState) - } - fn fetch_set_closed(cell: &AtomicUsize, ordering: Ordering) -> SemState { let value = cell.fetch_or(CLOSED_FLAG, ordering); SemState(value) @@ -1013,58 +1181,39 @@ impl fmt::Debug for SemState { } } -// ===== impl NodeState ===== +// ===== impl WaiterState ===== -impl NodeState { - fn new() -> NodeState { - NodeState::Idle +impl WaiterState { + fn permits_to_acquire(self) -> usize { + self.0 >> PERMIT_SHIFT } - fn from_usize(value: usize) -> NodeState { - use self::NodeState::*; - - match value { - 0 => Idle, - 1 => Queued, - 2 => QueuedWaiting, - 3 => Assigned, - 4 => Closed, - _ => panic!(), - } + fn set_permits_to_acquire(&mut self, val: usize) { + self.0 = (val << PERMIT_SHIFT) | (self.0 & !PERMIT_MASK) } - fn load(cell: &AtomicUsize, ordering: Ordering) -> NodeState { - NodeState::from_usize(cell.load(ordering)) - } - - /// Store a value - fn store(cell: &AtomicUsize, value: NodeState, ordering: Ordering) { - cell.store(value.to_usize(), ordering); - } - - fn compare_exchange( - self, - cell: &AtomicUsize, - prev: NodeState, - success: Ordering, - failure: Ordering, - ) -> Result { - cell.compare_exchange(prev.to_usize(), self.to_usize(), success, failure) - .map(NodeState::from_usize) - .map_err(NodeState::from_usize) - } - - /// Returns `true` if `self` represents a queued state. fn is_queued(self) -> bool { - use self::NodeState::*; - - match self { - Queued | QueuedWaiting => true, - _ => false, - } + self.0 & QUEUED == QUEUED } - fn to_usize(self) -> usize { - self as usize + fn set_queued(&mut self) { + self.0 |= QUEUED; + } + + fn is_closed(self) -> bool { + self.0 & CLOSED == CLOSED + } + + fn set_closed(&mut self) { + self.0 |= CLOSED; + } + + fn unset_queued(&mut self) { + assert!(self.is_queued()); + self.0 -= QUEUED; + } + + fn is_dropped(self) -> bool { + self.0 & DROPPED == DROPPED } } diff --git a/tokio/src/sync/tests/loom_semaphore_ll.rs b/tokio/src/sync/tests/loom_semaphore_ll.rs index cd4314ca0..b5e5efba8 100644 --- a/tokio/src/sync/tests/loom_semaphore_ll.rs +++ b/tokio/src/sync/tests/loom_semaphore_ll.rs @@ -31,7 +31,7 @@ fn basic_usage() { fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<()> { let me = &mut *self; - ready!(me.waiter.poll_acquire(cx, &me.shared.semaphore)).unwrap(); + ready!(me.waiter.poll_acquire(cx, 1, &me.shared.semaphore)).unwrap(); let actual = me.shared.active.fetch_add(1, SeqCst); assert!(actual <= NUM - 1); @@ -39,7 +39,7 @@ fn basic_usage() { let actual = me.shared.active.fetch_sub(1, SeqCst); assert!(actual <= NUM); - me.waiter.release(&me.shared.semaphore); + me.waiter.release(1, &me.shared.semaphore); Ready(()) } @@ -79,17 +79,17 @@ fn release() { thread::spawn(move || { let mut permit = Permit::new(); - block_on(poll_fn(|cx| permit.poll_acquire(cx, &semaphore))).unwrap(); + block_on(poll_fn(|cx| permit.poll_acquire(cx, 1, &semaphore))).unwrap(); - permit.release(&semaphore); + permit.release(1, &semaphore); }); } let mut permit = Permit::new(); - block_on(poll_fn(|cx| permit.poll_acquire(cx, &semaphore))).unwrap(); + block_on(poll_fn(|cx| permit.poll_acquire(cx, 1, &semaphore))).unwrap(); - permit.release(&semaphore); + permit.release(1, &semaphore); }); } @@ -108,10 +108,10 @@ fn basic_closing() { for _ in 0..2 { block_on(poll_fn(|cx| { - permit.poll_acquire(cx, &semaphore).map_err(|_| ()) + permit.poll_acquire(cx, 1, &semaphore).map_err(|_| ()) }))?; - permit.release(&semaphore); + permit.release(1, &semaphore); } Ok::<(), ()>(()) @@ -136,10 +136,10 @@ fn concurrent_close() { let mut permit = Permit::new(); block_on(poll_fn(|cx| { - permit.poll_acquire(cx, &semaphore).map_err(|_| ()) + permit.poll_acquire(cx, 1, &semaphore).map_err(|_| ()) }))?; - permit.release(&semaphore); + permit.release(1, &semaphore); semaphore.close(); @@ -148,3 +148,45 @@ fn concurrent_close() { } }); } + +#[test] +fn batch() { + let mut b = loom::model::Builder::new(); + b.preemption_bound = Some(1); + + b.check(|| { + let semaphore = Arc::new(Semaphore::new(10)); + let active = Arc::new(AtomicUsize::new(0)); + let mut ths = vec![]; + + for _ in 0..2 { + let semaphore = semaphore.clone(); + let active = active.clone(); + + ths.push(thread::spawn(move || { + let mut permit = Permit::new(); + + for n in &[4, 10, 8] { + block_on(poll_fn(|cx| permit.poll_acquire(cx, *n, &semaphore))).unwrap(); + + active.fetch_add(*n as usize, SeqCst); + + let num_active = active.load(SeqCst); + assert!(num_active <= 10); + + thread::yield_now(); + + active.fetch_sub(*n as usize, SeqCst); + + permit.release(*n, &semaphore); + } + })); + } + + for th in ths.into_iter() { + th.join().unwrap(); + } + + assert_eq!(10, semaphore.available_permits()); + }); +} diff --git a/tokio/src/sync/tests/semaphore_ll.rs b/tokio/src/sync/tests/semaphore_ll.rs index 8dd56c858..bfb075780 100644 --- a/tokio/src/sync/tests/semaphore_ll.rs +++ b/tokio/src/sync/tests/semaphore_ll.rs @@ -1,9 +1,8 @@ use crate::sync::semaphore_ll::{Permit, Semaphore}; -use tokio_test::task; -use tokio_test::{assert_pending, assert_ready_err, assert_ready_ok}; +use tokio_test::*; #[test] -fn available_permits() { +fn poll_acquire_one_available() { let s = Semaphore::new(100); assert_eq!(s.available_permits(), 100); @@ -11,44 +10,234 @@ fn available_permits() { let mut permit = task::spawn(Permit::new()); assert!(!permit.is_acquired()); - assert_ready_ok!(permit.enter(|cx, mut p| p.poll_acquire(cx, &s))); + assert_ready_ok!(permit.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); assert_eq!(s.available_permits(), 99); assert!(permit.is_acquired()); // Polling again on the same waiter does not claim a new permit - assert_ready_ok!(permit.enter(|cx, mut p| p.poll_acquire(cx, &s))); + assert_ready_ok!(permit.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); assert_eq!(s.available_permits(), 99); assert!(permit.is_acquired()); } #[test] -fn unavailable_permits() { +fn poll_acquire_many_available() { + let s = Semaphore::new(100); + assert_eq!(s.available_permits(), 100); + + // Polling for a permit succeeds immediately + let mut permit = task::spawn(Permit::new()); + assert!(!permit.is_acquired()); + + assert_ready_ok!(permit.enter(|cx, mut p| p.poll_acquire(cx, 5, &s))); + assert_eq!(s.available_permits(), 95); + assert!(permit.is_acquired()); + + // Polling again on the same waiter does not claim a new permit + assert_ready_ok!(permit.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); + assert_eq!(s.available_permits(), 95); + assert!(permit.is_acquired()); + + assert_ready_ok!(permit.enter(|cx, mut p| p.poll_acquire(cx, 5, &s))); + assert_eq!(s.available_permits(), 95); + assert!(permit.is_acquired()); + + // Polling for a larger number of permits acquires more + assert_ready_ok!(permit.enter(|cx, mut p| p.poll_acquire(cx, 8, &s))); + assert_eq!(s.available_permits(), 92); + assert!(permit.is_acquired()); +} + +#[test] +fn try_acquire_one_available() { + let s = Semaphore::new(100); + assert_eq!(s.available_permits(), 100); + + // Polling for a permit succeeds immediately + let mut permit = Permit::new(); + assert!(!permit.is_acquired()); + + assert_ok!(permit.try_acquire(1, &s)); + assert_eq!(s.available_permits(), 99); + assert!(permit.is_acquired()); + + // Polling again on the same waiter does not claim a new permit + assert_ok!(permit.try_acquire(1, &s)); + assert_eq!(s.available_permits(), 99); + assert!(permit.is_acquired()); +} + +#[test] +fn try_acquire_many_available() { + let s = Semaphore::new(100); + assert_eq!(s.available_permits(), 100); + + // Polling for a permit succeeds immediately + let mut permit = Permit::new(); + assert!(!permit.is_acquired()); + + assert_ok!(permit.try_acquire(5, &s)); + assert_eq!(s.available_permits(), 95); + assert!(permit.is_acquired()); + + // Polling again on the same waiter does not claim a new permit + assert_ok!(permit.try_acquire(5, &s)); + assert_eq!(s.available_permits(), 95); + assert!(permit.is_acquired()); +} + +#[test] +fn poll_acquire_one_unavailable() { let s = Semaphore::new(1); let mut permit_1 = task::spawn(Permit::new()); let mut permit_2 = task::spawn(Permit::new()); // Acquire the first permit - assert_ready_ok!(permit_1.enter(|cx, mut p| p.poll_acquire(cx, &s))); + assert_ready_ok!(permit_1.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); assert_eq!(s.available_permits(), 0); permit_2.enter(|cx, mut p| { // Try to acquire the second permit - assert_pending!(p.poll_acquire(cx, &s)); + assert_pending!(p.poll_acquire(cx, 1, &s)); }); - permit_1.release(&s); + permit_1.release(1, &s); assert_eq!(s.available_permits(), 0); assert!(permit_2.is_woken()); - assert_ready_ok!(permit_2.enter(|cx, mut p| p.poll_acquire(cx, &s))); + assert_ready_ok!(permit_2.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); - permit_2.release(&s); + permit_2.release(1, &s); assert_eq!(s.available_permits(), 1); } #[test] -fn zero_permits() { +fn forget_acquired() { + let s = Semaphore::new(1); + + // Polling for a permit succeeds immediately + let mut permit = task::spawn(Permit::new()); + + assert_ready_ok!(permit.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); + + assert_eq!(s.available_permits(), 0); + + permit.forget(1); + assert_eq!(s.available_permits(), 0); +} + +#[test] +fn forget_waiting() { + let s = Semaphore::new(0); + + // Polling for a permit succeeds immediately + let mut permit = task::spawn(Permit::new()); + + assert_pending!(permit.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); + + assert_eq!(s.available_permits(), 0); + + permit.forget(1); + + s.add_permits(1); + + assert!(!permit.is_woken()); + assert_eq!(s.available_permits(), 1); +} + +#[test] +fn poll_acquire_many_unavailable() { + let s = Semaphore::new(5); + + let mut permit_1 = task::spawn(Permit::new()); + let mut permit_2 = task::spawn(Permit::new()); + let mut permit_3 = task::spawn(Permit::new()); + + // Acquire the first permit + assert_ready_ok!(permit_1.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); + assert_eq!(s.available_permits(), 4); + + permit_2.enter(|cx, mut p| { + // Try to acquire the second permit + assert_pending!(p.poll_acquire(cx, 5, &s)); + }); + + assert_eq!(s.available_permits(), 0); + + permit_3.enter(|cx, mut p| { + // Try to acquire the third permit + assert_pending!(p.poll_acquire(cx, 3, &s)); + }); + + permit_1.release(1, &s); + + assert_eq!(s.available_permits(), 0); + assert!(permit_2.is_woken()); + assert_ready_ok!(permit_2.enter(|cx, mut p| p.poll_acquire(cx, 5, &s))); + + assert!(!permit_3.is_woken()); + assert_eq!(s.available_permits(), 0); + + permit_2.release(1, &s); + assert!(!permit_3.is_woken()); + assert_eq!(s.available_permits(), 0); + + permit_2.release(2, &s); + assert!(permit_3.is_woken()); + + assert_ready_ok!(permit_3.enter(|cx, mut p| p.poll_acquire(cx, 3, &s))); +} + +#[test] +fn try_acquire_one_unavailable() { + let s = Semaphore::new(1); + + let mut permit_1 = Permit::new(); + let mut permit_2 = Permit::new(); + + // Acquire the first permit + assert_ok!(permit_1.try_acquire(1, &s)); + assert_eq!(s.available_permits(), 0); + + assert_err!(permit_2.try_acquire(1, &s)); + + permit_1.release(1, &s); + + assert_eq!(s.available_permits(), 1); + assert_ok!(permit_2.try_acquire(1, &s)); + + permit_2.release(1, &s); + assert_eq!(s.available_permits(), 1); +} + +#[test] +fn try_acquire_many_unavailable() { + let s = Semaphore::new(5); + + let mut permit_1 = Permit::new(); + let mut permit_2 = Permit::new(); + + // Acquire the first permit + assert_ok!(permit_1.try_acquire(1, &s)); + assert_eq!(s.available_permits(), 4); + + assert_err!(permit_2.try_acquire(5, &s)); + + permit_1.release(1, &s); + assert_eq!(s.available_permits(), 5); + + assert_ok!(permit_2.try_acquire(5, &s)); + + permit_2.release(1, &s); + assert_eq!(s.available_permits(), 1); + + permit_2.release(1, &s); + assert_eq!(s.available_permits(), 2); +} + +#[test] +fn poll_acquire_one_zero_permits() { let s = Semaphore::new(0); assert_eq!(s.available_permits(), 0); @@ -56,13 +245,13 @@ fn zero_permits() { // Try to acquire the permit permit.enter(|cx, mut p| { - assert_pending!(p.poll_acquire(cx, &s)); + assert_pending!(p.poll_acquire(cx, 1, &s)); }); s.add_permits(1); assert!(permit.is_woken()); - assert_ready_ok!(permit.enter(|cx, mut p| p.poll_acquire(cx, &s))); + assert_ready_ok!(permit.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); } #[test] @@ -74,15 +263,19 @@ fn validates_max_permits() { #[test] fn close_semaphore_prevents_acquire() { - let s = Semaphore::new(1); + let s = Semaphore::new(5); s.close(); - assert_eq!(1, s.available_permits()); + assert_eq!(5, s.available_permits()); - let mut permit = task::spawn(Permit::new()); + let mut permit_1 = task::spawn(Permit::new()); + let mut permit_2 = task::spawn(Permit::new()); - assert_ready_err!(permit.enter(|cx, mut p| p.poll_acquire(cx, &s))); - assert_eq!(1, s.available_permits()); + assert_ready_err!(permit_1.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); + assert_eq!(5, s.available_permits()); + + assert_ready_err!(permit_2.enter(|cx, mut p| p.poll_acquire(cx, 2, &s))); + assert_eq!(5, s.available_permits()); } #[test] @@ -90,12 +283,12 @@ fn close_semaphore_notifies_permit1() { let s = Semaphore::new(0); let mut permit = task::spawn(Permit::new()); - assert_pending!(permit.enter(|cx, mut p| p.poll_acquire(cx, &s))); + assert_pending!(permit.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); s.close(); assert!(permit.is_woken()); - assert_ready_err!(permit.enter(|cx, mut p| p.poll_acquire(cx, &s))); + assert_ready_err!(permit.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); } #[test] @@ -108,29 +301,170 @@ fn close_semaphore_notifies_permit2() { let mut permit4 = task::spawn(Permit::new()); // Acquire a couple of permits - assert_ready_ok!(permit1.enter(|cx, mut p| p.poll_acquire(cx, &s))); - assert_ready_ok!(permit2.enter(|cx, mut p| p.poll_acquire(cx, &s))); + assert_ready_ok!(permit1.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); + assert_ready_ok!(permit2.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); - assert_pending!(permit3.enter(|cx, mut p| p.poll_acquire(cx, &s))); - assert_pending!(permit4.enter(|cx, mut p| p.poll_acquire(cx, &s))); + assert_pending!(permit3.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); + assert_pending!(permit4.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); s.close(); assert!(permit3.is_woken()); assert!(permit4.is_woken()); - assert_ready_err!(permit3.enter(|cx, mut p| p.poll_acquire(cx, &s))); - assert_ready_err!(permit4.enter(|cx, mut p| p.poll_acquire(cx, &s))); + assert_ready_err!(permit3.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); + assert_ready_err!(permit4.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); assert_eq!(0, s.available_permits()); - permit1.release(&s); + permit1.release(1, &s); assert_eq!(1, s.available_permits()); - assert_ready_err!(permit1.enter(|cx, mut p| p.poll_acquire(cx, &s))); + assert_ready_err!(permit1.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); - permit2.release(&s); + permit2.release(1, &s); assert_eq!(2, s.available_permits()); } + +#[test] +fn poll_acquire_additional_permits_while_waiting_before_assigned() { + let s = Semaphore::new(1); + + let mut permit = task::spawn(Permit::new()); + + assert_pending!(permit.enter(|cx, mut p| p.poll_acquire(cx, 2, &s))); + assert_pending!(permit.enter(|cx, mut p| p.poll_acquire(cx, 3, &s))); + + s.add_permits(1); + assert!(!permit.is_woken()); + + s.add_permits(1); + assert!(permit.is_woken()); + + assert_ready_ok!(permit.enter(|cx, mut p| p.poll_acquire(cx, 3, &s))); +} + +#[test] +fn try_acquire_additional_permits_while_waiting_before_assigned() { + let s = Semaphore::new(1); + + let mut permit = task::spawn(Permit::new()); + + assert_pending!(permit.enter(|cx, mut p| p.poll_acquire(cx, 2, &s))); + + assert_err!(permit.enter(|_, mut p| p.try_acquire(3, &s))); + + s.add_permits(1); + assert!(permit.is_woken()); + + assert_ok!(permit.enter(|_, mut p| p.try_acquire(2, &s))); +} + +#[test] +fn poll_acquire_additional_permits_while_waiting_after_assigned_success() { + let s = Semaphore::new(1); + + let mut permit = task::spawn(Permit::new()); + + assert_pending!(permit.enter(|cx, mut p| p.poll_acquire(cx, 2, &s))); + + s.add_permits(2); + + assert!(permit.is_woken()); + assert_ready_ok!(permit.enter(|cx, mut p| p.poll_acquire(cx, 3, &s))); +} + +#[test] +fn poll_acquire_additional_permits_while_waiting_after_assigned_requeue() { + let s = Semaphore::new(1); + + let mut permit = task::spawn(Permit::new()); + + assert_pending!(permit.enter(|cx, mut p| p.poll_acquire(cx, 2, &s))); + + s.add_permits(2); + + assert!(permit.is_woken()); + assert_pending!(permit.enter(|cx, mut p| p.poll_acquire(cx, 4, &s))); + + s.add_permits(1); + + assert!(permit.is_woken()); + assert_ready_ok!(permit.enter(|cx, mut p| p.poll_acquire(cx, 4, &s))); +} + +#[test] +fn poll_acquire_fewer_permits_while_waiting() { + let s = Semaphore::new(1); + + let mut permit = task::spawn(Permit::new()); + + assert_pending!(permit.enter(|cx, mut p| p.poll_acquire(cx, 2, &s))); + assert_eq!(s.available_permits(), 0); + + assert_ready_ok!(permit.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); + assert_eq!(s.available_permits(), 0); +} + +#[test] +fn poll_acquire_fewer_permits_after_assigned() { + let s = Semaphore::new(1); + + let mut permit1 = task::spawn(Permit::new()); + let mut permit2 = task::spawn(Permit::new()); + + assert_pending!(permit1.enter(|cx, mut p| p.poll_acquire(cx, 5, &s))); + assert_eq!(s.available_permits(), 0); + + assert_pending!(permit2.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); + + s.add_permits(4); + assert!(permit1.is_woken()); + assert!(!permit2.is_woken()); + + assert_ready_ok!(permit1.enter(|cx, mut p| p.poll_acquire(cx, 3, &s))); + + assert!(permit2.is_woken()); + assert_eq!(s.available_permits(), 1); + + assert_ready_ok!(permit2.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); +} + +#[test] +fn forget_partial_1() { + let s = Semaphore::new(0); + + let mut permit = task::spawn(Permit::new()); + + assert_pending!(permit.enter(|cx, mut p| p.poll_acquire(cx, 2, &s))); + s.add_permits(1); + + assert_eq!(0, s.available_permits()); + + permit.release(1, &s); + + assert_ready_ok!(permit.enter(|cx, mut p| p.poll_acquire(cx, 1, &s))); + + assert_eq!(s.available_permits(), 0); +} + +#[test] +fn forget_partial_2() { + let s = Semaphore::new(0); + + let mut permit = task::spawn(Permit::new()); + + assert_pending!(permit.enter(|cx, mut p| p.poll_acquire(cx, 2, &s))); + s.add_permits(1); + + assert_eq!(0, s.available_permits()); + + permit.release(1, &s); + + s.add_permits(1); + + assert_ready_ok!(permit.enter(|cx, mut p| p.poll_acquire(cx, 2, &s))); + assert_eq!(s.available_permits(), 0); +}