use anyhow::Result;
use parking_lot::RwLock;
use serde::{Deserialize, Serialize};
use std::sync::atomic::{AtomicU32, AtomicU64, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use thiserror::Error;
#[derive(Debug, Error)]
pub enum CircuitBreakerError {
#[error("Circuit breaker is open")]
CircuitOpen,
#[error("Service call failed: {0}")]
ServiceError(String),
#[error("Timeout after {0:?}")]
Timeout(Duration),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum CircuitState {
Closed,
Open,
HalfOpen,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CircuitBreakerConfig {
pub failure_threshold: u32,
pub failure_window: Duration,
pub success_threshold: f64,
pub recovery_timeout: Duration,
pub request_timeout: Duration,
pub half_open_max_requests: u32,
}
impl Default for CircuitBreakerConfig {
fn default() -> Self {
Self {
failure_threshold: 5,
failure_window: Duration::from_secs(60),
success_threshold: 0.8,
recovery_timeout: Duration::from_secs(30),
request_timeout: Duration::from_secs(10),
half_open_max_requests: 3,
}
}
}
pub struct CircuitBreaker {
config: CircuitBreakerConfig,
state: Arc<RwLock<CircuitState>>,
failure_count: AtomicU32,
success_count: AtomicU32,
total_requests: AtomicU64,
last_failure_time: Arc<RwLock<Option<Instant>>>,
last_state_change: Arc<RwLock<Instant>>,
half_open_requests: AtomicU32,
name: String,
}
impl CircuitBreaker {
pub fn new(name: impl Into<String>, config: CircuitBreakerConfig) -> Self {
Self {
config,
state: Arc::new(RwLock::new(CircuitState::Closed)),
failure_count: AtomicU32::new(0),
success_count: AtomicU32::new(0),
total_requests: AtomicU64::new(0),
last_failure_time: Arc::new(RwLock::new(None)),
last_state_change: Arc::new(RwLock::new(Instant::now())),
half_open_requests: AtomicU32::new(0),
name: name.into(),
}
}
pub fn state(&self) -> CircuitState {
*self.state.read()
}
pub fn stats(&self) -> CircuitBreakerStats {
CircuitBreakerStats {
state: self.state(),
failure_count: self.failure_count.load(Ordering::Relaxed),
success_count: self.success_count.load(Ordering::Relaxed),
total_requests: self.total_requests.load(Ordering::Relaxed),
}
}
pub async fn call<F, T, Fut>(&self, f: F) -> Result<T, CircuitBreakerError>
where
F: FnOnce() -> Fut,
Fut: std::future::Future<Output = Result<T>>,
{
self.total_requests.fetch_add(1, Ordering::Relaxed);
if !self.should_allow_request() {
return Err(CircuitBreakerError::CircuitOpen);
}
let result = tokio::time::timeout(self.config.request_timeout, f()).await;
match result {
Ok(Ok(value)) => {
self.on_success();
Ok(value)
},
Ok(Err(e)) => {
self.on_failure();
Err(CircuitBreakerError::ServiceError(e.to_string()))
},
Err(_) => {
self.on_failure();
Err(CircuitBreakerError::Timeout(self.config.request_timeout))
},
}
}
fn should_allow_request(&self) -> bool {
let state = *self.state.read();
match state {
CircuitState::Closed => true,
CircuitState::Open => {
let last_change = *self.last_state_change.read();
if last_change.elapsed() >= self.config.recovery_timeout {
self.transition_to(CircuitState::HalfOpen);
true
} else {
false
}
},
CircuitState::HalfOpen => {
let current = self.half_open_requests.fetch_add(1, Ordering::SeqCst);
current < self.config.half_open_max_requests
},
}
}
fn on_success(&self) {
self.success_count.fetch_add(1, Ordering::Relaxed);
let state = *self.state.read();
if state == CircuitState::HalfOpen {
let success = self.success_count.load(Ordering::Relaxed);
let failure = self.failure_count.load(Ordering::Relaxed);
let total = success + failure;
if total > 0 {
let success_rate = f64::from(success) / f64::from(total);
if success_rate >= self.config.success_threshold {
self.transition_to(CircuitState::Closed);
}
}
}
}
fn on_failure(&self) {
self.failure_count.fetch_add(1, Ordering::Relaxed);
*self.last_failure_time.write() = Some(Instant::now());
let state = *self.state.read();
match state {
CircuitState::Closed => {
let failures = self.failure_count.load(Ordering::Relaxed);
if failures >= self.config.failure_threshold {
if let Some(first_failure) = *self.last_failure_time.read() {
if first_failure.elapsed() <= self.config.failure_window {
self.transition_to(CircuitState::Open);
} else {
self.reset_counters();
}
}
}
},
CircuitState::HalfOpen => {
self.transition_to(CircuitState::Open);
},
_ => {},
}
}
fn transition_to(&self, new_state: CircuitState) {
let mut state = self.state.write();
if *state != new_state {
tracing::info!(
"Circuit breaker '{}' transitioning from {:?} to {:?}",
self.name,
*state,
new_state
);
*state = new_state;
*self.last_state_change.write() = Instant::now();
if new_state == CircuitState::Closed {
self.reset_counters();
} else if new_state == CircuitState::HalfOpen {
self.half_open_requests.store(0, Ordering::SeqCst);
self.success_count.store(0, Ordering::SeqCst);
self.failure_count.store(0, Ordering::SeqCst);
}
}
}
fn reset_counters(&self) {
self.failure_count.store(0, Ordering::SeqCst);
self.success_count.store(0, Ordering::SeqCst);
self.half_open_requests.store(0, Ordering::SeqCst);
*self.last_failure_time.write() = None;
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CircuitBreakerStats {
pub state: CircuitState,
pub failure_count: u32,
pub success_count: u32,
pub total_requests: u64,
}
pub struct CircuitBreakerRegistry {
breakers: RwLock<std::collections::HashMap<String, Arc<CircuitBreaker>>>,
default_config: CircuitBreakerConfig,
}
impl CircuitBreakerRegistry {
pub fn new(default_config: CircuitBreakerConfig) -> Self {
Self {
breakers: RwLock::new(std::collections::HashMap::new()),
default_config,
}
}
pub fn get_or_create(&self, name: &str) -> Arc<CircuitBreaker> {
let breakers = self.breakers.read();
if let Some(breaker) = breakers.get(name) {
return breaker.clone();
}
drop(breakers);
let mut breakers = self.breakers.write();
breakers
.entry(name.to_string())
.or_insert_with(|| Arc::new(CircuitBreaker::new(name, self.default_config.clone())))
.clone()
}
pub fn all_stats(&self) -> Vec<(String, CircuitBreakerStats)> {
self.breakers
.read()
.iter()
.map(|(name, breaker)| (name.clone(), breaker.stats()))
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_circuit_breaker_normal_operation() {
let config = CircuitBreakerConfig {
failure_threshold: 3,
..Default::default()
};
let breaker = CircuitBreaker::new("test", config);
assert_eq!(breaker.state(), CircuitState::Closed);
let result = breaker
.call(|| async { Ok::<_, anyhow::Error>("success") })
.await;
assert!(result.is_ok());
assert_eq!(result.unwrap(), "success");
}
#[tokio::test]
async fn test_circuit_breaker_opens_on_failures() {
let config = CircuitBreakerConfig {
failure_threshold: 2,
..Default::default()
};
let breaker = CircuitBreaker::new("test", config);
let _ = breaker
.call(|| async { Err::<String, _>(anyhow::anyhow!("failure")) })
.await;
assert_eq!(breaker.state(), CircuitState::Closed);
let _ = breaker
.call(|| async { Err::<String, _>(anyhow::anyhow!("failure")) })
.await;
assert_eq!(breaker.state(), CircuitState::Open);
let result = breaker
.call(|| async { Ok::<_, anyhow::Error>("should not execute") })
.await;
assert!(matches!(result, Err(CircuitBreakerError::CircuitOpen)));
}
#[tokio::test]
async fn test_circuit_breaker_half_open_recovery() {
let config = CircuitBreakerConfig {
failure_threshold: 1,
recovery_timeout: Duration::from_millis(100),
half_open_max_requests: 1,
..Default::default()
};
let breaker = CircuitBreaker::new("test", config);
let _ = breaker
.call(|| async { Err::<String, _>(anyhow::anyhow!("failure")) })
.await;
assert_eq!(breaker.state(), CircuitState::Open);
tokio::time::sleep(Duration::from_millis(150)).await;
let result = breaker
.call(|| async { Ok::<_, anyhow::Error>("recovered") })
.await;
assert!(result.is_ok());
assert_eq!(breaker.state(), CircuitState::Closed);
}
}