use std::{
cell::Cell,
collections::VecDeque,
error::Error,
pin::Pin,
sync::{Arc, Mutex},
task::{Context, Poll, Wake, Waker},
};
use super::process::{self, ProcessID};
struct Process<'a> {
future: Pin<Box<dyn Future<Output = process::ProcessResult> + 'a>>,
waker: Waker,
}
struct Inbox {
woken: Mutex<VecDeque<ProcessID>>,
}
struct ProcessWaker {
id: ProcessID,
inbox: Arc<Inbox>,
}
impl Wake for ProcessWaker {
fn wake(self: Arc<Self>) {
self.wake_by_ref();
}
fn wake_by_ref(self: &Arc<Self>) {
self.inbox.woken.lock().unwrap().push_back(self.id);
}
}
pub(crate) struct Executor<'a> {
processes: Vec<Option<Process<'a>>>,
queue: VecDeque<ProcessID>,
inbox: Arc<Inbox>,
live: usize,
flush_buf: Vec<ProcessID>,
}
impl<'a> Executor<'a> {
pub(crate) fn schedule(
&mut self,
code: impl Future<Output = process::ProcessResult> + 'a,
) -> ProcessID {
let id = self.processes.len();
let waker = Waker::from(Arc::new(ProcessWaker {
id,
inbox: Arc::clone(&self.inbox),
}));
self.processes.push(Some(Process {
future: Box::pin(code),
waker,
}));
self.live += 1;
self.queue.push_back(id);
id
}
pub(crate) fn pending(&self) -> usize {
self.live
}
pub(crate) fn execute(&mut self) -> Result<(), RawProcessError> {
loop {
self.flush_wakes();
let Some(id) = self.queue.pop_front() else {
return Ok(());
};
match self.poll(id) {
Poll::Ready(Ok(())) => {
self.processes[id] = None;
self.live -= 1;
}
Poll::Ready(Err(e)) => return Err(RawProcessError { pid: id, source: e }),
Poll::Pending => {}
}
}
}
fn poll(&mut self, id: ProcessID) -> Poll<process::ProcessResult> {
let _guard = Guard::enter(id);
let process = self.processes[id].as_mut().unwrap();
process
.future
.as_mut()
.poll(&mut Context::from_waker(&process.waker))
}
fn flush_wakes(&mut self) {
{
let mut woken = self.inbox.woken.lock().unwrap();
if woken.is_empty() {
return;
}
self.flush_buf.extend(woken.drain(..));
}
if self.flush_buf.len() > 1 {
self.flush_buf.sort_unstable();
self.flush_buf.dedup();
}
for i in 0..self.flush_buf.len() {
let id = self.flush_buf[i];
let live = matches!(self.processes.get(id), Some(Some(_)));
if live && !self.queue.contains(&id) {
self.queue.push_back(id);
}
}
self.flush_buf.clear();
}
}
impl Default for Executor<'_> {
fn default() -> Self {
Self {
processes: Vec::new(),
queue: VecDeque::new(),
inbox: Arc::new(Inbox {
woken: Mutex::new(VecDeque::new()),
}),
live: 0,
flush_buf: Vec::new(),
}
}
}
thread_local! {
static CURRENT: Cell<Option<ProcessID>> = const { Cell::new(None) };
}
struct Guard;
impl Guard {
fn enter(id: ProcessID) -> Self {
CURRENT.with(|c| c.set(Some(id)));
Guard
}
}
impl Drop for Guard {
fn drop(&mut self) {
CURRENT.with(|c| c.set(None));
}
}
pub(crate) fn pid() -> ProcessID {
CURRENT
.with(Cell::get)
.expect("pid() called outside an executor poll")
}
#[derive(Debug)]
pub(crate) struct RawProcessError {
pub(crate) pid: ProcessID,
pub(crate) source: Box<dyn Error + Send + Sync>,
}
#[cfg(test)]
mod tests {
use super::*;
use std::future::poll_fn;
#[test]
fn completes_ok() {
let mut exec = Executor::default();
exec.schedule(async { Ok(()) });
assert!(exec.execute().is_ok());
}
#[test]
fn process_error() {
let mut exec = Executor::default();
exec.schedule(async { Err("boom".into()) });
let err = exec.execute().unwrap_err();
assert!(matches!(err, RawProcessError { pid: 0, .. }));
}
#[test]
fn self_yield() {
let mut exec = Executor::default();
let mut yielded = false;
exec.schedule(poll_fn(move |cx| {
if yielded {
Poll::Ready(Ok(()))
} else {
yielded = true;
cx.waker().wake_by_ref();
Poll::Pending
}
}));
assert!(exec.execute().is_ok());
}
}