use parking_lot::RwLock;
use std::sync::atomic::{AtomicU32, AtomicU64, Ordering};
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 open_duration: Duration,
pub half_open_successes: u32,
}
impl Default for CircuitBreakerConfig {
fn default() -> Self {
Self {
failure_threshold: 5,
open_duration: Duration::from_secs(30),
half_open_successes: 2,
}
}
}
#[derive(Debug)]
pub struct CircuitBreaker {
config: CircuitBreakerConfig,
state: RwLock<CircuitState>,
failure_count: AtomicU32,
success_count: AtomicU32,
last_failure_time: RwLock<Option<Instant>>,
total_requests: AtomicU64,
total_failures: AtomicU64,
}
impl CircuitBreaker {
pub fn new() -> Self {
Self::with_config(CircuitBreakerConfig::default())
}
pub fn with_config(config: CircuitBreakerConfig) -> Self {
Self {
config,
state: RwLock::new(CircuitState::Closed),
failure_count: AtomicU32::new(0),
success_count: AtomicU32::new(0),
last_failure_time: RwLock::new(None),
total_requests: AtomicU64::new(0),
total_failures: AtomicU64::new(0),
}
}
pub fn state(&self) -> CircuitState {
*self.state.read()
}
pub fn allow_request(&self) -> bool {
self.total_requests.fetch_add(1, Ordering::Relaxed);
let current_state = *self.state.read();
match current_state {
CircuitState::Closed => true,
CircuitState::Open => {
if let Some(last_failure) = *self.last_failure_time.read() {
if last_failure.elapsed() >= self.config.open_duration {
let mut state = self.state.write();
if *state == CircuitState::Open {
*state = CircuitState::HalfOpen;
self.success_count.store(0, Ordering::Relaxed);
drop(state);
return true;
}
}
}
false
}
CircuitState::HalfOpen => true,
}
}
pub fn record_success(&self) {
let current_state = *self.state.read();
match current_state {
CircuitState::Closed => {
self.failure_count.store(0, Ordering::Relaxed);
}
CircuitState::HalfOpen => {
let successes = self.success_count.fetch_add(1, Ordering::Relaxed) + 1;
if successes >= self.config.half_open_successes {
let mut state = self.state.write();
*state = CircuitState::Closed;
self.failure_count.store(0, Ordering::Relaxed);
self.success_count.store(0, Ordering::Relaxed);
}
}
CircuitState::Open => {
}
}
}
pub fn record_failure(&self) {
self.total_failures.fetch_add(1, Ordering::Relaxed);
*self.last_failure_time.write() = Some(Instant::now());
let current_state = *self.state.read();
match current_state {
CircuitState::Closed => {
let failures = self.failure_count.fetch_add(1, Ordering::Relaxed) + 1;
if failures >= self.config.failure_threshold {
let mut state = self.state.write();
*state = CircuitState::Open;
}
}
CircuitState::HalfOpen => {
let mut state = self.state.write();
*state = CircuitState::Open;
self.success_count.store(0, Ordering::Relaxed);
}
CircuitState::Open => {
}
}
}
pub fn reset(&self) {
let mut state = self.state.write();
*state = CircuitState::Closed;
self.failure_count.store(0, Ordering::Relaxed);
self.success_count.store(0, Ordering::Relaxed);
}
pub fn stats(&self) -> CircuitBreakerStats {
CircuitBreakerStats {
state: *self.state.read(),
failure_count: self.failure_count.load(Ordering::Relaxed),
total_requests: self.total_requests.load(Ordering::Relaxed),
total_failures: self.total_failures.load(Ordering::Relaxed),
}
}
}
impl Default for CircuitBreaker {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone)]
pub struct CircuitBreakerStats {
pub state: CircuitState,
pub failure_count: u32,
pub total_requests: u64,
pub total_failures: u64,
}
#[cfg(test)]
mod tests {
use super::*;
#[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();
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_success_resets_failure_count() {
let config = CircuitBreakerConfig {
failure_threshold: 3,
..Default::default()
};
let cb = CircuitBreaker::with_config(config);
cb.record_failure();
cb.record_failure();
cb.record_success();
cb.record_failure();
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Closed);
}
#[test]
fn test_manual_reset() {
let config = CircuitBreakerConfig {
failure_threshold: 1,
..Default::default()
};
let cb = CircuitBreaker::with_config(config);
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
cb.reset();
assert_eq!(cb.state(), CircuitState::Closed);
assert!(cb.allow_request());
}
}