Skip to main content

prism_mcp_rs/core/
retry.rs

1//! Retry and circuit-breaker primitives for MCP operations.
2//!
3//! The module provides:
4//! - Smart retry decisions based on error recoverability
5//! - Exponential backoff with jitter
6//! - Circuit breaker pattern for cascading failure protection
7//! - structured logging and in-process metrics hooks
8
9use 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/// Retry policy configuration
20#[derive(Debug, Clone)]
21pub struct RetryConfig {
22    /// Maximum number of retry attempts
23    pub max_attempts: u32,
24    /// Initial retry delay in milliseconds
25    pub initial_delay_ms: u64,
26    /// Maximum retry delay in milliseconds
27    pub max_delay_ms: u64,
28    /// Exponential backoff multiplier
29    pub backoff_multiplier: f64,
30    /// Whether to add random jitter to delays
31    pub enable_jitter: bool,
32    /// Maximum jitter factor (0.0 to 1.0)
33    pub jitter_factor: f64,
34    /// Whether to respect error recoverability
35    pub respect_recoverability: bool,
36    /// Custom timeout for individual attempts
37    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    /// Create a conservative retry config for production
57    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    /// Create an aggressive retry config for high-availability scenarios
71    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    /// Create a retry config for network operations
85    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/// Circuit breaker states
100#[derive(Debug, Clone, Copy, PartialEq, Eq)]
101pub enum CircuitState {
102    /// Circuit is closed, requests pass through normally
103    Closed,
104    /// Circuit is open, requests fail immediately
105    Open,
106    /// Circuit is half-open, testing if service has recovered
107    HalfOpen,
108}
109
110/// Circuit breaker configuration
111#[derive(Debug, Clone)]
112pub struct CircuitBreakerConfig {
113    /// Number of failures to trigger circuit opening
114    pub failure_threshold: u32,
115    /// Time to wait before attempting recovery
116    pub recovery_timeout: Duration,
117    /// Number of successful requests needed to close circuit
118    pub success_threshold: u32,
119    /// Maximum number of requests allowed in half-open state
120    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/// Circuit breaker for protecting against cascading failures
135#[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    /// Create a new circuit breaker
147    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    /// Get current circuit state
159    pub async fn state(&self) -> CircuitState {
160        *self.state.read().await
161    }
162
163    /// Execute an operation through the circuit breaker
164    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        // Check if circuit is open and if recovery timeout has passed
169        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                // Limit concurrent requests in half-open state
179                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                        // Reset failure count on success
209                        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    /// Update circuit state based on time and failure count
224    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    /// Handle successful operation
249    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    /// Handle failed operation
269    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    /// Get circuit breaker statistics
290    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/// Circuit breaker statistics
308#[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/// Retry policy with smart error-based decisions
318#[derive(Debug)]
319pub struct RetryPolicy {
320    config: RetryConfig,
321    circuit_breaker: Option<Arc<CircuitBreaker>>,
322}
323
324impl RetryPolicy {
325    /// Create a new retry policy
326    pub fn new(config: RetryConfig) -> Self {
327        Self {
328            config,
329            circuit_breaker: None,
330        }
331    }
332
333    /// Create a retry policy with circuit breaker
334    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    /// Execute an operation with smart retry logic
345    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            // Execute through circuit breaker if available
356            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                    // Success! Log if we had previous attempts
365                    if attempt > 1 {
366                        ErrorLogger::log_retry_success(
367                            &context.operation,
368                            attempt,
369                            context.clone(),
370                        )
371                        .await;
372                    }
373
374                    // Record successful operation metrics
375                    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                    // Determine if we should retry
392                    let should_retry = self.should_retry(&error, attempt).await;
393
394                    // Log the retry attempt
395                    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                        // Final failure, log and return error
406                        error.log_with_context(context.clone()).await;
407                        return Err(error);
408                    }
409
410                    // Calculate and apply retry delay
411                    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        // All retries exhausted
424        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    /// Determine if an error should trigger a retry
438    async fn should_retry(&self, error: &McpError, attempt: u32) -> bool {
439        // Don't retry if we've reached max attempts
440        if attempt >= self.config.max_attempts {
441            return false;
442        }
443
444        // Respect error recoverability if configured
445        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    /// Calculate retry delay with exponential backoff and jitter
458    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    /// Get circuit breaker statistics if available
476    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
485/// Get current time in milliseconds since epoch
486fn 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        // Should only attempt once for non-recoverable errors
570        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        // Cause failures to open the circuit
584        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        // Circuit should now be open
595        assert_eq!(circuit_breaker.state().await, CircuitState::Open);
596
597        // Next call should fail immediately
598        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        // Open the circuit
620        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        // Wait for recovery timeout
632        tokio::time::sleep(Duration::from_millis(60)).await;
633
634        // Should transition to half-open and allow requests
635        let result = circuit_breaker
636            .call(async { Ok::<(), McpError>(()) }, &context)
637            .await;
638        assert!(result.is_ok());
639
640        // After enough successes, should close
641        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        // This should fail and open the circuit breaker
669        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); // Should retry once
681
682        // Circuit breaker should have stats
683        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        // Test conservative config
692        let conservative = RetryConfig::conservative();
693        assert_eq!(conservative.max_attempts, 2);
694        assert_eq!(conservative.initial_delay_ms, 500);
695
696        // Test aggressive config
697        let aggressive = RetryConfig::aggressive();
698        assert_eq!(aggressive.max_attempts, 5);
699        assert_eq!(aggressive.initial_delay_ms, 100);
700
701        // Test network config
702        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        // Test exponential backoff without jitter
718        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        // Should be capped at max_delay_ms
741        assert_eq!(delay3, Duration::from_millis(3000));
742        assert_eq!(delay4, Duration::from_millis(3000));
743    }
744}