use std::sync::Arc;
use std::sync::atomic::{AtomicU8, AtomicU32, Ordering};
use std::time::{Duration, Instant};
pub trait CircuitBreaker: Send + Sync {
fn check(&self) -> Result<(), BreakerError>;
fn record_success(&self);
fn record_failure(&self);
}
#[derive(Debug, Clone, thiserror::Error)]
#[non_exhaustive]
pub enum BreakerError {
#[error("circuit open: too many consecutive failures")]
Open,
}
impl BreakerError {
pub fn is_retryable(&self) -> bool {
false
}
}
const STATE_CLOSED: u8 = 0;
const STATE_OPEN: u8 = 1;
const STATE_HALF_OPEN: u8 = 2;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BreakerState {
Closed,
Open,
HalfOpen,
}
pub struct DefaultCircuitBreaker {
failure_threshold: u32,
reset_timeout: Duration,
created_at: Instant,
state: AtomicU8,
failure_count: AtomicU32,
last_failure_ms: AtomicU32,
}
impl DefaultCircuitBreaker {
pub fn new(failure_threshold: u32, reset_timeout: Duration) -> Self {
Self {
failure_threshold,
reset_timeout,
created_at: Instant::now(),
state: AtomicU8::new(STATE_CLOSED),
failure_count: AtomicU32::new(0),
last_failure_ms: AtomicU32::new(0),
}
}
pub fn state(&self) -> BreakerState {
match self.state.load(Ordering::Acquire) {
STATE_CLOSED => BreakerState::Closed,
STATE_OPEN => BreakerState::Open,
STATE_HALF_OPEN => BreakerState::HalfOpen,
_ => BreakerState::Closed,
}
}
pub fn failure_count(&self) -> u32 {
self.failure_count.load(Ordering::Acquire)
}
}
impl CircuitBreaker for DefaultCircuitBreaker {
fn check(&self) -> Result<(), BreakerError> {
let state = self.state.load(Ordering::Acquire);
match state {
STATE_CLOSED | STATE_HALF_OPEN => Ok(()),
STATE_OPEN => {
let last_ms = self.last_failure_ms.load(Ordering::Acquire);
let elapsed_ms = self.created_at.elapsed().as_millis() as u64;
if elapsed_ms.saturating_sub(u64::from(last_ms))
>= self.reset_timeout.as_millis() as u64
{
self.state.store(STATE_HALF_OPEN, Ordering::Release);
Ok(())
} else {
Err(BreakerError::Open)
}
}
_ => Ok(()),
}
}
fn record_success(&self) {
self.failure_count.store(0, Ordering::Release);
self.state.store(STATE_CLOSED, Ordering::Release);
}
fn record_failure(&self) {
let count = self.failure_count.fetch_add(1, Ordering::AcqRel) + 1;
self.last_failure_ms.store(
self.created_at.elapsed().as_millis() as u32,
Ordering::Release,
);
if count >= self.failure_threshold {
self.state.store(STATE_OPEN, Ordering::Release);
}
}
}
pub type SharedBreaker = Arc<dyn CircuitBreaker>;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn breaker_starts_closed_allows_calls() {
let b = DefaultCircuitBreaker::new(3, Duration::from_secs(30));
assert_eq!(b.state(), BreakerState::Closed);
assert!(b.check().is_ok());
assert_eq!(b.failure_count(), 0);
}
#[test]
fn breaker_opens_after_threshold_failures() {
let b = DefaultCircuitBreaker::new(3, Duration::from_secs(30));
b.record_failure();
b.record_failure();
assert!(b.check().is_ok(), "2 < 3, still closed");
b.record_failure();
assert_eq!(b.state(), BreakerState::Open);
assert!(b.check().is_err(), "3 >= 3, now open");
}
#[test]
fn breaker_half_opens_after_timeout() {
let b = DefaultCircuitBreaker::new(1, Duration::from_millis(20));
b.record_failure();
assert_eq!(b.state(), BreakerState::Open);
assert!(b.check().is_err(), "still open immediately after trip");
std::thread::sleep(Duration::from_millis(30));
assert!(b.check().is_ok(), "half-open allows trial after timeout");
assert_eq!(b.state(), BreakerState::HalfOpen);
b.record_success();
assert_eq!(b.state(), BreakerState::Closed);
}
#[test]
fn success_resets_failure_count() {
let b = DefaultCircuitBreaker::new(3, Duration::from_secs(30));
b.record_failure();
b.record_failure();
b.record_success();
b.record_failure();
b.record_failure();
assert!(b.check().is_ok(), "only 2 since reset, still closed");
assert_eq!(b.failure_count(), 2);
}
#[test]
fn trait_object_dispatch_works() {
let b: SharedBreaker = Arc::new(DefaultCircuitBreaker::new(2, Duration::from_secs(1)));
b.record_failure();
b.record_failure();
assert!(b.check().is_err());
}
#[test]
fn open_error_is_not_retryable() {
let err = BreakerError::Open;
assert!(!err.is_retryable());
}
}