use crate::config::RuntimeConfig;
use crate::error::HealthError;
use crate::health::{HealthCheck, HealthChecker};
use crate::registry::ServiceRegistry;
use crate::signal::{CompositeSignal, CtrlCSignal, ShutdownSignal, UnixSignal, UnixSignalKind};
use crate::state::StateTracker;
use crate::task::{SpawnTask, Task, TaskManager};
use anyhow::Result;
use std::net::SocketAddr;
use std::sync::Arc;
use tokio::sync::{mpsc, oneshot};
use tokio::task::JoinHandle;
use tracing::{error, info, warn};
struct TaskFailureHealthCheck {
tracker: Arc<StateTracker>,
}
impl TaskFailureHealthCheck {
fn new(tracker: Arc<StateTracker>) -> Self {
Self { tracker }
}
}
impl HealthCheck for TaskFailureHealthCheck {
fn check(
&self,
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = std::result::Result<(), HealthError>> + Send + '_>,
> {
Box::pin(async move {
if self.tracker.has_failures().await {
let failed = self.tracker.get_failed_tasks().await;
return Err(HealthError::CheckFailed {
name: self.name().to_string(),
reason: format!("failed tasks detected: {:?}", failed),
});
}
Ok(())
})
}
fn name(&self) -> &str {
"task-failure-monitor"
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum HealthFailureAction {
LogOnly,
GracefulShutdown,
}
struct HealthMonitorHandle {
stop_tx: oneshot::Sender<()>,
join_handle: JoinHandle<()>,
failure_rx: Option<mpsc::UnboundedReceiver<String>>,
}
pub struct ServiceRuntime {
service_name: String,
service_address: Option<SocketAddr>,
task_manager: TaskManager,
registry: Option<Box<dyn ServiceRegistry>>,
config: RuntimeConfig,
health_checker: Option<HealthChecker>,
health_failure_action: HealthFailureAction,
}
impl ServiceRuntime {
pub fn new(service_name: impl Into<String>) -> Self {
Self {
service_name: service_name.into(),
service_address: None,
task_manager: TaskManager::new(),
registry: None,
config: RuntimeConfig::default(),
health_checker: None,
health_failure_action: HealthFailureAction::LogOnly,
}
}
pub fn simple() -> Self {
Self {
service_name: "simple-runtime".to_string(),
service_address: None,
task_manager: TaskManager::new(),
registry: None,
config: RuntimeConfig::default(),
health_checker: None,
health_failure_action: HealthFailureAction::LogOnly,
}
}
pub fn mq_consumer() -> Self {
Self::simple()
}
pub fn tasks() -> Self {
Self::simple()
}
pub fn with_address(mut self, address: SocketAddr) -> Self {
self.service_address = Some(address);
self
}
pub fn with_config(mut self, config: RuntimeConfig) -> Self {
self.config = config;
self
}
pub fn with_registry(mut self, registry: Box<dyn ServiceRegistry>) -> Self {
self.registry = Some(registry);
self
}
pub fn with_health_checker(mut self, checker: HealthChecker) -> Self {
self.health_checker = Some(checker);
self
}
pub fn add_health_check(mut self, check: Arc<dyn HealthCheck>) -> Self {
if let Some(checker) = &mut self.health_checker {
checker.add_check(check);
} else {
let mut checker = HealthChecker::new()
.with_failure_threshold(self.config.health_check.failure_threshold);
checker.add_check(check);
self.health_checker = Some(checker);
}
self
}
pub fn with_health_failure_action(mut self, action: HealthFailureAction) -> Self {
self.health_failure_action = action;
self
}
pub fn add_task(mut self, task: Box<dyn Task>) -> Self {
self.task_manager.add_task(task);
self
}
pub fn add_spawn<Fut>(mut self, name: impl Into<String>, future: Fut) -> Self
where
Fut: std::future::Future<Output = Result<(), Box<dyn std::error::Error + Send + Sync>>>
+ Send
+ 'static,
{
self.task_manager
.add_task(Box::new(SpawnTask::new(name, future)));
self
}
pub fn add_spawn_with_deps<Fut>(
mut self,
name: impl Into<String>,
future: Fut,
dependencies: Vec<String>,
) -> Self
where
Fut: std::future::Future<Output = Result<(), Box<dyn std::error::Error + Send + Sync>>>
+ Send
+ 'static,
{
self.task_manager.add_task(Box::new(
SpawnTask::new(name, future).with_dependencies(dependencies),
));
self
}
pub fn add_spawn_with_shutdown<F, Fut>(mut self, name: impl Into<String>, future_fn: F) -> Self
where
F: FnOnce(oneshot::Receiver<()>) -> Fut + Send + 'static,
Fut: std::future::Future<Output = Result<(), Box<dyn std::error::Error + Send + Sync>>>
+ Send
+ 'static,
{
self.task_manager
.add_task(Box::new(SpawnTask::with_shutdown(name, future_fn)));
self
}
pub fn state_tracker(&self) -> Arc<StateTracker> {
self.task_manager.state_tracker()
}
fn start_health_monitor(&mut self) -> Option<HealthMonitorHandle> {
if !self.config.health_check.enabled {
return None;
}
let mut checker = self.health_checker.take().unwrap_or_else(|| {
let mut default_checker = HealthChecker::new()
.with_failure_threshold(self.config.health_check.failure_threshold);
default_checker.add_check(Arc::new(TaskFailureHealthCheck::new(
self.task_manager.state_tracker(),
)));
default_checker
});
if checker.check_count() == 0 {
checker.add_check(Arc::new(TaskFailureHealthCheck::new(
self.task_manager.state_tracker(),
)));
}
let mut failure_rx = None;
if self.health_failure_action == HealthFailureAction::GracefulShutdown {
let (failure_tx, rx) = mpsc::unbounded_channel::<String>();
checker = checker.with_on_failure(Arc::new(move |check_name: &str| {
let _ = failure_tx.send(check_name.to_string());
}));
failure_rx = Some(rx);
}
let service_name = self.service_name.clone();
let interval = self.config.health_check.interval;
let timeout = self.config.health_check.timeout;
let (stop_tx, mut stop_rx) = oneshot::channel::<()>();
let handle = tokio::spawn(async move {
let mut ticker = tokio::time::interval(interval);
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
info!(
service_name = %service_name,
interval_ms = interval.as_millis() as u64,
timeout_ms = timeout.as_millis() as u64,
check_count = checker.check_count(),
"Health monitor started"
);
loop {
tokio::select! {
_ = &mut stop_rx => {
info!(service_name = %service_name, "Health monitor stopped");
break;
}
_ = ticker.tick() => {
match tokio::time::timeout(timeout, checker.check_all()).await {
Ok(results) => {
let unhealthy: Vec<_> = results.into_iter().filter(|r| !r.healthy).collect();
if !unhealthy.is_empty() {
let names: Vec<_> = unhealthy.into_iter().map(|r| r.name).collect();
warn!(
service_name = %service_name,
unhealthy_checks = ?names,
"Health monitor detected unhealthy checks"
);
}
}
Err(_) => {
warn!(
service_name = %service_name,
timeout_ms = timeout.as_millis() as u64,
"Health monitor round timed out"
);
}
}
}
}
}
});
Some(HealthMonitorHandle {
stop_tx,
join_handle: handle,
failure_rx,
})
}
pub async fn run(self) -> Result<()> {
self.run_with_signals(vec![]).await
}
pub async fn run_with_signals(
mut self,
mut signals: Vec<Box<dyn ShutdownSignal>>,
) -> Result<()> {
info!(
service_name = %self.service_name,
task_count = self.task_manager.task_count(),
"🚀 Starting service runtime"
);
if signals.is_empty() {
signals.push(Box::new(CtrlCSignal::new()));
#[cfg(target_family = "unix")]
signals.push(Box::new(UnixSignal::new(UnixSignalKind::Terminate)));
}
let mut shutdown_signal = CompositeSignal::from_signals(signals);
let (join_set, shutdown_txs) = self
.task_manager
.start_all()
.await
.map_err(|e| anyhow::anyhow!("Failed to start tasks: {}", e))?;
self.task_manager
.wait_for_ready()
.await
.map_err(|e| anyhow::anyhow!("Failed to wait for tasks ready: {}", e))?;
let mut health_monitor = self.start_health_monitor();
info!("Waiting for shutdown signal...");
if let Some(monitor) = health_monitor.as_mut() {
if let Some(failure_rx) = monitor.failure_rx.as_mut() {
tokio::select! {
_ = shutdown_signal.wait() => {
info!("Shutdown signal received");
}
failed = failure_rx.recv() => {
warn!(failed_check = ?failed, "Health check threshold exceeded, triggering graceful shutdown");
}
}
} else {
shutdown_signal.wait().await;
info!("Shutdown signal received");
}
} else {
shutdown_signal.wait().await;
info!("Shutdown signal received");
}
if let Some(monitor) = health_monitor {
let _ = monitor.stop_tx.send(());
let _ = monitor.join_handle.await;
}
self.task_manager.stop_all(join_set, shutdown_txs).await;
info!(service_name = %self.service_name, "Service runtime stopped");
Ok(())
}
pub async fn run_with_registration<F, Fut>(self, register_fn: F) -> Result<()>
where
F: FnOnce(SocketAddr) -> Fut,
Fut: std::future::Future<
Output = Result<
Option<Box<dyn ServiceRegistry>>,
Box<dyn std::error::Error + Send + Sync>,
>,
> + Send,
{
self.run_with_registration_and_signals(register_fn, vec![])
.await
}
pub async fn run_with_registration_and_signals<F, Fut>(
mut self,
register_fn: F,
mut signals: Vec<Box<dyn ShutdownSignal>>,
) -> Result<()>
where
F: FnOnce(SocketAddr) -> Fut,
Fut: std::future::Future<
Output = Result<
Option<Box<dyn ServiceRegistry>>,
Box<dyn std::error::Error + Send + Sync>,
>,
> + Send,
{
let service_name = self.service_name.clone();
let service_address = self.service_address.ok_or_else(|| {
anyhow::anyhow!(
"Service address is required for service registration. \
Use `with_address()` to set the address."
)
})?;
info!(
service_name = %service_name,
address = %service_address,
task_count = self.task_manager.task_count(),
"🚀 Starting service runtime with registration"
);
if signals.is_empty() {
signals.push(Box::new(CtrlCSignal::new()));
#[cfg(target_family = "unix")]
signals.push(Box::new(UnixSignal::new(UnixSignalKind::Terminate)));
}
let mut shutdown_signal = CompositeSignal::from_signals(signals);
let (join_set, shutdown_txs) = self
.task_manager
.start_all()
.await
.map_err(|e| anyhow::anyhow!("Failed to start tasks: {}", e))?;
self.task_manager
.wait_for_ready()
.await
.map_err(|e| anyhow::anyhow!("Failed to wait for tasks ready: {}", e))?;
info!("Registering service...");
let registry = match register_fn(service_address).await {
Ok(Some(reg)) => {
info!("✅ Service registered: {}", service_name);
Some(reg)
}
Ok(None) => {
info!("Service registration skipped");
None
}
Err(e) => {
error!(error = %e, "❌ Service registration failed");
self.task_manager.stop_all(join_set, shutdown_txs).await;
return Err(anyhow::anyhow!("Service registration failed: {}", e));
}
};
let mut health_monitor = self.start_health_monitor();
info!("Waiting for shutdown signal...");
if let Some(monitor) = health_monitor.as_mut() {
if let Some(failure_rx) = monitor.failure_rx.as_mut() {
tokio::select! {
_ = shutdown_signal.wait() => {
info!("Shutdown signal received");
}
failed = failure_rx.recv() => {
warn!(failed_check = ?failed, "Health check threshold exceeded, triggering graceful shutdown");
}
}
} else {
shutdown_signal.wait().await;
info!("Shutdown signal received");
}
} else {
shutdown_signal.wait().await;
info!("Shutdown signal received");
}
if let Some(monitor) = health_monitor {
let _ = monitor.stop_tx.send(());
let _ = monitor.join_handle.await;
}
if let Some(mut reg) = registry {
info!("Deregistering service...");
if let Err(e) = reg.shutdown().await {
warn!(error = %e, "⚠️ Failed to deregister service gracefully");
} else {
info!("✅ Service deregistered");
}
}
self.task_manager.stop_all(join_set, shutdown_txs).await;
info!(service_name = %self.service_name, "Service runtime stopped");
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_service_runtime_new() {
let runtime = ServiceRuntime::new("test-service");
assert_eq!(runtime.service_name, "test-service");
}
#[test]
fn test_service_runtime_simple() {
let runtime = ServiceRuntime::simple();
assert_eq!(runtime.service_name, "simple-runtime");
assert!(runtime.service_address.is_none());
}
#[test]
fn test_service_runtime_mq_consumer() {
let runtime = ServiceRuntime::mq_consumer().add_spawn("kafka-consumer", async { Ok(()) });
assert_eq!(runtime.task_manager.task_count(), 1);
}
#[test]
fn test_service_runtime_tasks() {
let runtime = ServiceRuntime::tasks()
.add_spawn("task-1", async { Ok(()) })
.add_spawn("task-2", async { Ok(()) });
assert_eq!(runtime.task_manager.task_count(), 2);
}
#[test]
fn test_service_runtime_with_address() {
let addr: SocketAddr = "0.0.0.0:8080".parse().unwrap();
let runtime = ServiceRuntime::new("test-service").with_address(addr);
assert_eq!(runtime.service_address, Some(addr));
}
#[test]
fn test_service_runtime_add_spawn() {
let runtime = ServiceRuntime::new("test-service").add_spawn("task-1", async { Ok(()) });
assert_eq!(runtime.task_manager.task_count(), 1);
}
#[tokio::test]
async fn run_with_registration_accepts_custom_shutdown_signal() {
use crate::signal::ChannelSignal;
use tokio::sync::oneshot;
let service_address: SocketAddr = "127.0.0.1:0".parse().unwrap();
let (shutdown_tx, shutdown_rx) = oneshot::channel();
let runtime = ServiceRuntime::new("test-service")
.with_address(service_address)
.add_spawn_with_shutdown("wait-for-shutdown", |shutdown_rx| async move {
let _ = shutdown_rx.await;
Ok(())
});
let run = runtime.run_with_registration_and_signals(
|addr| async move {
assert_eq!(addr, service_address);
Ok(None)
},
vec![Box::new(ChannelSignal::new("test-shutdown", shutdown_rx))],
);
let stop = async move {
tokio::time::sleep(std::time::Duration::from_millis(25)).await;
shutdown_tx.send(()).unwrap();
};
let (result, _) = tokio::join!(run, stop);
result.unwrap();
}
}