use std::sync::Arc;
use std::time::Duration;
use tokio::sync::watch;
use tracing::Instrument;
#[derive(Clone)]
pub struct ShutdownSignal {
rx: watch::Receiver<bool>,
}
impl ShutdownSignal {
pub fn is_shutdown(&self) -> bool {
*self.rx.borrow()
}
pub async fn wait(&mut self) {
if *self.rx.borrow() {
return;
}
let _ = self.rx.changed().await;
}
}
#[derive(Clone)]
pub struct ShutdownManager {
inner: Arc<Inner>,
}
struct Inner {
tx: watch::Sender<bool>,
_keepalive: watch::Receiver<bool>,
tasks: std::sync::Mutex<Vec<tokio::task::JoinHandle<()>>>,
}
impl Default for ShutdownManager {
fn default() -> Self {
Self::new()
}
}
impl ShutdownManager {
pub fn new() -> Self {
let (tx, keepalive) = watch::channel(false);
Self {
inner: Arc::new(Inner {
tx,
_keepalive: keepalive,
tasks: std::sync::Mutex::new(Vec::new()),
}),
}
}
pub fn handle(&self) -> ShutdownSignal {
ShutdownSignal {
rx: self.inner.tx.subscribe(),
}
}
pub fn register(&self, handle: tokio::task::JoinHandle<()>) {
self.inner.tasks.lock().unwrap().push(handle);
}
pub fn is_shutdown(&self) -> bool {
*self.inner.tx.borrow()
}
pub async fn wait_for_shutdown(&self) {
let mut rx = self.inner.tx.subscribe();
if *rx.borrow() {
return;
}
let _ = rx.changed().await;
}
pub fn request(&self) {
let _ = self.inner.tx.send(true);
}
pub async fn drain(&self, grace: Duration) {
let span = tracing::info_span!("mytheclipse_shutdown_task");
self.wait_for_shutdown().instrument(span.clone()).await;
let tasks = {
let mut guard = self.inner.tasks.lock().unwrap();
std::mem::take(&mut *guard)
};
for task in tasks {
let _ = tokio::time::timeout(grace, task)
.instrument(span.clone())
.await;
}
}
pub async fn wait_for_os_signal(&self) {
let _ = os_signal().await;
}
}
#[cfg(unix)]
async fn os_signal() {
use tokio::signal::unix::{signal, SignalKind};
let mut sigint = signal(SignalKind::interrupt()).expect("failed to install SIGINT handler");
let mut sigterm = signal(SignalKind::terminate()).expect("failed to install SIGTERM handler");
tokio::select! {
_ = sigint.recv() => {},
_ = sigterm.recv() => {},
}
}
#[cfg(not(unix))]
async fn os_signal() {
use tokio::signal::ctrl_c;
let _ = ctrl_c().await;
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn signal_fires_on_request() {
let manager = ShutdownManager::new();
let mut sig = manager.handle();
assert!(!sig.is_shutdown());
assert!(!manager.is_shutdown());
let observer = tokio::spawn(async move {
sig.wait().await;
sig.is_shutdown()
});
manager.request();
assert!(manager.is_shutdown());
assert!(observer.await.unwrap());
}
#[tokio::test]
async fn request_is_idempotent() {
let manager = ShutdownManager::new();
manager.request();
manager.request();
let sig = manager.handle();
assert!(sig.is_shutdown());
}
#[tokio::test]
async fn drain_waits_for_registered_tasks() {
let manager = ShutdownManager::new();
let sig = manager.handle();
let handle = tokio::spawn(async move {
let mut sig = sig;
sig.wait().await;
});
manager.register(handle);
manager.request();
manager.drain(Duration::from_secs(5)).await;
}
#[tokio::test]
async fn drain_times_out_a_slow_task() {
let manager = ShutdownManager::new();
let _sig = manager.handle();
let slow = tokio::spawn(std::future::pending::<()>());
manager.register(slow);
manager.request();
manager.drain(Duration::from_millis(50)).await;
}
}