Skip to main content

tokio_prompt_orchestrator/
circuit_breaker.rs

1//! # Circuit Breaker State Machine
2//!
3//! Per-model circuit breaker with sliding-window failure tracking, three-state
4//! FSM, and a [`CircuitBreakerRegistry`] keyed by model name.
5//!
6//! ## States
7//!
8//! ```text
9//! Closed ──(failure_rate > threshold)──> Open
10//! Open   ──(timeout elapsed)──────────> HalfOpen
11//! HalfOpen ──(success_threshold met)──> Closed
12//! HalfOpen ──(any failure)────────────> Open
13//! ```
14//!
15//! ## Example
16//!
17//! ```rust
18//! use tokio_prompt_orchestrator::circuit_breaker::{CircuitBreaker, CircuitBreakerConfig};
19//!
20//! let cfg = CircuitBreakerConfig::default();
21//! let cb = CircuitBreaker::new(cfg);
22//!
23//! assert!(cb.is_allowed());
24//! cb.record_failure();
25//! ```
26
27use std::collections::VecDeque;
28use std::sync::{Arc, Mutex};
29use std::time::{Duration, Instant};
30
31use dashmap::DashMap;
32
33// ── State ─────────────────────────────────────────────────────────────────────
34
35/// Current state of a circuit breaker.
36#[derive(Debug, Clone, Copy, PartialEq, Eq)]
37pub enum CircuitState {
38    /// Normal operation — requests are allowed.
39    Closed,
40    /// Failure rate exceeded threshold — all requests are rejected.
41    Open,
42    /// Timeout elapsed — a single probe request is permitted.
43    HalfOpen,
44}
45
46// ── Config ────────────────────────────────────────────────────────────────────
47
48/// Configuration for a [`CircuitBreaker`].
49#[derive(Debug, Clone)]
50pub struct CircuitBreakerConfig {
51    /// Number of failures in the sliding window required to trip the breaker.
52    pub failure_threshold: usize,
53    /// Number of consecutive successes in HalfOpen required to close the
54    /// breaker again.
55    pub success_threshold: usize,
56    /// How long the breaker stays Open before moving to HalfOpen.
57    pub timeout_duration: Duration,
58    /// Width of the sliding window used for failure-rate calculation.
59    pub window_duration: Duration,
60    /// Minimum number of calls in the window before the failure rate is
61    /// evaluated (avoids tripping on the very first request).
62    pub min_calls: usize,
63}
64
65impl Default for CircuitBreakerConfig {
66    fn default() -> Self {
67        Self {
68            failure_threshold: 5,
69            success_threshold: 2,
70            timeout_duration: Duration::from_secs(30),
71            window_duration: Duration::from_secs(60),
72            min_calls: 5,
73        }
74    }
75}
76
77// ── Inner mutable state ───────────────────────────────────────────────────────
78
79struct Inner {
80    state: CircuitState,
81    /// Ring buffer of (timestamp, success) outcome pairs.
82    window: VecDeque<(Instant, bool)>,
83    /// When the breaker entered the Open state (used to compute timeout).
84    opened_at: Option<Instant>,
85    /// Consecutive successes recorded while in HalfOpen.
86    half_open_successes: usize,
87}
88
89impl Inner {
90    fn new() -> Self {
91        Self {
92            state: CircuitState::Closed,
93            window: VecDeque::new(),
94            opened_at: None,
95            half_open_successes: 0,
96        }
97    }
98
99    /// Remove entries older than `window_duration`.
100    fn evict_old(&mut self, window_duration: Duration) {
101        let cutoff = Instant::now() - window_duration;
102        while let Some(&(ts, _)) = self.window.front() {
103            if ts < cutoff {
104                self.window.pop_front();
105            } else {
106                break;
107            }
108        }
109    }
110
111    /// Number of failures currently in the sliding window.
112    fn failure_count(&self) -> usize {
113        self.window.iter().filter(|(_, ok)| !ok).count()
114    }
115
116    /// Total calls in the sliding window.
117    fn total_count(&self) -> usize {
118        self.window.len()
119    }
120}
121
122// ── CircuitBreaker ────────────────────────────────────────────────────────────
123
124/// A single per-model circuit breaker.
125pub struct CircuitBreaker {
126    config: CircuitBreakerConfig,
127    inner: Mutex<Inner>,
128}
129
130impl CircuitBreaker {
131    /// Create a new circuit breaker with the given configuration.
132    pub fn new(config: CircuitBreakerConfig) -> Self {
133        Self {
134            config,
135            inner: Mutex::new(Inner::new()),
136        }
137    }
138
139    /// Returns `true` when a request should be allowed through.
140    ///
141    /// - **Closed**: always allowed.
142    /// - **Open**: not allowed unless the timeout has elapsed, in which case
143    ///   the breaker transitions to HalfOpen and the probe request is allowed.
144    /// - **HalfOpen**: allowed (the probe).
145    pub fn is_allowed(&self) -> bool {
146        let mut g = self.inner.lock().unwrap_or_else(|e| e.into_inner());
147        match g.state {
148            CircuitState::Closed => true,
149            CircuitState::HalfOpen => true,
150            CircuitState::Open => {
151                if let Some(opened_at) = g.opened_at {
152                    if opened_at.elapsed() >= self.config.timeout_duration {
153                        g.state = CircuitState::HalfOpen;
154                        g.half_open_successes = 0;
155                        true
156                    } else {
157                        false
158                    }
159                } else {
160                    false
161                }
162            }
163        }
164    }
165
166    /// Record a successful call outcome and drive state transitions.
167    pub fn record_success(&self) {
168        let mut g = self.inner.lock().unwrap_or_else(|e| e.into_inner());
169        let now = Instant::now();
170        g.evict_old(self.config.window_duration);
171        g.window.push_back((now, true));
172
173        match g.state {
174            CircuitState::HalfOpen => {
175                g.half_open_successes += 1;
176                if g.half_open_successes >= self.config.success_threshold {
177                    g.state = CircuitState::Closed;
178                    g.opened_at = None;
179                    g.half_open_successes = 0;
180                }
181            }
182            CircuitState::Closed | CircuitState::Open => {}
183        }
184    }
185
186    /// Record a failed call outcome and drive state transitions.
187    pub fn record_failure(&self) {
188        let mut g = self.inner.lock().unwrap_or_else(|e| e.into_inner());
189        let now = Instant::now();
190        g.evict_old(self.config.window_duration);
191        g.window.push_back((now, false));
192
193        match g.state {
194            CircuitState::HalfOpen => {
195                // Any failure in HalfOpen immediately reopens the breaker.
196                g.state = CircuitState::Open;
197                g.opened_at = Some(now);
198                g.half_open_successes = 0;
199            }
200            CircuitState::Closed => {
201                let total = g.total_count();
202                let failures = g.failure_count();
203                if total >= self.config.min_calls
204                    && failures >= self.config.failure_threshold
205                {
206                    g.state = CircuitState::Open;
207                    g.opened_at = Some(now);
208                }
209            }
210            CircuitState::Open => {}
211        }
212    }
213
214    /// Return the current state of the circuit breaker.
215    pub fn state(&self) -> CircuitState {
216        let g = self.inner.lock().unwrap_or_else(|e| e.into_inner());
217        g.state
218    }
219
220    /// Return the failure rate (0.0–1.0) in the current sliding window.
221    ///
222    /// Returns `0.0` if no calls have been recorded yet.
223    pub fn failure_rate(&self) -> f64 {
224        let mut g = self.inner.lock().unwrap_or_else(|e| e.into_inner());
225        g.evict_old(self.config.window_duration);
226        let total = g.total_count();
227        if total == 0 {
228            return 0.0;
229        }
230        g.failure_count() as f64 / total as f64
231    }
232}
233
234// ── CircuitBreakerRegistry ────────────────────────────────────────────────────
235
236/// A registry of circuit breakers keyed by model name.
237///
238/// New entries are created lazily with the supplied default configuration.
239pub struct CircuitBreakerRegistry {
240    map: DashMap<String, Arc<CircuitBreaker>>,
241    default_config: CircuitBreakerConfig,
242}
243
244impl CircuitBreakerRegistry {
245    /// Create a registry; new breakers will be created with `default_config`.
246    pub fn new(default_config: CircuitBreakerConfig) -> Self {
247        Self {
248            map: DashMap::new(),
249            default_config,
250        }
251    }
252
253    /// Get (or lazily create) the circuit breaker for `model`.
254    pub fn get(&self, model: &str) -> Arc<CircuitBreaker> {
255        if let Some(cb) = self.map.get(model) {
256            return Arc::clone(&cb);
257        }
258        let cb = Arc::new(CircuitBreaker::new(self.default_config.clone()));
259        self.map.insert(model.to_string(), Arc::clone(&cb));
260        cb
261    }
262
263    /// Register a circuit breaker for `model` with a specific configuration.
264    pub fn register(&self, model: String, config: CircuitBreakerConfig) {
265        self.map
266            .insert(model, Arc::new(CircuitBreaker::new(config)));
267    }
268
269    /// Return the number of models tracked by this registry.
270    pub fn len(&self) -> usize {
271        self.map.len()
272    }
273
274    /// Returns `true` if no models have been registered yet.
275    pub fn is_empty(&self) -> bool {
276        self.map.is_empty()
277    }
278}
279
280// ── Unit tests ────────────────────────────────────────────────────────────────
281
282#[cfg(test)]
283mod tests {
284    use super::*;
285    use std::thread;
286
287    fn config_with(
288        failure_threshold: usize,
289        success_threshold: usize,
290        min_calls: usize,
291    ) -> CircuitBreakerConfig {
292        CircuitBreakerConfig {
293            failure_threshold,
294            success_threshold,
295            timeout_duration: Duration::from_millis(50),
296            window_duration: Duration::from_secs(60),
297            min_calls,
298        }
299    }
300
301    #[test]
302    fn starts_closed_and_allows_requests() {
303        let cb = CircuitBreaker::new(CircuitBreakerConfig::default());
304        assert_eq!(cb.state(), CircuitState::Closed);
305        assert!(cb.is_allowed());
306    }
307
308    #[test]
309    fn closed_to_open_on_failure_threshold() {
310        // threshold=3, min_calls=3
311        let cb = CircuitBreaker::new(config_with(3, 2, 3));
312        cb.record_failure();
313        cb.record_failure();
314        assert_eq!(cb.state(), CircuitState::Closed, "not enough failures yet");
315        cb.record_failure();
316        assert_eq!(cb.state(), CircuitState::Open);
317        assert!(!cb.is_allowed());
318    }
319
320    #[test]
321    fn open_to_half_open_after_timeout() {
322        let cb = CircuitBreaker::new(config_with(3, 2, 3));
323        for _ in 0..3 {
324            cb.record_failure();
325        }
326        assert_eq!(cb.state(), CircuitState::Open);
327        // Wait for timeout (50ms).
328        thread::sleep(Duration::from_millis(80));
329        assert!(cb.is_allowed(), "should probe after timeout");
330        assert_eq!(cb.state(), CircuitState::HalfOpen);
331    }
332
333    #[test]
334    fn half_open_to_closed_on_successes() {
335        let cb = CircuitBreaker::new(config_with(3, 2, 3));
336        for _ in 0..3 {
337            cb.record_failure();
338        }
339        thread::sleep(Duration::from_millis(80));
340        let _ = cb.is_allowed(); // transition to HalfOpen
341        cb.record_success();
342        assert_eq!(cb.state(), CircuitState::HalfOpen, "one success not enough");
343        cb.record_success();
344        assert_eq!(cb.state(), CircuitState::Closed);
345    }
346
347    #[test]
348    fn half_open_to_open_on_failure() {
349        let cb = CircuitBreaker::new(config_with(3, 2, 3));
350        for _ in 0..3 {
351            cb.record_failure();
352        }
353        thread::sleep(Duration::from_millis(80));
354        let _ = cb.is_allowed(); // transition to HalfOpen
355        cb.record_failure(); // immediate reopen
356        assert_eq!(cb.state(), CircuitState::Open);
357    }
358
359    #[test]
360    fn failure_rate_calculation() {
361        let cb = CircuitBreaker::new(CircuitBreakerConfig::default());
362        assert_eq!(cb.failure_rate(), 0.0);
363        cb.record_success();
364        cb.record_failure();
365        let rate = cb.failure_rate();
366        assert!((rate - 0.5).abs() < f64::EPSILON);
367    }
368
369    #[test]
370    fn successes_do_not_trip_breaker() {
371        let cb = CircuitBreaker::new(config_with(3, 2, 3));
372        for _ in 0..10 {
373            cb.record_success();
374        }
375        assert_eq!(cb.state(), CircuitState::Closed);
376        assert!(cb.is_allowed());
377    }
378
379    #[test]
380    fn registry_lazy_creation() {
381        let registry =
382            CircuitBreakerRegistry::new(CircuitBreakerConfig::default());
383        let cb = registry.get("gpt-4o");
384        assert!(cb.is_allowed());
385        assert_eq!(registry.len(), 1);
386        // Second get returns the same Arc.
387        let cb2 = registry.get("gpt-4o");
388        assert!(Arc::ptr_eq(&cb, &cb2));
389    }
390
391    #[test]
392    fn registry_register_custom_config() {
393        let registry =
394            CircuitBreakerRegistry::new(CircuitBreakerConfig::default());
395        registry.register("claude-3".to_string(), config_with(1, 1, 1));
396        let cb = registry.get("claude-3");
397        cb.record_failure();
398        assert_eq!(cb.state(), CircuitState::Open);
399    }
400}