PicoBot/src/task_supervisor.rs

203 lines
6.8 KiB
Rust

use std::future::Future;
use std::panic::AssertUnwindSafe;
use std::sync::{Arc, Mutex};
use std::time::Duration;
use futures_util::FutureExt;
use tokio::task::JoinHandle;
use tokio::time::{Instant, timeout_at};
use tokio_util::sync::CancellationToken;
#[derive(Clone)]
pub struct TaskSupervisor {
inner: Arc<Inner>,
}
struct Inner {
cancellation: CancellationToken,
state: Mutex<State>,
}
impl Drop for Inner {
fn drop(&mut self) {
self.cancellation.cancel();
}
}
#[derive(Default)]
struct State {
stopping: bool,
tasks: Vec<ManagedTask>,
}
struct ManagedTask {
name: String,
handle: JoinHandle<()>,
}
impl Default for TaskSupervisor {
fn default() -> Self {
Self::new()
}
}
impl TaskSupervisor {
pub fn new() -> Self {
Self {
inner: Arc::new(Inner {
cancellation: CancellationToken::new(),
state: Mutex::new(State::default()),
}),
}
}
pub fn cancellation_token(&self) -> CancellationToken {
self.inner.cancellation.clone()
}
/// Register a task before shutdown begins. Cancellation drops the task
/// future, so task code should keep externally visible state transactional.
pub fn spawn<F>(&self, name: impl Into<String>, future: F) -> bool
where
F: Future<Output = ()> + Send + 'static,
{
let name = name.into();
let cancellation = self.inner.cancellation.clone();
let mut state = self.inner.state.lock().unwrap_or_else(|e| e.into_inner());
if state.stopping {
return false;
}
// Completed handles no longer need to occupy the registry. Panics are
// observed and logged inside the wrapper below.
state.tasks.retain(|task| !task.handle.is_finished());
let task_name = name.clone();
let handle = tokio::spawn(async move {
tracing::debug!(task = %task_name, "Background task started");
let outcome = tokio::select! {
_ = cancellation.cancelled() => None,
outcome = AssertUnwindSafe(future).catch_unwind() => Some(outcome),
};
match outcome {
Some(Ok(())) => tracing::debug!(task = %task_name, "Background task finished"),
Some(Err(_)) => tracing::error!(task = %task_name, "Background task panicked"),
None => tracing::debug!(task = %task_name, "Background task cancelled"),
}
});
state.tasks.push(ManagedTask { name, handle });
true
}
/// Register a task that performs its own cooperative cancellation and
/// cleanup. The supervisor broadcasts cancellation during shutdown, but
/// does not drop this future until the grace period expires.
pub fn spawn_graceful<F>(&self, name: impl Into<String>, future: F) -> bool
where
F: Future<Output = ()> + Send + 'static,
{
let name = name.into();
let mut state = self.inner.state.lock().unwrap_or_else(|e| e.into_inner());
if state.stopping {
return false;
}
state.tasks.retain(|task| !task.handle.is_finished());
let task_name = name.clone();
let handle = tokio::spawn(async move {
tracing::debug!(task = %task_name, "Graceful background task started");
match AssertUnwindSafe(future).catch_unwind().await {
Ok(()) => tracing::debug!(task = %task_name, "Graceful background task finished"),
Err(_) => tracing::error!(task = %task_name, "Graceful background task panicked"),
}
});
state.tasks.push(ManagedTask { name, handle });
true
}
pub fn cancel(&self) {
let mut state = self.inner.state.lock().unwrap_or_else(|e| e.into_inner());
state.stopping = true;
self.inner.cancellation.cancel();
}
/// Stop accepting tasks, broadcast cancellation, and wait up to `grace`.
/// Remaining tasks are aborted so shutdown has a deterministic upper bound.
pub async fn shutdown(&self, grace: Duration) {
let mut tasks = {
let mut state = self.inner.state.lock().unwrap_or_else(|e| e.into_inner());
state.stopping = true;
self.inner.cancellation.cancel();
std::mem::take(&mut state.tasks)
};
let deadline = Instant::now() + grace;
for index in 0..tasks.len() {
let result = timeout_at(deadline, &mut tasks[index].handle).await;
match result {
Ok(Ok(())) => {}
Ok(Err(error)) if error.is_cancelled() => {}
Ok(Err(error)) => {
tracing::error!(task = %tasks[index].name, error = %error, "Background task join failed");
}
Err(_) => {
for task in &tasks[index..] {
if !task.handle.is_finished() {
tracing::warn!(task = %task.name, "Aborting background task after shutdown grace period");
task.handle.abort();
}
}
for task in &mut tasks[index..] {
let _ = (&mut task.handle).await;
}
break;
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicBool, Ordering};
#[tokio::test]
async fn shutdown_cancels_registered_task() {
let supervisor = TaskSupervisor::new();
let dropped = Arc::new(AtomicBool::new(false));
let marker = dropped.clone();
supervisor.spawn("pending", async move {
struct DropMarker(Arc<AtomicBool>);
impl Drop for DropMarker {
fn drop(&mut self) {
self.0.store(true, Ordering::SeqCst);
}
}
let _marker = DropMarker(marker);
std::future::pending::<()>().await;
});
tokio::task::yield_now().await;
supervisor.cancel();
assert!(!supervisor.spawn("late", async {}));
supervisor.shutdown(Duration::from_secs(1)).await;
assert!(dropped.load(Ordering::SeqCst));
}
#[tokio::test]
async fn graceful_task_observes_cancellation_before_shutdown_returns() {
let supervisor = TaskSupervisor::new();
let cancellation = supervisor.cancellation_token();
let cleaned_up = Arc::new(std::sync::atomic::AtomicBool::new(false));
let task_cleaned_up = cleaned_up.clone();
assert!(supervisor.spawn_graceful("graceful", async move {
cancellation.cancelled().await;
task_cleaned_up.store(true, std::sync::atomic::Ordering::SeqCst);
}));
supervisor.shutdown(Duration::from_secs(1)).await;
assert!(cleaned_up.load(std::sync::atomic::Ordering::SeqCst));
}
}