Skip to main content

crawlkit_engine/
circuit_breaker.rs

1use std::sync::atomic::{AtomicU32, AtomicU64, AtomicU8, Ordering};
2use std::sync::Arc;
3use std::time::Duration;
4
5use serde::{Deserialize, Serialize};
6
7/// Circuit breaker state.
8#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
9pub enum CircuitState {
10    /// Circuit is closed — requests are allowed.
11    Closed,
12    /// Circuit is open — requests are blocked.
13    Open,
14    /// Circuit is testing recovery — requests are allowed.
15    HalfOpen,
16}
17
18/// Per-domain circuit breaker for failure isolation.
19///
20/// Prevents cascade failures by breaking the circuit after consecutive failures.
21/// Transitions to Half-Open after cooldown to test recovery.
22///
23/// # States
24///
25/// - **Closed**: Normal operation. Failures increment counter.
26/// - **Open**: Requests are blocked. After cooldown, transitions to Half-Open.
27/// - **Half-Open**: Limited requests allowed. If they succeed, close circuit.
28pub struct CircuitBreaker {
29    state: AtomicU8,
30    failure_count: AtomicU32,
31    success_count: AtomicU32,
32    last_failure_time: AtomicU64,
33    config: CircuitBreakerConfig,
34}
35
36/// Configuration for circuit breaker behavior.
37#[derive(Debug, Clone)]
38pub struct CircuitBreakerConfig {
39    /// Number of consecutive failures before opening circuit.
40    pub failure_threshold: u32,
41    /// Number of successes in Half-Open state before closing circuit.
42    pub success_threshold: u32,
43    /// Duration to wait before transitioning from Open to Half-Open.
44    pub cooldown: Duration,
45}
46
47impl Default for CircuitBreakerConfig {
48    fn default() -> Self {
49        Self {
50            failure_threshold: 5,
51            success_threshold: 3,
52            cooldown: Duration::from_secs(60),
53        }
54    }
55}
56
57impl CircuitBreaker {
58    /// Create a new circuit breaker.
59    #[must_use]
60    pub fn new(config: CircuitBreakerConfig) -> Self {
61        Self {
62            state: AtomicU8::new(CircuitState::Closed as u8),
63            failure_count: AtomicU32::new(0),
64            success_count: AtomicU32::new(0),
65            last_failure_time: AtomicU64::new(0),
66            config,
67        }
68    }
69
70    /// Create with default configuration.
71    #[must_use]
72    pub fn with_default_config() -> Self {
73        Self::new(CircuitBreakerConfig::default())
74    }
75
76    /// Get current state.
77    #[must_use]
78    pub fn state(&self) -> CircuitState {
79        let state = self.state.load(Ordering::Acquire);
80        match state {
81            0 => CircuitState::Closed,
82            1 => {
83                // Check if cooldown expired
84                let last_failure = self.last_failure_time.load(Ordering::Acquire);
85                let now = now_millis();
86                if now.saturating_sub(last_failure) >= self.config.cooldown.as_millis() as u64 {
87                    self.state
88                        .store(CircuitState::HalfOpen as u8, Ordering::Release);
89                    CircuitState::HalfOpen
90                } else {
91                    CircuitState::Open
92                }
93            }
94            2 => CircuitState::HalfOpen,
95            _ => CircuitState::Closed,
96        }
97    }
98
99    /// Check if request is allowed.
100    #[must_use]
101    pub fn is_allowed(&self) -> bool {
102        match self.state() {
103            CircuitState::Closed => true,
104            CircuitState::HalfOpen => true,
105            CircuitState::Open => false,
106        }
107    }
108
109    /// Record a successful request.
110    pub fn record_success(&self) {
111        match self.state() {
112            CircuitState::Closed => {
113                self.failure_count.store(0, Ordering::Release);
114                self.success_count.store(0, Ordering::Release);
115            }
116            CircuitState::HalfOpen => {
117                let successes = self.success_count.fetch_add(1, Ordering::AcqRel) + 1;
118                if successes >= self.config.success_threshold {
119                    self.state
120                        .store(CircuitState::Closed as u8, Ordering::Release);
121                    self.failure_count.store(0, Ordering::Release);
122                    self.success_count.store(0, Ordering::Release);
123                }
124            }
125            CircuitState::Open => {}
126        }
127    }
128
129    /// Record a failed request.
130    pub fn record_failure(&self) {
131        let now = now_millis();
132        self.last_failure_time.store(now, Ordering::Release);
133
134        match self.state() {
135            CircuitState::Closed => {
136                let failures = self.failure_count.fetch_add(1, Ordering::AcqRel) + 1;
137                if failures >= self.config.failure_threshold {
138                    self.state
139                        .store(CircuitState::Open as u8, Ordering::Release);
140                }
141            }
142            CircuitState::HalfOpen => {
143                self.state
144                    .store(CircuitState::Open as u8, Ordering::Release);
145                self.success_count.store(0, Ordering::Release);
146            }
147            CircuitState::Open => {}
148        }
149    }
150
151    /// Reset circuit breaker to Closed state.
152    pub fn reset(&self) {
153        self.state
154            .store(CircuitState::Closed as u8, Ordering::Release);
155        self.failure_count.store(0, Ordering::Release);
156        self.success_count.store(0, Ordering::Release);
157    }
158
159    /// Get failure count.
160    #[must_use]
161    pub fn failure_count(&self) -> u32 {
162        self.failure_count.load(Ordering::Acquire)
163    }
164}
165
166/// Collection of per-domain circuit breakers.
167pub struct CircuitBreakerRegistry {
168    breakers: dashmap::DashMap<String, Arc<CircuitBreaker>>,
169    config: CircuitBreakerConfig,
170}
171
172impl CircuitBreakerRegistry {
173    /// Create registry with shared configuration.
174    #[must_use]
175    pub fn new(config: CircuitBreakerConfig) -> Self {
176        Self {
177            breakers: dashmap::DashMap::new(),
178            config,
179        }
180    }
181
182    /// Create with default configuration.
183    #[must_use]
184    pub fn with_default_config() -> Self {
185        Self::new(CircuitBreakerConfig::default())
186    }
187
188    /// Get or create circuit breaker for domain.
189    #[must_use]
190    pub fn get_or_create(&self, domain: &str) -> CircuitBreakerRef {
191        let entry = self
192            .breakers
193            .entry(domain.to_string())
194            .or_insert_with(|| Arc::new(CircuitBreaker::new(self.config.clone())));
195        CircuitBreakerRef {
196            inner: entry.value().clone(),
197        }
198    }
199}
200
201/// Reference to a domain's circuit breaker.
202pub struct CircuitBreakerRef {
203    inner: Arc<CircuitBreaker>,
204}
205
206impl CircuitBreakerRef {
207    /// Check if request is allowed.
208    #[must_use]
209    pub fn is_allowed(&self) -> bool {
210        self.inner.is_allowed()
211    }
212
213    /// Record success.
214    pub fn record_success(&self) {
215        self.inner.record_success();
216    }
217
218    /// Record failure.
219    pub fn record_failure(&self) {
220        self.inner.record_failure();
221    }
222
223    /// Get current state.
224    #[must_use]
225    pub fn state(&self) -> CircuitState {
226        self.inner.state()
227    }
228
229    /// Get failure count.
230    #[must_use]
231    pub fn failure_count(&self) -> u32 {
232        self.inner.failure_count()
233    }
234}
235
236/// Get current time in milliseconds since epoch.
237fn now_millis() -> u64 {
238    std::time::SystemTime::now()
239        .duration_since(std::time::UNIX_EPOCH)
240        .unwrap_or_default()
241        .as_millis() as u64
242}
243
244// ---------------------------------------------------------------------------
245// Tests
246// ---------------------------------------------------------------------------
247
248#[cfg(test)]
249mod tests {
250    use super::*;
251
252    #[test]
253    fn test_circuit_breaker_closed_by_default() {
254        let cb = CircuitBreaker::with_default_config();
255        assert_eq!(cb.state(), CircuitState::Closed);
256        assert!(cb.is_allowed());
257    }
258
259    #[test]
260    fn test_circuit_breaker_opens_after_failures() {
261        let config = CircuitBreakerConfig {
262            failure_threshold: 3,
263            ..Default::default()
264        };
265        let cb = CircuitBreaker::new(config);
266
267        cb.record_failure();
268        assert_eq!(cb.state(), CircuitState::Closed);
269
270        cb.record_failure();
271        assert_eq!(cb.state(), CircuitState::Closed);
272
273        cb.record_failure();
274        assert_eq!(cb.state(), CircuitState::Open);
275        assert!(!cb.is_allowed());
276    }
277
278    #[test]
279    fn test_circuit_breaker_success_resets_count() {
280        let config = CircuitBreakerConfig {
281            failure_threshold: 3,
282            ..Default::default()
283        };
284        let cb = CircuitBreaker::new(config);
285
286        cb.record_failure();
287        cb.record_failure();
288        cb.record_success(); // Resets failure count
289        cb.record_failure();
290        assert_eq!(cb.state(), CircuitState::Closed); // Not open
291    }
292
293    #[test]
294    fn test_circuit_breaker_registry() {
295        let registry = CircuitBreakerRegistry::with_default_config();
296        let cb = registry.get_or_create("example.com");
297        assert!(cb.is_allowed());
298        cb.record_failure();
299        assert_eq!(cb.failure_count(), 1);
300    }
301}