1use std::sync::atomic::{AtomicU32, AtomicU64, Ordering};
10use std::sync::Arc;
11use std::time::{Duration, Instant};
12use tokio::time::sleep;
13use tracing::{debug, error, warn};
14
15use crate::core::error::{McpError, McpResult};
16use crate::core::logging::{ErrorContext, ErrorLogger};
17use crate::core::metrics::global_metrics;
18
19#[derive(Debug, Clone)]
21pub struct RetryConfig {
22 pub max_attempts: u32,
24 pub initial_delay_ms: u64,
26 pub max_delay_ms: u64,
28 pub backoff_multiplier: f64,
30 pub enable_jitter: bool,
32 pub jitter_factor: f64,
34 pub respect_recoverability: bool,
36 pub attempt_timeout: Option<Duration>,
38}
39
40impl Default for RetryConfig {
41 fn default() -> Self {
42 Self {
43 max_attempts: 3,
44 initial_delay_ms: 1000,
45 max_delay_ms: 30000,
46 backoff_multiplier: 2.0,
47 enable_jitter: true,
48 jitter_factor: 0.1,
49 respect_recoverability: true,
50 attempt_timeout: None,
51 }
52 }
53}
54
55impl RetryConfig {
56 pub fn conservative() -> Self {
58 Self {
59 max_attempts: 2,
60 initial_delay_ms: 500,
61 max_delay_ms: 5000,
62 backoff_multiplier: 1.5,
63 enable_jitter: true,
64 jitter_factor: 0.05,
65 respect_recoverability: true,
66 attempt_timeout: Some(Duration::from_secs(30)),
67 }
68 }
69
70 pub fn aggressive() -> Self {
72 Self {
73 max_attempts: 5,
74 initial_delay_ms: 100,
75 max_delay_ms: 60000,
76 backoff_multiplier: 2.5,
77 enable_jitter: true,
78 jitter_factor: 0.15,
79 respect_recoverability: true,
80 attempt_timeout: Some(Duration::from_secs(60)),
81 }
82 }
83
84 pub fn network() -> Self {
86 Self {
87 max_attempts: 4,
88 initial_delay_ms: 200,
89 max_delay_ms: 15000,
90 backoff_multiplier: 2.0,
91 enable_jitter: true,
92 jitter_factor: 0.1,
93 respect_recoverability: true,
94 attempt_timeout: Some(Duration::from_secs(45)),
95 }
96 }
97}
98
99#[derive(Debug, Clone, Copy, PartialEq, Eq)]
101pub enum CircuitState {
102 Closed,
104 Open,
106 HalfOpen,
108}
109
110#[derive(Debug, Clone)]
112pub struct CircuitBreakerConfig {
113 pub failure_threshold: u32,
115 pub recovery_timeout: Duration,
117 pub success_threshold: u32,
119 pub half_open_max_requests: u32,
121}
122
123impl Default for CircuitBreakerConfig {
124 fn default() -> Self {
125 Self {
126 failure_threshold: 5,
127 recovery_timeout: Duration::from_secs(60),
128 success_threshold: 3,
129 half_open_max_requests: 3,
130 }
131 }
132}
133
134#[derive(Debug)]
136pub struct CircuitBreaker {
137 config: CircuitBreakerConfig,
138 failure_count: AtomicU32,
139 success_count: AtomicU32,
140 last_failure_time: AtomicU64,
141 half_open_requests: AtomicU32,
142 state: Arc<tokio::sync::RwLock<CircuitState>>,
143}
144
145impl CircuitBreaker {
146 pub fn new(config: CircuitBreakerConfig) -> Self {
148 Self {
149 config,
150 failure_count: AtomicU32::new(0),
151 success_count: AtomicU32::new(0),
152 last_failure_time: AtomicU64::new(0),
153 half_open_requests: AtomicU32::new(0),
154 state: Arc::new(tokio::sync::RwLock::new(CircuitState::Closed)),
155 }
156 }
157
158 pub async fn state(&self) -> CircuitState {
160 *self.state.read().await
161 }
162
163 pub async fn call<F, T>(&self, operation: F, context: &ErrorContext) -> McpResult<T>
165 where
166 F: std::future::Future<Output = McpResult<T>>,
167 {
168 let current_state = self.update_state_if_needed().await;
170
171 match current_state {
172 CircuitState::Open => {
173 let error = McpError::connection("Circuit breaker is open");
174 error.log_with_context(context.clone()).await;
175 Err(error)
176 }
177 CircuitState::HalfOpen => {
178 let current_requests = self.half_open_requests.fetch_add(1, Ordering::SeqCst);
180 if current_requests >= self.config.half_open_max_requests {
181 self.half_open_requests.fetch_sub(1, Ordering::SeqCst);
182 let error = McpError::connection(
183 "Circuit breaker is half-open with max concurrent requests",
184 );
185 error.log_with_context(context.clone()).await;
186 return Err(error);
187 }
188
189 let result = operation.await;
190 self.half_open_requests.fetch_sub(1, Ordering::SeqCst);
191
192 match &result {
193 Ok(_) => self.on_success().await,
194 Err(error) => {
195 if error.is_recoverable() {
196 self.on_failure().await;
197 }
198 }
199 }
200
201 result
202 }
203 CircuitState::Closed => {
204 let result = operation.await;
205
206 match &result {
207 Ok(_) => {
208 self.failure_count.store(0, Ordering::SeqCst);
210 }
211 Err(error) => {
212 if error.is_recoverable() {
213 self.on_failure().await;
214 }
215 }
216 }
217
218 result
219 }
220 }
221 }
222
223 async fn update_state_if_needed(&self) -> CircuitState {
225 let current_state = *self.state.read().await;
226
227 match current_state {
228 CircuitState::Open => {
229 let last_failure = self.last_failure_time.load(Ordering::SeqCst);
230 let now = current_time_millis();
231
232 if now.saturating_sub(last_failure)
233 >= self.config.recovery_timeout.as_millis() as u64
234 {
235 let mut state = self.state.write().await;
236 *state = CircuitState::HalfOpen;
237 self.success_count.store(0, Ordering::SeqCst);
238 debug!("Circuit breaker transitioned to HalfOpen state");
239 CircuitState::HalfOpen
240 } else {
241 CircuitState::Open
242 }
243 }
244 _ => current_state,
245 }
246 }
247
248 async fn on_success(&self) {
250 let current_state = *self.state.read().await;
251
252 if current_state == CircuitState::HalfOpen {
253 let success_count = self.success_count.fetch_add(1, Ordering::SeqCst) + 1;
254
255 if success_count >= self.config.success_threshold {
256 let mut state = self.state.write().await;
257 *state = CircuitState::Closed;
258 self.failure_count.store(0, Ordering::SeqCst);
259 self.success_count.store(0, Ordering::SeqCst);
260 debug!(
261 "Circuit breaker transitioned to Closed state after {} successes",
262 success_count
263 );
264 }
265 }
266 }
267
268 async fn on_failure(&self) {
270 let failure_count = self.failure_count.fetch_add(1, Ordering::SeqCst) + 1;
271 self.last_failure_time
272 .store(current_time_millis(), Ordering::SeqCst);
273
274 if failure_count >= self.config.failure_threshold {
275 let mut state = self.state.write().await;
276 if *state == CircuitState::Closed {
277 *state = CircuitState::Open;
278 warn!(
279 "Circuit breaker opened after {} failures, recovery timeout: {:?}",
280 failure_count, self.config.recovery_timeout
281 );
282 } else if *state == CircuitState::HalfOpen {
283 *state = CircuitState::Open;
284 warn!("Circuit breaker reopened during half-open state");
285 }
286 }
287 }
288
289 pub async fn stats(&self) -> CircuitBreakerStats {
291 CircuitBreakerStats {
292 state: self.state().await,
293 failure_count: self.failure_count.load(Ordering::SeqCst),
294 success_count: self.success_count.load(Ordering::SeqCst),
295 last_failure_time: self.last_failure_time.load(Ordering::SeqCst),
296 half_open_requests: self.half_open_requests.load(Ordering::SeqCst),
297 }
298 }
299}
300
301impl Default for CircuitBreaker {
302 fn default() -> Self {
303 Self::new(CircuitBreakerConfig::default())
304 }
305}
306
307#[derive(Debug, Clone)]
309pub struct CircuitBreakerStats {
310 pub state: CircuitState,
311 pub failure_count: u32,
312 pub success_count: u32,
313 pub last_failure_time: u64,
314 pub half_open_requests: u32,
315}
316
317#[derive(Debug)]
319pub struct RetryPolicy {
320 config: RetryConfig,
321 circuit_breaker: Option<Arc<CircuitBreaker>>,
322}
323
324impl RetryPolicy {
325 pub fn new(config: RetryConfig) -> Self {
327 Self {
328 config,
329 circuit_breaker: None,
330 }
331 }
332
333 pub fn with_circuit_breaker(
335 config: RetryConfig,
336 circuit_breaker_config: CircuitBreakerConfig,
337 ) -> Self {
338 Self {
339 config,
340 circuit_breaker: Some(Arc::new(CircuitBreaker::new(circuit_breaker_config))),
341 }
342 }
343
344 pub async fn execute<F, T>(&self, mut operation: F, context: ErrorContext) -> McpResult<T>
346 where
347 F: FnMut() -> std::pin::Pin<Box<dyn std::future::Future<Output = McpResult<T>> + Send>>,
348 {
349 let mut last_error = None;
350 let start_time = Instant::now();
351
352 for attempt in 1..=self.config.max_attempts {
353 let attempt_start = Instant::now();
354
355 let result = if let Some(ref circuit_breaker) = self.circuit_breaker {
357 circuit_breaker.call(operation(), &context).await
358 } else {
359 operation().await
360 };
361
362 match result {
363 Ok(value) => {
364 if attempt > 1 {
366 ErrorLogger::log_retry_success(
367 &context.operation,
368 attempt,
369 context.clone(),
370 )
371 .await;
372 }
373
374 let metrics = global_metrics();
376 if let Some(ref method) = context.method {
377 metrics
378 .record_request(
379 method,
380 context.transport.as_deref().unwrap_or("unknown"),
381 )
382 .await;
383 }
384
385 return Ok(value);
386 }
387 Err(error) => {
388 let attempt_duration = attempt_start.elapsed();
389 last_error = Some(error.clone());
390
391 let should_retry = self.should_retry(&error, attempt).await;
393
394 ErrorLogger::log_retry_attempt(
396 &error,
397 attempt,
398 self.config.max_attempts,
399 should_retry,
400 context.clone(),
401 )
402 .await;
403
404 if !should_retry {
405 error.log_with_context(context.clone()).await;
407 return Err(error);
408 }
409
410 if attempt < self.config.max_attempts {
412 let delay = self.calculate_delay(attempt, attempt_duration);
413 debug!(
414 "Retrying {} in {:?} (attempt {}/{})",
415 context.operation, delay, attempt, self.config.max_attempts
416 );
417 sleep(delay).await;
418 }
419 }
420 }
421 }
422
423 let total_duration = start_time.elapsed();
425 let final_error = last_error
426 .unwrap_or_else(|| McpError::internal("Retry logic failed without capturing error"));
427
428 error!(
429 "Operation '{}' failed after {} attempts in {:?}",
430 context.operation, self.config.max_attempts, total_duration
431 );
432
433 final_error.log_with_context(context).await;
434 Err(final_error)
435 }
436
437 async fn should_retry(&self, error: &McpError, attempt: u32) -> bool {
439 if attempt >= self.config.max_attempts {
441 return false;
442 }
443
444 if self.config.respect_recoverability && !error.is_recoverable() {
446 debug!(
447 "Not retrying non-recoverable error: {} (category: {})",
448 error,
449 error.category()
450 );
451 return false;
452 }
453
454 true
455 }
456
457 fn calculate_delay(&self, attempt: u32, _last_attempt_duration: Duration) -> Duration {
459 let base_delay = self.config.initial_delay_ms as f64
460 * self.config.backoff_multiplier.powi(attempt as i32 - 1);
461
462 let capped_delay = base_delay.min(self.config.max_delay_ms as f64);
463
464 let final_delay = if self.config.enable_jitter {
465 let jitter_range = capped_delay * self.config.jitter_factor;
466 let jitter = (fastrand::f64() - 0.5) * 2.0 * jitter_range;
467 (capped_delay + jitter).max(0.0)
468 } else {
469 capped_delay
470 };
471
472 Duration::from_millis(final_delay as u64)
473 }
474
475 pub async fn circuit_breaker_stats(&self) -> Option<CircuitBreakerStats> {
477 if let Some(ref circuit_breaker) = self.circuit_breaker {
478 Some(circuit_breaker.stats().await)
479 } else {
480 None
481 }
482 }
483}
484
485fn current_time_millis() -> u64 {
487 std::time::SystemTime::now()
488 .duration_since(std::time::UNIX_EPOCH)
489 .unwrap_or_default()
490 .as_millis() as u64
491}
492
493#[cfg(test)]
494mod tests {
495 use super::*;
496 use std::sync::atomic::AtomicU32;
497 use std::sync::Arc;
498 use tokio::time::Duration;
499
500 #[tokio::test]
501 async fn test_retry_policy_success_immediate() {
502 let policy = RetryPolicy::new(RetryConfig::default());
503 let context = ErrorContext::new("test_operation");
504
505 let result = policy
506 .execute(|| Box::pin(async { Ok::<i32, McpError>(42) }), context)
507 .await;
508
509 assert!(result.is_ok());
510 assert_eq!(result.unwrap(), 42);
511 }
512
513 #[tokio::test]
514 async fn test_retry_policy_success_after_retries() {
515 let policy = RetryPolicy::new(RetryConfig {
516 max_attempts: 3,
517 initial_delay_ms: 10,
518 ..Default::default()
519 });
520 let context = ErrorContext::new("test_retry_operation");
521
522 let attempt_count = Arc::new(AtomicU32::new(0));
523 let attempt_count_clone = attempt_count.clone();
524
525 let result = policy
526 .execute(
527 move || {
528 let count = attempt_count_clone.fetch_add(1, Ordering::SeqCst) + 1;
529 Box::pin(async move {
530 if count < 3 {
531 Err(McpError::connection("Temporary failure"))
532 } else {
533 Ok::<i32, McpError>(42)
534 }
535 })
536 },
537 context,
538 )
539 .await;
540
541 assert!(result.is_ok());
542 assert_eq!(result.unwrap(), 42);
543 assert_eq!(attempt_count.load(Ordering::SeqCst), 3);
544 }
545
546 #[tokio::test]
547 async fn test_retry_policy_non_recoverable_error() {
548 let policy = RetryPolicy::new(RetryConfig {
549 max_attempts: 3,
550 respect_recoverability: true,
551 ..Default::default()
552 });
553 let context = ErrorContext::new("test_non_recoverable");
554
555 let attempt_count = Arc::new(AtomicU32::new(0));
556 let attempt_count_clone = attempt_count.clone();
557
558 let result = policy
559 .execute(
560 move || {
561 attempt_count_clone.fetch_add(1, Ordering::SeqCst);
562 Box::pin(async { Err::<i32, McpError>(McpError::validation("Invalid input")) })
563 },
564 context,
565 )
566 .await;
567
568 assert!(result.is_err());
569 assert_eq!(attempt_count.load(Ordering::SeqCst), 1);
571 }
572
573 #[tokio::test]
574 async fn test_circuit_breaker_opens_after_failures() {
575 let circuit_breaker = CircuitBreaker::new(CircuitBreakerConfig {
576 failure_threshold: 3,
577 recovery_timeout: Duration::from_millis(100),
578 ..Default::default()
579 });
580
581 let context = ErrorContext::new("test_circuit_breaker");
582
583 for _ in 0..3 {
585 let result = circuit_breaker
586 .call(
587 async { Err::<(), McpError>(McpError::connection("Service down")) },
588 &context,
589 )
590 .await;
591 assert!(result.is_err());
592 }
593
594 assert_eq!(circuit_breaker.state().await, CircuitState::Open);
596
597 let result = circuit_breaker
599 .call(async { Ok::<(), McpError>(()) }, &context)
600 .await;
601 assert!(result.is_err());
602 assert!(result
603 .unwrap_err()
604 .to_string()
605 .contains("Circuit breaker is open"));
606 }
607
608 #[tokio::test]
609 async fn test_circuit_breaker_half_open_recovery() {
610 let circuit_breaker = CircuitBreaker::new(CircuitBreakerConfig {
611 failure_threshold: 2,
612 recovery_timeout: Duration::from_millis(50),
613 success_threshold: 2,
614 ..Default::default()
615 });
616
617 let context = ErrorContext::new("test_recovery");
618
619 for _ in 0..2 {
621 let _ = circuit_breaker
622 .call(
623 async { Err::<(), McpError>(McpError::connection("Failure")) },
624 &context,
625 )
626 .await;
627 }
628
629 assert_eq!(circuit_breaker.state().await, CircuitState::Open);
630
631 tokio::time::sleep(Duration::from_millis(60)).await;
633
634 let result = circuit_breaker
636 .call(async { Ok::<(), McpError>(()) }, &context)
637 .await;
638 assert!(result.is_ok());
639
640 let result = circuit_breaker
642 .call(async { Ok::<(), McpError>(()) }, &context)
643 .await;
644 assert!(result.is_ok());
645
646 assert_eq!(circuit_breaker.state().await, CircuitState::Closed);
647 }
648
649 #[tokio::test]
650 async fn test_retry_with_circuit_breaker() {
651 let policy = RetryPolicy::with_circuit_breaker(
652 RetryConfig {
653 max_attempts: 2,
654 initial_delay_ms: 10,
655 ..Default::default()
656 },
657 CircuitBreakerConfig {
658 failure_threshold: 2,
659 recovery_timeout: Duration::from_millis(100),
660 ..Default::default()
661 },
662 );
663
664 let context = ErrorContext::new("test_combined");
665 let attempt_count = Arc::new(AtomicU32::new(0));
666 let attempt_count_clone = attempt_count.clone();
667
668 let result = policy
670 .execute(
671 move || {
672 attempt_count_clone.fetch_add(1, Ordering::SeqCst);
673 Box::pin(async { Err::<(), McpError>(McpError::connection("Service down")) })
674 },
675 context,
676 )
677 .await;
678
679 assert!(result.is_err());
680 assert_eq!(attempt_count.load(Ordering::SeqCst), 2); let stats = policy.circuit_breaker_stats().await;
684 assert!(stats.is_some());
685 let stats = stats.unwrap();
686 assert_eq!(stats.failure_count, 2);
687 }
688
689 #[tokio::test]
690 async fn test_retry_config_variations() {
691 let conservative = RetryConfig::conservative();
693 assert_eq!(conservative.max_attempts, 2);
694 assert_eq!(conservative.initial_delay_ms, 500);
695
696 let aggressive = RetryConfig::aggressive();
698 assert_eq!(aggressive.max_attempts, 5);
699 assert_eq!(aggressive.initial_delay_ms, 100);
700
701 let network = RetryConfig::network();
703 assert_eq!(network.max_attempts, 4);
704 assert_eq!(network.initial_delay_ms, 200);
705 }
706
707 #[tokio::test]
708 async fn test_delay_calculation() {
709 let policy = RetryPolicy::new(RetryConfig {
710 initial_delay_ms: 1000,
711 max_delay_ms: 10000,
712 backoff_multiplier: 2.0,
713 enable_jitter: false,
714 ..Default::default()
715 });
716
717 let delay1 = policy.calculate_delay(1, Duration::from_millis(100));
719 let delay2 = policy.calculate_delay(2, Duration::from_millis(100));
720 let delay3 = policy.calculate_delay(3, Duration::from_millis(100));
721
722 assert_eq!(delay1, Duration::from_millis(1000));
723 assert_eq!(delay2, Duration::from_millis(2000));
724 assert_eq!(delay3, Duration::from_millis(4000));
725 }
726
727 #[tokio::test]
728 async fn test_delay_calculation_with_cap() {
729 let policy = RetryPolicy::new(RetryConfig {
730 initial_delay_ms: 1000,
731 max_delay_ms: 3000,
732 backoff_multiplier: 2.0,
733 enable_jitter: false,
734 ..Default::default()
735 });
736
737 let delay3 = policy.calculate_delay(3, Duration::from_millis(100));
738 let delay4 = policy.calculate_delay(4, Duration::from_millis(100));
739
740 assert_eq!(delay3, Duration::from_millis(3000));
742 assert_eq!(delay4, Duration::from_millis(3000));
743 }
744}