1use std::time::{Duration, Instant};
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
13pub enum CircuitState {
14 Closed,
16 Open,
18 HalfOpen,
20}
21
22pub trait CircuitBreaker: Send + Sync {
27 fn can_execute(&mut self) -> bool;
29 fn record_success(&mut self);
31 fn record_failure(&mut self);
33 fn state(&self) -> CircuitState;
35 fn reset(&mut self) -> bool;
37}
38
39pub struct DefaultCircuitBreaker {
42 failure_threshold: usize,
43 reset_timeout: Duration,
44 state: CircuitState,
45 consecutive_failures: usize,
46 last_failure_at: Option<Instant>,
47 total_trips: u64,
49}
50
51impl DefaultCircuitBreaker {
52 pub fn new(failure_threshold: usize, reset_timeout: Duration) -> Self {
57 Self {
58 failure_threshold,
59 reset_timeout,
60 state: CircuitState::Closed,
61 consecutive_failures: 0,
62 last_failure_at: None,
63 total_trips: 0,
64 }
65 }
66
67 #[cfg(feature = "prod-circuit-tuning")]
69 pub fn stats(&self) -> CircuitBreakerStats {
70 CircuitBreakerStats {
71 state: self.state,
72 consecutive_failures: self.consecutive_failures,
73 total_trips: self.total_trips,
74 }
75 }
76}
77
78impl CircuitBreaker for DefaultCircuitBreaker {
79 fn can_execute(&mut self) -> bool {
80 match self.state {
81 CircuitState::Closed => true,
82 CircuitState::HalfOpen => true,
83 CircuitState::Open => {
84 let elapsed = self
85 .last_failure_at
86 .map(|t| t.elapsed())
87 .unwrap_or(Duration::ZERO);
88 if elapsed >= self.reset_timeout {
89 self.state = CircuitState::HalfOpen;
90 true
91 } else {
92 false
93 }
94 }
95 }
96 }
97
98 fn record_success(&mut self) {
99 self.consecutive_failures = 0;
100 self.state = CircuitState::Closed;
101 self.last_failure_at = None;
102 }
103
104 fn record_failure(&mut self) {
105 self.consecutive_failures += 1;
106 self.last_failure_at = Some(Instant::now());
107 if self.consecutive_failures >= self.failure_threshold && self.state != CircuitState::Open {
108 self.state = CircuitState::Open;
109 self.total_trips += 1;
110 }
111 }
112
113 fn state(&self) -> CircuitState {
114 self.state
115 }
116
117 fn reset(&mut self) -> bool {
118 let changed = self.state != CircuitState::Closed || self.consecutive_failures != 0;
119 self.state = CircuitState::Closed;
120 self.consecutive_failures = 0;
121 self.last_failure_at = None;
122 changed
123 }
124}
125
126#[cfg(feature = "prod-circuit-tuning")]
131mod prod {
132 use super::CircuitState;
133 use serde::{Deserialize, Serialize};
134 use std::time::Duration;
135
136 #[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
138 pub enum CircuitBreakerProdError {
139 #[error("circuit breaker failure_threshold must be positive")]
141 FailureThresholdNotPositive,
142 #[error("circuit breaker reset_timeout must be positive")]
144 ResetTimeoutNotPositive,
145 }
146
147 #[derive(Debug, Clone, Serialize, Deserialize)]
149 pub struct CircuitBreakerProdConfig {
150 pub failure_threshold: u32,
152 pub reset_timeout: Duration,
154 }
155
156 impl Default for CircuitBreakerProdConfig {
157 fn default() -> Self {
158 Self {
159 failure_threshold: 5,
160 reset_timeout: Duration::from_secs(30),
161 }
162 }
163 }
164
165 impl CircuitBreakerProdConfig {
166 pub fn new(failure_threshold: u32, reset_timeout: Duration) -> Self {
168 Self {
169 failure_threshold,
170 reset_timeout,
171 }
172 }
173
174 pub fn validate(&self) -> Result<(), CircuitBreakerProdError> {
176 if self.failure_threshold == 0 {
177 return Err(CircuitBreakerProdError::FailureThresholdNotPositive);
178 }
179 if self.reset_timeout.is_zero() {
180 return Err(CircuitBreakerProdError::ResetTimeoutNotPositive);
181 }
182 Ok(())
183 }
184 }
185
186 #[derive(Debug, Clone, Serialize, Deserialize)]
188 pub struct CircuitBreakerStats {
189 pub state: CircuitState,
191 pub consecutive_failures: usize,
193 pub total_trips: u64,
195 }
196}
197
198#[cfg(feature = "prod-circuit-tuning")]
199pub use prod::{CircuitBreakerProdConfig, CircuitBreakerProdError, CircuitBreakerStats};
200
201#[cfg(test)]
202mod tests {
203 use super::*;
204
205 #[test]
206 fn test_circuit_breaker_starts_closed() {
207 let mut cb = DefaultCircuitBreaker::new(3, Duration::from_secs(60));
208 assert_eq!(cb.state(), CircuitState::Closed);
209 assert!(cb.can_execute());
210 }
211
212 #[test]
213 fn test_circuit_breaker_trips_after_threshold() {
214 let mut cb = DefaultCircuitBreaker::new(3, Duration::from_secs(60));
215 cb.record_failure();
216 cb.record_failure();
217 assert_eq!(cb.state(), CircuitState::Closed);
218 cb.record_failure();
219 assert_eq!(cb.state(), CircuitState::Open);
220 assert!(!cb.can_execute());
221 }
222
223 #[test]
224 fn test_circuit_breaker_success_resets() {
225 let mut cb = DefaultCircuitBreaker::new(2, Duration::from_secs(60));
226 cb.record_failure();
227 cb.record_success();
228 assert_eq!(cb.state(), CircuitState::Closed);
229 cb.record_failure();
231 assert_eq!(cb.state(), CircuitState::Closed);
232 cb.record_failure();
233 assert_eq!(cb.state(), CircuitState::Open);
234 }
235
236 #[test]
237 fn test_circuit_breaker_half_open_after_timeout() {
238 let mut cb = DefaultCircuitBreaker::new(1, Duration::from_millis(10));
239 cb.record_failure();
240 assert_eq!(cb.state(), CircuitState::Open);
241 std::thread::sleep(Duration::from_millis(20));
243 assert!(cb.can_execute());
244 assert_eq!(cb.state(), CircuitState::HalfOpen);
245 }
246
247 #[test]
248 fn test_circuit_breaker_half_open_success_closes() {
249 let mut cb = DefaultCircuitBreaker::new(1, Duration::from_millis(10));
250 cb.record_failure();
251 std::thread::sleep(Duration::from_millis(20));
252 assert!(cb.can_execute()); cb.record_success();
254 assert_eq!(cb.state(), CircuitState::Closed);
255 }
256
257 #[test]
258 fn test_circuit_breaker_half_open_failure_reopens() {
259 let mut cb = DefaultCircuitBreaker::new(1, Duration::from_millis(10));
260 cb.record_failure();
261 std::thread::sleep(Duration::from_millis(20));
262 assert!(cb.can_execute()); cb.record_failure();
264 assert_eq!(cb.state(), CircuitState::Open);
265 assert!(!cb.can_execute());
266 }
267
268 #[test]
269 fn test_circuit_breaker_reset() {
270 let mut cb = DefaultCircuitBreaker::new(1, Duration::from_secs(60));
271 cb.record_failure();
272 assert_eq!(cb.state(), CircuitState::Open);
273 assert!(cb.reset());
274 assert_eq!(cb.state(), CircuitState::Closed);
275 assert!(cb.can_execute());
276 assert!(!cb.reset());
278 }
279}
280
281#[cfg(all(test, feature = "prod-circuit-tuning"))]
282mod prod_tests {
283 use super::*;
284
285 #[test]
286 fn test_circuit_breaker_prod_config_validate_ok() {
287 let config = CircuitBreakerProdConfig::new(10, Duration::from_secs(60));
288 assert!(config.validate().is_ok());
289 }
290
291 #[test]
292 fn test_circuit_breaker_prod_config_threshold_zero_rejected() {
293 let config = CircuitBreakerProdConfig::new(0, Duration::from_secs(60));
294 let err = config.validate().unwrap_err();
295 assert!(err
296 .to_string()
297 .contains("failure_threshold must be positive"));
298 }
299
300 #[test]
301 fn test_circuit_breaker_prod_config_timeout_zero_rejected() {
302 let config = CircuitBreakerProdConfig::new(10, Duration::ZERO);
303 assert!(config.validate().is_err());
304 }
305
306 #[test]
307 fn test_circuit_breaker_stats_after_trips() {
308 let mut cb = DefaultCircuitBreaker::new(3, Duration::from_secs(60));
309 cb.record_failure();
310 cb.record_failure();
311 cb.record_failure();
312 let stats = cb.stats();
313 assert_eq!(stats.state, CircuitState::Open);
314 assert_eq!(stats.consecutive_failures, 3);
315 assert_eq!(stats.total_trips, 1);
316 }
317
318 #[test]
319 fn test_circuit_breaker_stats_no_trips() {
320 let cb = DefaultCircuitBreaker::new(5, Duration::from_secs(60));
321 let stats = cb.stats();
322 assert_eq!(stats.state, CircuitState::Closed);
323 assert_eq!(stats.total_trips, 0);
324 }
325
326 #[test]
327 fn test_circuit_breaker_total_trips_increments() {
328 let mut cb = DefaultCircuitBreaker::new(1, Duration::from_secs(60));
329 cb.record_failure();
330 assert_eq!(cb.stats().total_trips, 1);
331 cb.record_success();
332 cb.record_failure();
333 assert_eq!(cb.stats().total_trips, 2);
334 }
335}
336pub struct ErrorRateCircuitBreaker {
348 error_threshold: f64,
350 reset_timeout: Duration,
352 half_open_probes: u32,
354 state: CircuitState,
356 window: std::collections::VecDeque<bool>,
358 window_size: usize,
360 probes_in_half_open: u32,
362 probe_successes: u32,
364 last_failure_at: Option<Instant>,
366 total_trips: u64,
368}
369
370impl ErrorRateCircuitBreaker {
371 pub fn new(
378 error_threshold: f64,
379 reset_timeout: Duration,
380 half_open_probes: u32,
381 window_size: usize,
382 ) -> Self {
383 Self {
384 error_threshold,
385 reset_timeout,
386 half_open_probes: half_open_probes.max(1),
387 state: CircuitState::Closed,
388 window: std::collections::VecDeque::with_capacity(window_size),
389 window_size,
390 probes_in_half_open: 0,
391 probe_successes: 0,
392 last_failure_at: None,
393 total_trips: 0,
394 }
395 }
396
397 pub fn error_rate(&self) -> f64 {
399 if self.window.is_empty() {
400 return 0.0;
401 }
402 let failures = self.window.iter().filter(|&&s| !s).count() as f64;
403 failures / self.window.len() as f64
404 }
405
406 pub fn total_trips(&self) -> u64 {
408 self.total_trips
409 }
410
411 pub fn window_samples(&self) -> usize {
413 self.window.len()
414 }
415
416 fn record_result(&mut self, success: bool) {
418 if self.window.len() >= self.window_size {
420 self.window.pop_front();
421 }
422 self.window.push_back(success);
423
424 if !success {
425 self.last_failure_at = Some(Instant::now());
426 }
427
428 match self.state {
429 CircuitState::Closed => {
430 if self.window.len() >= self.window_size && self.error_rate() > self.error_threshold
432 {
433 self.state = CircuitState::Open;
434 self.total_trips += 1;
435 }
436 }
437 CircuitState::HalfOpen => {
438 if success {
440 self.probe_successes += 1;
441 }
442 if !success {
444 self.state = CircuitState::Open;
445 self.probes_in_half_open = 0;
446 self.probe_successes = 0;
447 } else if self.probes_in_half_open >= self.half_open_probes {
448 self.state = CircuitState::Closed;
450 self.probes_in_half_open = 0;
451 self.probe_successes = 0;
452 self.window.clear();
453 }
454 }
455 CircuitState::Open => {}
456 }
457 }
458}
459
460impl CircuitBreaker for ErrorRateCircuitBreaker {
461 fn can_execute(&mut self) -> bool {
462 match self.state {
463 CircuitState::Closed => true,
464 CircuitState::Open => {
465 let elapsed = self
466 .last_failure_at
467 .map(|t| t.elapsed())
468 .unwrap_or(Duration::ZERO);
469 if elapsed >= self.reset_timeout {
470 self.state = CircuitState::HalfOpen;
471 self.probes_in_half_open = 1;
473 self.probe_successes = 0;
474 true
475 } else {
476 false
477 }
478 }
479 CircuitState::HalfOpen => {
480 if self.probes_in_half_open < self.half_open_probes {
482 self.probes_in_half_open += 1;
483 true
484 } else {
485 false
486 }
487 }
488 }
489 }
490
491 fn record_success(&mut self) {
492 self.record_result(true);
493 }
494
495 fn record_failure(&mut self) {
496 self.record_result(false);
497 }
498
499 fn state(&self) -> CircuitState {
500 self.state
501 }
502
503 fn reset(&mut self) -> bool {
504 let changed = self.state != CircuitState::Closed || !self.window.is_empty();
505 self.state = CircuitState::Closed;
506 self.window.clear();
507 self.probes_in_half_open = 0;
508 self.probe_successes = 0;
509 self.last_failure_at = None;
510 changed
511 }
512}