use std::{
convert::Infallible,
fmt::Write,
time::{Duration, Instant},
};
use non_non_full::NonEmptyVec;
use rand::Rng;
use tap::TapOptional;
use tokio::{select, time::sleep};
use crate::{
error::{Error, FetcherError},
job::{ErrorChainDisplay, Trigger, cancel_wait},
maybe_send::MaybeSync,
};
use super::{HandleError, HandleErrorContext, HandleErrorResult};
#[derive(Clone, Debug)]
pub struct ExponentialBackoff {
pub max_attempts: u32,
pub use_jitter: bool,
pub pause_duration_net_error: Duration,
last_error_info: Option<ErrorInfo>,
}
#[derive(Clone, Debug)]
struct ErrorInfo {
attempt: u32,
happened_at: Instant,
must_sleep_for: Duration,
}
impl<Tr> HandleError<Tr> for ExponentialBackoff
where
Tr: Trigger,
{
type HandlerErr = Infallible;
async fn handle_errors(
&mut self,
errors: NonEmptyVec<FetcherError>,
cx: HandleErrorContext<'_, Tr>,
) -> HandleErrorResult<Self::HandlerErr> {
if self.resume_job(&errors, cx).await {
HandleErrorResult::ResumeJob {
wait_for_trigger: false,
}
} else {
HandleErrorResult::StopWithErrors(errors)
}
}
}
impl ExponentialBackoff {
#[expect(missing_docs, reason = "self-explanatory")]
pub const DEFAULT_MAX_ATTEMPT_COUNT: u32 = 15;
#[expect(missing_docs, reason = "self-explanatory")]
pub const DEFAULT_NETWORK_ERROR_PAUSE_DURATION: Duration =
Duration::from_secs(5 * 60 );
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn new_with_conf(
max_attempts: u32,
use_jitter: bool,
pause_duration_net_error: Duration,
) -> Self {
Self {
max_attempts,
use_jitter,
pause_duration_net_error,
..Default::default()
}
}
#[must_use]
pub fn next_attempt<Tr: Trigger>(&mut self, job_trigger: &Tr) -> u32 {
self.reset_error_count(job_trigger);
match self.check_limit_reached() {
AttemptLimitReached::No { current_attempt } => current_attempt,
AttemptLimitReached::Yes => self.max_attempts,
}
}
}
#[derive(Clone, Copy, Debug)]
enum AttemptLimitReached {
No { current_attempt: u32 },
Yes,
}
impl ExponentialBackoff {
async fn resume_job<Tr: Trigger>(
&mut self,
errors: &[FetcherError],
cx: HandleErrorContext<'_, Tr>,
) -> bool {
self.reset_error_count(cx.job_trigger);
let fatal_errors = errors.iter().filter(|e| {
e.is_network_related()
.tap_some(|net_err| {
tracing::warn!("Network error: {}", ErrorChainDisplay(net_err));
})
.is_none()
});
if fatal_errors.clone().count() == 0 {
return pause_job(self.pause_duration_net_error, cx).await;
}
let attempt_limit_reached = self.check_limit_reached();
self.log(attempt_limit_reached, fatal_errors, cx.job_name);
let AttemptLimitReached::No { current_attempt } = attempt_limit_reached else {
return false;
};
let pause_duration =
exponential_backoff_duration(current_attempt, self.use_jitter, rand::rng());
self.last_error_info = Some(ErrorInfo {
attempt: current_attempt,
happened_at: Instant::now(),
must_sleep_for: pause_duration,
});
pause_job(pause_duration, cx).await
}
fn reset_error_count<Tr: Trigger>(&mut self, job_trigger: &Tr) {
let Some(last_error) = self.last_error_info.as_ref() else {
return;
};
if last_error.happened_at.elapsed()
> last_error.must_sleep_for + job_trigger.twice_as_duration()
{
self.reset();
}
}
fn check_limit_reached(&self) -> AttemptLimitReached {
let prev_attempt = self
.last_error_info
.as_ref()
.map(|info| info.attempt)
.unwrap_or(0);
let current_attempt = prev_attempt + 1;
if current_attempt >= self.max_attempts {
AttemptLimitReached::Yes
} else {
AttemptLimitReached::No { current_attempt }
}
}
fn log<'a>(
&self,
attempt: AttemptLimitReached,
fatal_errors: impl Iterator<Item = &'a FetcherError>,
job_name: &str,
) {
let current_attempt = match attempt {
AttemptLimitReached::Yes => {
tracing::warn!(
"Maximum error limit reached ({max} out of {max}) for job {job_name}. Stopping retrying...",
max = self.max_attempts,
);
return;
}
AttemptLimitReached::No { current_attempt } => current_attempt,
};
let mut err_msg = format!(
"Job {job_name} finished {current_attempt}/{max} times in an error ",
max = self.max_attempts,
);
for (i, err) in fatal_errors.enumerate() {
_ = write!(
err_msg,
"\nError #{err_num}:\n{e}\n",
err_num = i + 1,
e = ErrorChainDisplay(err)
);
}
tracing::error!("{}", err_msg);
}
fn reset(&mut self) {
self.last_error_info = None;
}
}
impl Default for ExponentialBackoff {
fn default() -> Self {
Self {
max_attempts: Self::DEFAULT_MAX_ATTEMPT_COUNT,
use_jitter: true,
pause_duration_net_error: Self::DEFAULT_NETWORK_ERROR_PAUSE_DURATION,
last_error_info: None,
}
}
}
fn exponential_backoff_duration(attempt: u32, use_jitter: bool, mut rng: impl Rng) -> Duration {
let base_duration_min = 2u64.saturating_pow(attempt.saturating_sub(1));
let base_duration_sec = base_duration_min * 60;
let final_duration = if use_jitter {
#[expect(clippy::cast_precision_loss, reason = "what other way is there?")]
let duration_secs_f64 = base_duration_sec as f64 * (rng.random::<f64>() + 0.5);
#[expect(clippy::cast_possible_truncation, reason = "what other way is there?")]
#[expect(clippy::cast_sign_loss, reason = "always positive")]
let duration_secs = duration_secs_f64.round() as u64;
tracing::debug!(
"Calculated exponential backoff duration: base = {base_duration_min}m ({base_duration_sec}s), with jitter = {duration_secs_f64}s (rounded to {duration_secs}s, ~{}m)",
duration_secs / 60
);
duration_secs
} else {
tracing::debug!("Calculated exponential backoff duration: {base_duration_min}m");
base_duration_sec
};
Duration::from_secs(final_duration)
}
async fn pause_job<Tr: MaybeSync>(dur: Duration, cx: HandleErrorContext<'_, Tr>) -> bool {
tracing::info!("Pausing job {} for {}m", cx.job_name, dur.as_secs() / 60);
select! {
() = sleep(dur) => {
true
}
() = cancel_wait(cx.cancel_token) => {
tracing::debug!("Job terminated mid exponential backoff pause");
false
}
}
}
#[cfg(test)]
mod tests {
#![expect(clippy::cast_precision_loss, clippy::unimplemented)]
use std::time::Duration;
use rand::Rng;
use super::exponential_backoff_duration;
fn check_exp_backoff_duration(
attempt: u32,
use_jitter: bool,
expected_result: Duration,
rng: impl Rng,
) {
let dur = exponential_backoff_duration(attempt, use_jitter, rng);
if use_jitter {
let dur = dur.as_secs() as f64;
let expected_dur = expected_result.as_secs() as f64;
assert!(dur <= (expected_dur * 1.5) && dur >= (expected_dur / 2.0));
} else {
assert_eq!(dur, expected_result, "attempt: {attempt}");
}
}
fn m(mins: u64) -> Duration {
Duration::from_secs(mins * 60 )
}
#[test]
fn exponential_backoff_duration_no_jitter() {
for i in 0u32..=15 {
let expected_mins = 2u64.pow(i.saturating_sub(1));
check_exp_backoff_duration(i, false, m(expected_mins), rand::rng());
}
}
#[test]
fn exponential_backoff_duration_with_jitter() {
for i in 0u32..=15 {
let expected_mins = 2u64.pow(i.saturating_sub(1));
check_exp_backoff_duration(i, true, m(expected_mins), rand::rng());
}
}
#[test]
fn exponential_backoff_duration_with_fake_jitter() {
struct AlwaysExtremes(bool);
impl rand::RngCore for AlwaysExtremes {
fn next_u64(&mut self) -> u64 {
let min_or_max = self.0;
self.0 = !self.0;
if min_or_max { u64::MIN } else { u64::MAX }
}
fn next_u32(&mut self) -> u32 {
unimplemented!()
}
fn fill_bytes(&mut self, _dst: &mut [u8]) {
unimplemented!()
}
}
let mut always_extremes = AlwaysExtremes(true);
for i in 0u32..=15 {
let expected_mins = 2u64.pow(i.saturating_sub(1));
check_exp_backoff_duration(i, true, m(expected_mins), &mut always_extremes);
check_exp_backoff_duration(i, true, m(expected_mins), &mut always_extremes);
}
}
}