Skip to main content

tokio_prompt_orchestrator/enhanced/
circuit_breaker.rs

1//! Circuit Breaker
2//!
3//! Prevents cascading failures by stopping requests to failing services.
4//!
5//! ## States
6//! - **Closed**: Normal operation, requests flow through
7//! - **Open**: Service failing, requests rejected immediately
8//! - **Half-Open**: Testing if service recovered
9//!
10//! ## Usage
11//!
12//! ```no_run
13//! use std::time::Duration;
14//! use tokio_prompt_orchestrator::enhanced::CircuitBreaker;
15//! use tokio_prompt_orchestrator::enhanced::circuit_breaker::CircuitBreakerError;
16//! # #[tokio::main]
17//! # async fn main() {
18//! let breaker = CircuitBreaker::new(5, 0.5, Duration::from_secs(60));
19//!
20//! match breaker.call(|| async {
21//!     // Your operation  -  replace with a real async call
22//!     Ok::<&str, &str>("inference result")
23//! }).await {
24//!     Ok(result) => println!("{result}"), // Success
25//!     Err(CircuitBreakerError::Open) => {
26//!         // Circuit open, fail fast
27//!     }
28//!     Err(CircuitBreakerError::Failed(e)) => {
29//!         eprintln!("Operation failed: {e}");
30//!     }
31//! }
32//! # }
33//! ```
34
35use std::collections::VecDeque;
36use std::sync::Arc;
37use std::time::{Duration, Instant};
38use tokio::sync::RwLock;
39use tracing::{debug, info, warn};
40
41/// A circuit breaker that prevents cascading failures by stopping requests to
42/// a failing downstream service.
43///
44/// # State machine
45///
46/// ```text
47/// Closed ──(failures ≥ threshold)──► Open
48///   ▲                                  │
49///   │                                  │ (timeout elapsed)
50///   │                                  ▼
51///   └──(success_rate ≥ threshold)── HalfOpen
52/// ```
53///
54/// - **Closed** — normal operation; all requests flow through.
55/// - **Open** — all requests are rejected immediately with
56///   [`CircuitBreakerError::Open`] without calling the wrapped operation.
57/// - **HalfOpen** — a single probe request is allowed; on success the circuit
58///   closes; on failure it reopens.
59///
60/// # Cloning
61///
62/// `CircuitBreaker` is `Clone + Send + Sync`.  All clones share the same
63/// underlying `Arc<RwLock<CircuitState>>`.
64///
65/// # Examples
66///
67/// ```no_run
68/// use std::time::Duration;
69/// use tokio_prompt_orchestrator::enhanced::CircuitBreaker;
70/// use tokio_prompt_orchestrator::enhanced::circuit_breaker::CircuitBreakerError;
71///
72/// # #[tokio::main]
73/// # async fn main() {
74/// let cb = CircuitBreaker::new(5, 0.8, Duration::from_secs(30));
75///
76/// let result = cb.call(|| async {
77///     reqwest::get("https://api.example.com/infer")
78///         .await
79///         .map_err(|e| e.to_string())
80/// }).await;
81///
82/// match result {
83///     Ok(resp) => { /* handle response */ }
84///     Err(CircuitBreakerError::Open) => { /* fail fast */ }
85///     Err(CircuitBreakerError::Failed(e)) => { /* handle error */ }
86/// }
87/// # }
88/// ```
89#[derive(Clone)]
90pub struct CircuitBreaker {
91    state: Arc<RwLock<CircuitState>>,
92    config: CircuitBreakerConfig,
93}
94
95#[derive(Debug, Clone)]
96struct CircuitBreakerConfig {
97    /// Number of failures before opening circuit
98    failure_threshold: usize,
99    /// Success rate threshold (0.0 - 1.0) to close circuit
100    success_threshold: f64,
101    /// How long to wait before testing if service recovered
102    timeout: Duration,
103    /// Window size for tracking metrics
104    window_size: usize,
105}
106
107#[derive(Debug)]
108struct CircuitState {
109    status: CircuitStatus,
110    failures: usize,
111    successes: usize,
112    last_failure_time: Option<Instant>,
113    last_state_change: Instant,
114    /// Recent results (true = success, false = failure)
115    recent_results: VecDeque<bool>,
116    /// Number of consecutive half-open probe failures.
117    ///
118    /// Drives exponential backoff: the probe interval after the Nth failed
119    /// probe is `timeout * 2^min(probe_failures, 6)` (capped at 64×).
120    /// Resets to 0 when the circuit successfully closes.
121    probe_failures: usize,
122}
123
124/// Current state of a circuit breaker.
125#[derive(Debug, Clone, PartialEq)]
126pub enum CircuitStatus {
127    /// Circuit is closed  -  requests flow through normally.
128    Closed,
129    /// Circuit is open  -  requests are rejected immediately without calling the operation.
130    Open,
131    /// Circuit is half-open  -  one probe request is allowed through to test recovery.
132    HalfOpen,
133}
134
135/// Circuit breaker errors
136#[derive(Debug)]
137pub enum CircuitBreakerError<E> {
138    /// Circuit is open, request rejected
139    Open,
140    /// Operation failed
141    Failed(E),
142}
143
144impl CircuitBreaker {
145    /// Create a new `CircuitBreaker` in the `Closed` state.
146    ///
147    /// # Arguments
148    ///
149    /// * `failure_threshold` — Number of consecutive failures required to open
150    ///   the circuit.  A value of `1` opens immediately on the first error.
151    /// * `success_threshold` — Required success rate (`0.0..=1.0`) over the
152    ///   recent-results window before the circuit closes from `HalfOpen`.
153    ///   Typical values: `0.8` (80 %) or `1.0` (100 %).
154    /// * `timeout` — How long to stay in the `Open` state before transitioning
155    ///   to `HalfOpen` and allowing one probe request through.
156    ///
157    /// The default rolling-window size is 100 results.  Use
158    /// [`with_window_size`](Self::with_window_size) to customise it.
159    ///
160    /// # Examples
161    ///
162    /// ```
163    /// use std::time::Duration;
164    /// use tokio_prompt_orchestrator::enhanced::CircuitBreaker;
165    ///
166    /// // Open after 5 failures; require 80 % success rate to close; probe after 60 s.
167    /// let cb = CircuitBreaker::new(5, 0.8, Duration::from_secs(60));
168    /// ```
169    pub fn new(failure_threshold: usize, success_threshold: f64, timeout: Duration) -> Self {
170        Self {
171            state: Arc::new(RwLock::new(CircuitState {
172                status: CircuitStatus::Closed,
173                failures: 0,
174                successes: 0,
175                last_failure_time: None,
176                last_state_change: Instant::now(),
177                recent_results: VecDeque::new(),
178                probe_failures: 0,
179            })),
180            config: CircuitBreakerConfig {
181                failure_threshold,
182                success_threshold,
183                timeout,
184                window_size: 100,
185            },
186        }
187    }
188
189    /// Set the rolling-window size used to calculate the success rate.
190    ///
191    /// The window keeps the most recent `size` call outcomes. A smaller window
192    /// reacts faster to bursts of failures but is more sensitive to noise. The
193    /// default is `100`.
194    ///
195    /// # Examples
196    ///
197    /// ```
198    /// use std::time::Duration;
199    /// use tokio_prompt_orchestrator::enhanced::CircuitBreaker;
200    ///
201    /// // Tighter window — reacts faster to short failure bursts.
202    /// let cb = CircuitBreaker::new(5, 0.8, Duration::from_secs(30))
203    ///     .with_window_size(20);
204    /// ```
205    ///
206    /// # Panics
207    ///
208    /// This function does not panic.
209    #[must_use]
210    pub fn with_window_size(mut self, size: usize) -> Self {
211        self.config.window_size = size.max(1);
212        self
213    }
214
215    /// Execute a fallible async operation through the circuit breaker.
216    ///
217    /// If the circuit is `Open` the operation is **not called** and
218    /// `Err(CircuitBreakerError::Open)` is returned immediately.
219    /// Otherwise the closure is invoked; its `Ok`/`Err` outcome is recorded
220    /// and may cause a state transition.
221    ///
222    /// # Arguments
223    ///
224    /// * `f` — A `FnOnce` that returns a `Future<Output = Result<T, E>>`.
225    ///   The closure is called at most once per `call` invocation.
226    ///
227    /// # Returns
228    ///
229    /// - `Ok(value)` — operation succeeded.
230    /// - `Err(CircuitBreakerError::Open)` — circuit is open; request rejected.
231    /// - `Err(CircuitBreakerError::Failed(e))` — operation returned `Err(e)`.
232    ///
233    /// # Examples
234    ///
235    /// ```no_run
236    /// use std::time::Duration;
237    /// use tokio_prompt_orchestrator::enhanced::CircuitBreaker;
238    /// use tokio_prompt_orchestrator::enhanced::circuit_breaker::CircuitBreakerError;
239    ///
240    /// # #[tokio::main]
241    /// # async fn main() {
242    /// let cb = CircuitBreaker::new(3, 0.8, Duration::from_secs(10));
243    /// let result = cb.call(|| async { Ok::<_, String>("ok") }).await;
244    /// assert!(matches!(result, Ok("ok")));
245    /// # }
246    /// ```
247    pub async fn call<F, Fut, T, E>(&self, f: F) -> Result<T, CircuitBreakerError<E>>
248    where
249        F: FnOnce() -> Fut,
250        Fut: std::future::Future<Output = Result<T, E>>,
251    {
252        // Check if we should allow the request
253        {
254            let mut state = self.state.write().await;
255
256            match state.status {
257                CircuitStatus::Open => {
258                    // Exponential backoff: each failed half-open probe doubles the
259                    // wait, capped at 64× the base timeout.  This prevents flapping
260                    // when a downstream service recovers slowly.
261                    if let Some(last_failure) = state.last_failure_time {
262                        let backoff_factor = 1u32 << state.probe_failures.min(6);
263                        let probe_interval = self.config.timeout * backoff_factor;
264                        if last_failure.elapsed() >= probe_interval {
265                            // Try half-open  -  clear window so success-rate is
266                            // calculated only from post-recovery requests.
267                            state.status = CircuitStatus::HalfOpen;
268                            state.recent_results.clear();
269                            state.failures = 0;
270                            state.last_state_change = Instant::now();
271                            info!(
272                                probe_failures = state.probe_failures,
273                                backoff_factor = backoff_factor,
274                                "circuit breaker: transitioning to half-open"
275                            );
276                            crate::metrics::inc_cb_transition("half_open");
277                        } else {
278                            // Still open, reject
279                            debug!("circuit breaker: request rejected (open)");
280                            crate::metrics::inc_cb_rejected();
281                            return Err(CircuitBreakerError::Open);
282                        }
283                    }
284                }
285                CircuitStatus::HalfOpen | CircuitStatus::Closed => {
286                    // Allow request
287                }
288            }
289        }
290
291        // Execute operation
292        let result = f().await;
293
294        // Record result
295        match &result {
296            Ok(_) => self.record_success().await,
297            Err(_) => self.record_failure().await,
298        }
299
300        result.map_err(CircuitBreakerError::Failed)
301    }
302
303    async fn record_success(&self) {
304        let mut state = self.state.write().await;
305
306        state.successes += 1;
307        state.recent_results.push_back(true);
308        if state.recent_results.len() > self.config.window_size {
309            state.recent_results.pop_front();
310        }
311
312        debug!(
313            status = ?state.status,
314            successes = state.successes,
315            failures = state.failures,
316            "circuit breaker: success recorded"
317        );
318
319        match state.status {
320            CircuitStatus::HalfOpen => {
321                // Check if we can close the circuit
322                let success_rate = self.calculate_success_rate(&state);
323                if success_rate >= self.config.success_threshold {
324                    state.status = CircuitStatus::Closed;
325                    state.failures = 0;
326                    state.successes = 0;
327                    state.probe_failures = 0; // Service recovered — reset backoff
328                    state.last_state_change = Instant::now();
329                    info!(
330                        success_rate = success_rate,
331                        "circuit breaker: closing (service recovered)"
332                    );
333                    crate::metrics::inc_cb_transition("closed");
334                }
335            }
336            CircuitStatus::Closed
337                // Reset failure count on success
338                if state.failures > 0 => {
339                    state.failures = 0;
340                }
341            _ => {}
342        }
343    }
344
345    async fn record_failure(&self) {
346        let mut state = self.state.write().await;
347
348        state.failures += 1;
349        state.last_failure_time = Some(Instant::now());
350        state.recent_results.push_back(false);
351        if state.recent_results.len() > self.config.window_size {
352            state.recent_results.pop_front();
353        }
354
355        warn!(
356            status = ?state.status,
357            failures = state.failures,
358            threshold = self.config.failure_threshold,
359            "circuit breaker: failure recorded"
360        );
361
362        match state.status {
363            CircuitStatus::Closed => {
364                if state.failures >= self.config.failure_threshold {
365                    state.status = CircuitStatus::Open;
366                    state.last_state_change = Instant::now();
367                    warn!(
368                        failures = state.failures,
369                        threshold = self.config.failure_threshold,
370                        "circuit breaker: opening (threshold exceeded)"
371                    );
372                    crate::metrics::inc_cb_transition("open");
373                }
374            }
375            CircuitStatus::HalfOpen => {
376                // Failure in half-open — go back to open and back off longer.
377                state.probe_failures = state.probe_failures.saturating_add(1);
378                state.status = CircuitStatus::Open;
379                state.last_state_change = Instant::now();
380                let next_backoff = 1u32 << state.probe_failures.min(6);
381                warn!(
382                    probe_failures = state.probe_failures,
383                    next_wait_factor = next_backoff,
384                    "circuit breaker: reopening (half-open test failed — next probe in {}× timeout)",
385                    next_backoff
386                );
387                crate::metrics::inc_cb_transition("open");
388            }
389            _ => {}
390        }
391    }
392
393    fn calculate_success_rate(&self, state: &CircuitState) -> f64 {
394        if state.recent_results.is_empty() {
395            return 0.0;
396        }
397
398        let successes = state.recent_results.iter().filter(|&&x| x).count();
399        successes as f64 / state.recent_results.len() as f64
400    }
401
402    /// Returns `true` if the circuit is currently closed (normal operation).
403    ///
404    /// This is a best-effort synchronous read — use `status()` for authoritative state.
405    pub fn is_closed_sync(&self) -> bool {
406        self.state
407            .try_read()
408            .map(|s| matches!(s.status, CircuitStatus::Closed))
409            .unwrap_or(false)
410    }
411
412    /// Returns `true` if the circuit is currently open (rejecting requests).
413    pub fn is_open_sync(&self) -> bool {
414        self.state
415            .try_read()
416            .map(|s| matches!(s.status, CircuitStatus::Open))
417            .unwrap_or(false)
418    }
419
420    /// Returns `true` if the circuit is in half-open state (testing recovery).
421    pub fn is_half_open_sync(&self) -> bool {
422        self.state
423            .try_read()
424            .map(|s| matches!(s.status, CircuitStatus::HalfOpen))
425            .unwrap_or(false)
426    }
427
428    /// Get current circuit status
429    pub async fn status(&self) -> CircuitStatus {
430        self.state.read().await.status.clone()
431    }
432
433    /// Get circuit breaker statistics
434    pub async fn stats(&self) -> CircuitBreakerStats {
435        let state = self.state.read().await;
436
437        CircuitBreakerStats {
438            status: state.status.clone(),
439            failures: state.failures,
440            successes: state.successes,
441            success_rate: self.calculate_success_rate(&state),
442            time_in_current_state: state.last_state_change.elapsed(),
443            probe_failures: state.probe_failures,
444        }
445    }
446
447    /// Manually reset circuit breaker to closed state
448    pub async fn reset(&self) {
449        let mut state = self.state.write().await;
450        state.status = CircuitStatus::Closed;
451        state.failures = 0;
452        state.successes = 0;
453        state.probe_failures = 0;
454        state.recent_results.clear();
455        state.last_state_change = Instant::now();
456        info!("circuit breaker: manually reset to closed");
457    }
458
459    /// Force circuit to open state (for testing/maintenance)
460    pub async fn trip(&self) {
461        let mut state = self.state.write().await;
462        state.status = CircuitStatus::Open;
463        state.last_failure_time = Some(Instant::now());
464        state.last_state_change = Instant::now();
465        warn!("circuit breaker: manually tripped to open");
466    }
467}
468
469/// A point-in-time snapshot of [`CircuitBreaker`] metrics.
470///
471/// Obtain via [`CircuitBreaker::stats`].
472#[derive(Debug)]
473pub struct CircuitBreakerStats {
474    /// Current state of the circuit breaker.
475    pub status: CircuitStatus,
476    /// Total failures recorded in the current window.
477    pub failures: usize,
478    /// Total successes recorded in the current window.
479    pub successes: usize,
480    /// Fraction of recent requests that succeeded (0.0  -  1.0).
481    pub success_rate: f64,
482    /// Wall-clock time spent in the current state.
483    pub time_in_current_state: Duration,
484    /// Consecutive half-open probe failures.  The next probe interval is
485    /// `timeout * 2^min(probe_failures, 6)`.  Zero when the circuit is
486    /// closed or has not yet attempted any half-open probes.
487    pub probe_failures: usize,
488}
489
490#[cfg(test)]
491mod tests {
492    use super::*;
493
494    #[tokio::test]
495    async fn test_circuit_opens_on_failures() {
496        let breaker = CircuitBreaker::new(3, 0.8, Duration::from_secs(5));
497
498        // Record 3 failures
499        for _ in 0..3 {
500            let result: Result<(), CircuitBreakerError<()>> =
501                breaker.call(|| async { Err(()) }).await;
502            assert!(result.is_err());
503        }
504
505        // Circuit should be open now
506        assert_eq!(breaker.status().await, CircuitStatus::Open);
507
508        // Next request should be rejected
509        let result: Result<(), CircuitBreakerError<()>> = breaker.call(|| async { Ok(()) }).await;
510        assert!(matches!(result, Err(CircuitBreakerError::Open)));
511    }
512
513    #[tokio::test]
514    async fn test_circuit_closes_on_recovery() {
515        let breaker = CircuitBreaker::new(2, 0.8, Duration::from_millis(100));
516
517        // Open circuit
518        for _ in 0..2 {
519            let _: Result<(), CircuitBreakerError<()>> = breaker.call(|| async { Err(()) }).await;
520        }
521        assert_eq!(breaker.status().await, CircuitStatus::Open);
522
523        // Wait for timeout
524        tokio::time::sleep(Duration::from_millis(150)).await;
525
526        // Should transition to half-open and allow test request
527        let result: Result<(), CircuitBreakerError<()>> = breaker.call(|| async { Ok(()) }).await;
528        assert!(result.is_ok());
529
530        // Record more successes to close circuit
531        for _ in 0..5 {
532            let _: Result<(), CircuitBreakerError<()>> = breaker.call(|| async { Ok(()) }).await;
533        }
534
535        assert_eq!(breaker.status().await, CircuitStatus::Closed);
536    }
537
538    #[tokio::test]
539    async fn test_manual_reset() {
540        let breaker = CircuitBreaker::new(2, 0.8, Duration::from_secs(60));
541
542        // Open circuit
543        for _ in 0..2 {
544            let _: Result<(), CircuitBreakerError<()>> = breaker.call(|| async { Err(()) }).await;
545        }
546        assert_eq!(breaker.status().await, CircuitStatus::Open);
547
548        // Manual reset
549        breaker.reset().await;
550        assert_eq!(breaker.status().await, CircuitStatus::Closed);
551    }
552
553    #[tokio::test]
554    async fn test_circuit_breaker_clears_history_on_half_open() {
555        // Open the breaker with 2 failures, then wait for the timeout.
556        let breaker = CircuitBreaker::new(2, 0.8, Duration::from_millis(50));
557
558        for _ in 0..2 {
559            let _: Result<(), CircuitBreakerError<()>> = breaker.call(|| async { Err(()) }).await;
560        }
561        assert_eq!(breaker.status().await, CircuitStatus::Open);
562
563        // The recent_results window should have 2 failures recorded.
564        {
565            let state = breaker.state.read().await;
566            assert!(
567                !state.recent_results.is_empty(),
568                "should have failure history before half-open"
569            );
570        }
571
572        // Wait for the open timeout to elapse.
573        tokio::time::sleep(Duration::from_millis(100)).await;
574
575        // First successful probe transitions to HalfOpen and clears history.
576        let _: Result<(), CircuitBreakerError<()>> = breaker.call(|| async { Ok(()) }).await;
577
578        // After the transition the history must contain only the single probe
579        // result (the success above), not the old failures.
580        {
581            let state = breaker.state.read().await;
582            assert!(
583                !state.recent_results.contains(&false),
584                "old failures must be cleared on half-open transition; results: {:?}",
585                state.recent_results
586            );
587        }
588    }
589
590    #[tokio::test]
591    async fn test_circuit_breaker_half_open_failure_reopens() {
592        let cb = CircuitBreaker::new(1, 0.8, Duration::from_millis(10));
593
594        // Trip the circuit
595        let _ = cb.call(|| async { Err::<(), &str>("fail") }).await;
596
597        // Wait for timeout
598        tokio::time::sleep(Duration::from_millis(20)).await;
599
600        // Should be half-open now — a failure here should re-open
601        let result = cb.call(|| async { Err::<(), &str>("fail again") }).await;
602        assert!(result.is_err());
603
604        // Should be open again
605        let result = cb.call(|| async { Ok::<(), &str>(()) }).await;
606        assert!(matches!(result, Err(CircuitBreakerError::Open)));
607    }
608
609    #[tokio::test]
610    async fn test_stats() {
611        let breaker = CircuitBreaker::new(5, 0.8, Duration::from_secs(60));
612
613        // Record some results
614        let _: Result<(), CircuitBreakerError<()>> = breaker.call(|| async { Ok(()) }).await;
615        let _: Result<(), CircuitBreakerError<()>> = breaker.call(|| async { Err(()) }).await;
616        let _: Result<(), CircuitBreakerError<()>> = breaker.call(|| async { Ok(()) }).await;
617
618        let stats = breaker.stats().await;
619        assert_eq!(stats.successes, 2);
620        // record_success() resets the failure counter when status is Closed,
621        // so after the 3rd call (a success) the failure count is back to 0.
622        assert_eq!(stats.failures, 0);
623        assert_eq!(stats.status, CircuitStatus::Closed);
624    }
625}