//! Circuit breaker pattern implementation for fault tolerance use parking_lot::RwLock; use serde::{Deserialize, Serialize}; use std::sync::Arc; use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering}; use std::time::{Duration, Instant}; /// Circuit breaker states #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum CircuitState { /// Circuit is closed - normal operation Closed, /// Circuit is open - blocking requests Open, /// Circuit is half-open - testing recovery HalfOpen, } /// Circuit breaker configuration #[derive(Debug, Clone, Serialize, Deserialize)] pub struct CircuitBreakerConfig { /// Failure threshold to open the circuit pub failure_threshold: usize, /// Success threshold to close from half-open pub success_threshold: usize, /// Time window for counting failures pub window_duration: Duration, /// Time to wait before transitioning from open to half-open pub timeout_duration: Duration, /// Maximum requests allowed in half-open state pub half_open_max_requests: usize, } impl Default for CircuitBreakerConfig { fn default() -> Self { Self { failure_threshold: 5, success_threshold: 3, window_duration: Duration::from_secs(60), timeout_duration: Duration::from_secs(30), half_open_max_requests: 3, } } } /// Circuit breaker implementation pub struct CircuitBreaker { config: CircuitBreakerConfig, state: Arc>, failure_count: AtomicUsize, success_count: AtomicUsize, half_open_requests: AtomicUsize, last_failure_time: Arc>>, last_open_time: Arc>>, total_requests: AtomicU64, failed_requests: AtomicU64, } impl CircuitBreaker { /// Create a new circuit breaker pub fn new(config: CircuitBreakerConfig) -> Self { Self { config, state: Arc::new(RwLock::new(CircuitState::Closed)), failure_count: AtomicUsize::new(0), success_count: AtomicUsize::new(0), half_open_requests: AtomicUsize::new(0), last_failure_time: Arc::new(RwLock::new(None)), last_open_time: Arc::new(RwLock::new(None)), total_requests: AtomicU64::new(0), failed_requests: AtomicU64::new(0), } } /// Check if a request is allowed pub fn allow_request(&self) -> bool { let state = *self.state.read(); match state { CircuitState::Closed => true, CircuitState::Open => { // Check if we should transition to half-open if let Some(open_time) = *self.last_open_time.read() { if open_time.elapsed() >= self.config.timeout_duration { *self.state.write() = CircuitState::HalfOpen; self.half_open_requests.store(0, Ordering::SeqCst); true } else { false } } else { false } } CircuitState::HalfOpen => { let requests = self.half_open_requests.fetch_add(1, Ordering::SeqCst); requests < self.config.half_open_max_requests } } } /// Record a successful request pub fn record_success(&self) { self.total_requests.fetch_add(1, Ordering::SeqCst); let state = *self.state.read(); match state { CircuitState::HalfOpen => { let successes = self.success_count.fetch_add(1, Ordering::SeqCst) + 1; if successes >= self.config.success_threshold { *self.state.write() = CircuitState::Closed; self.failure_count.store(0, Ordering::SeqCst); self.success_count.store(0, Ordering::SeqCst); } } CircuitState::Closed => { // Reset failure count on success in closed state self.failure_count.store(0, Ordering::SeqCst); } CircuitState::Open => {} } } /// Record a failed request pub fn record_failure(&self) { self.total_requests.fetch_add(1, Ordering::SeqCst); self.failed_requests.fetch_add(1, Ordering::SeqCst); let now = Instant::now(); // Check if we're within the window let should_count = if let Some(last_failure) = *self.last_failure_time.read() { last_failure.elapsed() < self.config.window_duration } else { true }; if should_count { *self.last_failure_time.write() = Some(now); let state = *self.state.read(); match state { CircuitState::Closed => { let failures = self.failure_count.fetch_add(1, Ordering::SeqCst) + 1; if failures >= self.config.failure_threshold { *self.state.write() = CircuitState::Open; *self.last_open_time.write() = Some(now); self.failure_count.store(0, Ordering::SeqCst); } } CircuitState::HalfOpen => { // Any failure in half-open state reopens the circuit *self.state.write() = CircuitState::Open; *self.last_open_time.write() = Some(now); self.success_count.store(0, Ordering::SeqCst); self.half_open_requests.store(0, Ordering::SeqCst); } CircuitState::Open => {} } } else { // Reset window *self.last_failure_time.write() = Some(now); self.failure_count.store(1, Ordering::SeqCst); } } /// Get current state pub fn state(&self) -> CircuitState { *self.state.read() } /// Get statistics pub fn stats(&self) -> CircuitBreakerStats { CircuitBreakerStats { state: self.state(), total_requests: self.total_requests.load(Ordering::SeqCst), failed_requests: self.failed_requests.load(Ordering::SeqCst), failure_count: self.failure_count.load(Ordering::SeqCst), success_count: self.success_count.load(Ordering::SeqCst), } } /// Reset the circuit breaker pub fn reset(&self) { *self.state.write() = CircuitState::Closed; self.failure_count.store(0, Ordering::SeqCst); self.success_count.store(0, Ordering::SeqCst); self.half_open_requests.store(0, Ordering::SeqCst); *self.last_failure_time.write() = None; *self.last_open_time.write() = None; } /// Check if the circuit breaker is open pub fn is_open(&self) -> bool { matches!(*self.state.read(), CircuitState::Open) } /// Enable fast fail mode pub fn enable_fast_fail(&self) { *self.state.write() = CircuitState::Open; *self.last_open_time.write() = Some(Instant::now()); } /// Get circuit breaker statistics pub fn get_stats(&self) -> CircuitBreakerStats { CircuitBreakerStats { state: *self.state.read(), total_requests: self.total_requests.load(Ordering::SeqCst), failed_requests: self.failed_requests.load(Ordering::SeqCst), failure_count: self.failure_count.load(Ordering::SeqCst), success_count: self.success_count.load(Ordering::SeqCst), } } } /// Circuit breaker statistics #[derive(Debug, Clone, Serialize, Deserialize)] pub struct CircuitBreakerStats { pub state: CircuitState, pub total_requests: u64, pub failed_requests: u64, pub failure_count: usize, pub success_count: usize, } impl Serialize for CircuitState { fn serialize(&self, serializer: S) -> Result where S: serde::Serializer, { serializer.serialize_str(match self { Self::Closed => "closed", Self::Open => "open", Self::HalfOpen => "half_open", }) } } impl<'de> Deserialize<'de> for CircuitState { fn deserialize(deserializer: D) -> Result where D: serde::Deserializer<'de>, { let s = String::deserialize(deserializer)?; match s.as_str() { "closed" => Ok(Self::Closed), "open" => Ok(Self::Open), "half_open" => Ok(Self::HalfOpen), _ => Err(serde::de::Error::custom("invalid circuit state")), } } } #[cfg(test)] mod tests { use super::*; use std::thread; #[test] fn test_circuit_breaker_opens_on_failures() { let config = CircuitBreakerConfig { failure_threshold: 3, ..Default::default() }; let breaker = CircuitBreaker::new(config); assert_eq!(breaker.state(), CircuitState::Closed); // Record failures for _ in 0..3 { assert!(breaker.allow_request()); breaker.record_failure(); } // Circuit should be open now assert_eq!(breaker.state(), CircuitState::Open); assert!(!breaker.allow_request()); } #[test] fn test_circuit_breaker_half_open_transition() { let config = CircuitBreakerConfig { failure_threshold: 1, timeout_duration: Duration::from_millis(100), ..Default::default() }; let breaker = CircuitBreaker::new(config); // Open the circuit breaker.record_failure(); assert_eq!(breaker.state(), CircuitState::Open); // Wait for timeout thread::sleep(Duration::from_millis(150)); // Should allow request and transition to half-open assert!(breaker.allow_request()); assert_eq!(breaker.state(), CircuitState::HalfOpen); } #[test] fn test_circuit_breaker_closes_on_success() { let config = CircuitBreakerConfig { failure_threshold: 1, success_threshold: 2, timeout_duration: Duration::from_millis(100), ..Default::default() }; let breaker = CircuitBreaker::new(config); // Open the circuit breaker.record_failure(); thread::sleep(Duration::from_millis(150)); // Transition to half-open assert!(breaker.allow_request()); // Record successes breaker.record_success(); breaker.record_success(); // Circuit should be closed assert_eq!(breaker.state(), CircuitState::Closed); } }