use std::future::Future;
use std::sync::{Mutex, PoisonError};
use std::time::{Duration, Instant};
use async_trait::async_trait;
use crate::backend::{Backend, HealthStatus, LockableBackend};
use crate::client::SharedBackend;
use crate::error::BackendError;
use crate::random_unit;
#[derive(Debug, Clone, PartialEq)]
pub struct RetryConfig {
pub max_attempts: u32,
pub base_delay: Duration,
pub max_delay: Duration,
pub jitter: bool,
}
impl Default for RetryConfig {
fn default() -> Self {
Self {
max_attempts: 3,
base_delay: Duration::from_millis(100),
max_delay: Duration::from_secs(5),
jitter: true,
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct CircuitBreakerConfig {
pub failure_threshold: u32,
pub success_threshold: u32,
pub open_timeout: Duration,
pub half_open_max_calls: u32,
pub rolling_window: Duration,
}
impl Default for CircuitBreakerConfig {
fn default() -> Self {
Self {
failure_threshold: 5,
success_threshold: 3,
open_timeout: Duration::from_secs(5),
half_open_max_calls: 3,
rolling_window: Duration::from_secs(60),
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct BackpressureConfig {
pub max_concurrent: usize,
pub max_queue: usize,
pub acquire_timeout: Duration,
}
impl Default for BackpressureConfig {
fn default() -> Self {
Self {
max_concurrent: 100,
max_queue: 1000,
acquire_timeout: Duration::from_millis(100),
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct ReliabilityConfig {
pub retry: Option<RetryConfig>,
pub circuit_breaker: Option<CircuitBreakerConfig>,
pub backpressure: Option<BackpressureConfig>,
}
impl Default for ReliabilityConfig {
fn default() -> Self {
Self {
retry: Some(RetryConfig::default()),
circuit_breaker: Some(CircuitBreakerConfig::default()),
backpressure: Some(BackpressureConfig::default()),
}
}
}
impl ReliabilityConfig {
#[must_use]
pub fn disabled() -> Self {
Self {
retry: None,
circuit_breaker: None,
backpressure: None,
}
}
#[must_use]
pub fn is_disabled(&self) -> bool {
self.retry.is_none() && self.circuit_breaker.is_none() && self.backpressure.is_none()
}
}
#[derive(Debug)]
pub(crate) struct RetryPolicy {
config: RetryConfig,
}
impl RetryPolicy {
pub(crate) fn new(config: RetryConfig) -> Self {
Self { config }
}
fn delay(&self, attempt: u32) -> Duration {
let exp = self
.config
.base_delay
.saturating_mul(2u32.saturating_pow(attempt));
let capped = exp.min(self.config.max_delay);
if self.config.jitter {
capped.mul_f64(0.5 + random_unit())
} else {
capped
}
}
pub(crate) async fn execute<T, F, Fut>(&self, f: F) -> Result<T, BackendError>
where
F: Fn() -> Fut,
Fut: Future<Output = Result<T, BackendError>>,
{
let mut attempt: u32 = 0;
loop {
match f().await {
Ok(v) => return Ok(v),
Err(e) if e.kind.is_retryable() && attempt + 1 < self.config.max_attempts => {
tokio::time::sleep(self.delay(attempt)).await;
attempt += 1;
}
Err(e) => return Err(e),
}
}
}
}
#[cfg(test)]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum CircuitState {
Closed,
Open,
HalfOpen,
}
#[derive(Debug)]
enum State {
Closed,
Open { since: Instant },
HalfOpen,
}
#[derive(Debug)]
struct BreakerInner {
state: State,
failures: Vec<Instant>,
half_open_successes: u32,
half_open_calls: u32,
}
enum Outcome {
Success,
Failure,
Neutral,
}
#[derive(Debug)]
pub(crate) struct CircuitBreaker {
config: CircuitBreakerConfig,
inner: Mutex<BreakerInner>,
}
impl CircuitBreaker {
pub(crate) fn new(config: CircuitBreakerConfig) -> Self {
Self {
config,
inner: Mutex::new(BreakerInner {
state: State::Closed,
failures: Vec::new(),
half_open_successes: 0,
half_open_calls: 0,
}),
}
}
fn lock(&self) -> std::sync::MutexGuard<'_, BreakerInner> {
self.inner.lock().unwrap_or_else(PoisonError::into_inner)
}
#[cfg(test)]
pub(crate) fn state(&self) -> CircuitState {
let mut inner = self.lock();
self.maybe_half_open(&mut inner);
match inner.state {
State::Closed => CircuitState::Closed,
State::Open { .. } => CircuitState::Open,
State::HalfOpen => CircuitState::HalfOpen,
}
}
fn maybe_half_open(&self, inner: &mut BreakerInner) {
if let State::Open { since } = inner.state {
if since.elapsed() >= self.config.open_timeout {
inner.state = State::HalfOpen;
inner.half_open_successes = 0;
inner.half_open_calls = 0;
}
}
}
fn try_acquire(&self) -> Result<ProbePermit<'_>, BackendError> {
let mut inner = self.lock();
self.maybe_half_open(&mut inner);
match inner.state {
State::Closed => Ok(ProbePermit {
breaker: self,
took_slot: false,
}),
State::Open { .. } => Err(BackendError::circuit_open(
"circuit breaker is open: backend calls are failing fast",
)),
State::HalfOpen => {
if inner.half_open_calls >= self.config.half_open_max_calls {
Err(BackendError::circuit_open(
"circuit breaker is half-open and the probe limit is reached",
))
} else {
inner.half_open_calls += 1;
Ok(ProbePermit {
breaker: self,
took_slot: true,
})
}
}
}
}
fn record(&self, outcome: &Outcome) {
let mut inner = self.lock();
match outcome {
Outcome::Success => {
if matches!(inner.state, State::HalfOpen) {
inner.half_open_successes += 1;
if inner.half_open_successes >= self.config.success_threshold {
inner.state = State::Closed;
inner.failures.clear();
inner.half_open_successes = 0;
inner.half_open_calls = 0;
} else {
inner.half_open_calls = inner.half_open_calls.saturating_sub(1);
}
}
}
Outcome::Failure => match inner.state {
State::HalfOpen => {
inner.state = State::Open {
since: Instant::now(),
};
inner.half_open_successes = 0;
inner.half_open_calls = 0;
}
State::Closed => {
let now = Instant::now();
inner.failures.push(now);
let window = self.config.rolling_window;
inner.failures.retain(|t| now.duration_since(*t) <= window);
if inner.failures.len() >= self.config.failure_threshold as usize {
inner.state = State::Open { since: now };
inner.failures.clear();
}
}
State::Open { .. } => {}
},
Outcome::Neutral => {
if matches!(inner.state, State::HalfOpen) {
inner.half_open_calls = inner.half_open_calls.saturating_sub(1);
}
}
}
}
}
#[derive(Debug)]
struct ProbePermit<'a> {
breaker: &'a CircuitBreaker,
took_slot: bool,
}
impl ProbePermit<'_> {
fn complete(mut self, outcome: &Outcome) {
self.took_slot = false;
self.breaker.record(outcome);
}
}
impl Drop for ProbePermit<'_> {
fn drop(&mut self) {
if !self.took_slot {
return;
}
let mut inner = self.breaker.lock();
if matches!(inner.state, State::HalfOpen) {
inner.half_open_calls = inner.half_open_calls.saturating_sub(1);
}
}
}
#[derive(Debug)]
pub(crate) struct ConcurrencyLimiter {
semaphore: tokio::sync::Semaphore,
waiting: std::sync::atomic::AtomicUsize,
config: BackpressureConfig,
}
struct QueueSlot<'a> {
waiting: &'a std::sync::atomic::AtomicUsize,
}
impl Drop for QueueSlot<'_> {
fn drop(&mut self) {
self.waiting
.fetch_sub(1, std::sync::atomic::Ordering::AcqRel);
}
}
impl ConcurrencyLimiter {
pub(crate) fn new(config: BackpressureConfig) -> Self {
Self {
semaphore: tokio::sync::Semaphore::new(
config
.max_concurrent
.clamp(1, tokio::sync::Semaphore::MAX_PERMITS),
),
waiting: std::sync::atomic::AtomicUsize::new(0),
config,
}
}
async fn acquire(&self) -> Result<tokio::sync::SemaphorePermit<'_>, BackendError> {
use std::sync::atomic::Ordering;
if let Ok(permit) = self.semaphore.try_acquire() {
return Ok(permit);
}
if self.waiting.fetch_add(1, Ordering::AcqRel) >= self.config.max_queue {
self.waiting.fetch_sub(1, Ordering::AcqRel);
return Err(BackendError::backpressure(format!(
"backpressure: waiting queue is full (max_queue={}), call shed without reaching the backend",
self.config.max_queue
)));
}
let _slot = QueueSlot {
waiting: &self.waiting,
};
match tokio::time::timeout(self.config.acquire_timeout, self.semaphore.acquire()).await {
Ok(Ok(permit)) => Ok(permit),
Ok(Err(_closed)) => Err(BackendError::backpressure(
"backpressure: limiter unavailable, call shed without reaching the backend",
)),
Err(_elapsed) => Err(BackendError::backpressure(format!(
"backpressure: timed out waiting for a permit after {:?}, call shed without reaching the backend",
self.config.acquire_timeout
))),
}
}
}
pub(crate) struct ReliableBackend {
inner: SharedBackend,
retry: Option<RetryPolicy>,
breaker: Option<CircuitBreaker>,
limiter: Option<ConcurrencyLimiter>,
}
impl ReliableBackend {
pub(crate) fn new(inner: SharedBackend, config: ReliabilityConfig) -> Self {
Self {
inner,
retry: config.retry.map(RetryPolicy::new),
breaker: config.circuit_breaker.map(CircuitBreaker::new),
limiter: config.backpressure.map(ConcurrencyLimiter::new),
}
}
async fn guarded<T, F, Fut>(&self, f: F) -> Result<T, BackendError>
where
F: Fn() -> Fut,
Fut: Future<Output = Result<T, BackendError>>,
{
let _permit = match &self.limiter {
Some(limiter) => Some(limiter.acquire().await?),
None => None,
};
let permit = match &self.breaker {
Some(cb) => Some(cb.try_acquire()?),
None => None,
};
let result = match &self.retry {
Some(retry) => retry.execute(f).await,
None => f().await,
};
if let Some(permit) = permit {
let outcome = match &result {
Ok(_) => Outcome::Success,
Err(e) if e.kind.is_retryable() => Outcome::Failure,
Err(_) => Outcome::Neutral,
};
permit.complete(&outcome);
}
result
}
}
#[cfg_attr(not(feature = "unsync"), async_trait)]
#[cfg_attr(feature = "unsync", async_trait(?Send))]
impl Backend for ReliableBackend {
async fn get(&self, key: &str) -> Result<Option<Vec<u8>>, BackendError> {
self.guarded(|| self.inner.get(key)).await
}
async fn set(
&self,
key: &str,
value: Vec<u8>,
ttl: Option<Duration>,
) -> Result<(), BackendError> {
self.guarded(|| self.inner.set(key, value.clone(), ttl))
.await
}
async fn delete(&self, key: &str) -> Result<bool, BackendError> {
self.guarded(|| self.inner.delete(key)).await
}
async fn exists(&self, key: &str) -> Result<bool, BackendError> {
self.guarded(|| self.inner.exists(key)).await
}
async fn health(&self) -> Result<HealthStatus, BackendError> {
self.inner.health().await
}
fn as_lockable(&self) -> Option<&dyn LockableBackend> {
self.inner.as_lockable()
}
}
#[cfg(not(feature = "unsync"))]
pub(crate) fn wrap_reliable(inner: SharedBackend, config: ReliabilityConfig) -> SharedBackend {
std::sync::Arc::new(ReliableBackend::new(inner, config))
}
#[cfg(feature = "unsync")]
pub(crate) fn wrap_reliable(inner: SharedBackend, config: ReliabilityConfig) -> SharedBackend {
std::rc::Rc::new(ReliableBackend::new(inner, config))
}
#[cfg(test)]
#[allow(clippy::expect_used)] mod tests {
use super::*;
use crate::error::BackendErrorKind;
fn breaker(failure_threshold: u32, open_timeout: Duration) -> CircuitBreaker {
CircuitBreaker::new(CircuitBreakerConfig {
failure_threshold,
success_threshold: 2,
open_timeout,
half_open_max_calls: 2,
rolling_window: Duration::from_secs(60),
})
}
fn admit_and(cb: &CircuitBreaker, outcome: &Outcome) {
let permit = cb.try_acquire().expect("breaker admits the call");
permit.complete(outcome);
}
#[test]
fn breaker_opens_after_threshold_and_fails_fast() {
let cb = breaker(3, Duration::from_secs(60));
for _ in 0..3 {
admit_and(&cb, &Outcome::Failure);
}
assert_eq!(cb.state(), CircuitState::Open);
let err = cb.try_acquire().expect_err("open breaker fails fast");
assert_eq!(err.kind, BackendErrorKind::CircuitOpen);
assert!(!err.kind.is_retryable());
}
#[test]
fn breaker_ignores_permanent_errors() {
let cb = breaker(2, Duration::from_secs(60));
for _ in 0..10 {
admit_and(&cb, &Outcome::Neutral);
}
assert_eq!(cb.state(), CircuitState::Closed);
}
#[test]
fn breaker_half_open_recovers_on_successes() {
let cb = breaker(1, Duration::from_millis(0));
admit_and(&cb, &Outcome::Failure);
assert_eq!(cb.state(), CircuitState::HalfOpen);
for _ in 0..2 {
admit_and(&cb, &Outcome::Success);
}
assert_eq!(cb.state(), CircuitState::Closed);
}
#[test]
fn breaker_half_open_reopens_on_failure() {
let cb = breaker(1, Duration::from_millis(0));
admit_and(&cb, &Outcome::Failure);
assert_eq!(cb.state(), CircuitState::HalfOpen);
admit_and(&cb, &Outcome::Failure);
assert!(matches!(cb.lock().state, State::Open { .. }));
}
#[test]
fn breaker_half_open_slot_released_by_neutral_outcome() {
let cb = breaker(1, Duration::from_millis(0));
admit_and(&cb, &Outcome::Failure);
assert_eq!(cb.state(), CircuitState::HalfOpen);
admit_and(&cb, &Outcome::Neutral);
admit_and(&cb, &Outcome::Neutral);
let permit = cb
.try_acquire()
.expect("neutral outcomes release their probe slots");
permit.complete(&Outcome::Neutral);
}
#[test]
fn breaker_half_open_closes_when_success_threshold_exceeds_probe_cap() {
let cb = CircuitBreaker::new(CircuitBreakerConfig {
failure_threshold: 1,
success_threshold: 3,
open_timeout: Duration::from_millis(0),
half_open_max_calls: 1,
rolling_window: Duration::from_secs(60),
});
admit_and(&cb, &Outcome::Failure);
assert_eq!(cb.state(), CircuitState::HalfOpen);
for _ in 0..3 {
let permit = cb
.try_acquire()
.expect("a non-closing success must release its probe slot");
permit.complete(&Outcome::Success);
}
assert_eq!(cb.state(), CircuitState::Closed);
}
#[test]
fn breaker_dropped_permit_releases_probe_slot() {
let cb = breaker(1, Duration::from_millis(0));
admit_and(&cb, &Outcome::Failure);
assert_eq!(cb.state(), CircuitState::HalfOpen);
for _ in 0..2 {
let permit = cb.try_acquire().expect("half-open admits a probe");
drop(permit); }
let permit = cb
.try_acquire()
.expect("dropped permits release their probe slots");
permit.complete(&Outcome::Success);
}
#[test]
fn breaker_closed_permit_drop_does_not_touch_half_open_accounting() {
let cb = breaker(1, Duration::from_millis(0));
let closed_permit = cb.try_acquire().expect("closed breaker admits calls");
admit_and(&cb, &Outcome::Failure);
assert_eq!(cb.state(), CircuitState::HalfOpen);
let p1 = cb.try_acquire().expect("probe slot 1");
let p2 = cb.try_acquire().expect("probe slot 2");
drop(closed_permit); assert!(
cb.try_acquire().is_err(),
"probe cap must still be enforced after a closed-state permit drops"
);
p1.complete(&Outcome::Success);
p2.complete(&Outcome::Success);
assert_eq!(cb.state(), CircuitState::Closed);
}
#[test]
fn reliability_default_enables_backpressure_with_python_parity_defaults() {
let config = ReliabilityConfig::default();
let bp = config.backpressure.expect("backpressure is on by default");
assert_eq!(bp.max_concurrent, 100);
assert_eq!(bp.max_queue, 1000);
assert_eq!(bp.acquire_timeout, Duration::from_millis(100));
}
#[tokio::test]
async fn limiter_clamps_zero_max_concurrent_to_one() {
let limiter = ConcurrencyLimiter::new(BackpressureConfig {
max_concurrent: 0,
max_queue: 0,
acquire_timeout: Duration::from_millis(10),
});
let permit = limiter
.acquire()
.await
.expect("0 behaves as 1 — one permit exists");
drop(permit);
}
#[tokio::test]
async fn limiter_clamps_huge_max_concurrent_instead_of_panicking() {
let limiter = ConcurrencyLimiter::new(BackpressureConfig {
max_concurrent: usize::MAX,
max_queue: 0,
acquire_timeout: Duration::from_millis(10),
});
let permit = limiter.acquire().await.expect("clamped limiter admits");
drop(permit);
}
#[tokio::test]
async fn limiter_sheds_immediately_when_queue_disabled() {
let limiter = ConcurrencyLimiter::new(BackpressureConfig {
max_concurrent: 1,
max_queue: 0,
acquire_timeout: Duration::from_secs(5),
});
let _held = limiter.acquire().await.expect("first permit");
let start = Instant::now();
let err = limiter
.acquire()
.await
.expect_err("saturated with no waiting queue");
assert_eq!(err.kind, BackendErrorKind::Backpressure);
assert!(!err.kind.is_retryable());
assert!(
start.elapsed() < Duration::from_millis(500),
"queue-full sheds immediately, not after acquire_timeout"
);
}
#[tokio::test]
async fn limiter_waiting_slot_released_on_cancelled_wait() {
let limiter = ConcurrencyLimiter::new(BackpressureConfig {
max_concurrent: 1,
max_queue: 1,
acquire_timeout: Duration::from_millis(100),
});
let _held = limiter.acquire().await.expect("first permit");
let cancelled = tokio::time::timeout(Duration::from_millis(20), limiter.acquire()).await;
assert!(cancelled.is_err(), "waiter cancelled from outside");
let start = Instant::now();
let err = limiter
.acquire()
.await
.expect_err("permit never frees, waiter times out");
assert_eq!(err.kind, BackendErrorKind::Backpressure);
assert!(
start.elapsed() >= Duration::from_millis(80),
"must join the queue and wait out acquire_timeout — an instant \
queue-full shed means the cancelled waiter leaked its slot"
);
}
#[tokio::test]
async fn limiter_sheds_queue_full_at_nonzero_boundary() {
let limiter = ConcurrencyLimiter::new(BackpressureConfig {
max_concurrent: 1,
max_queue: 1,
acquire_timeout: Duration::from_millis(200),
});
let _held = limiter.acquire().await.expect("first permit");
let waiter = async {
limiter.acquire().await
};
let third = async {
tokio::time::sleep(Duration::from_millis(50)).await; let start = Instant::now();
let err = limiter.acquire().await.expect_err("queue of 1 is full");
assert_eq!(err.kind, BackendErrorKind::Backpressure);
assert!(
start.elapsed() < Duration::from_millis(100),
"queue-full sheds instantly, not after the wait timeout"
);
};
let (waited, ()) = tokio::join!(waiter, third);
waited.expect_err("the parked waiter itself times out");
}
#[test]
fn retry_delay_is_capped_and_jittered() {
let policy = RetryPolicy::new(RetryConfig {
max_attempts: 5,
base_delay: Duration::from_millis(100),
max_delay: Duration::from_millis(300),
jitter: true,
});
for attempt in 0..10 {
let d = policy.delay(attempt);
assert!(d < Duration::from_millis(450), "attempt {attempt}: {d:?}");
}
let no_jitter = RetryPolicy::new(RetryConfig {
jitter: false,
..RetryConfig::default()
});
assert_eq!(no_jitter.delay(0), Duration::from_millis(100));
assert_eq!(no_jitter.delay(1), Duration::from_millis(200));
assert_eq!(no_jitter.delay(20), Duration::from_secs(5));
}
}