use std::collections::HashMap;
use std::future::Future;
use std::time::Duration;
use tokio::task::{Id, JoinError, JoinSet};
use tokio::time::Instant;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Default)]
pub enum ManagedTaskPolicy {
#[default]
FaultOnExit,
AllowExit,
}
struct ManagedTaskInfo {
name: String,
policy: ManagedTaskPolicy,
}
pub(crate) struct ManagedTaskExit {
pub(crate) name: String,
pub(crate) panic_message: Option<String>,
}
#[derive(Default)]
pub(crate) struct ManagedTasks {
join_set: JoinSet<()>,
info: HashMap<Id, ManagedTaskInfo>,
}
impl ManagedTasks {
pub(crate) fn spawn<F>(&mut self, name: impl Into<String>, policy: ManagedTaskPolicy, future: F)
where
F: Future<Output = ()> + Send + 'static,
{
let abort = self.join_set.spawn(future);
self.info.insert(
abort.id(),
ManagedTaskInfo {
name: name.into(),
policy,
},
);
}
pub(crate) async fn next_unexpected_exit(&mut self) -> ManagedTaskExit {
loop {
let Some(result) = self.join_set.join_next_with_id().await else {
return std::future::pending().await;
};
match result {
Ok((id, ())) => {
let Some(info) = self.info.remove(&id) else {
continue;
};
if info.policy == ManagedTaskPolicy::AllowExit {
continue;
}
return ManagedTaskExit {
name: info.name,
panic_message: None,
};
}
Err(join_error) => {
let id = join_error.id();
let Some(info) = self.info.remove(&id) else {
continue;
};
if info.policy == ManagedTaskPolicy::AllowExit || join_error.is_cancelled() {
continue;
}
let panic_message = join_error
.is_panic()
.then(|| panic_message(join_error.into_panic()));
return ManagedTaskExit {
name: info.name,
panic_message,
};
}
};
}
}
pub(crate) fn cancel(&mut self) {
self.join_set.abort_all();
}
pub(crate) async fn shutdown_within(mut self, grace: Duration) -> Vec<String> {
self.cancel();
let deadline = Instant::now() + grace;
self.join_until(deadline).await
}
pub(crate) async fn join_until(mut self, deadline: Instant) -> Vec<String> {
loop {
self.drain_ready();
if self.info.is_empty() {
break;
}
let remaining = deadline.saturating_duration_since(Instant::now());
if remaining.is_zero() {
break;
}
match tokio::time::timeout(remaining, self.join_set.join_next_with_id()).await {
Ok(Some(result)) => self.forget_finished(result),
Ok(None) => break,
Err(_elapsed) => {
self.drain_ready();
break;
}
}
}
self.info.into_values().map(|info| info.name).collect()
}
fn drain_ready(&mut self) {
while let Some(result) = self.join_set.try_join_next_with_id() {
self.forget_finished(result);
}
}
fn forget_finished(&mut self, result: Result<(Id, ()), JoinError>) {
match result {
Ok((id, ())) => {
self.info.remove(&id);
}
Err(join_error) => {
self.info.remove(&join_error.id());
}
}
}
}
fn panic_message(payload: Box<dyn std::any::Any + Send>) -> String {
if let Some(message) = payload.downcast_ref::<&str>() {
(*message).to_string()
} else if let Some(message) = payload.downcast_ref::<String>() {
message.clone()
} else {
"managed task panicked with a non-string payload".to_string()
}
}