use std::future::Future;
use std::pin::Pin;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::task::{Context, Poll, Waker};
use std::collections::VecDeque;
pub struct MkAsyncBarrier {
inner: Arc<BarrierInner>,
}
struct BarrierInner {
count: AtomicUsize,
target: usize,
}
impl MkAsyncBarrier {
pub fn new(n: usize) -> Self {
Self {
inner: Arc::new(BarrierInner {
count: AtomicUsize::new(0),
target: n,
}),
}
}
pub async fn wait(&self) {
let prev = self.inner.count.fetch_add(1, Ordering::SeqCst);
if prev + 1 >= self.inner.target {
self.inner.count.store(0, Ordering::SeqCst);
return;
}
WaitBarrier::new(&self.inner.count, self.inner.target).await
}
pub fn waiting(&self) -> usize {
self.inner.count.load(Ordering::Relaxed)
}
}
impl Clone for MkAsyncBarrier {
fn clone(&self) -> Self {
Self {
inner: Arc::clone(&self.inner),
}
}
}
struct WaitBarrier<'a> {
count: &'a AtomicUsize,
target: usize,
}
impl<'a> WaitBarrier<'a> {
fn new(count: &'a AtomicUsize, target: usize) -> Self {
Self { count, target }
}
}
impl<'a> Future for WaitBarrier<'a> {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let current = self.count.load(Ordering::Acquire);
if current >= self.target || current == 0 {
Poll::Ready(())
} else {
cx.waker().wake_by_ref();
Poll::Pending
}
}
}
pub struct MkAsyncSemaphore {
permits: AtomicUsize,
max_permits: usize,
}
impl MkAsyncSemaphore {
pub fn new(permits: usize) -> Self {
Self {
permits: AtomicUsize::new(permits),
max_permits: permits,
}
}
pub async fn acquire(&self) -> SemaphorePermit<'_> {
loop {
let current = self.permits.load(Ordering::Acquire);
if current > 0 {
match self.permits.compare_exchange_weak(
current,
current - 1,
Ordering::AcqRel,
Ordering::Relaxed,
) {
Ok(_) => return SemaphorePermit { semaphore: self },
Err(_) => continue,
}
}
YieldOnce::new().await;
}
}
pub fn try_acquire(&self) -> Option<SemaphorePermit<'_>> {
loop {
let current = self.permits.load(Ordering::Acquire);
if current == 0 {
return None;
}
match self.permits.compare_exchange_weak(
current,
current - 1,
Ordering::AcqRel,
Ordering::Relaxed,
) {
Ok(_) => return Some(SemaphorePermit { semaphore: self }),
Err(_) => continue,
}
}
}
pub fn available(&self) -> usize {
self.permits.load(Ordering::Relaxed)
}
}
pub struct SemaphorePermit<'a> {
semaphore: &'a MkAsyncSemaphore,
}
impl<'a> Drop for SemaphorePermit<'a> {
fn drop(&mut self) {
self.semaphore.permits.fetch_add(1, Ordering::Release);
}
}
struct YieldOnce(bool);
impl YieldOnce {
fn new() -> Self {
Self(false)
}
}
impl Future for YieldOnce {
type Output = ();
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
if self.0 {
Poll::Ready(())
} else {
self.0 = true;
cx.waker().wake_by_ref();
Poll::Pending
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_semaphore_sync() {
let sem = MkAsyncSemaphore::new(3);
assert_eq!(sem.available(), 3);
let _p1 = sem.try_acquire().unwrap();
assert_eq!(sem.available(), 2);
let _p2 = sem.try_acquire().unwrap();
let _p3 = sem.try_acquire().unwrap();
assert_eq!(sem.available(), 0);
assert!(sem.try_acquire().is_none());
}
}