use std::sync::{Arc, Mutex};
use std::time::{Duration as StdDuration, Instant};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CircuitState {
Closed,
Open,
HalfOpen,
}
#[derive(Debug, Clone)]
pub struct CircuitBreakerConfig {
pub failure_threshold: u32,
pub success_threshold: u32,
pub timeout: StdDuration,
}
impl Default for CircuitBreakerConfig {
fn default() -> Self {
Self {
failure_threshold: 5,
success_threshold: 3,
timeout: StdDuration::from_secs(30),
}
}
}
#[derive(Debug)]
pub struct CircuitBreaker {
state: Arc<Mutex<CircuitState>>,
failure_count: Arc<Mutex<u32>>,
success_count: Arc<Mutex<u32>>,
last_failure: Arc<Mutex<Option<Instant>>>,
config: CircuitBreakerConfig,
}
impl CircuitBreaker {
pub fn new(failure_threshold: u32, timeout: StdDuration, success_threshold: u32) -> Self {
Self {
state: Arc::new(Mutex::new(CircuitState::Closed)),
failure_count: Arc::new(Mutex::new(0)),
success_count: Arc::new(Mutex::new(0)),
last_failure: Arc::new(Mutex::new(None)),
config: CircuitBreakerConfig {
failure_threshold,
success_threshold,
timeout,
},
}
}
pub fn with_config(config: CircuitBreakerConfig) -> Self {
Self {
state: Arc::new(Mutex::new(CircuitState::Closed)),
failure_count: Arc::new(Mutex::new(0)),
success_count: Arc::new(Mutex::new(0)),
last_failure: Arc::new(Mutex::new(None)),
config,
}
}
pub fn state(&self) -> CircuitState {
self.state
.lock()
.map(|guard| *guard)
.unwrap_or(CircuitState::Closed)
}
pub fn can_execute(&self) -> bool {
let state = self
.state
.lock()
.map(|guard| *guard)
.unwrap_or(CircuitState::Closed);
match state {
CircuitState::Closed => true,
CircuitState::Open => {
let last_failure = self.last_failure.lock().ok().and_then(|guard| *guard);
if let Some(time) = last_failure {
if time.elapsed() >= self.config.timeout {
if let Ok(mut guard) = self.state.lock() {
*guard = CircuitState::HalfOpen;
}
if let Ok(mut guard) = self.success_count.lock() {
*guard = 0;
}
return true;
}
}
false
}
CircuitState::HalfOpen => true,
}
}
pub fn record_success(&mut self) {
if let Ok(mut state_guard) = self.state.lock() {
if let Ok(mut success_count_guard) = self.success_count.lock() {
match *state_guard {
CircuitState::HalfOpen => {
*success_count_guard += 1;
if *success_count_guard >= self.config.success_threshold {
*state_guard = CircuitState::Closed;
if let Ok(mut failure_count_guard) = self.failure_count.lock() {
*failure_count_guard = 0;
}
}
}
CircuitState::Open => {
*state_guard = CircuitState::Closed;
if let Ok(mut failure_count_guard) = self.failure_count.lock() {
*failure_count_guard = 0;
}
}
CircuitState::Closed => {
if let Ok(mut failure_count_guard) = self.failure_count.lock() {
*failure_count_guard = 0;
}
}
}
}
}
}
pub fn record_failure(&mut self) {
if let Ok(mut state_guard) = self.state.lock() {
if let Ok(mut failure_count_guard) = self.failure_count.lock() {
if let Ok(mut last_failure_guard) = self.last_failure.lock() {
*last_failure_guard = Some(Instant::now());
*failure_count_guard += 1;
match *state_guard {
CircuitState::HalfOpen => {
*state_guard = CircuitState::Open;
}
CircuitState::Closed => {
if *failure_count_guard >= self.config.failure_threshold {
*state_guard = CircuitState::Open;
}
}
CircuitState::Open => {
}
}
}
}
}
}
pub fn reset(&mut self) {
if let Ok(mut guard) = self.state.lock() {
*guard = CircuitState::Closed;
}
if let Ok(mut guard) = self.failure_count.lock() {
*guard = 0;
}
if let Ok(mut guard) = self.success_count.lock() {
*guard = 0;
}
if let Ok(mut guard) = self.last_failure.lock() {
*guard = None;
}
}
pub fn failure_count(&self) -> u32 {
self.failure_count.lock().map(|guard| *guard).unwrap_or(0)
}
pub fn config(&self) -> &CircuitBreakerConfig {
&self.config
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_circuit_breaker_initial_state() {
let cb = CircuitBreaker::new(3, StdDuration::from_secs(1), 3);
assert_eq!(cb.state(), CircuitState::Closed);
assert_eq!(cb.failure_count(), 0);
assert!(cb.can_execute());
}
#[test]
fn test_circuit_breaker_open_after_failures() {
let mut cb = CircuitBreaker::new(3, StdDuration::from_secs(1), 3);
assert!(cb.can_execute());
cb.record_failure();
assert!(cb.can_execute());
assert_eq!(cb.state(), CircuitState::Closed);
cb.record_failure();
assert!(cb.can_execute());
cb.record_failure();
assert!(!cb.can_execute());
assert_eq!(cb.state(), CircuitState::Open);
}
#[test]
fn test_circuit_breaker_half_open_after_timeout() {
let mut cb = CircuitBreaker::new(2, StdDuration::from_millis(100), 3);
cb.record_failure();
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
assert!(!cb.can_execute());
std::thread::sleep(std::time::Duration::from_millis(150));
assert!(cb.can_execute());
assert_eq!(cb.state(), CircuitState::HalfOpen);
}
#[test]
fn test_circuit_breaker_close_after_successes() {
let mut cb = CircuitBreaker::new(2, StdDuration::from_millis(100), 3);
cb.record_failure();
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
std::thread::sleep(std::time::Duration::from_millis(150));
assert!(cb.can_execute());
assert_eq!(cb.state(), CircuitState::HalfOpen);
cb.record_success();
assert_eq!(cb.state(), CircuitState::HalfOpen);
cb.record_success();
assert_eq!(cb.state(), CircuitState::HalfOpen);
cb.record_success();
assert_eq!(cb.state(), CircuitState::Closed);
}
#[test]
fn test_circuit_breaker_reset() {
let mut cb = CircuitBreaker::new(2, StdDuration::from_secs(1), 3);
cb.record_failure();
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
cb.reset();
assert_eq!(cb.state(), CircuitState::Closed);
assert_eq!(cb.failure_count(), 0);
assert!(cb.can_execute());
}
#[test]
fn test_circuit_breaker_with_config() {
let config = CircuitBreakerConfig {
failure_threshold: 10,
success_threshold: 5,
timeout: StdDuration::from_secs(60),
};
let cb = CircuitBreaker::with_config(config.clone());
assert_eq!(cb.config().failure_threshold, 10);
assert_eq!(cb.config().success_threshold, 5);
}
#[test]
fn test_circuit_breaker_config_default() {
let config = CircuitBreakerConfig::default();
assert_eq!(config.failure_threshold, 5);
assert_eq!(config.success_threshold, 3);
assert_eq!(config.timeout, StdDuration::from_secs(30));
}
#[test]
fn test_record_success_on_open_state() {
let mut cb = CircuitBreaker::new(2, StdDuration::from_secs(60), 3);
cb.record_failure();
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
cb.record_success();
assert_eq!(cb.state(), CircuitState::Closed);
assert_eq!(cb.failure_count(), 0);
}
#[test]
fn test_record_failure_on_half_open_state() {
let mut cb = CircuitBreaker::new(2, StdDuration::from_millis(100), 3);
cb.record_failure();
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
std::thread::sleep(std::time::Duration::from_millis(150));
assert!(cb.can_execute()); assert_eq!(cb.state(), CircuitState::HalfOpen);
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
}
#[test]
fn test_record_failure_on_open_state() {
let mut cb = CircuitBreaker::new(2, StdDuration::from_secs(60), 3);
cb.record_failure();
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
assert_eq!(cb.failure_count(), 2);
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
assert_eq!(cb.failure_count(), 3);
}
#[test]
fn test_record_success_on_closed_state() {
let mut cb = CircuitBreaker::new(3, StdDuration::from_secs(60), 3);
cb.record_failure();
assert_eq!(cb.failure_count(), 1);
cb.record_success();
assert_eq!(cb.state(), CircuitState::Closed);
assert_eq!(cb.failure_count(), 0);
}
#[test]
fn test_half_open_can_execute() {
let mut cb = CircuitBreaker::new(1, StdDuration::from_millis(50), 2);
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
std::thread::sleep(std::time::Duration::from_millis(100));
assert!(cb.can_execute()); assert_eq!(cb.state(), CircuitState::HalfOpen);
assert!(cb.can_execute());
}
#[test]
fn test_open_state_can_execute_before_timeout() {
let mut cb = CircuitBreaker::new(1, StdDuration::from_secs(60), 2);
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
assert!(!cb.can_execute());
}
}