use std::sync::Arc;
use tokio::sync::watch;
#[derive(Clone, Debug)]
pub struct Shutdown(Arc<watch::Sender<bool>>);
impl Default for Shutdown {
fn default() -> Self {
Self::new()
}
}
impl Shutdown {
pub fn new() -> Self {
Self(Arc::new(watch::channel(false).0))
}
pub fn trigger(&self) {
self.0.send_replace(true);
}
pub fn is_triggered(&self) -> bool {
*self.0.borrow()
}
pub async fn wait(&self) {
let mut rx = self.0.subscribe();
let _ = rx.wait_for(|fired| *fired).await;
}
pub async fn wait_owned(self) {
self.wait().await
}
}
pub async fn os_signal() {
let ctrl_c = async {
if let Err(error) = tokio::signal::ctrl_c().await {
tracing::warn!(%error, "cannot listen for Ctrl-C");
std::future::pending::<()>().await;
}
};
#[cfg(unix)]
let terminate = async {
match tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate()) {
Ok(mut signal) => {
signal.recv().await;
}
Err(error) => {
tracing::warn!(%error, "cannot listen for SIGTERM");
std::future::pending::<()>().await;
}
}
};
#[cfg(not(unix))]
let terminate = std::future::pending::<()>();
tokio::select! {
_ = ctrl_c => {}
_ = terminate => {}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn trigger_wakes_every_waiter() {
let shutdown = Shutdown::new();
assert!(!shutdown.is_triggered());
let waiter = tokio::spawn(shutdown.clone().wait_owned());
shutdown.trigger();
shutdown.trigger();
assert!(tokio::time::timeout(std::time::Duration::from_secs(2), waiter).await.is_ok());
assert!(shutdown.is_triggered());
shutdown.wait().await;
}
}