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