tokio_prompt_orchestrator/
circuit_breaker.rs1use std::collections::VecDeque;
28use std::sync::{Arc, Mutex};
29use std::time::{Duration, Instant};
30
31use dashmap::DashMap;
32
33#[derive(Debug, Clone, Copy, PartialEq, Eq)]
37pub enum CircuitState {
38 Closed,
40 Open,
42 HalfOpen,
44}
45
46#[derive(Debug, Clone)]
50pub struct CircuitBreakerConfig {
51 pub failure_threshold: usize,
53 pub success_threshold: usize,
56 pub timeout_duration: Duration,
58 pub window_duration: Duration,
60 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
77struct Inner {
80 state: CircuitState,
81 window: VecDeque<(Instant, bool)>,
83 opened_at: Option<Instant>,
85 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 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 fn failure_count(&self) -> usize {
113 self.window.iter().filter(|(_, ok)| !ok).count()
114 }
115
116 fn total_count(&self) -> usize {
118 self.window.len()
119 }
120}
121
122pub struct CircuitBreaker {
126 config: CircuitBreakerConfig,
127 inner: Mutex<Inner>,
128}
129
130impl CircuitBreaker {
131 pub fn new(config: CircuitBreakerConfig) -> Self {
133 Self {
134 config,
135 inner: Mutex::new(Inner::new()),
136 }
137 }
138
139 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 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 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 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 pub fn state(&self) -> CircuitState {
216 let g = self.inner.lock().unwrap_or_else(|e| e.into_inner());
217 g.state
218 }
219
220 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
234pub struct CircuitBreakerRegistry {
240 map: DashMap<String, Arc<CircuitBreaker>>,
241 default_config: CircuitBreakerConfig,
242}
243
244impl CircuitBreakerRegistry {
245 pub fn new(default_config: CircuitBreakerConfig) -> Self {
247 Self {
248 map: DashMap::new(),
249 default_config,
250 }
251 }
252
253 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 pub fn register(&self, model: String, config: CircuitBreakerConfig) {
265 self.map
266 .insert(model, Arc::new(CircuitBreaker::new(config)));
267 }
268
269 pub fn len(&self) -> usize {
271 self.map.len()
272 }
273
274 pub fn is_empty(&self) -> bool {
276 self.map.is_empty()
277 }
278}
279
280#[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 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 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(); 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(); cb.record_failure(); 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 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}