use std::time::{Duration, Instant};
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub enum CircuitState {
Closed,
Open,
HalfOpen,
}
pub trait CircuitBreaker: Send + Sync {
fn can_execute(&mut self) -> bool;
fn record_success(&mut self);
fn record_failure(&mut self);
fn state(&self) -> CircuitState;
fn reset(&mut self) -> bool;
}
pub struct DefaultCircuitBreaker {
failure_threshold: usize,
reset_timeout: Duration,
state: CircuitState,
consecutive_failures: usize,
last_failure_at: Option<Instant>,
total_trips: u64,
}
impl DefaultCircuitBreaker {
pub fn new(failure_threshold: usize, reset_timeout: Duration) -> Self {
Self {
failure_threshold,
reset_timeout,
state: CircuitState::Closed,
consecutive_failures: 0,
last_failure_at: None,
total_trips: 0,
}
}
#[cfg(feature = "prod-circuit-tuning")]
pub fn stats(&self) -> CircuitBreakerStats {
CircuitBreakerStats {
state: self.state,
consecutive_failures: self.consecutive_failures,
total_trips: self.total_trips,
}
}
}
impl CircuitBreaker for DefaultCircuitBreaker {
fn can_execute(&mut self) -> bool {
match self.state {
CircuitState::Closed => true,
CircuitState::HalfOpen => true,
CircuitState::Open => {
let elapsed = self
.last_failure_at
.map(|t| t.elapsed())
.unwrap_or(Duration::ZERO);
if elapsed >= self.reset_timeout {
self.state = CircuitState::HalfOpen;
true
} else {
false
}
}
}
}
fn record_success(&mut self) {
self.consecutive_failures = 0;
self.state = CircuitState::Closed;
self.last_failure_at = None;
}
fn record_failure(&mut self) {
self.consecutive_failures += 1;
self.last_failure_at = Some(Instant::now());
if self.consecutive_failures >= self.failure_threshold && self.state != CircuitState::Open {
self.state = CircuitState::Open;
self.total_trips += 1;
}
}
fn state(&self) -> CircuitState {
self.state
}
fn reset(&mut self) -> bool {
let changed = self.state != CircuitState::Closed || self.consecutive_failures != 0;
self.state = CircuitState::Closed;
self.consecutive_failures = 0;
self.last_failure_at = None;
changed
}
}
#[cfg(feature = "prod-circuit-tuning")]
mod prod {
use super::CircuitState;
use serde::{Deserialize, Serialize};
use std::time::Duration;
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub enum CircuitBreakerProdError {
#[error("circuit breaker failure_threshold must be positive")]
FailureThresholdNotPositive,
#[error("circuit breaker reset_timeout must be positive")]
ResetTimeoutNotPositive,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CircuitBreakerProdConfig {
pub failure_threshold: u32,
pub reset_timeout: Duration,
}
impl Default for CircuitBreakerProdConfig {
fn default() -> Self {
Self {
failure_threshold: 5,
reset_timeout: Duration::from_secs(30),
}
}
}
impl CircuitBreakerProdConfig {
pub fn new(failure_threshold: u32, reset_timeout: Duration) -> Self {
Self {
failure_threshold,
reset_timeout,
}
}
pub fn validate(&self) -> Result<(), CircuitBreakerProdError> {
if self.failure_threshold == 0 {
return Err(CircuitBreakerProdError::FailureThresholdNotPositive);
}
if self.reset_timeout.is_zero() {
return Err(CircuitBreakerProdError::ResetTimeoutNotPositive);
}
Ok(())
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CircuitBreakerStats {
pub state: CircuitState,
pub consecutive_failures: usize,
pub total_trips: u64,
}
}
#[cfg(feature = "prod-circuit-tuning")]
pub use prod::{CircuitBreakerProdConfig, CircuitBreakerProdError, CircuitBreakerStats};
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_circuit_breaker_starts_closed() {
let mut cb = DefaultCircuitBreaker::new(3, Duration::from_secs(60));
assert_eq!(cb.state(), CircuitState::Closed);
assert!(cb.can_execute());
}
#[test]
fn test_circuit_breaker_trips_after_threshold() {
let mut cb = DefaultCircuitBreaker::new(3, Duration::from_secs(60));
cb.record_failure();
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Closed);
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
assert!(!cb.can_execute());
}
#[test]
fn test_circuit_breaker_success_resets() {
let mut cb = DefaultCircuitBreaker::new(2, Duration::from_secs(60));
cb.record_failure();
cb.record_success();
assert_eq!(cb.state(), CircuitState::Closed);
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Closed);
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
}
#[test]
fn test_circuit_breaker_half_open_after_timeout() {
let mut cb = DefaultCircuitBreaker::new(1, Duration::from_millis(10));
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
std::thread::sleep(Duration::from_millis(20));
assert!(cb.can_execute());
assert_eq!(cb.state(), CircuitState::HalfOpen);
}
#[test]
fn test_circuit_breaker_half_open_success_closes() {
let mut cb = DefaultCircuitBreaker::new(1, Duration::from_millis(10));
cb.record_failure();
std::thread::sleep(Duration::from_millis(20));
assert!(cb.can_execute()); cb.record_success();
assert_eq!(cb.state(), CircuitState::Closed);
}
#[test]
fn test_circuit_breaker_half_open_failure_reopens() {
let mut cb = DefaultCircuitBreaker::new(1, Duration::from_millis(10));
cb.record_failure();
std::thread::sleep(Duration::from_millis(20));
assert!(cb.can_execute()); cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
assert!(!cb.can_execute());
}
#[test]
fn test_circuit_breaker_reset() {
let mut cb = DefaultCircuitBreaker::new(1, Duration::from_secs(60));
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
assert!(cb.reset());
assert_eq!(cb.state(), CircuitState::Closed);
assert!(cb.can_execute());
assert!(!cb.reset());
}
}
#[cfg(all(test, feature = "prod-circuit-tuning"))]
mod prod_tests {
use super::*;
#[test]
fn test_circuit_breaker_prod_config_validate_ok() {
let config = CircuitBreakerProdConfig::new(10, Duration::from_secs(60));
assert!(config.validate().is_ok());
}
#[test]
fn test_circuit_breaker_prod_config_threshold_zero_rejected() {
let config = CircuitBreakerProdConfig::new(0, Duration::from_secs(60));
let err = config.validate().unwrap_err();
assert!(err
.to_string()
.contains("failure_threshold must be positive"));
}
#[test]
fn test_circuit_breaker_prod_config_timeout_zero_rejected() {
let config = CircuitBreakerProdConfig::new(10, Duration::ZERO);
assert!(config.validate().is_err());
}
#[test]
fn test_circuit_breaker_stats_after_trips() {
let mut cb = DefaultCircuitBreaker::new(3, Duration::from_secs(60));
cb.record_failure();
cb.record_failure();
cb.record_failure();
let stats = cb.stats();
assert_eq!(stats.state, CircuitState::Open);
assert_eq!(stats.consecutive_failures, 3);
assert_eq!(stats.total_trips, 1);
}
#[test]
fn test_circuit_breaker_stats_no_trips() {
let cb = DefaultCircuitBreaker::new(5, Duration::from_secs(60));
let stats = cb.stats();
assert_eq!(stats.state, CircuitState::Closed);
assert_eq!(stats.total_trips, 0);
}
#[test]
fn test_circuit_breaker_total_trips_increments() {
let mut cb = DefaultCircuitBreaker::new(1, Duration::from_secs(60));
cb.record_failure();
assert_eq!(cb.stats().total_trips, 1);
cb.record_success();
cb.record_failure();
assert_eq!(cb.stats().total_trips, 2);
}
}
pub struct ErrorRateCircuitBreaker {
error_threshold: f64,
reset_timeout: Duration,
half_open_probes: u32,
state: CircuitState,
window: std::collections::VecDeque<bool>,
window_size: usize,
probes_in_half_open: u32,
probe_successes: u32,
last_failure_at: Option<Instant>,
total_trips: u64,
}
impl ErrorRateCircuitBreaker {
pub fn new(
error_threshold: f64,
reset_timeout: Duration,
half_open_probes: u32,
window_size: usize,
) -> Self {
Self {
error_threshold,
reset_timeout,
half_open_probes: half_open_probes.max(1),
state: CircuitState::Closed,
window: std::collections::VecDeque::with_capacity(window_size),
window_size,
probes_in_half_open: 0,
probe_successes: 0,
last_failure_at: None,
total_trips: 0,
}
}
pub fn error_rate(&self) -> f64 {
if self.window.is_empty() {
return 0.0;
}
let failures = self.window.iter().filter(|&&s| !s).count() as f64;
failures / self.window.len() as f64
}
pub fn total_trips(&self) -> u64 {
self.total_trips
}
pub fn window_samples(&self) -> usize {
self.window.len()
}
fn record_result(&mut self, success: bool) {
if self.window.len() >= self.window_size {
self.window.pop_front();
}
self.window.push_back(success);
if !success {
self.last_failure_at = Some(Instant::now());
}
match self.state {
CircuitState::Closed => {
if self.window.len() >= self.window_size && self.error_rate() > self.error_threshold
{
self.state = CircuitState::Open;
self.total_trips += 1;
}
}
CircuitState::HalfOpen => {
if success {
self.probe_successes += 1;
}
if !success {
self.state = CircuitState::Open;
self.probes_in_half_open = 0;
self.probe_successes = 0;
} else if self.probes_in_half_open >= self.half_open_probes {
self.state = CircuitState::Closed;
self.probes_in_half_open = 0;
self.probe_successes = 0;
self.window.clear();
}
}
CircuitState::Open => {}
}
}
}
impl CircuitBreaker for ErrorRateCircuitBreaker {
fn can_execute(&mut self) -> bool {
match self.state {
CircuitState::Closed => true,
CircuitState::Open => {
let elapsed = self
.last_failure_at
.map(|t| t.elapsed())
.unwrap_or(Duration::ZERO);
if elapsed >= self.reset_timeout {
self.state = CircuitState::HalfOpen;
self.probes_in_half_open = 1;
self.probe_successes = 0;
true
} else {
false
}
}
CircuitState::HalfOpen => {
if self.probes_in_half_open < self.half_open_probes {
self.probes_in_half_open += 1;
true
} else {
false
}
}
}
}
fn record_success(&mut self) {
self.record_result(true);
}
fn record_failure(&mut self) {
self.record_result(false);
}
fn state(&self) -> CircuitState {
self.state
}
fn reset(&mut self) -> bool {
let changed = self.state != CircuitState::Closed || !self.window.is_empty();
self.state = CircuitState::Closed;
self.window.clear();
self.probes_in_half_open = 0;
self.probe_successes = 0;
self.last_failure_at = None;
changed
}
}