Skip to main content

confium_coordinator/
circuit_breaker.rs

1//! Circuit breaker pattern — fail fast on downstream failures.
2
3use 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()); // still closed
134    }
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}