1use async_trait::async_trait;
11use serde::{Deserialize, Deserializer, Serialize, Serializer};
12use std::collections::HashMap;
13use std::time::{Duration, Instant};
14use tokio::time::timeout;
15use tracing::debug;
16
17use crate::core::error::McpResult;
18use crate::core::retry::CircuitBreakerStats;
19
20mod instant_serde {
22 use super::*;
23 use std::time::SystemTime;
24
25 pub fn serialize<S>(_instant: &Instant, serializer: S) -> Result<S::Ok, S::Error>
26 where
27 S: Serializer,
28 {
29 let system_time = SystemTime::now();
31 let duration_since_epoch = system_time
32 .duration_since(std::time::UNIX_EPOCH)
33 .unwrap_or_default();
34 serializer.serialize_u64(duration_since_epoch.as_millis() as u64)
35 }
36
37 pub fn deserialize<'de, D>(deserializer: D) -> Result<Instant, D::Error>
38 where
39 D: Deserializer<'de>,
40 {
41 let _millis = u64::deserialize(deserializer)?;
43 Ok(Instant::now())
44 }
45}
46
47mod duration_serde {
49 use super::*;
50
51 pub fn serialize<S>(duration: &Duration, serializer: S) -> Result<S::Ok, S::Error>
52 where
53 S: Serializer,
54 {
55 serializer.serialize_u64(duration.as_millis() as u64)
56 }
57
58 pub fn deserialize<'de, D>(deserializer: D) -> Result<Duration, D::Error>
59 where
60 D: Deserializer<'de>,
61 {
62 let millis = u64::deserialize(deserializer)?;
63 Ok(Duration::from_millis(millis))
64 }
65}
66
67#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
69pub enum HealthStatus {
70 Healthy,
72 Degraded,
74 Unhealthy,
76 Unknown,
78}
79
80impl HealthStatus {
81 pub fn is_operational(&self) -> bool {
83 matches!(self, HealthStatus::Healthy | HealthStatus::Degraded)
84 }
85
86 pub fn score(&self) -> u8 {
88 match self {
89 HealthStatus::Healthy => 100,
90 HealthStatus::Degraded => 75,
91 HealthStatus::Unhealthy => 25,
92 HealthStatus::Unknown => 0,
93 }
94 }
95
96 pub fn combine(self, other: HealthStatus) -> HealthStatus {
98 if self.score() < other.score() {
99 self
100 } else {
101 other
102 }
103 }
104}
105
106impl std::fmt::Display for HealthStatus {
107 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
108 match self {
109 HealthStatus::Healthy => write!(f, "healthy"),
110 HealthStatus::Degraded => write!(f, "degraded"),
111 HealthStatus::Unhealthy => write!(f, "unhealthy"),
112 HealthStatus::Unknown => write!(f, "unknown"),
113 }
114 }
115}
116
117#[derive(Debug, Clone, Serialize, Deserialize)]
119pub struct HealthResult {
120 pub status: HealthStatus,
122 pub message: String,
124 pub metadata: HashMap<String, serde_json::Value>,
126 #[serde(with = "instant_serde")]
128 pub timestamp: Instant,
129 #[serde(with = "duration_serde")]
131 pub duration: Duration,
132}
133
134impl HealthResult {
135 pub fn healthy(message: impl Into<String>) -> Self {
137 Self::new(HealthStatus::Healthy, message)
138 }
139
140 pub fn degraded(message: impl Into<String>) -> Self {
142 Self::new(HealthStatus::Degraded, message)
143 }
144
145 pub fn unhealthy(message: impl Into<String>) -> Self {
147 Self::new(HealthStatus::Unhealthy, message)
148 }
149
150 pub fn unknown(message: impl Into<String>) -> Self {
152 Self::new(HealthStatus::Unknown, message)
153 }
154
155 pub fn new(status: HealthStatus, message: impl Into<String>) -> Self {
157 Self {
158 status,
159 message: message.into(),
160 metadata: HashMap::new(),
161 timestamp: Instant::now(),
162 duration: Duration::from_millis(0),
163 }
164 }
165
166 pub fn with_metadata(mut self, key: impl Into<String>, value: serde_json::Value) -> Self {
168 self.metadata.insert(key.into(), value);
169 self
170 }
171
172 pub fn with_duration(mut self, duration: Duration) -> Self {
174 self.duration = duration;
175 self
176 }
177}
178
179#[async_trait]
181pub trait HealthCheck: Send + Sync {
182 fn name(&self) -> &str;
184
185 async fn check(&self) -> HealthResult;
187
188 fn timeout(&self) -> Duration {
190 Duration::from_secs(5)
191 }
192
193 fn is_critical(&self) -> bool {
195 true
196 }
197}
198
199pub struct TransportHealthCheck {
201 name: String,
202 transport_type: String,
203 connection_test: Box<
204 dyn Fn() -> std::pin::Pin<Box<dyn std::future::Future<Output = bool> + Send>> + Send + Sync,
205 >,
206}
207
208impl TransportHealthCheck {
209 pub fn new<F, Fut>(
211 name: impl Into<String>,
212 transport_type: impl Into<String>,
213 connection_test: F,
214 ) -> Self
215 where
216 F: Fn() -> Fut + Send + Sync + 'static,
217 Fut: std::future::Future<Output = bool> + Send + 'static,
218 {
219 Self {
220 name: name.into(),
221 transport_type: transport_type.into(),
222 connection_test: Box::new(move || Box::pin(connection_test())),
223 }
224 }
225}
226
227#[async_trait]
228impl HealthCheck for TransportHealthCheck {
229 fn name(&self) -> &str {
230 &self.name
231 }
232
233 async fn check(&self) -> HealthResult {
234 let start = Instant::now();
235
236 match timeout(self.timeout(), (self.connection_test)()).await {
237 Ok(true) => {
238 HealthResult::healthy(format!("{} transport is connected", self.transport_type))
239 .with_duration(start.elapsed())
240 .with_metadata(
241 "transport_type",
242 serde_json::Value::String(self.transport_type.clone()),
243 )
244 }
245 Ok(false) => HealthResult::unhealthy(format!(
246 "{} transport connection failed",
247 self.transport_type
248 ))
249 .with_duration(start.elapsed())
250 .with_metadata(
251 "transport_type",
252 serde_json::Value::String(self.transport_type.clone()),
253 ),
254 Err(_) => HealthResult::unhealthy(format!(
255 "{} transport health check timed out",
256 self.transport_type
257 ))
258 .with_duration(start.elapsed())
259 .with_metadata(
260 "transport_type",
261 serde_json::Value::String(self.transport_type.clone()),
262 )
263 .with_metadata("timeout", serde_json::Value::Bool(true)),
264 }
265 }
266}
267
268type ProtocolTestFn = Box<
271 dyn Fn() -> std::pin::Pin<Box<dyn std::future::Future<Output = McpResult<()>> + Send>>
272 + Send
273 + Sync,
274>;
275
276pub struct ProtocolHealthCheck {
277 name: String,
278 protocol_test: ProtocolTestFn,
279}
280
281impl ProtocolHealthCheck {
282 pub fn new<F, Fut>(name: impl Into<String>, protocol_test: F) -> Self
284 where
285 F: Fn() -> Fut + Send + Sync + 'static,
286 Fut: std::future::Future<Output = McpResult<()>> + Send + 'static,
287 {
288 Self {
289 name: name.into(),
290 protocol_test: Box::new(move || Box::pin(protocol_test())),
291 }
292 }
293}
294
295#[async_trait]
296impl HealthCheck for ProtocolHealthCheck {
297 fn name(&self) -> &str {
298 &self.name
299 }
300
301 async fn check(&self) -> HealthResult {
302 let start = Instant::now();
303
304 match timeout(self.timeout(), (self.protocol_test)()).await {
305 Ok(Ok(())) => HealthResult::healthy("Protocol communication successful")
306 .with_duration(start.elapsed()),
307 Ok(Err(error)) => {
308 let status = if error.is_recoverable() {
309 HealthStatus::Degraded
310 } else {
311 HealthStatus::Unhealthy
312 };
313
314 HealthResult::new(status, format!("Protocol error: {error}"))
315 .with_duration(start.elapsed())
316 .with_metadata(
317 "error_category",
318 serde_json::Value::String(error.category().to_string()),
319 )
320 .with_metadata(
321 "error_recoverable",
322 serde_json::Value::Bool(error.is_recoverable()),
323 )
324 }
325 Err(_) => HealthResult::unhealthy("Protocol health check timed out")
326 .with_duration(start.elapsed())
327 .with_metadata("timeout", serde_json::Value::Bool(true)),
328 }
329 }
330}
331
332pub struct ResourceHealthCheck {
334 name: String,
335 resource_name: String,
336 availability_test: Box<
337 dyn Fn() -> std::pin::Pin<Box<dyn std::future::Future<Output = bool> + Send>> + Send + Sync,
338 >,
339}
340
341impl ResourceHealthCheck {
342 pub fn new<F, Fut>(
344 name: impl Into<String>,
345 resource_name: impl Into<String>,
346 availability_test: F,
347 ) -> Self
348 where
349 F: Fn() -> Fut + Send + Sync + 'static,
350 Fut: std::future::Future<Output = bool> + Send + 'static,
351 {
352 Self {
353 name: name.into(),
354 resource_name: resource_name.into(),
355 availability_test: Box::new(move || Box::pin(availability_test())),
356 }
357 }
358}
359
360#[async_trait]
361impl HealthCheck for ResourceHealthCheck {
362 fn name(&self) -> &str {
363 &self.name
364 }
365
366 async fn check(&self) -> HealthResult {
367 let start = Instant::now();
368
369 match timeout(self.timeout(), (self.availability_test)()).await {
370 Ok(true) => {
371 HealthResult::healthy(format!("Resource '{}' is available", self.resource_name))
372 .with_duration(start.elapsed())
373 .with_metadata(
374 "resource_name",
375 serde_json::Value::String(self.resource_name.clone()),
376 )
377 }
378 Ok(false) => {
379 HealthResult::unhealthy(format!("Resource '{}' is unavailable", self.resource_name))
380 .with_duration(start.elapsed())
381 .with_metadata(
382 "resource_name",
383 serde_json::Value::String(self.resource_name.clone()),
384 )
385 }
386 Err(_) => HealthResult::unknown(format!(
387 "Resource '{}' health check timed out",
388 self.resource_name
389 ))
390 .with_duration(start.elapsed())
391 .with_metadata(
392 "resource_name",
393 serde_json::Value::String(self.resource_name.clone()),
394 )
395 .with_metadata("timeout", serde_json::Value::Bool(true)),
396 }
397 }
398
399 fn is_critical(&self) -> bool {
400 false }
402}
403
404type StatsGetterFn = Box<
407 dyn Fn() -> std::pin::Pin<
408 Box<dyn std::future::Future<Output = Option<CircuitBreakerStats>> + Send>,
409 > + Send
410 + Sync,
411>;
412
413pub struct CircuitBreakerHealthCheck {
414 name: String,
415 get_stats: StatsGetterFn,
416}
417
418impl CircuitBreakerHealthCheck {
419 pub fn new<F, Fut>(name: impl Into<String>, get_stats: F) -> Self
421 where
422 F: Fn() -> Fut + Send + Sync + 'static,
423 Fut: std::future::Future<Output = Option<CircuitBreakerStats>> + Send + 'static,
424 {
425 Self {
426 name: name.into(),
427 get_stats: Box::new(move || Box::pin(get_stats())),
428 }
429 }
430}
431
432#[async_trait]
433impl HealthCheck for CircuitBreakerHealthCheck {
434 fn name(&self) -> &str {
435 &self.name
436 }
437
438 async fn check(&self) -> HealthResult {
439 let start = Instant::now();
440
441 match (self.get_stats)().await {
442 Some(stats) => {
443 let status = match stats.state {
444 crate::core::retry::CircuitState::Closed => HealthStatus::Healthy,
445 crate::core::retry::CircuitState::HalfOpen => HealthStatus::Degraded,
446 crate::core::retry::CircuitState::Open => HealthStatus::Unhealthy,
447 };
448
449 let message = format!("Circuit breaker state: {:?}", stats.state);
450
451 HealthResult::new(status, message)
452 .with_duration(start.elapsed())
453 .with_metadata(
454 "circuit_state",
455 serde_json::Value::String(format!("{:?}", stats.state)),
456 )
457 .with_metadata(
458 "failure_count",
459 serde_json::Value::from(stats.failure_count),
460 )
461 .with_metadata(
462 "success_count",
463 serde_json::Value::from(stats.success_count),
464 )
465 .with_metadata(
466 "half_open_requests",
467 serde_json::Value::from(stats.half_open_requests),
468 )
469 }
470 None => HealthResult::unknown("Circuit breaker stats not available")
471 .with_duration(start.elapsed()),
472 }
473 }
474
475 fn is_critical(&self) -> bool {
476 false }
478}
479
480#[derive(Debug, Clone, Serialize, Deserialize)]
482pub struct OverallHealth {
483 pub status: HealthStatus,
485 pub checks: HashMap<String, HealthResult>,
487 #[serde(with = "instant_serde")]
489 pub timestamp: Instant,
490 #[serde(with = "duration_serde")]
492 pub total_duration: Duration,
493}
494
495impl OverallHealth {
496 pub fn from_results(
498 results: Vec<(&str, Result<HealthResult, tokio::time::error::Elapsed>)>,
499 ) -> Self {
500 let start = Instant::now();
501 let mut checks = HashMap::new();
502 let mut overall_status = HealthStatus::Healthy;
503
504 for (name, result) in results {
505 let health_result = match result {
506 Ok(result) => result,
507 Err(_) => HealthResult::unknown("Health check timed out"),
508 };
509
510 overall_status = overall_status.combine(health_result.status);
512 checks.insert(name.to_string(), health_result);
513 }
514
515 Self {
516 status: overall_status,
517 checks,
518 timestamp: start,
519 total_duration: start.elapsed(),
520 }
521 }
522
523 pub fn is_operational(&self) -> bool {
525 self.status.is_operational()
526 }
527
528 pub fn healthy_count(&self) -> usize {
530 self.checks
531 .values()
532 .filter(|r| r.status == HealthStatus::Healthy)
533 .count()
534 }
535
536 pub fn unhealthy_count(&self) -> usize {
538 self.checks
539 .values()
540 .filter(|r| r.status == HealthStatus::Unhealthy)
541 .count()
542 }
543
544 pub fn degraded_count(&self) -> usize {
546 self.checks
547 .values()
548 .filter(|r| r.status == HealthStatus::Degraded)
549 .count()
550 }
551}
552
553pub struct HealthChecker {
555 checks: Vec<Box<dyn HealthCheck>>,
556 timeout: Duration,
557}
558
559impl Default for HealthChecker {
560 fn default() -> Self {
561 Self::new()
562 }
563}
564
565impl HealthChecker {
566 pub fn new() -> Self {
568 Self {
569 checks: Vec::new(),
570 timeout: Duration::from_secs(30),
571 }
572 }
573
574 pub fn with_timeout(timeout: Duration) -> Self {
576 Self {
577 checks: Vec::new(),
578 timeout,
579 }
580 }
581
582 pub fn add_check<T: HealthCheck + 'static>(mut self, check: T) -> Self {
584 self.checks.push(Box::new(check));
585 self
586 }
587
588 pub fn add_check_ref<T: HealthCheck + 'static>(&mut self, check: T) {
590 self.checks.push(Box::new(check));
591 }
592
593 pub async fn check_all(&self) -> OverallHealth {
595 let mut results = Vec::new();
596
597 for check in &self.checks {
598 let name = check.name();
599 let check_timeout = check.timeout().min(self.timeout);
600
601 debug!("Running health check: {}", name);
602
603 let result = timeout(check_timeout, check.check()).await;
604 results.push((name, result));
605 }
606
607 OverallHealth::from_results(results)
608 }
609
610 pub async fn check_critical(&self) -> OverallHealth {
612 let mut results = Vec::new();
613
614 for check in &self.checks {
615 if !check.is_critical() {
616 continue;
617 }
618
619 let name = check.name();
620 let check_timeout = check.timeout().min(self.timeout);
621
622 debug!("Running critical health check: {}", name);
623
624 let result = timeout(check_timeout, check.check()).await;
625 results.push((name, result));
626 }
627
628 OverallHealth::from_results(results)
629 }
630
631 pub fn check_count(&self) -> usize {
633 self.checks.len()
634 }
635
636 pub fn check_names(&self) -> Vec<&str> {
638 self.checks.iter().map(|c| c.name()).collect()
639 }
640}
641
642#[cfg(test)]
643mod tests {
644 use super::*;
645 use tokio::time::sleep;
646
647 struct TestHealthCheck {
648 name: String,
649 result: HealthResult,
650 delay: Duration,
651 }
652
653 impl TestHealthCheck {
654 fn new(name: &str, status: HealthStatus, delay: Duration) -> Self {
655 Self {
656 name: name.to_string(),
657 result: HealthResult::new(status, format!("{name} test result")),
658 delay,
659 }
660 }
661 }
662
663 #[async_trait]
664 impl HealthCheck for TestHealthCheck {
665 fn name(&self) -> &str {
666 &self.name
667 }
668
669 async fn check(&self) -> HealthResult {
670 sleep(self.delay).await;
671 self.result.clone()
672 }
673
674 fn timeout(&self) -> Duration {
675 Duration::from_millis(100)
676 }
677 }
678
679 #[tokio::test]
680 async fn test_health_status_operations() {
681 assert!(HealthStatus::Healthy.is_operational());
682 assert!(HealthStatus::Degraded.is_operational());
683 assert!(!HealthStatus::Unhealthy.is_operational());
684 assert!(!HealthStatus::Unknown.is_operational());
685
686 assert_eq!(HealthStatus::Healthy.score(), 100);
687 assert_eq!(HealthStatus::Degraded.score(), 75);
688 assert_eq!(HealthStatus::Unhealthy.score(), 25);
689 assert_eq!(HealthStatus::Unknown.score(), 0);
690
691 assert_eq!(
692 HealthStatus::Healthy.combine(HealthStatus::Degraded),
693 HealthStatus::Degraded
694 );
695 assert_eq!(
696 HealthStatus::Degraded.combine(HealthStatus::Unhealthy),
697 HealthStatus::Unhealthy
698 );
699 }
700
701 #[tokio::test]
702 async fn test_health_result_creation() {
703 let result = HealthResult::healthy("All good")
704 .with_metadata("version", serde_json::Value::String("1.0.0".to_string()))
705 .with_duration(Duration::from_millis(50));
706
707 assert_eq!(result.status, HealthStatus::Healthy);
708 assert_eq!(result.message, "All good");
709 assert_eq!(result.duration, Duration::from_millis(50));
710 assert!(result.metadata.contains_key("version"));
711 }
712
713 #[tokio::test]
714 async fn test_transport_health_check() {
715 let check = TransportHealthCheck::new("test_transport", "http", || async { true });
716
717 let result = check.check().await;
718 assert_eq!(result.status, HealthStatus::Healthy);
719 assert!(result.message.contains("http transport is connected"));
720 }
721
722 #[tokio::test]
723 async fn test_protocol_health_check() {
724 let check = ProtocolHealthCheck::new("test_protocol", || async { Ok(()) });
725
726 let result = check.check().await;
727 assert_eq!(result.status, HealthStatus::Healthy);
728 assert_eq!(result.message, "Protocol communication successful");
729 }
730
731 #[tokio::test]
732 async fn test_resource_health_check() {
733 let check = ResourceHealthCheck::new("test_resource", "database", || async { true });
734
735 let result = check.check().await;
736 assert_eq!(result.status, HealthStatus::Healthy);
737 assert!(result.message.contains("Resource 'database' is available"));
738 assert!(!check.is_critical()); }
740
741 #[tokio::test]
742 async fn test_health_checker() {
743 let checker = HealthChecker::new()
744 .add_check(TestHealthCheck::new(
745 "test1",
746 HealthStatus::Healthy,
747 Duration::from_millis(10),
748 ))
749 .add_check(TestHealthCheck::new(
750 "test2",
751 HealthStatus::Degraded,
752 Duration::from_millis(20),
753 ))
754 .add_check(TestHealthCheck::new(
755 "test3",
756 HealthStatus::Unhealthy,
757 Duration::from_millis(5),
758 ));
759
760 assert_eq!(checker.check_count(), 3);
761 assert_eq!(checker.check_names(), vec!["test1", "test2", "test3"]);
762
763 let overall = checker.check_all().await;
764 assert_eq!(overall.status, HealthStatus::Unhealthy); assert_eq!(overall.checks.len(), 3);
766 assert_eq!(overall.healthy_count(), 1);
767 assert_eq!(overall.degraded_count(), 1);
768 assert_eq!(overall.unhealthy_count(), 1);
769 assert!(overall.total_duration.as_nanos() > 0);
771 }
772
773 #[tokio::test]
774 async fn test_health_check_timeout() {
775 let checker =
776 HealthChecker::with_timeout(Duration::from_millis(50)).add_check(TestHealthCheck::new(
777 "slow_check",
778 HealthStatus::Healthy,
779 Duration::from_millis(200),
780 ));
781
782 let overall = checker.check_all().await;
783
784 assert_eq!(overall.checks.len(), 1);
786 let result = overall.checks.get("slow_check").unwrap();
787 assert_eq!(result.status, HealthStatus::Unknown);
788 assert!(result.message.contains("timed out"));
789 }
790
791 #[tokio::test]
792 async fn test_overall_health_operations() {
793 let results = vec![
794 ("check1", Ok(HealthResult::healthy("OK"))),
795 ("check2", Ok(HealthResult::degraded("Minor issue"))),
796 (
797 "check3",
798 Err(tokio::time::timeout(Duration::from_millis(10), async {
799 tokio::time::sleep(Duration::from_millis(100)).await
800 })
801 .await
802 .unwrap_err()),
803 ),
804 ];
805
806 let overall = OverallHealth::from_results(results);
807
808 assert_eq!(overall.status, HealthStatus::Unknown); assert_eq!(overall.checks.len(), 3);
810 assert_eq!(overall.healthy_count(), 1);
811 assert_eq!(overall.degraded_count(), 1);
812 assert_eq!(overall.unhealthy_count(), 0);
813 assert!(!overall.is_operational()); }
815}