use parking_lot::Mutex;
use std::collections::VecDeque;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::task::{Context, Poll, Waker};
#[allow(async_fn_in_trait)]
pub trait BytePermits: Send + Sync {
async fn acquire(&self, n_bytes: usize) -> Permit;
}
struct WaiterSlot {
needed: usize,
granted: AtomicBool,
waker: Mutex<Option<Waker>>,
}
struct SemInner {
available: usize,
waiters: VecDeque<Arc<WaiterSlot>>,
}
fn grant_front(inner: &mut SemInner) -> Vec<Waker> {
let mut wakers = Vec::new();
loop {
let Some(front) = inner.waiters.front() else {
break;
};
if front.needed > inner.available {
break;
}
let slot = inner.waiters.pop_front().expect("front just checked");
inner.available -= slot.needed;
slot.granted.store(true, Ordering::Release);
let waker = slot.waker.lock().take();
if let Some(w) = waker {
wakers.push(w);
}
}
wakers
}
pub struct Permit {
inner: Option<PermitInner>,
}
enum PermitInner {
ByteSem(Arc<Mutex<SemInner>>, usize),
NoOp,
}
impl Drop for Permit {
fn drop(&mut self) {
if let Some(PermitInner::ByteSem(sem, n_bytes)) = self.inner.take() {
let wakers = {
let mut inner = sem.lock();
inner.available += n_bytes;
let wakers = grant_front(&mut inner);
drop(inner);
wakers
};
for w in wakers {
w.wake();
}
}
}
}
impl Permit {
pub(crate) const fn noop() -> Self {
Self {
inner: Some(PermitInner::NoOp),
}
}
fn byte_sem(sem: Arc<Mutex<SemInner>>, n_bytes: usize) -> Self {
Self {
inner: Some(PermitInner::ByteSem(sem, n_bytes)),
}
}
}
#[derive(Debug, Clone, Copy, Default)]
pub struct NoOpPermits;
impl BytePermits for NoOpPermits {
async fn acquire(&self, _n_bytes: usize) -> Permit {
Permit::noop()
}
}
#[derive(Clone)]
pub struct SemaphorePermits {
inner: Arc<Mutex<SemInner>>,
max_bytes: usize,
}
impl SemaphorePermits {
#[must_use]
pub fn new(max_bytes: usize) -> Self {
Self {
inner: Arc::new(Mutex::new(SemInner {
available: max_bytes,
waiters: VecDeque::new(),
})),
max_bytes,
}
}
}
impl BytePermits for SemaphorePermits {
async fn acquire(&self, n_bytes: usize) -> Permit {
if n_bytes == 0 {
return Permit::noop();
}
let needed = n_bytes.min(self.max_bytes);
Acquire {
sem: self.inner.clone(),
needed,
slot: None,
#[cfg(test)]
counted_slow: false,
}
.await
}
}
struct Acquire {
sem: Arc<Mutex<SemInner>>,
needed: usize,
slot: Option<Arc<WaiterSlot>>,
#[cfg(test)]
counted_slow: bool,
}
impl Future for Acquire {
type Output = Permit;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Permit> {
let this = self.get_mut();
if let Some(slot) = &this.slot {
if slot.granted.load(Ordering::Acquire) {
let sem = this.sem.clone();
this.slot = None; return Poll::Ready(Permit::byte_sem(sem, this.needed));
}
*slot.waker.lock() = Some(cx.waker().clone());
if slot.granted.load(Ordering::Acquire) {
let sem = this.sem.clone();
this.slot = None;
return Poll::Ready(Permit::byte_sem(sem, this.needed));
}
return Poll::Pending;
}
let mut inner = this.sem.lock();
if inner.waiters.is_empty() && inner.available >= this.needed {
inner.available -= this.needed;
drop(inner);
return Poll::Ready(Permit::byte_sem(this.sem.clone(), this.needed));
}
#[cfg(test)]
if !this.counted_slow {
SLOW_PATH_ENTRIES.fetch_add(1, Ordering::Relaxed);
this.counted_slow = true;
}
let slot = Arc::new(WaiterSlot {
needed: this.needed,
granted: AtomicBool::new(false),
waker: Mutex::new(Some(cx.waker().clone())),
});
inner.waiters.push_back(slot.clone());
drop(inner);
this.slot = Some(slot);
Poll::Pending
}
}
impl Drop for Acquire {
fn drop(&mut self) {
let Some(slot) = self.slot.take() else {
return;
};
let wakers = {
let mut inner = self.sem.lock();
if slot.granted.load(Ordering::Acquire) {
inner.available += self.needed;
} else {
inner.waiters.retain(|s| !Arc::ptr_eq(s, &slot));
}
let wakers = grant_front(&mut inner);
drop(inner);
wakers
};
for w in wakers {
w.wake();
}
}
}
#[cfg(test)]
static SLOW_PATH_ENTRIES: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::Ordering;
static SLOW_PATH_TEST_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
fn lock_slow_path_counter() -> std::sync::MutexGuard<'static, ()> {
SLOW_PATH_TEST_LOCK
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
#[test]
fn uncontended_acquire_takes_fast_path() {
let _guard = lock_slow_path_counter();
let permits = SemaphorePermits::new(1024 * 1024);
let rt = crate::rt::LocalRuntime::new().unwrap();
SLOW_PATH_ENTRIES.store(0, Ordering::Relaxed);
rt.block_on(async {
for _ in 0..100 {
let permit = permits.acquire(1024).await;
drop(permit);
}
});
assert_eq!(
SLOW_PATH_ENTRIES.load(Ordering::Relaxed),
0,
"uncontended acquires must not park a waiter"
);
}
#[test]
fn contended_acquire_parks_then_completes_on_release() {
let _guard = lock_slow_path_counter();
let permits = SemaphorePermits::new(1024);
let rt = crate::rt::LocalRuntime::new().unwrap();
SLOW_PATH_ENTRIES.store(0, Ordering::Relaxed);
rt.block_on(async {
let p1 = permits.acquire(1024).await;
let permits2 = permits.clone();
let waiter = crate::rt::spawn(async move {
let _p2 = permits2.acquire(1024).await;
});
crate::rt::sleep(std::time::Duration::from_millis(50)).await;
assert!(
SLOW_PATH_ENTRIES.load(Ordering::Relaxed) >= 1,
"the second acquire should have parked while the pool was full"
);
drop(p1);
crate::rt::join(waiter).await;
});
}
#[test]
fn waiters_are_granted_in_fifo_order() {
let _guard = lock_slow_path_counter();
use std::cell::RefCell;
use std::rc::Rc;
let permits = SemaphorePermits::new(1024);
let rt = crate::rt::LocalRuntime::new().unwrap();
let order = Rc::new(RefCell::new(Vec::new()));
rt.block_on(async {
let p = permits.acquire(1024).await;
let mut handles = Vec::new();
for id in 0..3 {
let permits_i = permits.clone();
let order_i = order.clone();
handles.push(crate::rt::spawn(async move {
let _permit = permits_i.acquire(1024).await;
order_i.borrow_mut().push(id);
crate::rt::sleep(std::time::Duration::from_millis(10)).await;
}));
}
crate::rt::sleep(std::time::Duration::from_millis(50)).await;
drop(p); for h in handles {
crate::rt::join(h).await;
}
});
assert_eq!(
*order.borrow(),
vec![0, 1, 2],
"waiters must be granted in the order they queued"
);
}
#[test]
fn cancelled_waiter_does_not_leak_capacity() {
let _guard = lock_slow_path_counter();
let permits = SemaphorePermits::new(1024);
let rt = crate::rt::LocalRuntime::new().unwrap();
rt.block_on(async {
let p1 = permits.acquire(1024).await;
{
let mut fut = Box::pin(permits.acquire(512));
let polled = futures::poll!(fut.as_mut());
assert!(polled.is_pending(), "waiter should park while pool is full");
drop(fut); }
drop(p1);
let _p2 = permits.acquire(1024).await;
});
}
#[test]
fn noop_permits_always_succeed() {
let permits = NoOpPermits;
let rt = crate::rt::LocalRuntime::new().unwrap();
rt.block_on(async {
let _p1 = permits.acquire(1024).await;
let _p2 = permits.acquire(1_000_000).await;
});
}
#[test]
fn semaphore_permits_enforce_limit() {
let permits = SemaphorePermits::new(1024);
let rt = crate::rt::LocalRuntime::new().unwrap();
rt.block_on(async {
let p1 = permits.acquire(1024).await;
drop(p1);
let _p2 = permits.acquire(512).await;
let _p3 = permits.acquire(512).await;
});
}
#[test]
fn semaphore_permits_release_on_drop() {
let permits = SemaphorePermits::new(1000);
let rt = crate::rt::LocalRuntime::new().unwrap();
rt.block_on(async {
{
let _p1 = permits.acquire(500).await;
let _p2 = permits.acquire(500).await;
}
let _p3 = permits.acquire(1000).await;
});
}
#[test]
fn semaphore_permits_oversized_acquire_does_not_deadlock() {
let permits = SemaphorePermits::new(1024);
let rt = crate::rt::LocalRuntime::new().unwrap();
rt.block_on(async {
let permit = permits.acquire(2048).await; drop(permit);
let _p = permits.acquire(1024).await;
});
}
#[test]
fn semaphore_permits_single_atomic_acquire() {
let permits = SemaphorePermits::new(1024 * 1024); let rt = crate::rt::LocalRuntime::new().unwrap();
rt.block_on(async {
let permit = permits.acquire(512 * 1024).await; drop(permit);
});
}
}