use spin::Mutex;
use crate::Semaphore;
pub struct Condvar {
semaphore: Semaphore,
num_waiters: Mutex<usize>,
}
impl Condvar {
pub const fn new() -> Self {
Self {
semaphore: Semaphore::new(0),
num_waiters: Mutex::new(0),
}
}
pub async fn wait(&self) {
let mut num_waiters = self.num_waiters.lock();
*num_waiters += 1;
drop(num_waiters);
self.semaphore.acquire(1).await;
}
pub fn notify_one(&self) {
let mut num_waiters = self.num_waiters.lock();
if *num_waiters > 0 {
*num_waiters -= 1;
self.semaphore.release(1);
}
}
pub fn notify_all(&self) {
let mut num_waiters = self.num_waiters.lock();
if *num_waiters > 0 {
let total_waiters = *num_waiters;
*num_waiters = 0;
self.semaphore.release(total_waiters);
}
}
}
#[cfg(test)]
mod tests {
use core::time::Duration;
use std::thread;
use super::*;
#[test]
fn notify_all() {
static CONDVAR: Condvar = Condvar::new();
let task1 = thread::spawn(|| pollster::block_on(CONDVAR.wait()));
let task2 = thread::spawn(|| pollster::block_on(CONDVAR.wait()));
let task3 = thread::spawn(|| pollster::block_on(CONDVAR.wait()));
thread::sleep(Duration::from_millis(100));
CONDVAR.notify_all();
thread::sleep(Duration::from_millis(100));
assert!(task1.is_finished());
assert!(task2.is_finished());
assert!(task3.is_finished());
}
}