crawlkit_engine/
circuit_breaker.rs1use std::sync::atomic::{AtomicU32, AtomicU64, AtomicU8, Ordering};
2use std::sync::Arc;
3use std::time::Duration;
4
5use serde::{Deserialize, Serialize};
6
7#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
9pub enum CircuitState {
10 Closed,
12 Open,
14 HalfOpen,
16}
17
18pub struct CircuitBreaker {
29 state: AtomicU8,
30 failure_count: AtomicU32,
31 success_count: AtomicU32,
32 last_failure_time: AtomicU64,
33 config: CircuitBreakerConfig,
34}
35
36#[derive(Debug, Clone)]
38pub struct CircuitBreakerConfig {
39 pub failure_threshold: u32,
41 pub success_threshold: u32,
43 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 #[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 #[must_use]
72 pub fn with_default_config() -> Self {
73 Self::new(CircuitBreakerConfig::default())
74 }
75
76 #[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 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 #[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 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 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 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 #[must_use]
161 pub fn failure_count(&self) -> u32 {
162 self.failure_count.load(Ordering::Acquire)
163 }
164}
165
166pub struct CircuitBreakerRegistry {
168 breakers: dashmap::DashMap<String, Arc<CircuitBreaker>>,
169 config: CircuitBreakerConfig,
170}
171
172impl CircuitBreakerRegistry {
173 #[must_use]
175 pub fn new(config: CircuitBreakerConfig) -> Self {
176 Self {
177 breakers: dashmap::DashMap::new(),
178 config,
179 }
180 }
181
182 #[must_use]
184 pub fn with_default_config() -> Self {
185 Self::new(CircuitBreakerConfig::default())
186 }
187
188 #[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
201pub struct CircuitBreakerRef {
203 inner: Arc<CircuitBreaker>,
204}
205
206impl CircuitBreakerRef {
207 #[must_use]
209 pub fn is_allowed(&self) -> bool {
210 self.inner.is_allowed()
211 }
212
213 pub fn record_success(&self) {
215 self.inner.record_success();
216 }
217
218 pub fn record_failure(&self) {
220 self.inner.record_failure();
221 }
222
223 #[must_use]
225 pub fn state(&self) -> CircuitState {
226 self.inner.state()
227 }
228
229 #[must_use]
231 pub fn failure_count(&self) -> u32 {
232 self.inner.failure_count()
233 }
234}
235
236fn 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#[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(); cb.record_failure();
290 assert_eq!(cb.state(), CircuitState::Closed); }
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}