327 lines
10 KiB
Rust
327 lines
10 KiB
Rust
//! 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<RwLock<CircuitState>>,
|
|
failure_count: AtomicUsize,
|
|
success_count: AtomicUsize,
|
|
half_open_requests: AtomicUsize,
|
|
last_failure_time: Arc<RwLock<Option<Instant>>>,
|
|
last_open_time: Arc<RwLock<Option<Instant>>>,
|
|
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<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
|
|
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<D>(deserializer: D) -> Result<Self, D::Error>
|
|
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);
|
|
}
|
|
}
|