Files
rustytorch/crates/production/rtx-streaming/src/circuit_breaker.rs
T
2026-03-04 00:08:42 +00:00

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);
}
}