use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU32, Ordering};
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll};
use std::time::{Duration, Instant};
use tower::{Layer, Service};
use super::types::{LlmRequest, LlmResponse};
use crate::client::BoxFuture;
use crate::error::{LiterLlmError, Result};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub enum CircuitState {
Closed = 0,
Open = 1,
HalfOpen = 2,
}
impl CircuitState {
fn from_u8(v: u8) -> Self {
match v {
1 => Self::Open,
2 => Self::HalfOpen,
_ => Self::Closed,
}
}
}
enum ProbeGuard<P: CircuitPolicy> {
None,
Half(Arc<P>),
}
impl<P: CircuitPolicy> ProbeGuard<P> {
fn disarm(&mut self) {
*self = Self::None;
}
}
impl<P: CircuitPolicy> Drop for ProbeGuard<P> {
fn drop(&mut self) {
if let Self::Half(policy) = self {
policy.release_probe_slot();
}
}
}
#[cfg_attr(alef, alef(skip))]
pub trait CircuitPolicy: Send + Sync + 'static {
fn record_success(&self);
fn record_failure(&self);
fn should_allow(&self) -> bool;
fn state(&self) -> CircuitState;
fn release_probe_slot(&self) {}
}
struct CircuitInner {
state: AtomicU8,
consecutive_failures: AtomicU32,
open_since: Mutex<Option<Instant>>,
probe_in_flight: AtomicBool,
}
#[cfg_attr(alef, alef(skip))]
pub struct ExponentialBackoffCircuit {
failure_threshold: u32,
base_backoff: Duration,
max_backoff: Duration,
inner: Arc<CircuitInner>,
open_count: AtomicU32,
}
impl ExponentialBackoffCircuit {
#[must_use]
pub fn new(failure_threshold: u32, base_backoff: Duration) -> Self {
Self {
failure_threshold,
base_backoff,
max_backoff: Duration::from_secs(120),
inner: Arc::new(CircuitInner {
state: AtomicU8::new(CircuitState::Closed as u8),
consecutive_failures: AtomicU32::new(0),
open_since: Mutex::new(None),
probe_in_flight: AtomicBool::new(false),
}),
open_count: AtomicU32::new(0),
}
}
fn current_backoff(&self) -> Duration {
let count = self.open_count.load(Ordering::Relaxed);
let shift = count.min(62) as u64;
let factor = 1u64.checked_shl(shift as u32).unwrap_or(u64::MAX);
let nanos = self.base_backoff.as_nanos().saturating_mul(factor as u128);
let computed = Duration::from_nanos(nanos.min(u64::MAX as u128) as u64);
computed.min(self.max_backoff)
}
fn maybe_half_open(&self) -> bool {
let backoff = self.current_backoff();
let guard = self.inner.open_since.lock().expect("open_since mutex poisoned");
if let Some(open_at) = *guard
&& open_at.elapsed() >= backoff
{
drop(guard);
self.inner.state.store(CircuitState::HalfOpen as u8, Ordering::Release);
tracing::info!(backoff = ?backoff, "circuit breaker entering half-open");
return true;
}
false
}
}
impl CircuitPolicy for ExponentialBackoffCircuit {
fn record_success(&self) {
self.inner.consecutive_failures.store(0, Ordering::Relaxed);
let prev = self.inner.state.swap(CircuitState::Closed as u8, Ordering::Release);
self.inner.probe_in_flight.store(false, Ordering::Release);
if CircuitState::from_u8(prev) != CircuitState::Closed {
tracing::info!("circuit breaker closed after successful probe");
}
}
fn record_failure(&self) {
let failures = self.inner.consecutive_failures.fetch_add(1, Ordering::AcqRel) + 1;
let current_u8 = self.inner.state.load(Ordering::Acquire);
let current = CircuitState::from_u8(current_u8);
if current == CircuitState::Open {
return;
}
let should_open = failures >= self.failure_threshold || current == CircuitState::HalfOpen;
if !should_open {
return;
}
let result = self.inner.state.compare_exchange(
current_u8,
CircuitState::Open as u8,
Ordering::AcqRel,
Ordering::Acquire,
);
if result.is_ok() {
let backoff = self.current_backoff();
let open_count = self.open_count.fetch_add(1, Ordering::Relaxed) + 1;
{
let mut guard = self.inner.open_since.lock().expect("open_since mutex poisoned");
*guard = Some(Instant::now());
}
self.inner.probe_in_flight.store(false, Ordering::Release);
tracing::warn!(
consecutive_failures = failures,
backoff = ?backoff,
open_count,
"circuit breaker opened"
);
}
}
fn should_allow(&self) -> bool {
match CircuitState::from_u8(self.inner.state.load(Ordering::Acquire)) {
CircuitState::Closed => true,
CircuitState::Open => {
if self.maybe_half_open() {
self.inner
.probe_in_flight
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
} else {
false
}
}
CircuitState::HalfOpen => {
self.inner
.probe_in_flight
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
}
}
}
fn state(&self) -> CircuitState {
CircuitState::from_u8(self.inner.state.load(Ordering::Acquire))
}
fn release_probe_slot(&self) {
self.inner.probe_in_flight.store(false, Ordering::Release);
}
}
#[cfg_attr(alef, alef(skip))]
pub struct CircuitLayer<P> {
policy: Arc<P>,
provider: String,
}
impl<P: CircuitPolicy> CircuitLayer<P> {
#[must_use]
pub fn new(policy: Arc<P>, provider: impl Into<String>) -> Self {
Self {
policy,
provider: provider.into(),
}
}
}
impl<P: CircuitPolicy, S> Layer<S> for CircuitLayer<P> {
type Service = CircuitService<P, S>;
fn layer(&self, inner: S) -> Self::Service {
CircuitService {
inner,
policy: Arc::clone(&self.policy),
provider: self.provider.clone(),
}
}
}
#[cfg_attr(alef, alef(skip))]
pub struct CircuitService<P, S> {
inner: S,
policy: Arc<P>,
provider: String,
}
impl<P: CircuitPolicy, S: Clone> Clone for CircuitService<P, S> {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
policy: Arc::clone(&self.policy),
provider: self.provider.clone(),
}
}
}
impl<P, S> Service<LlmRequest> for CircuitService<P, S>
where
P: CircuitPolicy + 'static,
S: Service<LlmRequest, Response = LlmResponse, Error = LiterLlmError> + Send + Clone + 'static,
S::Future: Send + 'static,
{
type Response = LlmResponse;
type Error = LiterLlmError;
type Future = BoxFuture<'static, Result<LlmResponse>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<()>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, req: LlmRequest) -> Self::Future {
let policy = Arc::clone(&self.policy);
let provider = self.provider.clone();
let model = req.model().unwrap_or("").to_owned();
let system = model.split_once('/').map(|(p, _)| p.to_owned()).unwrap_or_default();
let state = self.policy.state();
let clone = self.inner.clone();
let mut inner = std::mem::replace(&mut self.inner, clone);
Box::pin(async move {
let (allowed, is_probe) = {
let _span = tracing::debug_span!(
"circuit_breaker",
gen_ai.circuit.state = ?state,
provider = %provider,
)
.entered();
let probe = state != CircuitState::Closed;
(policy.should_allow(), probe)
};
if !allowed {
tracing::debug!(provider = %provider, "circuit open -- rejecting request");
super::metrics::record_circuit_trip(&system, &model);
return Err(LiterLlmError::ServiceUnavailable {
message: format!("circuit breaker open for provider '{provider}'"),
status: 503,
});
}
let mut probe_guard: ProbeGuard<P> = if is_probe {
ProbeGuard::Half(Arc::clone(&policy))
} else {
ProbeGuard::None
};
tracing::debug!(provider = %provider, state = ?policy.state(), "circuit allowing request through");
match inner.call(req).await {
Ok(resp) => {
probe_guard.disarm();
policy.record_success();
Ok(resp)
}
Err(e) => {
if e.is_transient() {
probe_guard.disarm();
policy.record_failure();
} else {
probe_guard.disarm();
if is_probe {
policy.release_probe_slot();
}
}
Err(e)
}
}
})
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::time::Duration;
use tower::{Layer as _, Service as _};
use super::*;
use crate::tower::service::LlmService;
use crate::tower::tests_common::{MockClient, chat_req};
use crate::tower::types::LlmRequest;
fn policy(n: u32) -> Arc<ExponentialBackoffCircuit> {
Arc::new(ExponentialBackoffCircuit::new(n, Duration::from_millis(50)))
}
#[tokio::test]
async fn circuit_starts_closed() {
let p = policy(3);
assert_eq!(p.state(), CircuitState::Closed);
}
#[tokio::test]
async fn circuit_breaker_opens_after_n_failures() {
let p = policy(3);
let layer = CircuitLayer::new(Arc::clone(&p), "test");
let mut svc = layer.layer(LlmService::new(MockClient::failing_timeout()));
for _ in 0..3 {
let _ = svc.call(LlmRequest::Chat(chat_req("openai/gpt-4"))).await;
}
assert_eq!(
p.state(),
CircuitState::Open,
"circuit should be open after threshold failures"
);
}
#[tokio::test]
async fn open_circuit_rejects_requests_without_calling_inner() {
let p = policy(1);
let mock = MockClient::failing_timeout();
let call_count = Arc::clone(&mock.call_count);
let layer = CircuitLayer::new(Arc::clone(&p), "test");
let mut svc = layer.layer(LlmService::new(mock));
let _ = svc.call(LlmRequest::Chat(chat_req("openai/gpt-4"))).await;
let before = call_count.load(std::sync::atomic::Ordering::SeqCst);
let err = svc
.call(LlmRequest::Chat(chat_req("openai/gpt-4")))
.await
.expect_err("should be rejected by open circuit");
assert!(
matches!(err, LiterLlmError::ServiceUnavailable { .. }),
"expected ServiceUnavailable from open circuit, got {err:?}"
);
assert_eq!(
call_count.load(std::sync::atomic::Ordering::SeqCst),
before,
"inner service should not be called while circuit is open"
);
}
#[tokio::test]
async fn circuit_breaker_half_opens_after_delay() {
let p = Arc::new(ExponentialBackoffCircuit::new(1, Duration::from_millis(20)));
let layer = CircuitLayer::new(Arc::clone(&p), "test");
let mut svc = layer.layer(LlmService::new(MockClient::failing_timeout()));
let _ = svc.call(LlmRequest::Chat(chat_req("openai/gpt-4"))).await;
assert_eq!(p.state(), CircuitState::Open);
tokio::time::sleep(Duration::from_millis(50)).await;
let allowed = p.maybe_half_open();
assert!(allowed, "should transition to half-open after backoff");
assert_eq!(p.state(), CircuitState::HalfOpen);
}
#[tokio::test]
async fn circuit_closes_after_successful_probe() {
let p = policy(1);
let failing = LlmService::new(MockClient::failing_timeout());
let layer = CircuitLayer::new(Arc::clone(&p), "test");
let mut svc = layer.layer(failing);
let _ = svc.call(LlmRequest::Chat(chat_req("openai/gpt-4"))).await;
assert_eq!(p.state(), CircuitState::Open);
p.inner.state.store(CircuitState::HalfOpen as u8, Ordering::Release);
let recovering = LlmService::new(MockClient::ok());
let layer2 = CircuitLayer::new(Arc::clone(&p), "test");
let mut svc2 = layer2.layer(recovering);
let resp = svc2
.call(LlmRequest::Chat(chat_req("openai/gpt-4")))
.await
.expect("probe should succeed");
assert!(matches!(resp, LlmResponse::Chat(_)));
assert_eq!(p.state(), CircuitState::Closed, "circuit should close after success");
}
#[tokio::test]
async fn non_transient_errors_do_not_trip_circuit() {
let p = policy(2);
let layer = CircuitLayer::new(Arc::clone(&p), "test");
let mut svc = layer.layer(LlmService::new(MockClient::failing_auth()));
for _ in 0..5 {
let _ = svc.call(LlmRequest::Chat(chat_req("openai/gpt-4"))).await;
}
assert_eq!(
p.state(),
CircuitState::Closed,
"non-transient errors should not open the circuit"
);
}
#[test]
fn circuit_concurrent_failures_open_count_increments_once() {
use std::thread;
let circuit = Arc::new(ExponentialBackoffCircuit::new(3, Duration::from_millis(50)));
let handles: Vec<_> = (0..10)
.map(|_| {
let c = Arc::clone(&circuit);
thread::spawn(move || c.record_failure())
})
.collect();
for h in handles {
h.join().expect("thread panicked");
}
let open_count = circuit.open_count.load(Ordering::Relaxed);
assert_eq!(
open_count, 1,
"open_count should be 1 regardless of how many concurrent callers hit the threshold; got {open_count}"
);
assert_eq!(
circuit.state(),
CircuitState::Open,
"circuit should be open after concurrent failures"
);
}
#[test]
fn circuit_record_failure_works_outside_tokio_runtime() {
let circuit = ExponentialBackoffCircuit::new(1, Duration::from_millis(50));
circuit.record_failure();
assert_eq!(
circuit.state(),
CircuitState::Open,
"state should be Open after one failure with threshold=1"
);
}
#[tokio::test]
async fn circuit_service_respects_inner_readiness() {
use std::pin::Pin;
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering as AtomicOrdering;
use tower::limit::ConcurrencyLimit;
#[derive(Clone)]
struct BlockingInner {
call_count: Arc<AtomicUsize>,
}
impl tower::Service<LlmRequest> for BlockingInner {
type Response = LlmResponse;
type Error = LiterLlmError;
type Future = crate::client::BoxFuture<'static, Result<LlmResponse>>;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<()>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, _req: LlmRequest) -> Self::Future {
self.call_count.fetch_add(1, AtomicOrdering::SeqCst);
Box::pin(std::future::pending())
}
}
let call_count = Arc::new(AtomicUsize::new(0));
let inner = BlockingInner {
call_count: Arc::clone(&call_count),
};
let limited: ConcurrencyLimit<BlockingInner> = ConcurrencyLimit::new(inner, 1);
let p = Arc::new(ExponentialBackoffCircuit::new(5, Duration::from_millis(50)));
let mut svc = CircuitService {
inner: limited,
policy: Arc::clone(&p),
provider: "test".into(),
};
futures_util::future::poll_fn(|cx| svc.poll_ready(cx))
.await
.expect("first poll_ready should succeed");
let mut held_fut = svc.call(LlmRequest::ListModels());
{
let mut noop_cx = std::task::Context::from_waker(futures_util::task::noop_waker_ref());
let _ = Pin::new(&mut held_fut).poll(&mut noop_cx);
}
assert_eq!(
call_count.load(AtomicOrdering::SeqCst),
1,
"inner service should have been called exactly once"
);
let mut noop_cx = std::task::Context::from_waker(futures_util::task::noop_waker_ref());
let poll = svc.poll_ready(&mut noop_cx);
assert!(
poll.is_pending(),
"second poll_ready must be Pending when the concurrency permit is exhausted"
);
}
#[tokio::test]
async fn circuit_half_open_after_cooldown() {
let p = Arc::new(ExponentialBackoffCircuit::new(1, Duration::from_millis(20)));
let layer = CircuitLayer::new(Arc::clone(&p), "test");
let mut svc_fail = layer.layer(LlmService::new(MockClient::failing_timeout()));
let _ = svc_fail.call(LlmRequest::Chat(chat_req("openai/gpt-4"))).await;
assert_eq!(p.state(), CircuitState::Open);
tokio::time::sleep(Duration::from_millis(50)).await;
let layer2 = CircuitLayer::new(Arc::clone(&p), "test");
let mut svc_ok = layer2.layer(LlmService::new(MockClient::ok()));
let resp = svc_ok
.call(LlmRequest::Chat(chat_req("openai/gpt-4")))
.await
.expect("probe after cooldown must succeed");
assert!(matches!(resp, LlmResponse::Chat(_)));
assert_eq!(p.state(), CircuitState::Closed);
}
#[tokio::test]
async fn circuit_half_open_single_probe() {
use std::sync::atomic::AtomicUsize;
let p = Arc::new(ExponentialBackoffCircuit::new(1, Duration::from_millis(20)));
p.record_failure();
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(p.maybe_half_open(), "should transition to HalfOpen");
assert_eq!(p.state(), CircuitState::HalfOpen);
let probe_count = Arc::new(AtomicUsize::new(0));
let rejected_count = Arc::new(AtomicUsize::new(0));
let handles: Vec<_> = (0..50)
.map(|_| {
let p2 = Arc::clone(&p);
let pc = Arc::clone(&probe_count);
let rc = Arc::clone(&rejected_count);
tokio::spawn(async move {
if p2.should_allow() {
pc.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
} else {
rc.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
}
})
})
.collect();
for h in handles {
h.await.unwrap();
}
assert_eq!(
probe_count.load(std::sync::atomic::Ordering::SeqCst),
1,
"exactly 1 probe"
);
assert_eq!(
rejected_count.load(std::sync::atomic::Ordering::SeqCst),
49,
"49 rejected"
);
}
#[tokio::test]
async fn circuit_probe_flag_cleared_on_cancel() {
use std::sync::atomic::Ordering as AO;
let p = Arc::new(ExponentialBackoffCircuit::new(1, Duration::from_millis(10)));
p.record_failure();
assert_eq!(p.state(), CircuitState::Open);
p.inner.state.store(CircuitState::HalfOpen as u8, AO::Release);
p.inner.probe_in_flight.store(false, AO::Release);
#[derive(Clone)]
struct BlockForever;
impl tower::Service<LlmRequest> for BlockForever {
type Response = LlmResponse;
type Error = LiterLlmError;
type Future = crate::client::BoxFuture<'static, Result<LlmResponse>>;
fn poll_ready(&mut self, _cx: &mut Context<'_>) -> Poll<Result<()>> {
Poll::Ready(Ok(()))
}
fn call(&mut self, _req: LlmRequest) -> Self::Future {
Box::pin(std::future::pending())
}
}
let layer = CircuitLayer::new(Arc::clone(&p), "test");
let mut svc = layer.layer(BlockForever);
{
let fut = svc.call(LlmRequest::Chat(chat_req("openai/gpt-4")));
drop(fut);
}
assert!(
!p.inner.probe_in_flight.load(AO::Acquire),
"probe_in_flight must be false after probe future is dropped"
);
assert!(
p.should_allow(),
"should allow another probe after cancelled probe slot was released"
);
}
#[tokio::test]
async fn circuit_half_open_failure_reopens() {
let p = Arc::new(ExponentialBackoffCircuit::new(1, Duration::from_millis(20)));
let layer = CircuitLayer::new(Arc::clone(&p), "test");
let mut svc_fail = layer.layer(LlmService::new(MockClient::failing_timeout()));
let _ = svc_fail.call(LlmRequest::Chat(chat_req("openai/gpt-4"))).await;
assert_eq!(p.state(), CircuitState::Open);
tokio::time::sleep(Duration::from_millis(50)).await;
let layer2 = CircuitLayer::new(Arc::clone(&p), "test");
let mut svc_fail2 = layer2.layer(LlmService::new(MockClient::failing_timeout()));
let err = svc_fail2
.call(LlmRequest::Chat(chat_req("openai/gpt-4")))
.await
.expect_err("failing probe must error");
assert!(matches!(err, LiterLlmError::Timeout));
assert_eq!(p.state(), CircuitState::Open, "must re-open after failing probe");
assert!(
!p.inner.probe_in_flight.load(Ordering::Acquire),
"probe_in_flight must be cleared after failure"
);
}
}