use tokio_util::sync::CancellationToken;
#[derive(Debug, Clone)]
pub struct ShutdownSignal(CancellationToken);
impl ShutdownSignal {
pub fn new() -> (Self, ShutdownCtl) {
let token = CancellationToken::new();
(Self(token.clone()), ShutdownCtl(token))
}
pub async fn wait_for_shutdown(&self) {
self.0.cancelled().await
}
}
#[derive(Debug)]
pub struct ShutdownCtl(CancellationToken);
impl ShutdownCtl {
pub fn shutdown_now(self) {
self.0.cancel();
}
pub fn get_signal(&self) -> ShutdownSignal {
ShutdownSignal(self.0.clone())
}
}
impl Drop for ShutdownCtl {
fn drop(&mut self) {
self.0.cancel();
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use tokio::{join, sync::mpsc};
use super::*;
#[tokio::test]
async fn test_explicit_shutdown() {
let (signal, ctl) = ShutdownSignal::new();
let (tx, mut rx) = mpsc::channel::<()>(1);
let a = tokio::spawn({
let signal = signal.clone();
let tx = tx.clone();
async move {
drop(tx);
signal.wait_for_shutdown().await;
}
});
let b = tokio::spawn({
async move {
drop(tx);
signal.wait_for_shutdown().await;
signal.wait_for_shutdown().await;
signal.wait_for_shutdown().await;
signal.wait_for_shutdown().await;
}
});
let _ = rx.recv().await;
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(!a.is_finished());
assert!(!b.is_finished());
ctl.shutdown_now();
let (a, b) = tokio::time::timeout(Duration::from_secs(5), async { join!(a, b) })
.await
.expect("timeout waiting for worker tasks to stop");
assert!(a.is_ok());
assert!(b.is_ok());
}
#[tokio::test]
async fn test_implicit_shutdown() {
let (signal, ctl) = ShutdownSignal::new();
let (tx, mut rx) = mpsc::channel::<()>(1);
let a = tokio::spawn({
let signal = signal.clone();
let tx = tx.clone();
async move {
drop(tx);
signal.wait_for_shutdown().await;
}
});
let b = tokio::spawn({
async move {
drop(tx);
signal.wait_for_shutdown().await;
signal.wait_for_shutdown().await;
signal.wait_for_shutdown().await;
signal.wait_for_shutdown().await;
}
});
let _ = rx.recv().await;
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(!a.is_finished());
assert!(!b.is_finished());
drop(ctl);
let (a, b) = tokio::time::timeout(Duration::from_secs(5), async { join!(a, b) })
.await
.expect("timeout waiting for worker tasks to stop");
assert!(a.is_ok());
assert!(b.is_ok());
}
}