#![allow(
clippy::single_call_fn,
reason = "retry/is_transient/jittered split is intentional for unit testability"
)]
use std::borrow::Cow;
use std::future::Future;
use std::sync::OnceLock;
use std::time::{Duration, Instant};
use tokio::time::sleep;
use tracing::warn;
use crate::types::error::{HorizonError, PostgresError};
const TRANSIENT_SQLSTATES: &[&str] = &["08000", "08003", "08004", "08006", "40001", "40P01"];
pub const DEFAULT_MAX_RETRIES: u32 = 3;
pub const DEFAULT_INITIAL_DELAY: Duration = Duration::from_millis(50);
pub const DEFAULT_MAX_DELAY: Duration = Duration::from_secs(2);
pub const DEFAULT_BACKOFF_FACTOR: u32 = 2;
#[derive(Debug, Clone, Copy)]
pub struct RetryPolicy {
pub backoff_factor: u32,
pub initial_delay: Duration,
pub max_delay: Duration,
pub max_retries: u32,
}
impl Default for RetryPolicy {
fn default() -> Self {
Self {
backoff_factor: DEFAULT_BACKOFF_FACTOR,
initial_delay: DEFAULT_INITIAL_DELAY,
max_delay: DEFAULT_MAX_DELAY,
max_retries: DEFAULT_MAX_RETRIES,
}
}
}
impl RetryPolicy {
#[must_use]
pub const fn with_max_retries(mut self, max_retries: u32) -> Self {
self.max_retries = max_retries;
self
}
}
pub async fn retry<T, F, Fut>(policy: RetryPolicy, mut op: F) -> Result<T, HorizonError>
where
F: FnMut() -> Fut,
Fut: Future<Output = Result<T, HorizonError>>,
{
let mut attempt = 0_u32;
let mut delay = policy.initial_delay;
loop {
match op().await {
Ok(value) => return Ok(value),
Err(error) => {
if !is_transient(&error) || attempt >= policy.max_retries {
return Err(error);
}
let sleep_for = jittered(delay);
let next_attempt = attempt.saturating_add(1_u32);
warn!(
attempt = next_attempt,
max_retries = policy.max_retries,
sleep_ms = u64::try_from(sleep_for.as_millis()).unwrap_or(u64::MAX),
error = %error,
"transient Postgres error; retrying"
);
sleep(sleep_for).await;
attempt = next_attempt;
delay = next_delay(delay, policy);
}
}
}
}
fn sqlstate(error: &HorizonError) -> Option<Cow<'_, str>> {
let HorizonError::Postgres(PostgresError::Query(sqlx_error)) = error else {
return None;
};
let database_error = sqlx_error.as_database_error()?;
database_error
.code()
.map(|code| Cow::Owned(code.into_owned()))
}
fn is_transient(error: &HorizonError) -> bool {
if let Some(code) = sqlstate(error) {
return TRANSIENT_SQLSTATES.iter().any(|known| *known == code);
}
let HorizonError::Postgres(PostgresError::Query(sqlx_error)) = error else {
return false;
};
matches!(
sqlx_error,
sqlx::Error::Io(_) | sqlx::Error::PoolClosed | sqlx::Error::PoolTimedOut
)
}
fn next_delay(delay: Duration, policy: RetryPolicy) -> Duration {
let scaled = delay.saturating_mul(policy.backoff_factor);
if scaled > policy.max_delay {
policy.max_delay
} else {
scaled
}
}
#[allow(
clippy::arithmetic_side_effects,
clippy::integer_division,
clippy::integer_division_remainder_used,
reason = "jitter math operates on bounded u64 nanos with explicit shr/saturating ops"
)]
fn jittered(delay: Duration) -> Duration {
let entropy = jitter_entropy();
let sample = i64::from(u32::try_from(entropy & 0xFFFF_u64).unwrap_or(0_u32));
let centered = sample - 0x8000_i64;
let half_range = i64::try_from(delay.as_nanos() >> 2_u32).unwrap_or(i64::MAX);
let offset_nanos = centered
.saturating_mul(half_range)
.checked_shr(15_u32)
.unwrap_or(0_i64);
let base_nanos = i64::try_from(delay.as_nanos()).unwrap_or(i64::MAX);
let total_nanos = base_nanos.saturating_add(offset_nanos).max(0_i64);
Duration::from_nanos(u64::try_from(total_nanos).unwrap_or(0_u64))
}
fn jitter_entropy() -> u64 {
static START: OnceLock<Instant> = OnceLock::new();
let start = START.get_or_init(Instant::now);
let elapsed = start.elapsed().as_nanos();
u64::try_from(elapsed & u128::from(u64::MAX)).unwrap_or(0_u64)
}
#[cfg(test)]
#[allow(
clippy::absolute_paths,
clippy::arbitrary_source_item_ordering,
clippy::assertions_on_result_states,
clippy::default_numeric_fallback,
clippy::expect_used,
clippy::little_endian_bytes,
clippy::missing_trait_methods,
clippy::panic,
clippy::tests_outside_test_module,
clippy::unnecessary_literal_bound,
clippy::unwrap_used,
reason = "test-only stub for sqlx::DatabaseError needs unscoped restriction allows"
)]
mod tests {
use std::error::Error as StdError;
use std::fmt;
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
use sqlx::error::{DatabaseError, ErrorKind};
use super::*;
fn database_error_with_sqlstate(code: &'static str) -> HorizonError {
#[derive(Debug)]
struct StubDatabaseError {
code: &'static str,
}
impl fmt::Display for StubDatabaseError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "stub sqlstate={}", self.code)
}
}
impl StdError for StubDatabaseError {}
impl DatabaseError for StubDatabaseError {
fn message(&self) -> &str {
"stub"
}
fn code(&self) -> Option<Cow<'_, str>> {
Some(Cow::Borrowed(self.code))
}
fn as_error(&self) -> &(dyn StdError + Send + Sync + 'static) {
self
}
fn as_error_mut(&mut self) -> &mut (dyn StdError + Send + Sync + 'static) {
self
}
fn into_error(self: Box<Self>) -> Box<dyn StdError + Send + Sync + 'static> {
self
}
fn kind(&self) -> ErrorKind {
ErrorKind::Other
}
}
HorizonError::Postgres(PostgresError::Query(sqlx::Error::Database(Box::new(
StubDatabaseError { code },
))))
}
fn fast_policy() -> RetryPolicy {
RetryPolicy {
backoff_factor: DEFAULT_BACKOFF_FACTOR,
initial_delay: Duration::from_millis(1_u64),
max_delay: Duration::from_millis(4_u64),
max_retries: 3_u32,
}
}
#[test]
fn sqlstate_extracts_known_code() {
let error = database_error_with_sqlstate("40001");
assert_eq!(sqlstate(&error).as_deref(), Some("40001"));
}
#[test]
fn sqlstate_returns_none_for_non_postgres_error() {
let error = HorizonError::Postgres(PostgresError::NotFound {
entity: "platform".to_owned(),
id: uuid::Uuid::nil().into(),
});
assert!(sqlstate(&error).is_none());
}
#[test]
fn is_transient_matches_serialization_failure() {
assert!(is_transient(&database_error_with_sqlstate("40001")));
}
#[test]
fn is_transient_matches_deadlock() {
assert!(is_transient(&database_error_with_sqlstate("40P01")));
}
#[test]
fn is_transient_matches_connection_class() {
for code in ["08000", "08003", "08004", "08006"] {
assert!(
is_transient(&database_error_with_sqlstate(code)),
"expected {code} to be transient"
);
}
}
#[test]
fn is_transient_rejects_unique_violation() {
assert!(!is_transient(&database_error_with_sqlstate("23505")));
}
#[test]
fn is_transient_rejects_not_found() {
let error = HorizonError::Postgres(PostgresError::NotFound {
entity: "platform".to_owned(),
id: uuid::Uuid::nil().into(),
});
assert!(!is_transient(&error));
}
#[test]
fn is_transient_matches_pool_timeout() {
let error = HorizonError::Postgres(PostgresError::Query(sqlx::Error::PoolTimedOut));
assert!(is_transient(&error));
}
#[test]
fn next_delay_caps_at_max() {
let policy = RetryPolicy {
backoff_factor: 2_u32,
initial_delay: Duration::from_millis(50_u64),
max_delay: Duration::from_millis(120_u64),
max_retries: 5_u32,
};
assert_eq!(
next_delay(Duration::from_millis(50_u64), policy),
Duration::from_millis(100_u64)
);
assert_eq!(
next_delay(Duration::from_millis(100_u64), policy),
Duration::from_millis(120_u64)
);
assert_eq!(
next_delay(Duration::from_millis(120_u64), policy),
Duration::from_millis(120_u64)
);
}
#[test]
fn jittered_stays_within_quarter_band() {
let delay = Duration::from_millis(100_u64);
for _ in 0_i32..32_i32 {
let actual = jittered(delay);
assert!(actual >= Duration::from_millis(75_u64));
assert!(actual <= Duration::from_millis(125_u64));
}
}
#[tokio::test]
async fn retry_returns_immediately_on_success() {
let calls = Arc::new(AtomicU32::new(0_u32));
let calls_for_op = Arc::clone(&calls);
let result: Result<i32, HorizonError> = retry(fast_policy(), || {
let calls_inner = Arc::clone(&calls_for_op);
async move {
calls_inner.fetch_add(1_u32, Ordering::SeqCst);
Ok(7_i32)
}
})
.await;
assert_eq!(result.unwrap(), 7_i32);
assert_eq!(calls.load(Ordering::SeqCst), 1_u32);
}
#[tokio::test]
async fn retry_succeeds_after_transient_failure() {
let calls = Arc::new(AtomicU32::new(0_u32));
let calls_for_op = Arc::clone(&calls);
let result: Result<i32, HorizonError> = retry(fast_policy(), || {
let calls_inner = Arc::clone(&calls_for_op);
async move {
let attempt = calls_inner.fetch_add(1_u32, Ordering::SeqCst);
if attempt < 2_u32 {
Err(database_error_with_sqlstate("40001"))
} else {
Ok(11_i32)
}
}
})
.await;
assert_eq!(result.unwrap(), 11_i32);
assert_eq!(calls.load(Ordering::SeqCst), 3_u32);
}
#[tokio::test]
async fn retry_exhausts_and_returns_last_error() {
let calls = Arc::new(AtomicU32::new(0_u32));
let calls_for_op = Arc::clone(&calls);
let result: Result<i32, HorizonError> = retry(fast_policy(), || {
let calls_inner = Arc::clone(&calls_for_op);
async move {
calls_inner.fetch_add(1_u32, Ordering::SeqCst);
Err(database_error_with_sqlstate("40P01"))
}
})
.await;
match result {
Err(HorizonError::Postgres(PostgresError::Query(sqlx::Error::Database(database)))) => {
assert_eq!(database.code().as_deref(), Some("40P01"));
}
other => panic!("expected deadlock SQLSTATE error, got {other:?}"),
}
assert_eq!(calls.load(Ordering::SeqCst), 4_u32);
}
#[tokio::test]
async fn retry_propagates_non_transient_error_immediately() {
let calls = Arc::new(AtomicU32::new(0_u32));
let calls_for_op = Arc::clone(&calls);
let result: Result<i32, HorizonError> = retry(fast_policy(), || {
let calls_inner = Arc::clone(&calls_for_op);
async move {
calls_inner.fetch_add(1_u32, Ordering::SeqCst);
Err(database_error_with_sqlstate("23505"))
}
})
.await;
assert!(matches!(
result,
Err(HorizonError::Postgres(PostgresError::Query(
sqlx::Error::Database(_)
)))
));
assert_eq!(calls.load(Ordering::SeqCst), 1_u32);
}
#[tokio::test]
async fn retry_with_zero_max_retries_runs_once() {
let calls = Arc::new(AtomicU32::new(0_u32));
let calls_for_op = Arc::clone(&calls);
let result: Result<i32, HorizonError> =
retry(RetryPolicy::default().with_max_retries(0_u32), || {
let calls_inner = Arc::clone(&calls_for_op);
async move {
calls_inner.fetch_add(1_u32, Ordering::SeqCst);
Err(database_error_with_sqlstate("40001"))
}
})
.await;
assert!(result.is_err());
assert_eq!(calls.load(Ordering::SeqCst), 1_u32);
}
#[test]
fn jittered_never_returns_negative_for_small_delay() {
for _ in 0_i32..32_i32 {
let actual = jittered(Duration::from_millis(1_u64));
assert!(actual >= Duration::ZERO);
}
}
}