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}