use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::Notify;
use tracing::info;
#[derive(Clone, Default)]
pub struct ShutdownSignal {
inner: Arc<Inner>,
}
#[derive(Default)]
struct Inner {
triggered: AtomicBool,
notify: Notify,
}
impl ShutdownSignal {
pub fn new() -> Self {
Self::default()
}
pub fn shutdown(&self) {
if !self.inner.triggered.swap(true, Ordering::AcqRel) {
info!("shutdown requested");
self.inner.notify.notify_waiters();
}
}
pub fn is_shutdown(&self) -> bool {
self.inner.triggered.load(Ordering::Acquire)
}
pub async fn cancelled(&self) {
let notified = self.inner.notify.notified();
tokio::pin!(notified);
notified.as_mut().enable();
if self.is_shutdown() {
return;
}
notified.await;
}
pub async fn wait_timeout(&self, timeout: Duration) -> bool {
tokio::time::timeout(timeout, self.cancelled())
.await
.is_ok()
}
}
impl std::fmt::Debug for ShutdownSignal {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ShutdownSignal")
.field("triggered", &self.is_shutdown())
.finish()
}
}
pub fn install_signal_handlers(signal: ShutdownSignal) {
#[cfg(unix)]
tokio::spawn(async move {
use tokio::signal::unix::{signal as unix_signal, SignalKind};
let mut streams: Vec<_> = [
("SIGINT", SignalKind::interrupt()),
("SIGTERM", SignalKind::terminate()),
("SIGQUIT", SignalKind::quit()),
]
.into_iter()
.filter_map(|(name, kind)| match unix_signal(kind) {
Ok(stream) => Some((name, stream)),
Err(e) => {
tracing::warn!(signal = name, error = %e, "could not install a signal handler");
None
}
})
.collect();
if streams.is_empty() {
return;
}
let received =
futures_util::future::select_all(streams.iter_mut().map(|(name, stream)| {
Box::pin(async move { stream.recv().await.map(|()| *name) })
}))
.await
.0;
if let Some(name) = received {
info!(signal = name, "received termination signal");
}
signal.shutdown();
});
#[cfg(windows)]
tokio::spawn(async move {
match tokio::signal::ctrl_c().await {
Ok(()) => info!("received Ctrl-C"),
Err(e) => {
tracing::warn!(error = %e, "could not install the Ctrl-C handler");
return;
}
}
signal.shutdown();
});
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn shutdown_is_observable_and_idempotent() {
let signal = ShutdownSignal::new();
assert!(!signal.is_shutdown());
signal.shutdown();
signal.shutdown();
assert!(signal.is_shutdown());
}
#[tokio::test]
async fn a_waiter_registered_first_is_woken() {
let signal = ShutdownSignal::new();
let waiter = tokio::spawn({
let signal = signal.clone();
async move { signal.cancelled().await }
});
tokio::task::yield_now().await;
signal.shutdown();
tokio::time::timeout(Duration::from_secs(1), waiter)
.await
.expect("waiter must be woken")
.unwrap();
}
#[tokio::test]
async fn cancellation_latches_for_late_waiters() {
let signal = ShutdownSignal::new();
signal.shutdown();
tokio::time::timeout(Duration::from_millis(50), signal.cancelled())
.await
.expect("a late waiter must return immediately");
}
#[tokio::test]
async fn clones_share_one_state() {
let signal = ShutdownSignal::new();
let a = signal.clone();
let b = signal.clone();
a.shutdown();
assert!(b.is_shutdown() && signal.is_shutdown());
}
#[tokio::test]
async fn wait_timeout_reports_which_happened() {
let signal = ShutdownSignal::new();
assert!(!signal.wait_timeout(Duration::from_millis(20)).await);
signal.shutdown();
assert!(signal.wait_timeout(Duration::from_millis(20)).await);
}
}