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.task_manager.set_config(config.clone());
235 self.config = config;
236 self
237 }
238
239 pub fn with_registry(mut self, registry: Box<dyn ServiceRegistry>) -> Self {
241 self.registry = Some(registry);
242 self
243 }
244
245 pub fn with_health_checker(mut self, checker: HealthChecker) -> Self {
247 self.health_checker = Some(checker);
248 self
249 }
250
251 pub fn add_health_check(mut self, check: Arc<dyn HealthCheck>) -> Self {
253 if let Some(checker) = &mut self.health_checker {
254 checker.add_check(check);
255 } else {
256 let mut checker = HealthChecker::new()
257 .with_failure_threshold(self.config.health_check.failure_threshold);
258 checker.add_check(check);
259 self.health_checker = Some(checker);
260 }
261 self
262 }
263
264 pub fn with_health_failure_action(mut self, action: HealthFailureAction) -> Self {
266 self.health_failure_action = action;
267 self
268 }
269
270 pub fn add_task(mut self, task: Box<dyn Task>) -> Self {
276 self.task_manager.add_task(task);
277 self
278 }
279
280 pub fn add_spawn<Fut>(mut self, name: impl Into<String>, future: Fut) -> Self
296 where
297 Fut: std::future::Future<Output = Result<(), Box<dyn std::error::Error + Send + Sync>>>
298 + Send
299 + 'static,
300 {
301 self.task_manager
302 .add_task(Box::new(SpawnTask::new(name, future)));
303 self
304 }
305
306 pub fn add_spawn_with_deps<Fut>(
308 mut self,
309 name: impl Into<String>,
310 future: Fut,
311 dependencies: Vec<String>,
312 ) -> Self
313 where
314 Fut: std::future::Future<Output = Result<(), Box<dyn std::error::Error + Send + Sync>>>
315 + Send
316 + 'static,
317 {
318 self.task_manager.add_task(Box::new(
319 SpawnTask::new(name, future).with_dependencies(dependencies),
320 ));
321 self
322 }
323
324 pub fn add_spawn_with_shutdown<F, Fut>(mut self, name: impl Into<String>, future_fn: F) -> Self
326 where
327 F: FnOnce(oneshot::Receiver<()>) -> Fut + Send + 'static,
328 Fut: std::future::Future<Output = Result<(), Box<dyn std::error::Error + Send + Sync>>>
329 + Send
330 + 'static,
331 {
332 self.task_manager
333 .add_task(Box::new(SpawnTask::with_shutdown(name, future_fn)));
334 self
335 }
336
337 pub fn state_tracker(&self) -> Arc<StateTracker> {
339 self.task_manager.state_tracker()
340 }
341
342 fn start_health_monitor(&mut self) -> Option<HealthMonitorHandle> {
344 if !self.config.health_check.enabled {
345 return None;
346 }
347
348 let mut checker = self.health_checker.take().unwrap_or_else(|| {
349 let mut default_checker = HealthChecker::new()
350 .with_failure_threshold(self.config.health_check.failure_threshold);
351 default_checker.add_check(Arc::new(TaskFailureHealthCheck::new(
352 self.task_manager.state_tracker(),
353 )));
354 default_checker
355 });
356
357 if checker.check_count() == 0 {
358 checker.add_check(Arc::new(TaskFailureHealthCheck::new(
359 self.task_manager.state_tracker(),
360 )));
361 }
362
363 let mut failure_rx = None;
364 if self.health_failure_action == HealthFailureAction::GracefulShutdown {
365 let (failure_tx, rx) = mpsc::unbounded_channel::<String>();
366 checker = checker.with_on_failure(Arc::new(move |check_name: &str| {
367 let _ = failure_tx.send(check_name.to_string());
368 }));
369 failure_rx = Some(rx);
370 }
371
372 let service_name = self.service_name.clone();
373 let interval = self.config.health_check.interval;
374 let timeout = self.config.health_check.timeout;
375 let (stop_tx, mut stop_rx) = oneshot::channel::<()>();
376
377 let handle = tokio::spawn(async move {
378 let mut ticker = tokio::time::interval(interval);
379 ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
380
381 info!(
382 service_name = %service_name,
383 interval_ms = interval.as_millis() as u64,
384 timeout_ms = timeout.as_millis() as u64,
385 check_count = checker.check_count(),
386 "Health monitor started"
387 );
388
389 loop {
390 tokio::select! {
391 _ = &mut stop_rx => {
392 info!(service_name = %service_name, "Health monitor stopped");
393 break;
394 }
395 _ = ticker.tick() => {
396 match tokio::time::timeout(timeout, checker.check_all()).await {
397 Ok(results) => {
398 let unhealthy: Vec<_> = results.into_iter().filter(|r| !r.healthy).collect();
399 if !unhealthy.is_empty() {
400 let names: Vec<_> = unhealthy.into_iter().map(|r| r.name).collect();
401 warn!(
402 service_name = %service_name,
403 unhealthy_checks = ?names,
404 "Health monitor detected unhealthy checks"
405 );
406 }
407 }
408 Err(_) => {
409 warn!(
410 service_name = %service_name,
411 timeout_ms = timeout.as_millis() as u64,
412 "Health monitor round timed out"
413 );
414 }
415 }
416 }
417 }
418 }
419 });
420
421 Some(HealthMonitorHandle {
422 stop_tx,
423 join_handle: handle,
424 failure_rx,
425 })
426 }
427
428 pub async fn run(self) -> Result<()> {
436 self.run_with_signals(vec![]).await
437 }
438
439 pub async fn run_with_signals(
445 mut self,
446 mut signals: Vec<Box<dyn ShutdownSignal>>,
447 ) -> Result<()> {
448 info!(
449 service_name = %self.service_name,
450 task_count = self.task_manager.task_count(),
451 "🚀 Starting service runtime"
452 );
453
454 if signals.is_empty() {
456 signals.push(Box::new(CtrlCSignal::new()));
458
459 #[cfg(target_family = "unix")]
460 signals.push(Box::new(UnixSignal::new(UnixSignalKind::Terminate)));
461 }
462
463 let mut shutdown_signal = CompositeSignal::from_signals(signals);
464
465 let (join_set, shutdown_txs) = self
467 .task_manager
468 .start_all()
469 .await
470 .map_err(|e| anyhow::anyhow!("Failed to start tasks: {}", e))?;
471
472 self.task_manager
474 .wait_for_ready()
475 .await
476 .map_err(|e| anyhow::anyhow!("Failed to wait for tasks ready: {}", e))?;
477
478 let mut health_monitor = self.start_health_monitor();
480
481 info!("Waiting for shutdown signal...");
483 if let Some(monitor) = health_monitor.as_mut() {
484 if let Some(failure_rx) = monitor.failure_rx.as_mut() {
485 tokio::select! {
486 _ = shutdown_signal.wait() => {
487 info!("Shutdown signal received");
488 }
489 failed = failure_rx.recv() => {
490 warn!(failed_check = ?failed, "Health check threshold exceeded, triggering graceful shutdown");
491 }
492 }
493 } else {
494 shutdown_signal.wait().await;
495 info!("Shutdown signal received");
496 }
497 } else {
498 shutdown_signal.wait().await;
499 info!("Shutdown signal received");
500 }
501
502 if let Some(monitor) = health_monitor {
504 let _ = monitor.stop_tx.send(());
505 let _ = monitor.join_handle.await;
506 }
507
508 self.task_manager.stop_all(join_set, shutdown_txs).await;
510
511 info!(service_name = %self.service_name, "Service runtime stopped");
512 Ok(())
513 }
514
515 pub async fn run_with_registration<F, Fut>(self, register_fn: F) -> Result<()>
525 where
526 F: FnOnce(SocketAddr) -> Fut,
527 Fut: std::future::Future<
528 Output = Result<
529 Option<Box<dyn ServiceRegistry>>,
530 Box<dyn std::error::Error + Send + Sync>,
531 >,
532 > + Send,
533 {
534 self.run_with_registration_and_signals(register_fn, vec![])
535 .await
536 }
537
538 pub async fn run_with_registration_and_signals<F, Fut>(
543 mut self,
544 register_fn: F,
545 mut signals: Vec<Box<dyn ShutdownSignal>>,
546 ) -> Result<()>
547 where
548 F: FnOnce(SocketAddr) -> Fut,
549 Fut: std::future::Future<
550 Output = Result<
551 Option<Box<dyn ServiceRegistry>>,
552 Box<dyn std::error::Error + Send + Sync>,
553 >,
554 > + Send,
555 {
556 let service_name = self.service_name.clone();
557 let service_address = self.service_address.ok_or_else(|| {
558 anyhow::anyhow!(
559 "Service address is required for service registration. \
560 Use `with_address()` to set the address."
561 )
562 })?;
563
564 info!(
565 service_name = %service_name,
566 address = %service_address,
567 task_count = self.task_manager.task_count(),
568 "🚀 Starting service runtime with registration"
569 );
570
571 if signals.is_empty() {
573 signals.push(Box::new(CtrlCSignal::new()));
574
575 #[cfg(target_family = "unix")]
576 signals.push(Box::new(UnixSignal::new(UnixSignalKind::Terminate)));
577 }
578 let mut shutdown_signal = CompositeSignal::from_signals(signals);
579
580 let (join_set, shutdown_txs) = self
582 .task_manager
583 .start_all()
584 .await
585 .map_err(|e| anyhow::anyhow!("Failed to start tasks: {}", e))?;
586
587 self.task_manager
589 .wait_for_ready()
590 .await
591 .map_err(|e| anyhow::anyhow!("Failed to wait for tasks ready: {}", e))?;
592
593 info!("Registering service...");
595 let registry = match register_fn(service_address).await {
596 Ok(Some(reg)) => {
597 info!("✅ Service registered: {}", service_name);
598 Some(reg)
599 }
600 Ok(None) => {
601 info!("Service registration skipped");
602 None
603 }
604 Err(e) => {
605 error!(error = %e, "❌ Service registration failed");
606
607 self.task_manager.stop_all(join_set, shutdown_txs).await;
609
610 return Err(anyhow::anyhow!("Service registration failed: {}", e));
611 }
612 };
613
614 let mut health_monitor = self.start_health_monitor();
616
617 info!("Waiting for shutdown signal...");
619 if let Some(monitor) = health_monitor.as_mut() {
620 if let Some(failure_rx) = monitor.failure_rx.as_mut() {
621 tokio::select! {
622 _ = shutdown_signal.wait() => {
623 info!("Shutdown signal received");
624 }
625 failed = failure_rx.recv() => {
626 warn!(failed_check = ?failed, "Health check threshold exceeded, triggering graceful shutdown");
627 }
628 }
629 } else {
630 shutdown_signal.wait().await;
631 info!("Shutdown signal received");
632 }
633 } else {
634 shutdown_signal.wait().await;
635 info!("Shutdown signal received");
636 }
637
638 if let Some(monitor) = health_monitor {
640 let _ = monitor.stop_tx.send(());
641 let _ = monitor.join_handle.await;
642 }
643
644 if let Some(mut reg) = registry {
646 info!("Deregistering service...");
647 if let Err(e) = reg.shutdown().await {
648 warn!(error = %e, "⚠️ Failed to deregister service gracefully");
649 } else {
650 info!("✅ Service deregistered");
651 }
652 }
653
654 self.task_manager.stop_all(join_set, shutdown_txs).await;
656
657 info!(service_name = %self.service_name, "Service runtime stopped");
658 Ok(())
659 }
660}
661
662#[cfg(test)]
663mod tests {
664 use super::*;
665 use std::time::Duration;
666
667 #[test]
668 fn test_service_runtime_new() {
669 let runtime = ServiceRuntime::new("test-service");
670 assert_eq!(runtime.service_name, "test-service");
671 }
672
673 #[test]
674 fn test_service_runtime_simple() {
675 let runtime = ServiceRuntime::simple();
676 assert_eq!(runtime.service_name, "simple-runtime");
677 assert!(runtime.service_address.is_none());
678 }
679
680 #[test]
681 fn test_service_runtime_mq_consumer() {
682 let runtime = ServiceRuntime::mq_consumer().add_spawn("kafka-consumer", async { Ok(()) });
683
684 assert_eq!(runtime.task_manager.task_count(), 1);
685 }
686
687 #[test]
688 fn test_service_runtime_tasks() {
689 let runtime = ServiceRuntime::tasks()
690 .add_spawn("task-1", async { Ok(()) })
691 .add_spawn("task-2", async { Ok(()) });
692
693 assert_eq!(runtime.task_manager.task_count(), 2);
694 }
695
696 #[test]
697 fn test_service_runtime_with_address() {
698 let addr: SocketAddr = "0.0.0.0:8080".parse().unwrap();
699 let runtime = ServiceRuntime::new("test-service").with_address(addr);
700
701 assert_eq!(runtime.service_address, Some(addr));
702 }
703
704 #[test]
705 fn test_service_runtime_with_config_updates_task_manager() {
706 let config = RuntimeConfig::new().with_shutdown_timeout(Duration::from_millis(25));
707 let runtime = ServiceRuntime::new("test-service")
708 .add_spawn("task-1", async { Ok(()) })
709 .with_config(config);
710
711 assert_eq!(runtime.task_manager.task_count(), 1);
712 assert_eq!(
713 runtime.task_manager.shutdown_timeout(),
714 Duration::from_millis(25)
715 );
716 }
717
718 #[test]
719 fn test_service_runtime_add_spawn() {
720 let runtime = ServiceRuntime::new("test-service").add_spawn("task-1", async { Ok(()) });
721
722 assert_eq!(runtime.task_manager.task_count(), 1);
723 }
724
725 #[tokio::test]
726 async fn run_with_registration_accepts_custom_shutdown_signal() {
727 use crate::signal::ChannelSignal;
728 use tokio::sync::oneshot;
729
730 let service_address: SocketAddr = "127.0.0.1:0".parse().unwrap();
731 let (shutdown_tx, shutdown_rx) = oneshot::channel();
732 let runtime = ServiceRuntime::new("test-service")
733 .with_address(service_address)
734 .add_spawn_with_shutdown("wait-for-shutdown", |shutdown_rx| async move {
735 let _ = shutdown_rx.await;
736 Ok(())
737 });
738
739 let run = runtime.run_with_registration_and_signals(
740 |addr| async move {
741 assert_eq!(addr, service_address);
742 Ok(None)
743 },
744 vec![Box::new(ChannelSignal::new("test-shutdown", shutdown_rx))],
745 );
746 let stop = async move {
747 tokio::time::sleep(std::time::Duration::from_millis(25)).await;
748 shutdown_tx.send(()).unwrap();
749 };
750
751 let (result, _) = tokio::join!(run, stop);
752 result.unwrap();
753 }
754}