#![cfg_attr(docsrs, feature(doc_cfg))]
#![warn(missing_docs)]
use std::{
cell::RefCell,
future::Future,
hint::unreachable_unchecked,
pin::Pin,
sync::Mutex,
task::{Context, Poll, Waker},
};
pub trait Runnable {
fn run();
}
impl Runnable for () {
fn run() {}
}
#[derive(Debug)]
enum WakerState<T> {
Inactive,
Active(Waker),
Signaled(T),
}
#[derive(Debug)]
pub struct Callback<T = ()>(RefCell<WakerState<T>>);
impl<T> Default for Callback<T> {
fn default() -> Self {
Self::new()
}
}
impl<T> Callback<T> {
pub fn new() -> Self {
Self(RefCell::new(WakerState::Inactive))
}
pub fn signal<R: Runnable>(&self, v: T) -> bool {
let mut state = self.0.borrow_mut();
match &*state {
WakerState::Inactive => return true,
WakerState::Signaled(_) => {
}
WakerState::Active(waker) => {
waker.wake_by_ref();
}
}
*state = WakerState::Signaled(v);
drop(state);
R::run();
false
}
pub(crate) fn register(&self, waker: &Waker) -> Poll<T> {
let mut state = self.0.borrow_mut();
match &*state {
WakerState::Signaled(_) => {
let state = std::mem::replace(&mut *state, WakerState::Inactive);
let v = if let WakerState::Signaled(v) = state {
v
} else {
unsafe { unreachable_unchecked() }
};
Poll::Ready(v)
}
_ => {
*state = WakerState::Active(waker.clone());
Poll::Pending
}
}
}
pub fn wait(&self) -> impl Future<Output = T> + '_ {
WaitFut(self)
}
}
struct WaitFut<'a, T>(&'a Callback<T>);
impl<T> Future for WaitFut<'_, T> {
type Output = T;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
self.0.register(cx.waker())
}
}
impl<T> Drop for WaitFut<'_, T> {
fn drop(&mut self) {
*self.0.0.borrow_mut() = WakerState::Inactive;
}
}
#[derive(Debug)]
pub struct SyncCallback<T = ()>(Mutex<WakerState<T>>);
impl<T> Default for SyncCallback<T> {
fn default() -> Self {
Self::new()
}
}
impl<T> SyncCallback<T> {
pub fn new() -> Self {
Self(Mutex::new(WakerState::Inactive))
}
pub fn signal(&self, v: T) -> bool {
let mut state = self.0.lock().unwrap();
match &*state {
WakerState::Inactive => return true,
WakerState::Signaled(_) => {
}
WakerState::Active(waker) => {
waker.wake_by_ref();
}
}
*state = WakerState::Signaled(v);
false
}
pub(crate) fn register(&self, waker: &Waker) -> Poll<T> {
let mut state = self.0.lock().unwrap();
match &*state {
WakerState::Signaled(_) => {
let state = std::mem::replace(&mut *state, WakerState::Inactive);
let v = if let WakerState::Signaled(v) = state {
v
} else {
unsafe { unreachable_unchecked() }
};
Poll::Ready(v)
}
_ => {
*state = WakerState::Active(waker.clone());
Poll::Pending
}
}
}
pub fn wait(&self) -> impl Future<Output = T> + '_ {
SyncWaitFut(self)
}
}
struct SyncWaitFut<'a, T>(&'a SyncCallback<T>);
impl<T> Future for SyncWaitFut<'_, T> {
type Output = T;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
self.0.register(cx.waker())
}
}
impl<T> Drop for SyncWaitFut<'_, T> {
fn drop(&mut self) {
*self.0.0.lock().unwrap() = WakerState::Inactive;
}
}