use parking_lot::RwLock;
use std::future::Future;
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, AtomicU64, Ordering};
use std::time::{Duration, Instant};
use tracing::{debug, info, warn};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CircuitState {
Closed,
Open,
HalfOpen,
}
impl std::fmt::Display for CircuitState {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Closed => write!(f, "Closed"),
Self::Open => write!(f, "Open"),
Self::HalfOpen => write!(f, "HalfOpen"),
}
}
}
#[derive(Debug, Clone)]
pub struct CircuitBreakerConfig {
pub name: String,
pub failure_threshold: u32,
pub success_threshold: u32,
pub reset_timeout: Duration,
pub half_open_requests: u32,
pub failure_window: Duration,
pub automatic_transitions: bool,
}
impl Default for CircuitBreakerConfig {
fn default() -> Self {
Self {
name: "default".to_string(),
failure_threshold: 5,
success_threshold: 3,
reset_timeout: Duration::from_secs(30),
half_open_requests: 3,
failure_window: Duration::from_secs(60),
automatic_transitions: true,
}
}
}
impl CircuitBreakerConfig {
pub fn new(name: impl Into<String>) -> Self {
Self {
name: name.into(),
..Default::default()
}
}
pub fn failure_threshold(mut self, threshold: u32) -> Self {
self.failure_threshold = threshold;
self
}
pub fn success_threshold(mut self, threshold: u32) -> Self {
self.success_threshold = threshold;
self
}
pub fn reset_timeout(mut self, timeout: Duration) -> Self {
self.reset_timeout = timeout;
self
}
pub fn half_open_requests(mut self, count: u32) -> Self {
self.half_open_requests = count;
self
}
pub fn failure_window(mut self, window: Duration) -> Self {
self.failure_window = window;
self
}
}
#[derive(Debug)]
pub enum CircuitBreakerError<E> {
Open,
Execution(E),
HalfOpenLimitReached,
}
impl<E: std::fmt::Display> std::fmt::Display for CircuitBreakerError<E> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Open => write!(f, "Circuit breaker is open"),
Self::Execution(e) => write!(f, "Execution failed: {}", e),
Self::HalfOpenLimitReached => write!(f, "Half-open request limit reached"),
}
}
}
impl<E: std::fmt::Debug + std::fmt::Display> std::error::Error for CircuitBreakerError<E> {}
struct CircuitBreakerState {
state: CircuitState,
opened_at: Option<Instant>,
half_open_at: Option<Instant>,
failure_timestamps: Vec<Instant>,
}
pub struct CircuitBreaker {
config: CircuitBreakerConfig,
inner: RwLock<CircuitBreakerState>,
failure_count: AtomicU32,
success_count: AtomicU32,
half_open_count: AtomicU32,
total_requests: AtomicU64,
total_failures: AtomicU64,
total_successes: AtomicU64,
total_rejections: AtomicU64,
}
impl CircuitBreaker {
pub fn new(mut config: CircuitBreakerConfig) -> Arc<Self> {
if config.success_threshold > config.half_open_requests {
warn!(
name = %config.name,
success_threshold = config.success_threshold,
half_open_requests = config.half_open_requests,
"success_threshold exceeds half_open_requests, clamping"
);
config.success_threshold = config.half_open_requests;
}
info!(
name = %config.name,
failure_threshold = config.failure_threshold,
reset_timeout = ?config.reset_timeout,
"Circuit breaker initialized"
);
Arc::new(Self {
config,
inner: RwLock::new(CircuitBreakerState {
state: CircuitState::Closed,
opened_at: None,
half_open_at: None,
failure_timestamps: Vec::new(),
}),
failure_count: AtomicU32::new(0),
success_count: AtomicU32::new(0),
half_open_count: AtomicU32::new(0),
total_requests: AtomicU64::new(0),
total_failures: AtomicU64::new(0),
total_successes: AtomicU64::new(0),
total_rejections: AtomicU64::new(0),
})
}
pub fn default_circuit() -> Arc<Self> {
Self::new(CircuitBreakerConfig::default())
}
pub fn state(&self) -> CircuitState {
self.maybe_transition_to_half_open();
self.inner.read().state
}
pub fn name(&self) -> &str {
&self.config.name
}
pub fn is_allowed(&self) -> bool {
self.maybe_transition_to_half_open();
let state = self.inner.read().state;
match state {
CircuitState::Closed => true,
CircuitState::Open => false,
CircuitState::HalfOpen => {
let count = self.half_open_count.fetch_add(1, Ordering::SeqCst);
count < self.config.half_open_requests
}
}
}
pub async fn call<F, Fut, T, E>(&self, f: F) -> Result<T, CircuitBreakerError<E>>
where
F: FnOnce() -> Fut,
Fut: Future<Output = Result<T, E>>,
{
self.total_requests.fetch_add(1, Ordering::Relaxed);
if !self.is_allowed() {
self.total_rejections.fetch_add(1, Ordering::Relaxed);
debug!(
name = %self.config.name,
state = %self.state(),
"Circuit breaker rejected request"
);
return Err(CircuitBreakerError::Open);
}
match f().await {
Ok(result) => {
self.record_success();
Ok(result)
}
Err(e) => {
self.record_failure();
Err(CircuitBreakerError::Execution(e))
}
}
}
pub async fn call_with_predicate<F, Fut, T, P>(
&self,
f: F,
is_failure: P,
) -> Result<T, CircuitBreakerError<()>>
where
F: FnOnce() -> Fut,
Fut: Future<Output = T>,
P: FnOnce(&T) -> bool,
{
self.total_requests.fetch_add(1, Ordering::Relaxed);
if !self.is_allowed() {
self.total_rejections.fetch_add(1, Ordering::Relaxed);
return Err(CircuitBreakerError::Open);
}
let result = f().await;
if is_failure(&result) {
self.record_failure();
Err(CircuitBreakerError::Execution(()))
} else {
self.record_success();
Ok(result)
}
}
pub fn record_success(&self) {
self.total_successes.fetch_add(1, Ordering::Relaxed);
let state = self.inner.read().state;
match state {
CircuitState::Closed => {
self.failure_count.store(0, Ordering::SeqCst);
let mut inner = self.inner.write();
inner.failure_timestamps.clear();
}
CircuitState::HalfOpen => {
let successes = self.success_count.fetch_add(1, Ordering::SeqCst) + 1;
if successes >= self.config.success_threshold {
self.close();
}
}
CircuitState::Open => {
debug!(name = %self.config.name, "Success recorded while circuit open");
}
}
}
pub fn record_failure(&self) {
self.total_failures.fetch_add(1, Ordering::Relaxed);
let now = Instant::now();
let state = self.inner.read().state;
match state {
CircuitState::Closed => {
let mut inner = self.inner.write();
let window_start = now - self.config.failure_window;
inner.failure_timestamps.retain(|&t| t > window_start);
inner.failure_timestamps.push(now);
let failure_count = inner.failure_timestamps.len() as u32;
self.failure_count.store(failure_count, Ordering::SeqCst);
if failure_count >= self.config.failure_threshold {
drop(inner); self.open();
}
}
CircuitState::HalfOpen => {
self.open();
}
CircuitState::Open => {
}
}
}
fn open(&self) {
let mut inner = self.inner.write();
if inner.state != CircuitState::Open {
warn!(
name = %self.config.name,
failures = self.failure_count.load(Ordering::SeqCst),
"Circuit breaker OPENED"
);
inner.state = CircuitState::Open;
inner.opened_at = Some(Instant::now());
inner.half_open_at = None;
self.half_open_count.store(0, Ordering::SeqCst);
self.success_count.store(0, Ordering::SeqCst);
}
}
fn close(&self) {
let mut inner = self.inner.write();
if inner.state != CircuitState::Closed {
info!(name = %self.config.name, "Circuit breaker CLOSED");
inner.state = CircuitState::Closed;
inner.opened_at = None;
inner.half_open_at = None;
inner.failure_timestamps.clear();
self.failure_count.store(0, Ordering::SeqCst);
self.success_count.store(0, Ordering::SeqCst);
self.half_open_count.store(0, Ordering::SeqCst);
}
}
fn maybe_transition_to_half_open(&self) {
if !self.config.automatic_transitions {
return;
}
let inner = self.inner.read();
match inner.state {
CircuitState::Closed => {}
CircuitState::Open => {
if let Some(opened_at) = inner.opened_at
&& opened_at.elapsed() >= self.config.reset_timeout
{
drop(inner);
let mut inner = self.inner.write();
if inner.state == CircuitState::Open {
debug!(name = %self.config.name, "Circuit breaker transitioning to HALF-OPEN");
inner.state = CircuitState::HalfOpen;
inner.half_open_at = Some(Instant::now());
self.half_open_count.store(0, Ordering::SeqCst);
self.success_count.store(0, Ordering::SeqCst);
}
}
}
CircuitState::HalfOpen => {
if let Some(half_open_at) = inner.half_open_at
&& half_open_at.elapsed() >= self.config.reset_timeout
{
drop(inner);
let mut inner = self.inner.write();
if inner.state == CircuitState::HalfOpen
&& inner
.half_open_at
.is_some_and(|at| at.elapsed() >= self.config.reset_timeout)
{
debug!(name = %self.config.name, "Circuit breaker re-arming HALF-OPEN probes");
inner.half_open_at = Some(Instant::now());
self.half_open_count.store(0, Ordering::SeqCst);
}
}
}
}
}
pub fn reset(&self) {
self.close();
}
pub fn force_open(&self) {
self.open();
}
pub fn failure_count(&self) -> u32 {
self.failure_count.load(Ordering::SeqCst)
}
pub fn success_count(&self) -> u32 {
self.success_count.load(Ordering::SeqCst)
}
pub fn total_requests(&self) -> u64 {
self.total_requests.load(Ordering::Relaxed)
}
pub fn total_successes(&self) -> u64 {
self.total_successes.load(Ordering::Relaxed)
}
pub fn total_failures(&self) -> u64 {
self.total_failures.load(Ordering::Relaxed)
}
pub fn total_rejections(&self) -> u64 {
self.total_rejections.load(Ordering::Relaxed)
}
pub fn stats(&self) -> CircuitBreakerStats {
CircuitBreakerStats {
name: self.config.name.clone(),
state: self.state(),
total_requests: self.total_requests(),
total_successes: self.total_successes(),
total_failures: self.total_failures(),
total_rejections: self.total_rejections(),
current_failure_count: self.failure_count(),
}
}
}
impl Clone for CircuitBreaker {
fn clone(&self) -> Self {
Self {
config: self.config.clone(),
inner: RwLock::new(CircuitBreakerState {
state: self.inner.read().state,
opened_at: self.inner.read().opened_at,
half_open_at: self.inner.read().half_open_at,
failure_timestamps: self.inner.read().failure_timestamps.clone(),
}),
failure_count: AtomicU32::new(self.failure_count.load(Ordering::SeqCst)),
success_count: AtomicU32::new(self.success_count.load(Ordering::SeqCst)),
half_open_count: AtomicU32::new(self.half_open_count.load(Ordering::SeqCst)),
total_requests: AtomicU64::new(self.total_requests.load(Ordering::Relaxed)),
total_failures: AtomicU64::new(self.total_failures.load(Ordering::Relaxed)),
total_successes: AtomicU64::new(self.total_successes.load(Ordering::Relaxed)),
total_rejections: AtomicU64::new(self.total_rejections.load(Ordering::Relaxed)),
}
}
}
#[derive(Debug, Clone)]
pub struct CircuitBreakerStats {
pub name: String,
pub state: CircuitState,
pub total_requests: u64,
pub total_successes: u64,
pub total_failures: u64,
pub total_rejections: u64,
pub current_failure_count: u32,
}
impl CircuitBreakerStats {
pub fn success_rate(&self) -> f64 {
if self.total_requests == 0 {
1.0
} else {
self.total_successes as f64 / self.total_requests as f64
}
}
pub fn failure_rate(&self) -> f64 {
if self.total_requests == 0 {
0.0
} else {
self.total_failures as f64 / self.total_requests as f64
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_circuit_breaker_opens_after_failures() {
let config = CircuitBreakerConfig {
failure_threshold: 3,
reset_timeout: Duration::from_secs(30),
..Default::default()
};
let cb = CircuitBreaker::new(config);
assert_eq!(cb.state(), CircuitState::Closed);
for _ in 0..3 {
let _: Result<(), CircuitBreakerError<&str>> = cb.call(|| async { Err("error") }).await;
}
assert_eq!(cb.state(), CircuitState::Open);
}
#[tokio::test]
async fn test_circuit_breaker_rejects_when_open() {
let config = CircuitBreakerConfig {
failure_threshold: 1,
..Default::default()
};
let cb = CircuitBreaker::new(config);
let _: Result<(), _> = cb.call(|| async { Err::<(), _>("error") }).await;
let result: Result<(), CircuitBreakerError<&str>> = cb.call(|| async { Ok(()) }).await;
assert!(matches!(result, Err(CircuitBreakerError::Open)));
}
#[tokio::test]
async fn test_circuit_breaker_success_resets_count() {
let config = CircuitBreakerConfig {
failure_threshold: 3,
..Default::default()
};
let cb = CircuitBreaker::new(config);
cb.record_failure();
cb.record_failure();
assert_eq!(cb.failure_count(), 2);
cb.record_success();
assert_eq!(cb.failure_count(), 0);
}
#[tokio::test]
async fn test_circuit_breaker_half_open() {
let config = CircuitBreakerConfig {
failure_threshold: 1,
success_threshold: 2,
reset_timeout: Duration::from_millis(50),
half_open_requests: 3,
..Default::default()
};
let cb = CircuitBreaker::new(config);
cb.record_failure();
assert_eq!(cb.state(), CircuitState::Open);
tokio::time::sleep(Duration::from_millis(100)).await;
assert_eq!(cb.state(), CircuitState::HalfOpen);
cb.record_success();
cb.record_success();
assert_eq!(cb.state(), CircuitState::Closed);
}
#[tokio::test]
async fn test_circuit_breaker_clamps_success_threshold() {
let config = CircuitBreakerConfig {
failure_threshold: 1,
success_threshold: 5,
reset_timeout: Duration::from_millis(50),
half_open_requests: 2,
..Default::default()
};
let cb = CircuitBreaker::new(config);
cb.record_failure();
tokio::time::sleep(Duration::from_millis(100)).await;
assert_eq!(cb.state(), CircuitState::HalfOpen);
for _ in 0..2 {
let result: Result<(), CircuitBreakerError<&str>> = cb.call(|| async { Ok(()) }).await;
assert!(result.is_ok());
}
assert_eq!(cb.state(), CircuitState::Closed);
}
#[tokio::test]
async fn test_circuit_breaker_half_open_rearms_after_stalled_probes() {
let config = CircuitBreakerConfig {
failure_threshold: 1,
success_threshold: 1,
reset_timeout: Duration::from_millis(50),
half_open_requests: 1,
..Default::default()
};
let cb = CircuitBreaker::new(config);
cb.record_failure();
tokio::time::sleep(Duration::from_millis(100)).await;
assert_eq!(cb.state(), CircuitState::HalfOpen);
assert!(cb.is_allowed());
assert!(!cb.is_allowed());
tokio::time::sleep(Duration::from_millis(100)).await;
assert!(cb.is_allowed());
cb.record_success();
assert_eq!(cb.state(), CircuitState::Closed);
}
}