use crate::error::{Result, TalosError};
use std::future::Future;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::time::{Duration, Instant};
use tokio::sync::RwLock;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CircuitState {
Closed,
Open,
HalfOpen,
}
#[derive(Debug, Clone)]
pub struct CircuitBreakerConfig {
pub failure_threshold: usize,
pub success_threshold: usize,
pub reset_timeout: Duration,
pub half_open_max_requests: usize,
}
impl Default for CircuitBreakerConfig {
fn default() -> Self {
Self {
failure_threshold: 5,
success_threshold: 2,
reset_timeout: Duration::from_secs(30),
half_open_max_requests: 3,
}
}
}
impl CircuitBreakerConfig {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn with_failure_threshold(mut self, threshold: usize) -> Self {
self.failure_threshold = threshold;
self
}
#[must_use]
pub fn with_success_threshold(mut self, threshold: usize) -> Self {
self.success_threshold = threshold;
self
}
#[must_use]
pub fn with_reset_timeout(mut self, timeout: Duration) -> Self {
self.reset_timeout = timeout;
self
}
#[must_use]
pub fn with_half_open_max_requests(mut self, max: usize) -> Self {
self.half_open_max_requests = max;
self
}
}
pub struct CircuitBreaker {
config: CircuitBreakerConfig,
state: RwLock<CircuitState>,
failure_count: AtomicUsize,
success_count: AtomicUsize,
half_open_requests: AtomicUsize,
last_failure_time: RwLock<Option<Instant>>,
opened_at: RwLock<Option<Instant>>,
total_calls: AtomicU64,
total_failures: AtomicU64,
total_rejections: AtomicU64,
}
impl CircuitBreaker {
#[must_use]
pub fn new(config: CircuitBreakerConfig) -> Self {
Self {
config,
state: RwLock::new(CircuitState::Closed),
failure_count: AtomicUsize::new(0),
success_count: AtomicUsize::new(0),
half_open_requests: AtomicUsize::new(0),
last_failure_time: RwLock::new(None),
opened_at: RwLock::new(None),
total_calls: AtomicU64::new(0),
total_failures: AtomicU64::new(0),
total_rejections: AtomicU64::new(0),
}
}
#[must_use]
pub fn with_defaults() -> Self {
Self::new(CircuitBreakerConfig::default())
}
pub async fn state(&self) -> CircuitState {
let current_state = *self.state.read().await;
if current_state == CircuitState::Open {
if let Some(opened_at) = *self.opened_at.read().await {
if opened_at.elapsed() >= self.config.reset_timeout {
let mut state = self.state.write().await;
if *state == CircuitState::Open {
*state = CircuitState::HalfOpen;
self.half_open_requests.store(0, Ordering::Relaxed);
self.success_count.store(0, Ordering::Relaxed);
}
return CircuitState::HalfOpen;
}
}
}
current_state
}
pub async fn can_execute(&self) -> bool {
match self.state().await {
CircuitState::Closed => true,
CircuitState::Open => false,
CircuitState::HalfOpen => {
let current = self.half_open_requests.load(Ordering::Relaxed);
current < self.config.half_open_max_requests
}
}
}
pub async fn call<F, Fut, T>(&self, operation: F) -> Result<T>
where
F: FnOnce() -> Fut,
Fut: Future<Output = Result<T>>,
{
self.total_calls.fetch_add(1, Ordering::Relaxed);
if !self.can_execute().await {
self.total_rejections.fetch_add(1, Ordering::Relaxed);
return Err(TalosError::CircuitOpen(format!(
"Circuit breaker is open, will retry after {:?}",
self.time_until_retry().await
)));
}
let current_state = self.state().await;
if current_state == CircuitState::HalfOpen {
self.half_open_requests.fetch_add(1, Ordering::Relaxed);
}
match operation().await {
Ok(result) => {
self.on_success().await;
Ok(result)
}
Err(e) => {
self.on_failure().await;
Err(e)
}
}
}
async fn on_success(&self) {
let state = *self.state.read().await;
match 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.success_threshold {
let mut state = self.state.write().await;
*state = CircuitState::Closed;
self.failure_count.store(0, Ordering::Relaxed);
self.success_count.store(0, Ordering::Relaxed);
}
}
CircuitState::Open => {
self.failure_count.store(0, Ordering::Relaxed);
}
}
}
async fn on_failure(&self) {
self.total_failures.fetch_add(1, Ordering::Relaxed);
*self.last_failure_time.write().await = Some(Instant::now());
let state = *self.state.read().await;
match state {
CircuitState::Closed => {
let failures = self.failure_count.fetch_add(1, Ordering::Relaxed) + 1;
if failures >= self.config.failure_threshold {
self.open_circuit().await;
}
}
CircuitState::HalfOpen => {
self.open_circuit().await;
}
CircuitState::Open => {
}
}
}
async fn open_circuit(&self) {
let mut state = self.state.write().await;
*state = CircuitState::Open;
*self.opened_at.write().await = Some(Instant::now());
}
pub async fn reset(&self) {
let mut state = self.state.write().await;
*state = CircuitState::Closed;
self.failure_count.store(0, Ordering::Relaxed);
self.success_count.store(0, Ordering::Relaxed);
self.half_open_requests.store(0, Ordering::Relaxed);
*self.opened_at.write().await = None;
}
pub async fn time_until_retry(&self) -> Option<Duration> {
if *self.state.read().await != CircuitState::Open {
return None;
}
self.opened_at.read().await.map(|opened| {
let elapsed = opened.elapsed();
if elapsed >= self.config.reset_timeout {
Duration::ZERO
} else {
self.config.reset_timeout - elapsed
}
})
}
#[must_use]
pub fn failure_count(&self) -> usize {
self.failure_count.load(Ordering::Relaxed)
}
#[must_use]
pub fn total_calls(&self) -> u64 {
self.total_calls.load(Ordering::Relaxed)
}
#[must_use]
pub fn total_failures(&self) -> u64 {
self.total_failures.load(Ordering::Relaxed)
}
#[must_use]
pub fn total_rejections(&self) -> u64 {
self.total_rejections.load(Ordering::Relaxed)
}
#[must_use]
pub fn failure_rate(&self) -> f64 {
let total = self.total_calls.load(Ordering::Relaxed);
if total == 0 {
return 0.0;
}
let failures = self.total_failures.load(Ordering::Relaxed);
failures as f64 / total as f64
}
#[must_use]
pub fn config(&self) -> &CircuitBreakerConfig {
&self.config
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_circuit_breaker_config_default() {
let config = CircuitBreakerConfig::default();
assert_eq!(config.failure_threshold, 5);
assert_eq!(config.success_threshold, 2);
assert_eq!(config.reset_timeout, Duration::from_secs(30));
assert_eq!(config.half_open_max_requests, 3);
}
#[test]
fn test_circuit_breaker_config_builder() {
let config = CircuitBreakerConfig::new()
.with_failure_threshold(10)
.with_success_threshold(5)
.with_reset_timeout(Duration::from_secs(60))
.with_half_open_max_requests(5);
assert_eq!(config.failure_threshold, 10);
assert_eq!(config.success_threshold, 5);
assert_eq!(config.reset_timeout, Duration::from_secs(60));
assert_eq!(config.half_open_max_requests, 5);
}
#[tokio::test]
async fn test_circuit_breaker_initial_state() {
let breaker = CircuitBreaker::with_defaults();
assert_eq!(breaker.state().await, CircuitState::Closed);
assert!(breaker.can_execute().await);
}
#[tokio::test]
async fn test_circuit_breaker_opens_on_failures() {
let config = CircuitBreakerConfig::new().with_failure_threshold(3);
let breaker = CircuitBreaker::new(config);
for _ in 0..3 {
let _ = breaker
.call(|| async { Err::<(), _>(TalosError::Connection("test".to_string())) })
.await;
}
assert_eq!(breaker.state().await, CircuitState::Open);
assert!(!breaker.can_execute().await);
}
#[tokio::test]
async fn test_circuit_breaker_rejects_when_open() {
let config = CircuitBreakerConfig::new()
.with_failure_threshold(2)
.with_reset_timeout(Duration::from_secs(60));
let breaker = CircuitBreaker::new(config);
for _ in 0..2 {
let _ = breaker
.call(|| async { Err::<(), _>(TalosError::Connection("test".to_string())) })
.await;
}
let result = breaker
.call(|| async { Ok::<_, TalosError>("success") })
.await;
assert!(matches!(result, Err(TalosError::CircuitOpen(_))));
assert_eq!(breaker.total_rejections(), 1);
}
#[tokio::test]
async fn test_circuit_breaker_success_resets_failures() {
let config = CircuitBreakerConfig::new().with_failure_threshold(3);
let breaker = CircuitBreaker::new(config);
for _ in 0..2 {
let _ = breaker
.call(|| async { Err::<(), _>(TalosError::Connection("test".to_string())) })
.await;
}
assert_eq!(breaker.failure_count(), 2);
let _ = breaker.call(|| async { Ok::<_, TalosError>("ok") }).await;
assert_eq!(breaker.failure_count(), 0);
}
#[tokio::test]
async fn test_circuit_breaker_reset() {
let config = CircuitBreakerConfig::new().with_failure_threshold(2);
let breaker = CircuitBreaker::new(config);
for _ in 0..2 {
let _ = breaker
.call(|| async { Err::<(), _>(TalosError::Connection("test".to_string())) })
.await;
}
assert_eq!(breaker.state().await, CircuitState::Open);
breaker.reset().await;
assert_eq!(breaker.state().await, CircuitState::Closed);
assert!(breaker.can_execute().await);
}
#[tokio::test]
async fn test_circuit_breaker_half_open_transition() {
let config = CircuitBreakerConfig::new()
.with_failure_threshold(2)
.with_reset_timeout(Duration::from_millis(50));
let breaker = CircuitBreaker::new(config);
for _ in 0..2 {
let _ = breaker
.call(|| async { Err::<(), _>(TalosError::Connection("test".to_string())) })
.await;
}
assert_eq!(breaker.state().await, CircuitState::Open);
tokio::time::sleep(Duration::from_millis(60)).await;
assert_eq!(breaker.state().await, CircuitState::HalfOpen);
}
#[tokio::test]
async fn test_circuit_breaker_closes_after_success_in_half_open() {
let config = CircuitBreakerConfig::new()
.with_failure_threshold(2)
.with_success_threshold(2)
.with_reset_timeout(Duration::from_millis(10));
let breaker = CircuitBreaker::new(config);
for _ in 0..2 {
let _ = breaker
.call(|| async { Err::<(), _>(TalosError::Connection("test".to_string())) })
.await;
}
tokio::time::sleep(Duration::from_millis(20)).await;
assert_eq!(breaker.state().await, CircuitState::HalfOpen);
for _ in 0..2 {
let _ = breaker.call(|| async { Ok::<_, TalosError>("ok") }).await;
}
assert_eq!(breaker.state().await, CircuitState::Closed);
}
#[tokio::test]
async fn test_circuit_breaker_failure_rate() {
let breaker = CircuitBreaker::with_defaults();
assert_eq!(breaker.failure_rate(), 0.0);
for _ in 0..4 {
let _ = breaker.call(|| async { Ok::<_, TalosError>("ok") }).await;
}
let _ = breaker
.call(|| async { Err::<(), _>(TalosError::Connection("test".to_string())) })
.await;
assert!((breaker.failure_rate() - 0.2).abs() < f64::EPSILON);
}
#[tokio::test]
async fn test_circuit_breaker_time_until_retry() {
let config = CircuitBreakerConfig::new()
.with_failure_threshold(2)
.with_reset_timeout(Duration::from_secs(30));
let breaker = CircuitBreaker::new(config);
assert!(breaker.time_until_retry().await.is_none());
for _ in 0..2 {
let _ = breaker
.call(|| async { Err::<(), _>(TalosError::Connection("test".to_string())) })
.await;
}
let retry_time = breaker.time_until_retry().await;
assert!(retry_time.is_some());
assert!(retry_time.unwrap() > Duration::ZERO);
}
#[test]
fn test_circuit_state_equality() {
assert_eq!(CircuitState::Closed, CircuitState::Closed);
assert_ne!(CircuitState::Closed, CircuitState::Open);
assert_ne!(CircuitState::Open, CircuitState::HalfOpen);
}
}