use std::time::Duration;
use tokio::task::JoinSet;
use tokio_util::sync::CancellationToken;
pub struct GracefulShutdown {
tasks: JoinSet<()>,
token: CancellationToken,
}
impl GracefulShutdown {
pub fn new() -> Self {
Self {
tasks: JoinSet::new(),
token: CancellationToken::new(),
}
}
pub fn token(&self) -> CancellationToken {
self.token.clone()
}
pub fn spawn<F>(&mut self, future: F)
where
F: std::future::Future<Output = ()> + Send + 'static,
{
self.tasks.spawn(future);
}
pub fn len(&self) -> usize {
self.tasks.len()
}
pub fn is_empty(&self) -> bool {
self.tasks.is_empty()
}
pub async fn shutdown(mut self, timeout: Duration) -> (bool, usize) {
self.token.cancel();
let deadline = tokio::time::Instant::now() + timeout;
let mut success = true;
let mut aborted = 0usize;
loop {
if self.tasks.is_empty() {
break;
}
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
if remaining.is_zero() {
aborted = self.tasks.len();
self.tasks.abort_all();
while self.tasks.join_next().await.is_some() {}
success = false;
break;
}
tokio::select! {
biased;
_ = tokio::time::sleep(remaining) => {
aborted = self.tasks.len();
self.tasks.abort_all();
while self.tasks.join_next().await.is_some() {}
success = false;
break;
}
res = self.tasks.join_next() => {
if res.is_none() {
break;
}
}
}
}
(success, aborted)
}
pub fn abort_now(mut self) -> usize {
let n = self.tasks.len();
self.tasks.abort_all();
n
}
}
impl Default for GracefulShutdown {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
#[tokio::test]
async fn test_new_empty() {
let gs = GracefulShutdown::new();
assert!(gs.is_empty());
assert_eq!(gs.len(), 0);
}
#[tokio::test]
async fn test_spawn_increments_len() {
let mut gs = GracefulShutdown::new();
let token = gs.token();
gs.spawn(async move {
token.cancelled().await;
});
assert_eq!(gs.len(), 1);
assert!(!gs.is_empty());
}
#[tokio::test]
async fn test_shutdown_success_all_tasks_complete() {
let mut gs = GracefulShutdown::new();
let token = gs.token();
let flag = Arc::new(AtomicBool::new(false));
let flag_clone = flag.clone();
gs.spawn(async move {
token.cancelled().await;
flag_clone.store(true, Ordering::SeqCst);
});
let (success, aborted) = gs.shutdown(Duration::from_secs(1)).await;
assert!(success);
assert_eq!(aborted, 0);
assert!(flag.load(Ordering::SeqCst));
}
#[tokio::test]
async fn test_shutdown_multiple_tasks() {
let mut gs = GracefulShutdown::new();
let token1 = gs.token();
let token2 = gs.token();
gs.spawn(async move {
token1.cancelled().await;
});
gs.spawn(async move {
token2.cancelled().await;
});
let (success, aborted) = gs.shutdown(Duration::from_secs(1)).await;
assert!(success);
assert_eq!(aborted, 0);
}
#[tokio::test]
async fn test_shutdown_timeout_aborts_tasks() {
let mut gs = GracefulShutdown::new();
gs.spawn(async {
std::future::pending::<()>().await;
});
let (success, aborted) = gs.shutdown(Duration::from_millis(50)).await;
assert!(!success);
assert_eq!(aborted, 1);
}
#[tokio::test]
async fn test_shutdown_with_mixed_tasks() {
let mut gs = GracefulShutdown::new();
let token = gs.token();
gs.spawn(async move {
token.cancelled().await;
});
gs.spawn(async {
std::future::pending::<()>().await;
});
let (success, aborted) = gs.shutdown(Duration::from_millis(50)).await;
assert!(!success);
assert_eq!(aborted, 1);
}
#[tokio::test]
async fn test_shutdown_no_tasks() {
let gs = GracefulShutdown::new();
let (success, aborted) = gs.shutdown(Duration::from_millis(100)).await;
assert!(success);
assert_eq!(aborted, 0);
}
#[tokio::test]
async fn test_abort_now() {
let mut gs = GracefulShutdown::new();
gs.spawn(async {
std::future::pending::<()>().await;
});
gs.spawn(async {
std::future::pending::<()>().await;
});
let aborted = gs.abort_now();
assert_eq!(aborted, 2);
}
#[tokio::test]
async fn test_abort_now_empty() {
let gs = GracefulShutdown::new();
assert_eq!(gs.abort_now(), 0);
}
#[tokio::test]
async fn test_token_cancellation_propagates() {
let mut gs = GracefulShutdown::new();
let token = gs.token();
let counter = Arc::new(std::sync::atomic::AtomicUsize::new(0));
for _ in 0..5 {
let t = token.clone();
let c = counter.clone();
gs.spawn(async move {
t.cancelled().await;
c.fetch_add(1, Ordering::SeqCst);
});
}
let (success, aborted) = gs.shutdown(Duration::from_secs(1)).await;
assert!(success);
assert_eq!(aborted, 0);
assert_eq!(counter.load(Ordering::SeqCst), 5);
}
#[tokio::test]
async fn test_default_impl() {
let gs = GracefulShutdown::default();
assert!(gs.is_empty());
}
#[tokio::test]
async fn test_shutdown_zero_timeout() {
let mut gs = GracefulShutdown::new();
let token = gs.token();
gs.spawn(async move {
token.cancelled().await;
});
let (success, aborted) = gs.shutdown(Duration::from_millis(0)).await;
let _ = (success, aborted);
}
#[tokio::test]
async fn test_shutdown_returns_correct_aborted_count() {
let mut gs = GracefulShutdown::new();
gs.spawn(async {
std::future::pending::<()>().await;
});
gs.spawn(async {
std::future::pending::<()>().await;
});
gs.spawn(async {
std::future::pending::<()>().await;
});
let (success, aborted) = gs.shutdown(Duration::from_millis(10)).await;
assert!(!success);
assert_eq!(aborted, 3);
}
#[tokio::test]
async fn test_task_completes_before_shutdown() {
let mut gs = GracefulShutdown::new();
let flag = Arc::new(AtomicBool::new(false));
let flag_clone = flag.clone();
gs.spawn(async move {
flag_clone.store(true, Ordering::SeqCst);
});
tokio::time::sleep(Duration::from_millis(50)).await;
let (success, aborted) = gs.shutdown(Duration::from_secs(1)).await;
assert!(success);
assert_eq!(aborted, 0);
assert!(flag.load(Ordering::SeqCst));
}
}