memkit-async 0.2.0-beta.1

Async-aware memory allocators for memkit
//! Async synchronization primitives.

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;

/// Async barrier for synchronizing multiple tasks.
pub struct MkAsyncBarrier {
    inner: Arc<BarrierInner>,
}

struct BarrierInner {
    count: AtomicUsize,
    target: usize,
}

impl MkAsyncBarrier {
    /// Create a new barrier for `n` tasks.
    pub fn new(n: usize) -> Self {
        Self {
            inner: Arc::new(BarrierInner {
                count: AtomicUsize::new(0),
                target: n,
            }),
        }
    }

    /// Wait at the barrier.
    pub async fn wait(&self) {
        // Increment count
        let prev = self.inner.count.fetch_add(1, Ordering::SeqCst);
        
        if prev + 1 >= self.inner.target {
            // We're the last one - reset for reuse
            self.inner.count.store(0, Ordering::SeqCst);
            return;
        }

        // Wait for others
        WaitBarrier::new(&self.inner.count, self.inner.target).await
    }

    /// Get the number of tasks currently waiting.
    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
        }
    }
}

/// Async semaphore for controlling concurrent access.
pub struct MkAsyncSemaphore {
    permits: AtomicUsize,
    max_permits: usize,
}

impl MkAsyncSemaphore {
    /// Create a new semaphore with the given number of permits.
    pub fn new(permits: usize) -> Self {
        Self {
            permits: AtomicUsize::new(permits),
            max_permits: permits,
        }
    }

    /// Acquire a permit.
    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,
                }
            }
            
            // Yield and retry
            YieldOnce::new().await;
        }
    }

    /// Try to acquire a permit without waiting.
    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,
            }
        }
    }

    /// Get the number of available permits.
    pub fn available(&self) -> usize {
        self.permits.load(Ordering::Relaxed)
    }
}

/// A permit from a semaphore.
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);
    }
}

/// Yield once to the runtime.
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());
    }
}