use parking_lot::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)]
struct CircuitBreakerInner {
state: CircuitState,
failure_count: u32,
success_count: u32,
last_failure: Option<Instant>,
}
#[derive(Debug)]
pub struct CircuitBreaker {
inner: Mutex<CircuitBreakerInner>,
config: CircuitBreakerConfig,
}
impl CircuitBreaker {
pub fn new(failure_threshold: u32, timeout: StdDuration, success_threshold: u32) -> Self {
Self {
inner: Mutex::new(CircuitBreakerInner {
state: CircuitState::Closed,
failure_count: 0,
success_count: 0,
last_failure: None,
}),
config: CircuitBreakerConfig {
failure_threshold,
success_threshold,
timeout,
},
}
}
pub fn with_config(config: CircuitBreakerConfig) -> Self {
Self {
inner: Mutex::new(CircuitBreakerInner {
state: CircuitState::Closed,
failure_count: 0,
success_count: 0,
last_failure: None,
}),
config,
}
}
pub fn state(&self) -> CircuitState {
self.inner.lock().state
}
pub fn can_execute(&self) -> bool {
let mut inner = self.inner.lock();
match inner.state {
CircuitState::Closed => true,
CircuitState::Open => {
if let Some(time) = inner.last_failure
&& time.elapsed() >= self.config.timeout
{
inner.state = CircuitState::HalfOpen;
inner.success_count = 0;
return true;
}
false
}
CircuitState::HalfOpen => true,
}
}
pub fn record_success(&self) {
let mut inner = self.inner.lock();
match inner.state {
CircuitState::HalfOpen => {
inner.success_count += 1;
if inner.success_count >= self.config.success_threshold {
inner.state = CircuitState::Closed;
inner.failure_count = 0;
}
}
CircuitState::Open => {
inner.state = CircuitState::Closed;
inner.failure_count = 0;
}
CircuitState::Closed => {
inner.failure_count = 0;
}
}
}
pub fn record_failure(&self) {
let mut inner = self.inner.lock();
inner.last_failure = Some(Instant::now());
inner.failure_count += 1;
match inner.state {
CircuitState::HalfOpen => {
inner.state = CircuitState::Open;
}
CircuitState::Closed => {
if inner.failure_count >= self.config.failure_threshold {
inner.state = CircuitState::Open;
}
}
CircuitState::Open => {
}
}
}
pub fn reset(&self) {
let mut inner = self.inner.lock();
inner.state = CircuitState::Closed;
inner.failure_count = 0;
inner.success_count = 0;
inner.last_failure = None;
}
pub fn failure_count(&self) -> u32 {
self.inner.lock().failure_count
}
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 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 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 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 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 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 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 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 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 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 cb = CircuitBreaker::new(1, StdDuration::from_secs(60), 2);
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
assert!(!cb.can_execute());
}
}