confium_coordinator/
circuit_breaker.rs1use std::sync::atomic::{AtomicU32, Ordering};
4use std::time::{Duration, Instant};
5
6#[derive(Debug, Clone, Copy, PartialEq, Eq)]
7pub enum CircuitState {
8 Closed,
9 Open,
10 HalfOpen,
11}
12
13pub struct CircuitBreaker {
14 state: std::sync::Mutex<CircuitState>,
15 failure_count: AtomicU32,
16 success_count: AtomicU32,
17 failure_threshold: u32,
18 success_threshold: u32,
19 timeout: Duration,
20 last_failure: std::sync::Mutex<Option<Instant>>,
21}
22
23impl CircuitBreaker {
24 pub fn new(failure_threshold: u32, success_threshold: u32, timeout: Duration) -> Self {
25 Self {
26 state: std::sync::Mutex::new(CircuitState::Closed),
27 failure_count: AtomicU32::new(0),
28 success_count: AtomicU32::new(0),
29 failure_threshold,
30 success_threshold,
31 timeout,
32 last_failure: std::sync::Mutex::new(None),
33 }
34 }
35
36 pub fn state(&self) -> CircuitState {
37 let state = *self.state.lock().unwrap();
38 match state {
39 CircuitState::Open => {
40 if let Some(last) = *self.last_failure.lock().unwrap() {
41 if last.elapsed() >= self.timeout {
42 return CircuitState::HalfOpen;
43 }
44 }
45 CircuitState::Open
46 }
47 _ => state,
48 }
49 }
50
51 pub fn allow_request(&self) -> bool {
52 match self.state() {
53 CircuitState::Closed => true,
54 CircuitState::HalfOpen => true,
55 CircuitState::Open => false,
56 }
57 }
58
59 pub fn record_success(&self) {
60 let observed = self.state();
61 let mut state = self.state.lock().unwrap();
62 match observed {
63 CircuitState::HalfOpen => {
64 let count = self.success_count.fetch_add(1, Ordering::SeqCst) + 1;
65 if count >= self.success_threshold {
66 *state = CircuitState::Closed;
67 self.failure_count.store(0, Ordering::SeqCst);
68 self.success_count.store(0, Ordering::SeqCst);
69 }
70 }
71 CircuitState::Closed => {
72 self.failure_count.store(0, Ordering::SeqCst);
73 }
74 _ => {}
75 }
76 }
77
78 pub fn record_failure(&self) {
79 let mut state = self.state.lock().unwrap();
80 *self.last_failure.lock().unwrap() = Some(Instant::now());
81 match *state {
82 CircuitState::Closed => {
83 let count = self.failure_count.fetch_add(1, Ordering::SeqCst) + 1;
84 if count >= self.failure_threshold {
85 *state = CircuitState::Open;
86 }
87 }
88 CircuitState::HalfOpen => {
89 *state = CircuitState::Open;
90 self.success_count.store(0, Ordering::SeqCst);
91 }
92 _ => {}
93 }
94 }
95
96 pub fn reset(&self) {
97 *self.state.lock().unwrap() = CircuitState::Closed;
98 self.failure_count.store(0, Ordering::SeqCst);
99 self.success_count.store(0, Ordering::SeqCst);
100 }
101}
102
103#[cfg(test)]
104mod tests {
105 use super::*;
106
107 #[test]
108 fn starts_closed() {
109 let cb = CircuitBreaker::new(3, 2, Duration::from_secs(1));
110 assert_eq!(cb.state(), CircuitState::Closed);
111 assert!(cb.allow_request());
112 }
113
114 #[test]
115 fn opens_after_threshold() {
116 let cb = CircuitBreaker::new(3, 2, Duration::from_secs(1));
117 cb.record_failure();
118 cb.record_failure();
119 assert!(cb.allow_request());
120 cb.record_failure();
121 assert_eq!(cb.state(), CircuitState::Open);
122 assert!(!cb.allow_request());
123 }
124
125 #[test]
126 fn success_resets_failure_count() {
127 let cb = CircuitBreaker::new(3, 2, Duration::from_secs(1));
128 cb.record_failure();
129 cb.record_failure();
130 cb.record_success();
131 cb.record_failure();
132 cb.record_failure();
133 assert!(cb.allow_request()); }
135
136 #[test]
137 fn reset_clears_state() {
138 let cb = CircuitBreaker::new(1, 1, Duration::from_secs(1));
139 cb.record_failure();
140 assert_eq!(cb.state(), CircuitState::Open);
141 cb.reset();
142 assert_eq!(cb.state(), CircuitState::Closed);
143 }
144
145 #[test]
146 fn half_open_after_timeout() {
147 let cb = CircuitBreaker::new(1, 1, Duration::from_millis(10));
148 cb.record_failure();
149 assert_eq!(cb.state(), CircuitState::Open);
150 std::thread::sleep(Duration::from_millis(20));
151 assert_eq!(cb.state(), CircuitState::HalfOpen);
152 }
153
154 #[test]
155 fn half_open_success_closes() {
156 let cb = CircuitBreaker::new(1, 1, Duration::from_millis(10));
157 cb.record_failure();
158 std::thread::sleep(Duration::from_millis(20));
159 assert_eq!(cb.state(), CircuitState::HalfOpen);
160 cb.record_success();
161 assert_eq!(cb.state(), CircuitState::Closed);
162 }
163
164 #[test]
165 fn half_open_failure_reopens() {
166 let cb = CircuitBreaker::new(1, 1, Duration::from_millis(10));
167 cb.record_failure();
168 std::thread::sleep(Duration::from_millis(20));
169 cb.record_failure();
170 assert_eq!(cb.state(), CircuitState::Open);
171 }
172}