use std::{collections::VecDeque, future::Future, pin::Pin, task::Poll};
#[cfg(not(target_family = "wasm"))]
pub(crate) type MaybeSendBox<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
#[cfg(target_family = "wasm")]
pub(crate) type MaybeSendBox<'a, T> = Pin<Box<dyn Future<Output = T> + 'a>>;
#[cfg(not(target_family = "wasm"))]
pub(crate) type MaybeSendTask = Box<dyn FnMut(&kio::Waiter) -> Poll<()> + Send>;
#[cfg(target_family = "wasm")]
pub(crate) type MaybeSendTask = Box<dyn FnMut(&kio::Waiter) -> Poll<()>>;
#[cfg(not(target_family = "wasm"))]
pub(crate) trait MaybeBoxedExt<'a>: Future + Send + Sized + 'a {
fn maybe_boxed(self) -> MaybeSendBox<'a, Self::Output> {
Box::pin(self)
}
}
#[cfg(not(target_family = "wasm"))]
impl<'a, F: Future + Send + 'a> MaybeBoxedExt<'a> for F {}
#[cfg(target_family = "wasm")]
pub(crate) trait MaybeBoxedExt<'a>: Future + Sized + 'a {
fn maybe_boxed(self) -> MaybeSendBox<'a, Self::Output> {
Box::pin(self)
}
}
#[cfg(target_family = "wasm")]
impl<'a, F: Future + 'a> MaybeBoxedExt<'a> for F {}
fn future_task(mut task: MaybeSendBox<'static, ()>) -> MaybeSendTask {
Box::new(move |waiter| waiter.poll_future(task.as_mut()))
}
pub(crate) async fn err_only<E>(fut: impl Future<Output = Result<(), E>>) -> E {
match fut.await {
Err(err) => err,
Ok(()) => std::future::pending().await,
}
}
#[derive(Clone)]
pub(crate) struct Tasks {
state: kio::Shared<Submissions>,
alive: kio::Producer<()>,
}
pub(crate) struct Keepalive {
_alive: kio::Producer<()>,
}
#[derive(Default)]
struct Submissions {
queued: VecDeque<MaybeSendBox<'static, ()>>,
closed: bool,
}
impl Tasks {
pub fn push(&self, task: impl MaybeBoxedExt<'static, Output = ()>) {
let mut state = self.state.lock();
if !state.closed {
state.queued.push_back(task.maybe_boxed());
}
}
}
impl Tasks {
pub fn downgrade(&self) -> TasksWeak {
TasksWeak {
state: self.state.clone(),
alive: self.alive.consume(),
}
}
pub fn keepalive(&self) -> Keepalive {
Keepalive {
_alive: self.alive.clone(),
}
}
}
#[derive(Clone)]
pub(crate) struct TasksWeak {
state: kio::Shared<Submissions>,
alive: kio::Consumer<()>,
}
impl TasksWeak {
pub fn push(&self, task: impl MaybeBoxedExt<'static, Output = ()>) {
let mut state = self.state.lock();
if !state.closed && !self.alive.is_closed() {
state.queued.push_back(task.maybe_boxed());
}
}
}
pub(crate) struct TaskSet {
state: kio::Shared<Submissions>,
active: kio::Tasks<MaybeSendTask>,
alive: kio::Consumer<()>,
}
impl TaskSet {
pub fn new() -> (Tasks, Self) {
let state = kio::Shared::<Submissions>::default();
let alive = kio::Producer::<()>::default();
let set = Self {
state: state.clone(),
active: kio::Tasks::new(),
alive: alive.consume(),
};
(Tasks { state, alive }, set)
}
pub fn owned() -> Self {
Self {
state: kio::Shared::default(),
active: kio::Tasks::new(),
alive: kio::Producer::<()>::default().consume(),
}
}
pub fn push(&mut self, task: impl MaybeBoxedExt<'static, Output = ()>) {
self.active.push(future_task(task.maybe_boxed()));
}
pub fn poll(&mut self, waiter: &kio::Waiter) -> Poll<()> {
let mut submissions_done = false;
loop {
let orphaned = self.alive.poll_closed(waiter).is_ready();
let Poll::Ready(mut state) = self.state.poll(waiter, |state| {
if state.queued.is_empty() && !orphaned {
Poll::Pending
} else {
Poll::Ready(())
}
}) else {
break;
};
while let Some(task) = state.queued.pop_front() {
self.active.push(future_task(task));
}
if orphaned {
submissions_done = true;
break;
}
}
match self.active.poll(waiter) {
Poll::Ready(()) if submissions_done => Poll::Ready(()),
_ => Poll::Pending,
}
}
pub async fn drive<T>(&mut self, mut f: impl FnMut(&kio::Waiter) -> Poll<T> + Unpin) -> T {
kio::wait(|waiter| {
if let Poll::Ready(output) = f(waiter) {
return Poll::Ready(output);
}
let _ = self.poll(waiter);
Poll::Pending
})
.await
}
#[cfg(test)]
pub async fn run(mut self) {
kio::wait(|waiter| self.poll(waiter)).await
}
}
impl Drop for TaskSet {
fn drop(&mut self) {
let queued = {
let mut state = self.state.lock();
state.closed = true;
std::mem::take(&mut state.queued)
};
drop(queued);
}
}
#[cfg(test)]
mod tests {
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};
use super::*;
#[test]
fn task_set_runs_nested_work_without_a_runtime() {
let (tasks, task_set) = TaskSet::new();
let completed = Arc::new(AtomicUsize::new(0));
let nested_tasks = tasks.clone();
let outer_completed = completed.clone();
tasks.push(async move {
outer_completed.fetch_add(1, Ordering::SeqCst);
let inner_completed = outer_completed.clone();
nested_tasks.push(async move {
inner_completed.fetch_add(1, Ordering::SeqCst);
});
});
drop(tasks);
futures::executor::block_on(task_set.run());
assert_eq!(completed.load(Ordering::SeqCst), 2);
}
#[test]
fn drive_polls_children_alongside_the_accept_future() {
let (tasks, mut set) = TaskSet::new();
let completed = Arc::new(AtomicUsize::new(0));
let child_completed = completed.clone();
tasks.push(async move {
child_completed.fetch_add(1, Ordering::SeqCst);
});
let gate = completed.clone();
let output = futures::executor::block_on(set.drive(move |waiter| {
if gate.load(Ordering::SeqCst) == 1 {
std::task::Poll::Ready(42)
} else {
waiter.waker().wake_by_ref();
std::task::Poll::Pending
}
}));
assert_eq!(output, 42);
}
}