use std::mem;
use std::sync::Mutex;
use crate::{RustFutureContinuationBoundCallback, RustFuturePoll};
#[derive(Debug)]
enum State<Callback> {
Empty,
Waked,
Cancelled,
Set(Callback),
}
pub struct Scheduler<Callback = RustFutureContinuationBoundCallback> {
state: Mutex<State<Callback>>,
}
pub trait RustFutureCallback {
fn invoke(self, poll: RustFuturePoll);
}
impl RustFutureCallback for RustFutureContinuationBoundCallback {
fn invoke(self, poll: RustFuturePoll) {
(self.callback)(self.data, poll)
}
}
impl<Callback: RustFutureCallback> Default for Scheduler<Callback> {
fn default() -> Self {
Self::new()
}
}
impl<Callback: RustFutureCallback> Scheduler<Callback> {
pub fn new() -> Self {
Self {
state: Mutex::new(State::Empty),
}
}
pub fn store(&self, callback: Callback) {
let to_invoke = {
let mut state = self.state.lock().unwrap();
match *state {
State::Empty => {
*state = State::Set(callback);
None
}
State::Set(_) => {
trace!(
"store: observed `Self::Set` state. Is poll() being called from multiple threads at once?"
);
let State::Set(old_callback) = mem::replace(&mut *state, State::Set(callback))
else {
unreachable!();
};
Some((old_callback, RustFuturePoll::Wake))
}
State::Waked => {
*state = State::Empty;
Some((callback, RustFuturePoll::Wake))
}
State::Cancelled => Some((callback, RustFuturePoll::Ready)),
}
};
if let Some((cb, poll)) = to_invoke {
cb.invoke(poll);
}
}
pub fn wake(&self) {
let callback = {
let mut state = self.state.lock().unwrap();
match *state {
State::Set(_) => {
let State::Set(callback) = mem::replace(&mut *state, State::Empty) else {
unreachable!();
};
Some(callback)
}
State::Empty => {
*state = State::Waked;
None
}
_ => None,
}
};
if let Some(cb) = callback {
cb.invoke(RustFuturePoll::Wake);
}
}
pub fn cancel(&self) {
let callback = {
let mut state = self.state.lock().unwrap();
match mem::replace(&mut *state, State::Cancelled) {
State::Set(cb) => Some(cb),
_ => None,
}
};
if let Some(cb) = callback {
cb.invoke(RustFuturePoll::Ready);
}
}
pub fn is_cancelled(&self) -> bool {
matches!(*self.state.lock().unwrap(), State::Cancelled)
}
}