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()
}
}
#[cfg(test)]
mod tests {
use super::{ManagedTaskPolicy, ManagedTasks};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
const NEVER: Duration = Duration::from_secs(3600);
#[tokio::test(start_paused = true)]
async fn an_early_return_faults_and_names_the_task() {
let mut tasks = ManagedTasks::default();
tasks.spawn("sensor-loop", ManagedTaskPolicy::FaultOnExit, async {});
let exit = tokio::time::timeout(NEVER, tasks.next_unexpected_exit())
.await
.expect("a FaultOnExit task that returns must be reported");
assert_eq!(exit.name, "sensor-loop");
assert_eq!(
exit.panic_message, None,
"a normal return is an unexpected exit, not a panic"
);
}
#[tokio::test(start_paused = true)]
async fn a_panic_is_reported_with_its_message() {
let mut tasks = ManagedTasks::default();
tasks.spawn("io-pump", ManagedTaskPolicy::FaultOnExit, async {
panic!("serial port vanished");
});
let exit = tokio::time::timeout(NEVER, tasks.next_unexpected_exit())
.await
.expect("a panicking FaultOnExit task must be reported");
assert_eq!(exit.name, "io-pump");
assert_eq!(exit.panic_message.as_deref(), Some("serial port vanished"));
}
#[tokio::test(start_paused = true)]
async fn allow_exit_suppresses_both_a_return_and_a_panic() {
let mut tasks = ManagedTasks::default();
tasks.spawn("cache-prime", ManagedTaskPolicy::AllowExit, async {});
tasks.spawn("warm-up", ManagedTaskPolicy::AllowExit, async {
panic!("best-effort work failed");
});
assert!(
tokio::time::timeout(NEVER, tasks.next_unexpected_exit())
.await
.is_err(),
"AllowExit completions must never surface as unexpected exits"
);
}
#[tokio::test(start_paused = true)]
async fn a_real_fault_is_not_masked_by_allow_exit_siblings() {
let mut tasks = ManagedTasks::default();
tasks.spawn("cache-prime", ManagedTaskPolicy::AllowExit, async {});
tasks.spawn("watchdog", ManagedTaskPolicy::FaultOnExit, async {
tokio::task::yield_now().await;
});
let exit = tokio::time::timeout(NEVER, tasks.next_unexpected_exit())
.await
.expect("the FaultOnExit task must still be reported");
assert_eq!(exit.name, "watchdog");
}
#[tokio::test]
async fn shutdown_cancels_and_joins_every_task() {
let cancelled = Arc::new(AtomicBool::new(false));
let observed = Arc::clone(&cancelled);
let started = Arc::new(AtomicBool::new(false));
let running = Arc::clone(&started);
let mut tasks = ManagedTasks::default();
tasks.spawn("forever", ManagedTaskPolicy::FaultOnExit, async move {
struct OnCancel(Arc<AtomicBool>);
impl Drop for OnCancel {
fn drop(&mut self) {
self.0.store(true, Ordering::Relaxed);
}
}
let _guard = OnCancel(observed);
running.store(true, Ordering::Relaxed);
std::future::pending::<()>().await;
});
while !started.load(Ordering::Relaxed) {
tokio::task::yield_now().await;
}
let unjoined = tasks.shutdown_within(Duration::from_secs(5)).await;
assert!(
unjoined.is_empty(),
"a cancellable task must join: {unjoined:?}"
);
assert!(
cancelled.load(Ordering::Relaxed),
"the task must observe cancellation"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn an_uncancellable_task_is_reported_rather_than_waited_for() {
let started = Arc::new(AtomicBool::new(false));
let running = Arc::clone(&started);
let finished = Arc::new(AtomicBool::new(false));
let completed = Arc::clone(&finished);
let mut tasks = ManagedTasks::default();
tasks.spawn(
"uncancellable",
ManagedTaskPolicy::FaultOnExit,
async move {
running.store(true, Ordering::Relaxed);
std::thread::sleep(Duration::from_millis(1500));
completed.store(true, Ordering::Relaxed);
},
);
while !started.load(Ordering::Relaxed) {
tokio::task::yield_now().await;
}
let began = std::time::Instant::now();
let unjoined = tasks.shutdown_within(Duration::from_millis(150)).await;
let elapsed = began.elapsed();
assert_eq!(
unjoined,
vec!["uncancellable".to_string()],
"a task still running at the grace deadline is reported by name"
);
assert!(
elapsed < Duration::from_millis(700),
"shutdown must be bounded by the grace budget, took {elapsed:?}"
);
assert!(
!finished.load(Ordering::Relaxed),
"shutdown must return before a cancellation-ignoring task finishes"
);
}
}