use std::future::Future;
use std::pin::Pin;
use std::sync::Mutex;
use std::task::{Context, Poll};
use crate::sync::wait_queue::{WaitQueue, WaiterPoll};
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
enum NotifyGrant {
One,
All,
}
pub struct Notify {
state: Mutex<NotifyState>,
}
struct NotifyState {
notified: bool,
waiters: WaitQueue<NotifyGrant>,
}
impl Notify {
pub fn new() -> Self {
Self {
state: Mutex::new(NotifyState {
notified: false,
waiters: WaitQueue::new(),
}),
}
}
pub fn notified(&self) -> NotifyFuture<'_> {
NotifyFuture {
notify: self,
id: None,
}
}
pub fn notify_one(&self) {
let mut state = self.state.lock().unwrap();
match state.waiters.grant_oldest(NotifyGrant::One) {
Some(waker) => waker.wake(),
None => state.notified = true,
}
}
pub fn notify_waiters(&self) {
let mut state = self.state.lock().unwrap();
let wakers = state.waiters.grant_all(NotifyGrant::All);
drop(state);
for waker in wakers {
waker.wake();
}
}
}
impl Default for Notify {
fn default() -> Self {
Self::new()
}
}
pub struct NotifyFuture<'a> {
notify: &'a Notify,
id: Option<u64>,
}
impl<'a> Future for NotifyFuture<'a> {
type Output = ();
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let mut state = self.notify.state.lock().unwrap();
if state.notified {
state.notified = false;
if let Some(id) = self.id.take() {
let _removed_grant = state.waiters.deregister(id);
}
return Poll::Ready(());
}
if let Some(id) = self.id {
match state.waiters.poll_waiter(id, cx.waker()) {
WaiterPoll::Granted(_) => {
self.id = None;
return Poll::Ready(());
}
WaiterPoll::Pending => return Poll::Pending,
WaiterPoll::NotRegistered => {}
}
}
self.id = Some(state.waiters.register(cx.waker().clone()));
Poll::Pending
}
}
impl<'a> Drop for NotifyFuture<'a> {
fn drop(&mut self) {
if let Some(id) = self.id {
if let Ok(mut state) = self.notify.state.lock() {
if state.waiters.deregister(id) == Some(NotifyGrant::One) {
match state.waiters.grant_oldest(NotifyGrant::One) {
Some(waker) => waker.wake(),
None => state.notified = true,
}
}
}
}
}
}