use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use tokio::sync::{oneshot, watch};
use tokio::task::JoinHandle;
#[derive(Clone)]
pub struct FatalSignal {
component: watch::Sender<Option<&'static str>>,
}
impl FatalSignal {
pub fn new() -> Self {
Self {
component: watch::Sender::new(None),
}
}
pub fn trip(&self, component: &'static str) {
self.component.send_if_modified(|current| {
if current.is_none() {
*current = Some(component);
true
} else {
false
}
});
}
pub fn check(&self) -> Option<&'static str> {
*self.component.borrow()
}
pub async fn tripped(&self) -> &'static str {
let mut watcher = self.component.subscribe();
loop {
if let Some(component) = *watcher.borrow_and_update() {
return component;
}
if watcher.changed().await.is_err() {
std::future::pending::<()>().await;
}
}
}
}
impl Default for FatalSignal {
fn default() -> Self {
Self::new()
}
}
pub struct TransportHandle {
cancel: Option<oneshot::Sender<()>>,
join: Option<JoinHandle<()>>,
}
impl TransportHandle {
pub fn noop() -> Self {
Self {
cancel: None,
join: None,
}
}
pub fn new(cancel: oneshot::Sender<()>, join: JoinHandle<()>) -> Self {
Self {
cancel: Some(cancel),
join: Some(join),
}
}
pub fn spawn_supervised<F, Fut, E>(
component: &'static str,
fatal: FatalSignal,
serve: F,
) -> Self
where
F: FnOnce(Pin<Box<dyn Future<Output = ()> + Send>>) -> Fut,
Fut: Future<Output = Result<(), E>> + Send + 'static,
E: std::fmt::Debug + Send + 'static,
{
let (cancel_tx, cancel_rx) = oneshot::channel::<()>();
let stop_requested = Arc::new(AtomicBool::new(false));
let stop_observed = stop_requested.clone();
let shutdown: Pin<Box<dyn Future<Output = ()> + Send>> = Box::pin(async move {
let _ = cancel_rx.await;
stop_observed.store(true, Ordering::Release);
});
let server = tokio::spawn(serve(shutdown));
let supervisor = tokio::spawn(async move {
match server.await {
Ok(Ok(())) if stop_requested.load(Ordering::Acquire) => {}
Ok(Ok(())) => {
tracing::error!(component, "server exited without a shutdown request");
fatal.trip(component);
}
Ok(Err(err)) => {
tracing::error!(error = ?err, component, "server died");
fatal.trip(component);
}
Err(join_error) => {
tracing::error!(error = ?join_error, component, "server task panicked or was aborted");
fatal.trip(component);
}
}
});
Self::new(cancel_tx, supervisor)
}
pub async fn shutdown(&mut self) {
if let Some(cancel) = self.cancel.take() {
let _ = cancel.send(());
}
if let Some(join) = self.join.take() {
if let Err(err) = join.await {
tracing::warn!(error = ?err, "peer transport task join error");
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn shutdown_signals_the_task_and_joins() {
let (cancel_tx, cancel_rx) = oneshot::channel::<()>();
let join = tokio::spawn(async move {
let _ = cancel_rx.await;
});
let mut handle = TransportHandle::new(cancel_tx, join);
handle.shutdown().await;
handle.shutdown().await;
}
#[tokio::test]
async fn noop_shutdown_is_harmless() {
let mut handle = TransportHandle::noop();
handle.shutdown().await;
}
#[tokio::test]
async fn shutdown_after_task_already_returned() {
let (cancel_tx, _cancel_rx) = oneshot::channel::<()>();
let join = tokio::spawn(async {});
tokio::task::yield_now().await;
let mut handle = TransportHandle::new(cancel_tx, join);
handle.shutdown().await;
}
#[test]
fn fatal_signal_starts_untripped() {
assert_eq!(FatalSignal::new().check(), None);
}
#[tokio::test]
async fn fatal_trip_records_first_component_only() {
let fatal = FatalSignal::new();
fatal.trip("peer server");
fatal.trip("admin server");
assert_eq!(fatal.check(), Some("peer server"));
assert_eq!(fatal.tripped().await, "peer server");
}
#[tokio::test]
async fn fatal_tripped_wakes_a_pending_waiter() {
let fatal = FatalSignal::new();
let waiter = tokio::spawn({
let fatal = fatal.clone();
async move { fatal.tripped().await }
});
tokio::task::yield_now().await;
fatal.trip("peer server");
assert_eq!(waiter.await.unwrap(), "peer server");
}
#[tokio::test]
async fn supervised_error_exit_trips_fatal() {
let fatal = FatalSignal::new();
let mut handle =
TransportHandle::spawn_supervised("test server", fatal.clone(), |_shutdown| async {
Err::<(), std::io::Error>(std::io::Error::other("listener torn down"))
});
assert_eq!(fatal.tripped().await, "test server");
handle.shutdown().await;
}
#[tokio::test]
async fn supervised_graceful_cancel_does_not_trip() {
let fatal = FatalSignal::new();
let mut handle = TransportHandle::spawn_supervised(
"test server",
fatal.clone(),
|shutdown| async move {
shutdown.await;
Ok::<(), std::io::Error>(())
},
);
handle.shutdown().await;
assert_eq!(fatal.check(), None);
}
#[tokio::test]
async fn supervised_premature_clean_exit_trips_fatal() {
let fatal = FatalSignal::new();
let mut handle =
TransportHandle::spawn_supervised("test server", fatal.clone(), |_shutdown| async {
Ok::<(), std::io::Error>(())
});
assert_eq!(fatal.tripped().await, "test server");
handle.shutdown().await;
}
#[tokio::test]
async fn supervised_panic_trips_fatal() {
let fatal = FatalSignal::new();
let mut handle =
TransportHandle::spawn_supervised("test server", fatal.clone(), |_shutdown| async {
panic!("server task blew up");
#[allow(unreachable_code)]
Ok::<(), std::io::Error>(())
});
assert_eq!(fatal.tripped().await, "test server");
handle.shutdown().await;
}
#[tokio::test]
async fn supervised_drop_without_shutdown_does_not_trip() {
let fatal = FatalSignal::new();
let handle = TransportHandle::spawn_supervised(
"test server",
fatal.clone(),
|shutdown| async move {
shutdown.await;
Ok::<(), std::io::Error>(())
},
);
drop(handle);
for _ in 0..32 {
tokio::task::yield_now().await;
assert_eq!(fatal.check(), None);
}
}
#[cfg(feature = "openraft")]
#[tokio::test]
async fn tonic_incoming_error_propagates_to_fatal() {
use std::sync::Arc;
use crate::admin::service::AdminServiceImpl;
use crate::admin::{MembershipAdmin, MembershipView, UnsupportedAdmin};
use crate::admin_proto::membership_admin_server::MembershipAdminServer;
let admin: Arc<dyn MembershipAdmin> = Arc::new(UnsupportedAdmin::new(MembershipView {
members: Vec::new(),
leader: None,
}));
let service = MembershipAdminServer::new(AdminServiceImpl::new(admin));
let incoming = futures::stream::iter(vec![Err::<tokio::net::TcpStream, std::io::Error>(
std::io::Error::other("accept failed"),
)]);
let fatal = FatalSignal::new();
let mut handle =
TransportHandle::spawn_supervised("admin server", fatal.clone(), move |shutdown| {
tonic::transport::Server::builder()
.add_service(service)
.serve_with_incoming_shutdown(incoming, shutdown)
});
assert_eq!(fatal.tripped().await, "admin server");
handle.shutdown().await;
}
}