use crate::get_task_from_context;
use crate::runtime::local_executor;
use crate::runtime::task::Task;
use crate::sync::mutexes::AsyncSubscribableMutex;
use crate::sync::{AsyncCondVar, AsyncMutex, AsyncMutexGuard, LocalMutex};
use std::cell::UnsafeCell;
use std::future::Future;
use std::marker::PhantomData;
use std::pin::Pin;
use std::task::{Context, Poll};
enum WaitState {
Sleep,
Wake,
Lock,
}
pub struct WaitLocalCondVar<'mutex, 'cond_var, T, Guard>
where
T: 'mutex + ?Sized,
Guard: AsyncMutexGuard<'mutex, T>,
Guard::Mutex: AsyncSubscribableMutex<T>,
{
state: WaitState,
cond_var: &'cond_var LocalCondVar,
mutex: &'mutex Guard::Mutex,
no_send_marker: PhantomData<*const ()>,
pd: PhantomData<T>,
}
impl<'mutex, 'cond_var, T, Guard> WaitLocalCondVar<'mutex, 'cond_var, T, Guard>
where
T: 'mutex + ?Sized,
Guard: AsyncMutexGuard<'mutex, T>,
Guard::Mutex: AsyncSubscribableMutex<T>,
{
#[inline(always)]
pub fn new(cond_var: &'cond_var LocalCondVar, mutex: &'mutex Guard::Mutex) -> Self {
WaitLocalCondVar {
state: WaitState::Sleep,
cond_var,
mutex,
no_send_marker: PhantomData,
pd: PhantomData,
}
}
}
impl<'mutex, T, Guard> Future for WaitLocalCondVar<'mutex, '_, T, Guard>
where
T: 'mutex + ?Sized,
Guard: AsyncMutexGuard<'mutex, T>,
Guard::Mutex: AsyncSubscribableMutex<T>,
{
type Output = <<Guard as AsyncMutexGuard<'mutex, T>>::Mutex as AsyncMutex<T>>::Guard<'mutex>;
fn poll(self: Pin<&mut Self>, cx: &mut Context) -> Poll<Self::Output> {
let this = unsafe { self.get_unchecked_mut() };
match this.state {
WaitState::Sleep => {
this.state = WaitState::Wake;
let task = unsafe { get_task_from_context!(cx) };
let wait_queue = unsafe { &mut *this.cond_var.wait_queue.get() };
wait_queue.push(task);
Poll::Pending
}
WaitState::Wake => {
if let Some(guard) = this.mutex.try_lock() {
Poll::Ready(guard)
} else {
this.state = WaitState::Lock;
this.mutex.low_level_subscribe(cx);
Poll::Pending
}
}
WaitState::Lock => Poll::Ready(unsafe { this.mutex.get_locked() }),
}
}
}
pub struct LocalCondVar {
wait_queue: UnsafeCell<Vec<Task>>,
no_send_marker: std::marker::PhantomData<*const ()>,
}
impl LocalCondVar {
#[inline(always)]
pub const fn new() -> Self {
Self {
wait_queue: UnsafeCell::new(Vec::new()),
no_send_marker: PhantomData,
}
}
}
impl AsyncCondVar for LocalCondVar {
type SubscribableMutex<T>
= LocalMutex<T>
where
T: ?Sized;
#[inline(always)]
#[allow(clippy::future_not_send, reason = "LocalCondVar is !Send")]
fn wait<'mutex, T>(
&self,
guard: <Self::SubscribableMutex<T> as AsyncMutex<T>>::Guard<'mutex>,
) -> impl Future<Output = <Self::SubscribableMutex<T> as AsyncMutex<T>>::Guard<'mutex>>
where
T: ?Sized + 'mutex,
{
WaitLocalCondVar::<
'mutex,
'_,
T,
<Self::SubscribableMutex<T> as AsyncMutex<T>>::Guard<'mutex>,
>::new(self, guard.mutex())
}
#[inline(always)]
fn notify_one(&self) {
let wait_queue = unsafe { &mut *self.wait_queue.get() };
if let Some(task) = wait_queue.pop() {
local_executor().exec_task(task);
}
}
#[inline(always)]
fn notify_all(&self) {
let executor = local_executor();
let wait_queue = unsafe { &mut *self.wait_queue.get() };
while let Some(task) = wait_queue.pop() {
executor.exec_task(task);
}
}
}
impl Default for LocalCondVar {
fn default() -> Self {
Self::new()
}
}
unsafe impl Sync for LocalCondVar {}
#[allow(dead_code, reason = "It is used only in compile tests")]
fn test_compile_local_cond_var() {}
#[cfg(test)]
mod tests {
use super::*;
use crate as orengine;
use crate::runtime::local_executor;
use crate::sleep::sleep;
use crate::sync::{AsyncMutex, AsyncWaitGroup, LocalMutex, LocalWaitGroup};
use std::rc::Rc;
use std::time::{Duration, Instant};
const TIME_TO_SLEEP: Duration = Duration::from_millis(1);
#[allow(clippy::future_not_send, reason = "It is local.")]
async fn test_notify_one(need_drop: bool) {
let start = Instant::now();
let pair = Rc::new((LocalMutex::new(false), LocalCondVar::new()));
let pair2 = pair.clone();
local_executor().spawn_local(async move {
let (lock, cvar) = &*pair2;
let mut started = lock.lock().await;
sleep(TIME_TO_SLEEP).await;
*started = true;
if need_drop {
drop(started);
}
cvar.notify_one();
});
let (lock, cvar) = &*pair;
let mut started = lock.lock().await;
while !*started {
started = cvar.wait(started).await;
}
assert!(start.elapsed() >= TIME_TO_SLEEP);
}
#[allow(clippy::future_not_send, reason = "It is local.")]
async fn test_notify_all(need_drop: bool) {
const NUMBER_OF_WAITERS: usize = 10;
let start = Instant::now();
let pair = Rc::new((LocalMutex::new(false), LocalCondVar::new()));
let pair2 = pair.clone();
local_executor().spawn_local(async move {
let (lock, cvar) = &*pair2;
let mut started = lock.lock().await;
sleep(TIME_TO_SLEEP).await;
*started = true;
if need_drop {
drop(started);
}
cvar.notify_all();
});
let wg = Rc::new(LocalWaitGroup::new());
for _ in 0..NUMBER_OF_WAITERS {
let pair = pair.clone();
let wg = wg.clone();
wg.add(1);
local_executor().spawn_local(async move {
let (lock, cvar) = &*pair;
let mut started = lock.lock().await;
while !*started {
started = cvar.wait(started).await;
}
wg.done();
});
}
wg.wait().await;
assert!(start.elapsed() >= TIME_TO_SLEEP);
}
#[orengine::test::test_local]
fn test_local_cond_var_notify_one_with_drop_guard() {
test_notify_one(true).await;
}
#[orengine::test::test_local]
fn test_local_cond_var_notify_all_with_drop_guard() {
test_notify_all(true).await;
}
#[orengine::test::test_local]
fn test_local_cond_var_notify_one_without_drop_guard() {
test_notify_one(false).await;
}
#[orengine::test::test_local]
fn test_local_cond_var_notify_all_without_drop_guard() {
test_notify_all(false).await;
}
}