use crate::metrics::MetricsCollector;
use parking_lot::Mutex;
use std::sync::Arc;
use std::time::{Duration, Instant};
use vtcode_commons::error_category::ErrorCategory;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum CircuitState {
Closed = 0,
Open = 1,
HalfOpen = 2,
}
use std::fs;
use std::path::PathBuf;
#[derive(Debug, Clone)]
struct InternalState {
status: CircuitState,
consecutive_failures: u32,
half_open_successes: u32,
last_failure_time: Option<Instant>,
blocked_requests: u32,
}
impl Default for InternalState {
fn default() -> Self {
Self {
status: CircuitState::Closed,
consecutive_failures: 0,
half_open_successes: 0,
last_failure_time: None,
blocked_requests: 0,
}
}
}
pub struct McpCircuitBreaker {
state: Mutex<InternalState>,
config: CircuitBreakerConfig,
persistence_path: Option<PathBuf>,
metrics: Option<Arc<MetricsCollector>>,
}
#[derive(Debug, Clone, Copy)]
pub struct CircuitBreakerConfig {
pub failure_threshold: u32,
pub success_threshold: u32,
pub base_timeout: Duration,
pub max_timeout: Duration,
}
impl Default for CircuitBreakerConfig {
fn default() -> Self {
Self {
failure_threshold: 3, success_threshold: 2, base_timeout: Duration::from_secs(10),
max_timeout: Duration::from_secs(60),
}
}
}
impl McpCircuitBreaker {
pub fn new() -> Self {
Self::with_config(CircuitBreakerConfig::default())
}
pub fn with_metrics(metrics: Arc<MetricsCollector>) -> Self {
Self::with_config_and_metrics(CircuitBreakerConfig::default(), metrics)
}
pub fn with_config(config: CircuitBreakerConfig) -> Self {
Self::build(config, None, None)
}
pub fn with_config_and_metrics(config: CircuitBreakerConfig, metrics: Arc<MetricsCollector>) -> Self {
Self::build(config, None, Some(metrics))
}
fn build(
config: CircuitBreakerConfig,
persistence_path: Option<PathBuf>,
metrics: Option<Arc<MetricsCollector>>,
) -> Self {
Self {
state: Mutex::new(InternalState::default()),
config,
persistence_path,
metrics,
}
}
#[allow(dead_code)]
pub fn with_persistence(path: PathBuf) -> Self {
let breaker = Self::build(CircuitBreakerConfig::default(), Some(path.clone()), None);
if let Ok(data) = fs::read_to_string(&path)
&& let Ok(persisted) = serde_json::from_str::<PersistedState>(&data)
{
let mut state = breaker.state.lock();
state.status = match persisted.state {
0 => CircuitState::Closed,
1 => CircuitState::Open,
2 => CircuitState::HalfOpen,
_ => CircuitState::Closed,
};
state.consecutive_failures = persisted.consecutive_failures;
if let Some(epoch) = persisted.last_failure_epoch_secs {
let now = Instant::now();
let elapsed = Duration::from_secs(epoch);
state.last_failure_time = Some(now.checked_sub(elapsed).unwrap_or(now));
}
}
breaker
}
#[inline]
fn record_half_open_metric(&self) {
MetricsCollector::record_circuit_breaker_metrics(&self.metrics, true, false, false);
}
#[inline]
fn record_breaker_denial_metric(&self) {
MetricsCollector::record_circuit_breaker_metrics(&self.metrics, false, true, false);
}
#[inline]
fn record_circuit_open_metric(&self) {
MetricsCollector::record_circuit_breaker_metrics(&self.metrics, false, false, true);
}
fn persist(&self, state: &InternalState) {
if let Some(path) = &self.persistence_path {
let epoch = state.last_failure_time.map(|t| {
t.elapsed().as_secs()
});
let persisted = PersistedState {
state: state.status as u8,
consecutive_failures: state.consecutive_failures,
last_failure_epoch_secs: epoch,
};
if let Ok(data) = serde_json::to_string(&persisted) {
let path = path.clone();
match tokio::runtime::Handle::try_current() {
Ok(handle) => {
let _persist_task = handle.spawn_blocking(move || {
let _ = fs::write(&path, data);
});
}
Err(_) => {
let _ = fs::write(&path, data);
}
}
}
}
}
#[cfg(test)]
pub fn state(&self) -> CircuitState {
self.state.lock().status
}
pub fn allow_request(&self) -> bool {
let mut state = self.state.lock();
let mut should_persist = false;
let result = match state.status {
CircuitState::Closed | CircuitState::HalfOpen => true,
CircuitState::Open => {
if let Some(last_failure) = state.last_failure_time {
let timeout = self.calculate_timeout(&state);
if last_failure.elapsed() >= timeout {
state.status = CircuitState::HalfOpen;
state.half_open_successes = 0;
self.record_half_open_metric();
should_persist = true;
true
} else {
state.blocked_requests += 1;
self.record_breaker_denial_metric();
should_persist = true;
false
}
} else {
state.blocked_requests += 1;
self.record_breaker_denial_metric();
should_persist = true;
false
}
}
};
if should_persist {
let state_clone = state.clone();
drop(state); self.persist(&state_clone);
}
result
}
pub fn record_success(&self) {
let mut state = self.state.lock();
let mut should_persist = false;
match state.status {
CircuitState::Closed => {
state.consecutive_failures = 0;
}
CircuitState::HalfOpen => {
state.half_open_successes += 1;
if state.half_open_successes >= self.config.success_threshold {
state.status = CircuitState::Closed;
state.consecutive_failures = 0;
state.half_open_successes = 0;
state.last_failure_time = None;
should_persist = true;
}
}
CircuitState::Open => {
state.status = CircuitState::HalfOpen;
state.half_open_successes = 1;
should_persist = true;
}
}
if should_persist {
let state_clone = state.clone();
drop(state); self.persist(&state_clone);
}
}
#[cfg(test)]
pub fn record_failure(&self) {
self.record_failure_category(ErrorCategory::ExecutionError);
}
pub fn record_failure_category(&self, category: ErrorCategory) {
if !category.should_trip_circuit_breaker() {
return;
}
let mut state = self.state.lock();
state.last_failure_time = Some(Instant::now());
match state.status {
CircuitState::Closed => {
state.consecutive_failures += 1;
if state.consecutive_failures >= self.config.failure_threshold {
state.status = CircuitState::Open;
self.record_circuit_open_metric();
}
}
CircuitState::HalfOpen => {
state.status = CircuitState::Open;
state.consecutive_failures += 1;
state.half_open_successes = 0;
self.record_circuit_open_metric();
}
CircuitState::Open => {
state.consecutive_failures += 1;
}
}
let state_clone = state.clone();
drop(state); self.persist(&state_clone);
}
fn calculate_timeout(&self, state: &InternalState) -> Duration {
let failures = state.consecutive_failures;
let multiplier = if failures > self.config.failure_threshold {
2u32.saturating_pow(failures.saturating_sub(self.config.failure_threshold))
} else {
1
};
let timeout = self.config.base_timeout.saturating_mul(multiplier);
timeout.min(self.config.max_timeout)
}
pub fn diagnostics(&self) -> CircuitBreakerDiagnostics {
let state = self.state.lock();
let timeout = self.calculate_timeout(&state);
let retry_after = if state.status == CircuitState::Open {
state.last_failure_time.and_then(|failure_time| {
let elapsed = failure_time.elapsed();
let timeout = self.calculate_timeout(&state);
timeout.checked_sub(elapsed)
})
} else {
None
};
CircuitBreakerDiagnostics {
status: state.status,
consecutive_failures: state.consecutive_failures,
half_open_successes: state.half_open_successes,
last_failure_time: state.last_failure_time,
current_timeout: timeout,
retry_after,
blocked_requests: state.blocked_requests,
is_blocking: state.status == CircuitState::Open,
}
}
}
impl Default for McpCircuitBreaker {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone)]
pub struct CircuitBreakerDiagnostics {
pub status: CircuitState,
pub consecutive_failures: u32,
#[allow(dead_code)]
pub half_open_successes: u32,
pub last_failure_time: Option<Instant>,
pub current_timeout: Duration,
pub retry_after: Option<Duration>,
#[allow(dead_code)]
pub blocked_requests: u32,
#[allow(dead_code)]
pub is_blocking: bool,
}
#[derive(serde::Serialize, serde::Deserialize)]
struct PersistedState {
state: u8,
consecutive_failures: u32,
last_failure_epoch_secs: Option<u64>,
}
#[cfg(test)]
mod tests {
use super::*;
use crate::metrics::MetricsCollector;
use std::thread;
#[test]
fn test_circuit_breaker_closed_state() {
let breaker = McpCircuitBreaker::new();
assert_eq!(breaker.state(), CircuitState::Closed);
assert!(breaker.allow_request());
}
#[test]
fn test_circuit_breaker_opens_after_threshold() {
let config = CircuitBreakerConfig { failure_threshold: 3, ..Default::default() };
let breaker = McpCircuitBreaker::with_config(config);
breaker.record_failure(); assert_eq!(breaker.state(), CircuitState::Closed);
breaker.record_failure(); assert_eq!(breaker.state(), CircuitState::Closed);
breaker.record_failure(); assert_eq!(breaker.state(), CircuitState::Open);
assert!(!breaker.allow_request()); assert!(breaker.diagnostics().blocked_requests > 0);
}
#[test]
fn test_circuit_breaker_half_open_transition() {
let config = CircuitBreakerConfig {
failure_threshold: 2,
base_timeout: Duration::from_millis(100),
..Default::default()
};
let breaker = McpCircuitBreaker::with_config(config);
breaker.record_failure();
breaker.record_failure();
assert_eq!(breaker.state(), CircuitState::Open);
thread::sleep(Duration::from_millis(150));
assert!(breaker.allow_request());
assert_eq!(breaker.state(), CircuitState::HalfOpen);
}
#[test]
fn test_circuit_breaker_closes_after_successes() {
let config = CircuitBreakerConfig {
failure_threshold: 2,
success_threshold: 2,
base_timeout: Duration::from_millis(50),
..Default::default()
};
let breaker = McpCircuitBreaker::with_config(config);
breaker.record_failure();
breaker.record_failure();
assert_eq!(breaker.state(), CircuitState::Open);
thread::sleep(Duration::from_millis(60));
assert!(breaker.allow_request());
assert_eq!(breaker.state(), CircuitState::HalfOpen);
breaker.record_success(); assert_eq!(breaker.state(), CircuitState::HalfOpen);
breaker.record_success(); assert_eq!(breaker.state(), CircuitState::Closed);
}
#[test]
fn test_circuit_breaker_failure_in_half_open() {
let config = CircuitBreakerConfig {
failure_threshold: 2,
base_timeout: Duration::from_millis(50),
..Default::default()
};
let breaker = McpCircuitBreaker::with_config(config);
breaker.record_failure();
breaker.record_failure();
thread::sleep(Duration::from_millis(60));
breaker.allow_request();
assert_eq!(breaker.state(), CircuitState::HalfOpen);
breaker.record_failure();
assert_eq!(breaker.state(), CircuitState::Open);
}
#[test]
fn test_exponential_backoff() {
let config = CircuitBreakerConfig {
failure_threshold: 2,
base_timeout: Duration::from_secs(10),
max_timeout: Duration::from_secs(60),
..Default::default()
};
let breaker = McpCircuitBreaker::with_config(config);
for _ in 0..5 {
breaker.record_failure();
}
let diag = breaker.diagnostics();
assert_eq!(diag.current_timeout, Duration::from_secs(60));
}
#[test]
fn authentication_failure_does_not_trip_breaker() {
let breaker = McpCircuitBreaker::new();
breaker.record_failure_category(ErrorCategory::Authentication);
assert_eq!(breaker.state(), CircuitState::Closed);
assert_eq!(breaker.diagnostics().consecutive_failures, 0);
}
#[test]
fn reliability_metrics_capture_open_half_open_and_denials() {
let metrics = Arc::new(MetricsCollector::new());
let breaker = McpCircuitBreaker::with_config_and_metrics(
CircuitBreakerConfig {
failure_threshold: 1,
base_timeout: Duration::from_millis(10),
..Default::default()
},
metrics.clone(),
);
breaker.record_failure_category(ErrorCategory::ExecutionError);
assert_eq!(breaker.state(), CircuitState::Open);
assert!(!breaker.allow_request());
thread::sleep(Duration::from_millis(20));
assert!(breaker.allow_request());
assert_eq!(breaker.state(), CircuitState::HalfOpen);
let execution = metrics.get_execution_metrics();
assert_eq!(execution.circuit_open_events, 1);
assert_eq!(execution.breaker_denials, 1);
assert_eq!(execution.half_open_events, 1);
}
#[test]
fn concurrent_requests_do_not_cause_inconsistent_state() {
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
let config = CircuitBreakerConfig {
failure_threshold: 1,
base_timeout: Duration::from_millis(50),
..Default::default()
};
let breaker = Arc::new(McpCircuitBreaker::with_config(config));
let request_count = Arc::new(AtomicUsize::new(0));
breaker.record_failure();
assert_eq!(breaker.state(), CircuitState::Open);
thread::sleep(Duration::from_millis(100));
let mut handles = vec![];
for _ in 0..10 {
let breaker = breaker.clone();
let request_count = request_count.clone();
handles.push(thread::spawn(move || {
if breaker.allow_request() {
request_count.fetch_add(1, Ordering::SeqCst);
}
}));
}
for handle in handles {
handle.join().unwrap();
}
let final_state = breaker.state();
assert!(
final_state == CircuitState::HalfOpen || final_state == CircuitState::Closed,
"Circuit should be in HalfOpen or Closed state, got {final_state:?}"
);
let count = request_count.load(Ordering::SeqCst);
assert!(count >= 1, "At least one thread should have succeeded");
}
}