use std::{
fmt,
future::Future,
io::Write as _,
sync::Arc,
time::{Duration, Instant, SystemTime},
};
use http::{HeaderMap, Method, Uri, header::RETRY_AFTER};
use crate::{
config::ZERO_TIMEOUT,
constants::RETRY_AFTER_MS_HEADER,
error::{Error, ErrorKind, parse_retry_after},
telemetry,
};
type Predicate = Arc<dyn Fn(&Error) -> bool + Send + Sync>;
#[derive(Clone)]
pub struct RetryPolicy {
max_retries: u32,
backoff_initial: Duration,
backoff_max: Duration,
backoff_jitter: f64,
http_statuses: StatusSet,
respect_retry_after: bool,
api_connection_error: bool,
api_timeout_error: bool,
predicate: Option<Predicate>,
timeout: Option<Duration>,
#[cfg(test)]
time: Option<Arc<tests::FakeTime>>,
}
impl RetryPolicy {
#[must_use]
pub const fn new() -> Self {
Self {
max_retries: 2,
backoff_initial: Duration::from_millis(500),
backoff_max: Duration::from_secs(5),
backoff_jitter: 0.25,
http_statuses: StatusSet::DEFAULT,
respect_retry_after: true,
api_connection_error: true,
api_timeout_error: true,
predicate: None,
timeout: Some(Duration::from_secs(30)),
#[cfg(test)]
time: None,
}
}
#[must_use]
pub fn max_retries(mut self, retries: u32) -> Self {
self.max_retries = retries;
self
}
#[must_use]
pub fn backoff_initial(mut self, delay: Duration) -> Self {
self.backoff_initial = delay;
self
}
#[must_use]
pub fn backoff_max(mut self, delay: Duration) -> Self {
self.backoff_max = delay;
self
}
pub fn backoff_jitter(mut self, jitter: f64) -> Result<Self, Error> {
if !(0.0..=1.0).contains(&jitter) {
return Err(Error::config("backoff_jitter must be between zero and one."));
}
self.backoff_jitter = jitter;
Ok(self)
}
#[must_use]
pub fn http_statuses(mut self, statuses: StatusSet) -> Self {
self.http_statuses = statuses;
self
}
#[must_use]
pub fn respect_retry_after(mut self, respect: bool) -> Self {
self.respect_retry_after = respect;
self
}
#[must_use]
pub fn api_connection_error(mut self, retry: bool) -> Self {
self.api_connection_error = retry;
self
}
#[must_use]
pub fn api_timeout_error(mut self, retry: bool) -> Self {
self.api_timeout_error = retry;
self
}
#[must_use]
pub fn predicate<F>(mut self, predicate: F) -> Self
where
F: Fn(&Error) -> bool + Send + Sync + 'static,
{
self.predicate = Some(Arc::new(predicate));
self
}
pub fn timeout(mut self, budget: Duration) -> Result<Self, Error> {
if budget.is_zero() {
return Err(Error::config(ZERO_TIMEOUT));
}
self.timeout = Some(budget);
Ok(self)
}
#[must_use]
pub fn no_timeout(mut self) -> Self {
self.timeout = None;
self
}
pub(crate) fn can_retry(&self) -> bool {
self.max_retries > 0
}
fn retryable(&self, error: &Error) -> bool {
let builtin = match error.kind() {
ErrorKind::Timeout { .. } => self.api_timeout_error,
ErrorKind::Connection => self.api_connection_error,
ErrorKind::Api(api) => self.http_statuses.contains(api.status().as_u16()),
ErrorKind::InvalidRequest
| ErrorKind::Config
| ErrorKind::ResponseValidation(_)
| ErrorKind::ResponseTooLarge { .. } => false,
};
builtin || self.predicate.as_ref().is_some_and(|predicate| predicate(error))
}
fn delay<T: Time>(&self, time: &T, attempts: u32, error: &Error) -> Duration {
if self.respect_retry_after {
let names_delay = |headers: &HeaderMap| {
headers.contains_key(RETRY_AFTER_MS_HEADER) || headers.contains_key(RETRY_AFTER)
};
let asked = match error.kind() {
ErrorKind::Api(api) if names_delay(api.headers()) => {
parse_retry_after(api.headers(), time.system_now())
}
ErrorKind::ResponseValidation(invalid) if names_delay(invalid.headers()) => {
parse_retry_after(invalid.headers(), time.system_now())
}
_ => None,
};
if let Some(asked) = asked {
return asked;
}
}
seconds_to_duration(backoff_seconds(
attempts,
self.backoff_initial.as_secs_f64(),
self.backoff_max.as_secs_f64(),
self.backoff_jitter,
time.draw(),
))
}
fn next_delay<T: Time>(
&self,
time: &T,
retries: u32,
error: &Error,
started: Option<Instant>,
) -> Option<Duration> {
if !self.retryable(error) || retries >= self.max_retries {
return None;
}
let attempts = retries + 1;
let delay = self.delay(time, attempts, error);
if let (Some(budget), Some(started)) = (self.timeout, started) {
let elapsed = time.now().saturating_duration_since(started);
if elapsed.saturating_add(delay) >= budget {
return None;
}
}
Some(delay)
}
}
impl Default for RetryPolicy {
fn default() -> Self {
Self::new()
}
}
impl fmt::Debug for RetryPolicy {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("RetryPolicy")
.field("max_retries", &self.max_retries)
.field("backoff_initial", &self.backoff_initial)
.field("backoff_max", &self.backoff_max)
.field("backoff_jitter", &self.backoff_jitter)
.field("http_statuses", &self.http_statuses)
.field("respect_retry_after", &self.respect_retry_after)
.field("api_connection_error", &self.api_connection_error)
.field("api_timeout_error", &self.api_timeout_error)
.field("predicate", &self.predicate.as_ref().map(|_| Opaque))
.field("timeout", &self.timeout)
.finish()
}
}
struct Opaque;
impl fmt::Debug for Opaque {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str("<predicate>")
}
}
const STATUS_LIMIT: u16 = 640;
#[derive(Clone, Copy, PartialEq, Eq, Hash)]
pub struct StatusSet([u64; 10]);
impl StatusSet {
pub const DEFAULT: Self = {
let mut set = Self::empty();
set.insert(408);
set.insert(429);
let mut status = 500;
while status < 600 {
set.insert(status);
status += 1;
}
set
};
#[must_use]
pub const fn empty() -> Self {
Self([0; 10])
}
#[must_use]
pub const fn contains(&self, status: u16) -> bool {
match Self::slot(status) {
Some((word, bit)) => self.0[word] & bit != 0,
None => false,
}
}
pub const fn insert(&mut self, status: u16) -> bool {
match Self::slot(status) {
Some((word, bit)) => {
let added = self.0[word] & bit == 0;
self.0[word] |= bit;
added
}
None => false,
}
}
pub const fn remove(&mut self, status: u16) -> bool {
match Self::slot(status) {
Some((word, bit)) => {
let removed = self.0[word] & bit != 0;
self.0[word] &= !bit;
removed
}
None => false,
}
}
#[must_use]
pub const fn is_empty(&self) -> bool {
let mut word = 0;
while word < self.0.len() {
if self.0[word] != 0 {
return false;
}
word += 1;
}
true
}
pub fn iter(&self) -> impl Iterator<Item = u16> + use<> {
let words = self.0;
(0_u16..).zip(words).flat_map(|(index, word)| {
let mut rest = word;
std::iter::from_fn(move || {
if rest == 0 {
return None;
}
let bit = rest.trailing_zeros() as u16;
rest &= rest - 1;
Some(index * 64 + bit)
})
})
}
const fn slot(status: u16) -> Option<(usize, u64)> {
if status >= STATUS_LIMIT {
return None;
}
Some(((status / 64) as usize, 1 << (status % 64)))
}
}
impl Default for StatusSet {
fn default() -> Self {
Self::DEFAULT
}
}
impl FromIterator<u16> for StatusSet {
fn from_iter<I: IntoIterator<Item = u16>>(statuses: I) -> Self {
let mut set = Self::empty();
set.extend(statuses);
set
}
}
impl Extend<u16> for StatusSet {
fn extend<I: IntoIterator<Item = u16>>(&mut self, statuses: I) {
for status in statuses {
self.insert(status);
}
}
}
impl fmt::Debug for StatusSet {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str("{")?;
let mut statuses = self.iter().peekable();
let mut first = true;
while let Some(start) = statuses.next() {
let mut end = start;
while statuses.next_if_eq(&(end + 1)).is_some() {
end += 1;
}
if !first {
formatter.write_str(", ")?;
}
first = false;
match end - start {
0 => write!(formatter, "{start}")?,
1 => write!(formatter, "{start}, {end}")?,
_ => write!(formatter, "{start}..={end}")?,
}
}
formatter.write_str("}")
}
}
trait Time {
fn now(&self) -> Instant;
fn system_now(&self) -> SystemTime;
fn sleep(&self, delay: Duration) -> impl Future<Output = ()> + Send;
fn draw(&self) -> f64;
}
struct Tokio;
impl Time for Tokio {
fn now(&self) -> Instant {
Instant::now()
}
fn system_now(&self) -> SystemTime {
SystemTime::now()
}
fn sleep(&self, delay: Duration) -> impl Future<Output = ()> + Send {
tokio::time::sleep(delay)
}
fn draw(&self) -> f64 {
fastrand::f64()
}
}
#[cfg(not(test))]
pub(crate) fn run<R, F, Fut>(
policy: &RetryPolicy,
method: &Method,
uri: &Uri,
attempt: F,
) -> impl Future<Output = Result<R, Error>>
where
F: FnMut(u32) -> Fut,
Fut: Future<Output = Result<R, Error>>,
{
run_on(&Tokio, policy, method, uri, attempt)
}
#[cfg(test)]
pub(crate) async fn run<R, F, Fut>(
policy: &RetryPolicy,
method: &Method,
uri: &Uri,
attempt: F,
) -> Result<R, Error>
where
F: FnMut(u32) -> Fut,
Fut: Future<Output = Result<R, Error>>,
{
if let Some(time) = &policy.time {
return run_on(&**time, policy, method, uri, attempt).await;
}
run_on(&Tokio, policy, method, uri, attempt).await
}
async fn run_on<T, R, F, Fut>(
time: &T,
policy: &RetryPolicy,
method: &Method,
uri: &Uri,
mut attempt: F,
) -> Result<R, Error>
where
T: Time,
F: FnMut(u32) -> Fut,
Fut: Future<Output = Result<R, Error>>,
{
let started = (policy.timeout.is_some() && policy.can_retry()).then(|| time.now());
let mut retry = 0_u32;
loop {
if retry > 0 {
telemetry::retrying(telemetry::Exchange::new(method, uri, retry));
}
let error = match attempt(retry).await {
Ok(done) => return Ok(done),
Err(error) => error,
};
let Some(delay) = policy.next_delay(time, retry, &error, started) else {
return Err(error);
};
drop(error);
time.sleep(delay).await;
retry += 1;
}
}
pub(crate) fn backoff_seconds(attempt: u32, initial: f64, max: f64, jitter: f64, draw: f64) -> f64 {
if initial == 0.0 || max == 0.0 {
return 0.0;
}
let exponential = if f64::from(attempt) - 1.0 >= max.log2() - initial.log2() {
max
} else {
let exponent = i32::try_from(i64::from(attempt) - 1).unwrap_or(i32::MAX);
scalbn(initial, exponent)
};
let delay = exponential * (1.0 - draw * jitter);
exponential.min(round_to_millis(delay))
}
fn seconds_to_duration(seconds: f64) -> Duration {
if seconds.is_nan() || seconds <= 0.0 {
return Duration::ZERO;
}
Duration::try_from_secs_f64(seconds).unwrap_or(Duration::MAX)
}
fn scalbn(x: f64, n: i32) -> f64 {
const TWO_POW_1023: f64 = f64::from_bits(0x7FE0_0000_0000_0000);
const TWO_POW_MINUS_969: f64 = f64::from_bits(0x0360_0000_0000_0000);
let mut y = x;
let mut n = n;
if n > 1023 {
y *= TWO_POW_1023;
n -= 1023;
if n > 1023 {
y *= TWO_POW_1023;
n = (n - 1023).min(1023);
}
} else if n < -1022 {
y *= TWO_POW_MINUS_969;
n += 1022 - 53;
if n < -1022 {
y *= TWO_POW_MINUS_969;
n = (n + 1022 - 53).max(-1022);
}
}
y * f64::from_bits(u64::from((0x3FF + n).unsigned_abs()) << 52)
}
fn round_to_millis(x: f64) -> f64 {
const INTEGRAL_FROM: f64 = 4_503_599_627_370_496.0; if x.is_nan() || x.abs() >= INTEGRAL_FROM {
return x;
}
let mut buf = [0_u8; 24];
let unused = {
let mut rest = &mut buf[..];
write!(rest, "{x:.3}")
.expect("invariant: a value below 2^52 needs at most 21 bytes at three decimals");
rest.len()
};
let written = buf.len() - unused;
std::str::from_utf8(&buf[..written])
.expect("invariant: formatted digits are ASCII")
.parse()
.expect("invariant: a formatted finite f64 parses back")
}
#[cfg(test)]
#[path = "retry_tests.rs"]
mod tests;