use parking_lot::RwLock;
use std::sync::Arc;
use std::time::{Duration, Instant};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CircuitState {
Closed,
Open,
HalfOpen,
}
#[derive(Debug, Clone)]
pub struct CircuitBreakerConfig {
pub failure_threshold: u32,
pub window_duration: Duration,
pub timeout_duration: Duration,
pub success_threshold: u32,
}
impl Default for CircuitBreakerConfig {
fn default() -> Self {
Self {
failure_threshold: 5,
window_duration: Duration::from_secs(60),
timeout_duration: Duration::from_secs(30),
success_threshold: 2,
}
}
}
#[derive(Debug)]
struct CircuitBreakerState {
state: CircuitState,
failure_count: u32,
success_count: u32,
last_failure_time: Option<Instant>,
last_state_change: Instant,
}
impl Default for CircuitBreakerState {
fn default() -> Self {
Self {
state: CircuitState::Closed,
failure_count: 0,
success_count: 0,
last_failure_time: None,
last_state_change: Instant::now(),
}
}
}
#[derive(Clone)]
pub struct CircuitBreaker {
config: CircuitBreakerConfig,
state: Arc<RwLock<CircuitBreakerState>>,
}
impl CircuitBreaker {
pub fn new() -> Self {
Self::with_config(CircuitBreakerConfig::default())
}
pub fn with_config(config: CircuitBreakerConfig) -> Self {
Self {
config,
state: Arc::new(RwLock::new(CircuitBreakerState::default())),
}
}
pub fn allow_request(&self) -> bool {
let state = self.state.read();
match state.state {
CircuitState::Closed => true,
CircuitState::Open => {
let elapsed = state.last_state_change.elapsed();
if elapsed >= self.config.timeout_duration {
drop(state);
self.transition_to_half_open();
true
} else {
false
}
}
CircuitState::HalfOpen => true,
}
}
pub fn record_success(&self) {
let mut state = self.state.write();
match state.state {
CircuitState::Closed => {
state.failure_count = 0;
state.last_failure_time = None;
}
CircuitState::HalfOpen => {
state.success_count += 1;
if state.success_count >= self.config.success_threshold {
state.state = CircuitState::Closed;
state.failure_count = 0;
state.success_count = 0;
state.last_failure_time = None;
state.last_state_change = Instant::now();
tracing::info!("Circuit breaker closed - service recovered");
}
}
CircuitState::Open => {
tracing::warn!("Received success while circuit is open");
}
}
}
pub fn record_failure(&self) {
let mut state = self.state.write();
let now = Instant::now();
match state.state {
CircuitState::Closed => {
if let Some(last_failure) = state.last_failure_time {
if now.duration_since(last_failure) > self.config.window_duration {
state.failure_count = 1;
state.last_failure_time = Some(now);
return;
}
} else {
state.last_failure_time = Some(now);
}
state.failure_count += 1;
if state.failure_count >= self.config.failure_threshold {
state.state = CircuitState::Open;
state.last_state_change = now;
tracing::warn!(
failure_count = state.failure_count,
"Circuit breaker opened - too many failures"
);
}
}
CircuitState::HalfOpen => {
state.state = CircuitState::Open;
state.failure_count = self.config.failure_threshold;
state.success_count = 0;
state.last_state_change = now;
tracing::warn!("Circuit breaker reopened - recovery attempt failed");
}
CircuitState::Open => {
state.last_failure_time = Some(now);
}
}
}
pub fn state(&self) -> CircuitState {
self.state.read().state
}
pub fn stats(&self) -> CircuitBreakerStats {
let state = self.state.read();
CircuitBreakerStats {
state: state.state,
failure_count: state.failure_count,
success_count: state.success_count,
time_in_current_state: state.last_state_change.elapsed(),
}
}
pub fn reset(&self) {
let mut state = self.state.write();
state.state = CircuitState::Closed;
state.failure_count = 0;
state.success_count = 0;
state.last_failure_time = None;
state.last_state_change = Instant::now();
tracing::info!("Circuit breaker manually reset");
}
fn transition_to_half_open(&self) {
let mut state = self.state.write();
if state.state == CircuitState::Open {
state.state = CircuitState::HalfOpen;
state.success_count = 0;
state.last_state_change = Instant::now();
tracing::info!("Circuit breaker transitioning to half-open - testing recovery");
}
}
}
impl Default for CircuitBreaker {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone)]
pub struct CircuitBreakerStats {
pub state: CircuitState,
pub failure_count: u32,
pub success_count: u32,
pub time_in_current_state: Duration,
}
#[cfg(test)]
mod tests {
use super::*;
use std::thread;
#[test]
fn test_circuit_breaker_starts_closed() {
let cb = CircuitBreaker::new();
assert_eq!(cb.state(), CircuitState::Closed);
assert!(cb.allow_request());
}
#[test]
fn test_circuit_opens_after_failures() {
let config = CircuitBreakerConfig {
failure_threshold: 3,
..Default::default()
};
let cb = CircuitBreaker::with_config(config);
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Closed);
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Closed);
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
assert!(!cb.allow_request());
}
#[test]
fn test_circuit_resets_on_success() {
let config = CircuitBreakerConfig {
failure_threshold: 3,
..Default::default()
};
let cb = CircuitBreaker::with_config(config);
cb.record_failure();
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Closed);
cb.record_success();
let stats = cb.stats();
assert_eq!(stats.failure_count, 0);
}
#[test]
fn test_circuit_transitions_to_half_open() {
let config = CircuitBreakerConfig {
failure_threshold: 2,
timeout_duration: Duration::from_millis(100),
..Default::default()
};
let cb = CircuitBreaker::with_config(config);
cb.record_failure();
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
thread::sleep(Duration::from_millis(150));
assert!(cb.allow_request());
assert_eq!(cb.state(), CircuitState::HalfOpen);
}
#[test]
fn test_circuit_closes_from_half_open() {
let config = CircuitBreakerConfig {
failure_threshold: 2,
timeout_duration: Duration::from_millis(100),
success_threshold: 2,
..Default::default()
};
let cb = CircuitBreaker::with_config(config);
cb.record_failure();
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
thread::sleep(Duration::from_millis(150));
assert!(cb.allow_request());
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_reopens_on_half_open_failure() {
let config = CircuitBreakerConfig {
failure_threshold: 2,
timeout_duration: Duration::from_millis(100),
..Default::default()
};
let cb = CircuitBreaker::with_config(config);
cb.record_failure();
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
thread::sleep(Duration::from_millis(150));
assert!(cb.allow_request());
assert_eq!(cb.state(), CircuitState::HalfOpen);
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
}
#[test]
fn test_manual_reset() {
let cb = CircuitBreaker::new();
cb.record_failure();
cb.record_failure();
cb.record_failure();
cb.record_failure();
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
cb.reset();
assert_eq!(cb.state(), CircuitState::Closed);
assert!(cb.allow_request());
}
#[test]
fn test_failure_window_expiration() {
let config = CircuitBreakerConfig {
failure_threshold: 3,
window_duration: Duration::from_millis(100),
..Default::default()
};
let cb = CircuitBreaker::with_config(config);
cb.record_failure();
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Closed);
thread::sleep(Duration::from_millis(150));
cb.record_failure();
let stats = cb.stats();
assert_eq!(stats.failure_count, 1);
assert_eq!(cb.state(), CircuitState::Closed);
}
#[test]
fn test_stats() {
let cb = CircuitBreaker::new();
let stats = cb.stats();
assert_eq!(stats.state, CircuitState::Closed);
assert_eq!(stats.failure_count, 0);
assert_eq!(stats.success_count, 0);
cb.record_failure();
let stats = cb.stats();
assert_eq!(stats.failure_count, 1);
}
}