1use crate::config::RuntimeConfig;
11use crate::error::HealthError;
12use crate::health::{HealthCheck, HealthChecker};
13use crate::registry::ServiceRegistry;
14use crate::signal::{CompositeSignal, CtrlCSignal, ShutdownSignal, UnixSignal, UnixSignalKind};
15use crate::state::StateTracker;
16use crate::task::{SpawnTask, Task, TaskManager};
17use anyhow::Result;
18use std::net::SocketAddr;
19use std::sync::Arc;
20use tokio::sync::{mpsc, oneshot};
21use tokio::task::JoinHandle;
22use tracing::{error, info, warn};
23
24struct TaskFailureHealthCheck {
26 tracker: Arc<StateTracker>,
27}
28
29impl TaskFailureHealthCheck {
30 fn new(tracker: Arc<StateTracker>) -> Self {
31 Self { tracker }
32 }
33}
34
35impl HealthCheck for TaskFailureHealthCheck {
36 fn check(
37 &self,
38 ) -> std::pin::Pin<
39 Box<dyn std::future::Future<Output = std::result::Result<(), HealthError>> + Send + '_>,
40 > {
41 Box::pin(async move {
42 if self.tracker.has_failures().await {
43 let failed = self.tracker.get_failed_tasks().await;
44 return Err(HealthError::CheckFailed {
45 name: self.name().to_string(),
46 reason: format!("failed tasks detected: {:?}", failed),
47 });
48 }
49 Ok(())
50 })
51 }
52
53 fn name(&self) -> &str {
54 "task-failure-monitor"
55 }
56}
57
58#[derive(Debug, Clone, Copy, PartialEq, Eq)]
60pub enum HealthFailureAction {
61 LogOnly,
63 GracefulShutdown,
65}
66
67struct HealthMonitorHandle {
68 stop_tx: oneshot::Sender<()>,
69 join_handle: JoinHandle<()>,
70 failure_rx: Option<mpsc::UnboundedReceiver<String>>,
71}
72
73pub struct ServiceRuntime {
121 service_name: String,
123 service_address: Option<SocketAddr>,
125 task_manager: TaskManager,
127 registry: Option<Box<dyn ServiceRegistry>>,
129 config: RuntimeConfig,
131 health_checker: Option<HealthChecker>,
133 health_failure_action: HealthFailureAction,
135}
136
137impl ServiceRuntime {
138 pub fn new(service_name: impl Into<String>) -> Self {
152 Self {
153 service_name: service_name.into(),
154 service_address: None,
155 task_manager: TaskManager::new(),
156 registry: None,
157 config: RuntimeConfig::default(),
158 health_checker: None,
159 health_failure_action: HealthFailureAction::LogOnly,
160 }
161 }
162
163 pub fn simple() -> Self {
177 Self {
178 service_name: "simple-runtime".to_string(),
179 service_address: None,
180 task_manager: TaskManager::new(),
181 registry: None,
182 config: RuntimeConfig::default(),
183 health_checker: None,
184 health_failure_action: HealthFailureAction::LogOnly,
185 }
186 }
187
188 pub fn mq_consumer() -> Self {
202 Self::simple()
203 }
204
205 pub fn tasks() -> Self {
219 Self::simple()
220 }
221
222 pub fn with_address(mut self, address: SocketAddr) -> Self {
228 self.service_address = Some(address);
229 self
230 }
231
232 pub fn with_config(mut self, config: RuntimeConfig) -> Self {
234 self.config = config;
235 self
236 }
237
238 pub fn with_registry(mut self, registry: Box<dyn ServiceRegistry>) -> Self {
240 self.registry = Some(registry);
241 self
242 }
243
244 pub fn with_health_checker(mut self, checker: HealthChecker) -> Self {
246 self.health_checker = Some(checker);
247 self
248 }
249
250 pub fn add_health_check(mut self, check: Arc<dyn HealthCheck>) -> Self {
252 if let Some(checker) = &mut self.health_checker {
253 checker.add_check(check);
254 } else {
255 let mut checker = HealthChecker::new()
256 .with_failure_threshold(self.config.health_check.failure_threshold);
257 checker.add_check(check);
258 self.health_checker = Some(checker);
259 }
260 self
261 }
262
263 pub fn with_health_failure_action(mut self, action: HealthFailureAction) -> Self {
265 self.health_failure_action = action;
266 self
267 }
268
269 pub fn add_task(mut self, task: Box<dyn Task>) -> Self {
275 self.task_manager.add_task(task);
276 self
277 }
278
279 pub fn add_spawn<Fut>(mut self, name: impl Into<String>, future: Fut) -> Self
295 where
296 Fut: std::future::Future<Output = Result<(), Box<dyn std::error::Error + Send + Sync>>>
297 + Send
298 + 'static,
299 {
300 self.task_manager
301 .add_task(Box::new(SpawnTask::new(name, future)));
302 self
303 }
304
305 pub fn add_spawn_with_deps<Fut>(
307 mut self,
308 name: impl Into<String>,
309 future: Fut,
310 dependencies: Vec<String>,
311 ) -> Self
312 where
313 Fut: std::future::Future<Output = Result<(), Box<dyn std::error::Error + Send + Sync>>>
314 + Send
315 + 'static,
316 {
317 self.task_manager.add_task(Box::new(
318 SpawnTask::new(name, future).with_dependencies(dependencies),
319 ));
320 self
321 }
322
323 pub fn add_spawn_with_shutdown<F, Fut>(mut self, name: impl Into<String>, future_fn: F) -> Self
325 where
326 F: FnOnce(oneshot::Receiver<()>) -> Fut + Send + 'static,
327 Fut: std::future::Future<Output = Result<(), Box<dyn std::error::Error + Send + Sync>>>
328 + Send
329 + 'static,
330 {
331 self.task_manager
332 .add_task(Box::new(SpawnTask::with_shutdown(name, future_fn)));
333 self
334 }
335
336 pub fn state_tracker(&self) -> Arc<StateTracker> {
338 self.task_manager.state_tracker()
339 }
340
341 fn start_health_monitor(&mut self) -> Option<HealthMonitorHandle> {
343 if !self.config.health_check.enabled {
344 return None;
345 }
346
347 let mut checker = self.health_checker.take().unwrap_or_else(|| {
348 let mut default_checker = HealthChecker::new()
349 .with_failure_threshold(self.config.health_check.failure_threshold);
350 default_checker.add_check(Arc::new(TaskFailureHealthCheck::new(
351 self.task_manager.state_tracker(),
352 )));
353 default_checker
354 });
355
356 if checker.check_count() == 0 {
357 checker.add_check(Arc::new(TaskFailureHealthCheck::new(
358 self.task_manager.state_tracker(),
359 )));
360 }
361
362 let mut failure_rx = None;
363 if self.health_failure_action == HealthFailureAction::GracefulShutdown {
364 let (failure_tx, rx) = mpsc::unbounded_channel::<String>();
365 checker = checker.with_on_failure(Arc::new(move |check_name: &str| {
366 let _ = failure_tx.send(check_name.to_string());
367 }));
368 failure_rx = Some(rx);
369 }
370
371 let service_name = self.service_name.clone();
372 let interval = self.config.health_check.interval;
373 let timeout = self.config.health_check.timeout;
374 let (stop_tx, mut stop_rx) = oneshot::channel::<()>();
375
376 let handle = tokio::spawn(async move {
377 let mut ticker = tokio::time::interval(interval);
378 ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
379
380 info!(
381 service_name = %service_name,
382 interval_ms = interval.as_millis() as u64,
383 timeout_ms = timeout.as_millis() as u64,
384 check_count = checker.check_count(),
385 "Health monitor started"
386 );
387
388 loop {
389 tokio::select! {
390 _ = &mut stop_rx => {
391 info!(service_name = %service_name, "Health monitor stopped");
392 break;
393 }
394 _ = ticker.tick() => {
395 match tokio::time::timeout(timeout, checker.check_all()).await {
396 Ok(results) => {
397 let unhealthy: Vec<_> = results.into_iter().filter(|r| !r.healthy).collect();
398 if !unhealthy.is_empty() {
399 let names: Vec<_> = unhealthy.into_iter().map(|r| r.name).collect();
400 warn!(
401 service_name = %service_name,
402 unhealthy_checks = ?names,
403 "Health monitor detected unhealthy checks"
404 );
405 }
406 }
407 Err(_) => {
408 warn!(
409 service_name = %service_name,
410 timeout_ms = timeout.as_millis() as u64,
411 "Health monitor round timed out"
412 );
413 }
414 }
415 }
416 }
417 }
418 });
419
420 Some(HealthMonitorHandle {
421 stop_tx,
422 join_handle: handle,
423 failure_rx,
424 })
425 }
426
427 pub async fn run(self) -> Result<()> {
435 self.run_with_signals(vec![]).await
436 }
437
438 pub async fn run_with_signals(
444 mut self,
445 mut signals: Vec<Box<dyn ShutdownSignal>>,
446 ) -> Result<()> {
447 info!(
448 service_name = %self.service_name,
449 task_count = self.task_manager.task_count(),
450 "🚀 Starting service runtime"
451 );
452
453 if signals.is_empty() {
455 signals.push(Box::new(CtrlCSignal::new()));
457
458 #[cfg(target_family = "unix")]
459 signals.push(Box::new(UnixSignal::new(UnixSignalKind::Terminate)));
460 }
461
462 let mut shutdown_signal = CompositeSignal::from_signals(signals);
463
464 let (join_set, shutdown_txs) = self
466 .task_manager
467 .start_all()
468 .await
469 .map_err(|e| anyhow::anyhow!("Failed to start tasks: {}", e))?;
470
471 self.task_manager
473 .wait_for_ready()
474 .await
475 .map_err(|e| anyhow::anyhow!("Failed to wait for tasks ready: {}", e))?;
476
477 let mut health_monitor = self.start_health_monitor();
479
480 info!("Waiting for shutdown signal...");
482 if let Some(monitor) = health_monitor.as_mut() {
483 if let Some(failure_rx) = monitor.failure_rx.as_mut() {
484 tokio::select! {
485 _ = shutdown_signal.wait() => {
486 info!("Shutdown signal received");
487 }
488 failed = failure_rx.recv() => {
489 warn!(failed_check = ?failed, "Health check threshold exceeded, triggering graceful shutdown");
490 }
491 }
492 } else {
493 shutdown_signal.wait().await;
494 info!("Shutdown signal received");
495 }
496 } else {
497 shutdown_signal.wait().await;
498 info!("Shutdown signal received");
499 }
500
501 if let Some(monitor) = health_monitor {
503 let _ = monitor.stop_tx.send(());
504 let _ = monitor.join_handle.await;
505 }
506
507 self.task_manager.stop_all(join_set, shutdown_txs).await;
509
510 info!(service_name = %self.service_name, "Service runtime stopped");
511 Ok(())
512 }
513
514 pub async fn run_with_registration<F, Fut>(self, register_fn: F) -> Result<()>
524 where
525 F: FnOnce(SocketAddr) -> Fut,
526 Fut: std::future::Future<
527 Output = Result<
528 Option<Box<dyn ServiceRegistry>>,
529 Box<dyn std::error::Error + Send + Sync>,
530 >,
531 > + Send,
532 {
533 self.run_with_registration_and_signals(register_fn, vec![])
534 .await
535 }
536
537 pub async fn run_with_registration_and_signals<F, Fut>(
542 mut self,
543 register_fn: F,
544 mut signals: Vec<Box<dyn ShutdownSignal>>,
545 ) -> Result<()>
546 where
547 F: FnOnce(SocketAddr) -> Fut,
548 Fut: std::future::Future<
549 Output = Result<
550 Option<Box<dyn ServiceRegistry>>,
551 Box<dyn std::error::Error + Send + Sync>,
552 >,
553 > + Send,
554 {
555 let service_name = self.service_name.clone();
556 let service_address = self.service_address.ok_or_else(|| {
557 anyhow::anyhow!(
558 "Service address is required for service registration. \
559 Use `with_address()` to set the address."
560 )
561 })?;
562
563 info!(
564 service_name = %service_name,
565 address = %service_address,
566 task_count = self.task_manager.task_count(),
567 "🚀 Starting service runtime with registration"
568 );
569
570 if signals.is_empty() {
572 signals.push(Box::new(CtrlCSignal::new()));
573
574 #[cfg(target_family = "unix")]
575 signals.push(Box::new(UnixSignal::new(UnixSignalKind::Terminate)));
576 }
577 let mut shutdown_signal = CompositeSignal::from_signals(signals);
578
579 let (join_set, shutdown_txs) = self
581 .task_manager
582 .start_all()
583 .await
584 .map_err(|e| anyhow::anyhow!("Failed to start tasks: {}", e))?;
585
586 self.task_manager
588 .wait_for_ready()
589 .await
590 .map_err(|e| anyhow::anyhow!("Failed to wait for tasks ready: {}", e))?;
591
592 info!("Registering service...");
594 let registry = match register_fn(service_address).await {
595 Ok(Some(reg)) => {
596 info!("✅ Service registered: {}", service_name);
597 Some(reg)
598 }
599 Ok(None) => {
600 info!("Service registration skipped");
601 None
602 }
603 Err(e) => {
604 error!(error = %e, "❌ Service registration failed");
605
606 self.task_manager.stop_all(join_set, shutdown_txs).await;
608
609 return Err(anyhow::anyhow!("Service registration failed: {}", e));
610 }
611 };
612
613 let mut health_monitor = self.start_health_monitor();
615
616 info!("Waiting for shutdown signal...");
618 if let Some(monitor) = health_monitor.as_mut() {
619 if let Some(failure_rx) = monitor.failure_rx.as_mut() {
620 tokio::select! {
621 _ = shutdown_signal.wait() => {
622 info!("Shutdown signal received");
623 }
624 failed = failure_rx.recv() => {
625 warn!(failed_check = ?failed, "Health check threshold exceeded, triggering graceful shutdown");
626 }
627 }
628 } else {
629 shutdown_signal.wait().await;
630 info!("Shutdown signal received");
631 }
632 } else {
633 shutdown_signal.wait().await;
634 info!("Shutdown signal received");
635 }
636
637 if let Some(monitor) = health_monitor {
639 let _ = monitor.stop_tx.send(());
640 let _ = monitor.join_handle.await;
641 }
642
643 if let Some(mut reg) = registry {
645 info!("Deregistering service...");
646 if let Err(e) = reg.shutdown().await {
647 warn!(error = %e, "⚠️ Failed to deregister service gracefully");
648 } else {
649 info!("✅ Service deregistered");
650 }
651 }
652
653 self.task_manager.stop_all(join_set, shutdown_txs).await;
655
656 info!(service_name = %self.service_name, "Service runtime stopped");
657 Ok(())
658 }
659}
660
661#[cfg(test)]
662mod tests {
663 use super::*;
664
665 #[test]
666 fn test_service_runtime_new() {
667 let runtime = ServiceRuntime::new("test-service");
668 assert_eq!(runtime.service_name, "test-service");
669 }
670
671 #[test]
672 fn test_service_runtime_simple() {
673 let runtime = ServiceRuntime::simple();
674 assert_eq!(runtime.service_name, "simple-runtime");
675 assert!(runtime.service_address.is_none());
676 }
677
678 #[test]
679 fn test_service_runtime_mq_consumer() {
680 let runtime = ServiceRuntime::mq_consumer().add_spawn("kafka-consumer", async { Ok(()) });
681
682 assert_eq!(runtime.task_manager.task_count(), 1);
683 }
684
685 #[test]
686 fn test_service_runtime_tasks() {
687 let runtime = ServiceRuntime::tasks()
688 .add_spawn("task-1", async { Ok(()) })
689 .add_spawn("task-2", async { Ok(()) });
690
691 assert_eq!(runtime.task_manager.task_count(), 2);
692 }
693
694 #[test]
695 fn test_service_runtime_with_address() {
696 let addr: SocketAddr = "0.0.0.0:8080".parse().unwrap();
697 let runtime = ServiceRuntime::new("test-service").with_address(addr);
698
699 assert_eq!(runtime.service_address, Some(addr));
700 }
701
702 #[test]
703 fn test_service_runtime_add_spawn() {
704 let runtime = ServiceRuntime::new("test-service").add_spawn("task-1", async { Ok(()) });
705
706 assert_eq!(runtime.task_manager.task_count(), 1);
707 }
708
709 #[tokio::test]
710 async fn run_with_registration_accepts_custom_shutdown_signal() {
711 use crate::signal::ChannelSignal;
712 use tokio::sync::oneshot;
713
714 let service_address: SocketAddr = "127.0.0.1:0".parse().unwrap();
715 let (shutdown_tx, shutdown_rx) = oneshot::channel();
716 let runtime = ServiceRuntime::new("test-service")
717 .with_address(service_address)
718 .add_spawn_with_shutdown("wait-for-shutdown", |shutdown_rx| async move {
719 let _ = shutdown_rx.await;
720 Ok(())
721 });
722
723 let run = runtime.run_with_registration_and_signals(
724 |addr| async move {
725 assert_eq!(addr, service_address);
726 Ok(None)
727 },
728 vec![Box::new(ChannelSignal::new("test-shutdown", shutdown_rx))],
729 );
730 let stop = async move {
731 tokio::time::sleep(std::time::Duration::from_millis(25)).await;
732 shutdown_tx.send(()).unwrap();
733 };
734
735 let (result, _) = tokio::join!(run, stop);
736 result.unwrap();
737 }
738}