use std::time::Duration;
use crate::error::{Error, Result};
use crate::registry::Endpoint;
use crate::transport::{Query, RawResponse, Transport};
#[cfg(feature = "async")]
use crate::transport::{AsyncTransport, BoxFuture};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RetryPolicy {
pub max_attempts: u32,
pub initial_backoff: Duration,
pub backoff_multiplier: u32,
pub max_backoff: Duration,
}
impl RetryPolicy {
pub const DEFAULT: RetryPolicy = RetryPolicy {
max_attempts: 3,
initial_backoff: Duration::from_millis(250),
backoff_multiplier: 2,
max_backoff: Duration::from_secs(4),
};
pub const NONE: RetryPolicy = RetryPolicy {
max_attempts: 1,
initial_backoff: Duration::ZERO,
backoff_multiplier: 1,
max_backoff: Duration::ZERO,
};
pub fn fixed(attempts: u32, delay: Duration) -> Self {
RetryPolicy {
max_attempts: attempts.max(1),
initial_backoff: delay,
backoff_multiplier: 1,
max_backoff: delay,
}
}
pub fn backoff_before(&self, attempt: u32) -> Duration {
if attempt <= 1 {
return Duration::ZERO;
}
let steps = attempt - 2;
let factor = self.backoff_multiplier.saturating_pow(steps);
self.initial_backoff
.saturating_mul(factor.max(1))
.min(self.max_backoff)
}
pub fn should_retry(&self, attempt: u32, error: &Error) -> bool {
attempt < self.max_attempts && error.is_transient()
}
}
impl Default for RetryPolicy {
fn default() -> Self {
RetryPolicy::DEFAULT
}
}
#[derive(Debug, Clone)]
pub struct RetryTransport<T> {
inner: T,
policy: RetryPolicy,
}
impl<T> RetryTransport<T> {
pub fn new(inner: T, policy: RetryPolicy) -> Self {
RetryTransport { inner, policy }
}
pub fn policy(&self) -> RetryPolicy {
self.policy
}
pub fn inner(&self) -> &T {
&self.inner
}
pub fn into_inner(self) -> T {
self.inner
}
}
impl<T: Transport> Transport for RetryTransport<T> {
fn supports(&self, endpoint: &Endpoint) -> bool {
self.inner.supports(endpoint)
}
fn fetch(&self, query: &Query) -> Result<RawResponse> {
let mut attempt = 1;
loop {
let backoff = self.policy.backoff_before(attempt);
if !backoff.is_zero() {
std::thread::sleep(backoff);
}
match self.inner.fetch(query) {
Ok(response) => return Ok(response),
Err(error) if self.policy.should_retry(attempt, &error) => {
attempt += 1;
}
Err(error) => return Err(error),
}
}
}
fn name(&self) -> String {
format!("retry({})", self.inner.name())
}
}
#[cfg(feature = "async")]
#[derive(Debug, Clone)]
pub struct AsyncRetryTransport<T> {
inner: T,
policy: RetryPolicy,
}
#[cfg(feature = "async")]
impl<T> AsyncRetryTransport<T> {
pub fn new(inner: T, policy: RetryPolicy) -> Self {
AsyncRetryTransport { inner, policy }
}
pub fn policy(&self) -> RetryPolicy {
self.policy
}
pub fn inner(&self) -> &T {
&self.inner
}
pub fn into_inner(self) -> T {
self.inner
}
}
#[cfg(feature = "async")]
impl<T: AsyncTransport> AsyncTransport for AsyncRetryTransport<T> {
fn supports(&self, endpoint: &Endpoint) -> bool {
self.inner.supports(endpoint)
}
fn fetch<'a>(&'a self, query: &'a Query) -> BoxFuture<'a, Result<RawResponse>> {
Box::pin(async move {
let mut attempt = 1;
loop {
let backoff = self.policy.backoff_before(attempt);
if !backoff.is_zero() {
tokio::time::sleep(backoff).await;
}
match self.inner.fetch(query).await {
Ok(response) => return Ok(response),
Err(error) if self.policy.should_retry(attempt, &error) => {
attempt += 1;
}
Err(error) => return Err(error),
}
}
})
}
fn name(&self) -> String {
format!("async-retry({})", self.inner.name())
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::error::Refusal;
use crate::transport::mock::{MockTransport, Scripted};
use crate::transport::ResponseKind;
fn query() -> Query {
Query::new(
Endpoint::whois("whois.example"),
"example.com",
crate::domain::Tld::parse("com").unwrap(),
)
}
fn timeout() -> Error {
Error::Timeout {
server: "whois.example".into(),
elapsed: Duration::ZERO,
}
}
#[test]
fn backoff_grows_then_caps() {
let policy = RetryPolicy::DEFAULT;
assert_eq!(policy.backoff_before(1), Duration::ZERO);
assert_eq!(policy.backoff_before(2), Duration::from_millis(250));
assert_eq!(policy.backoff_before(3), Duration::from_millis(500));
assert_eq!(policy.backoff_before(4), Duration::from_secs(1));
assert_eq!(policy.backoff_before(u32::MAX), policy.max_backoff);
}
#[test]
fn fixed_policy_does_not_grow() {
let policy = RetryPolicy::fixed(4, Duration::from_millis(10));
assert_eq!(policy.backoff_before(2), Duration::from_millis(10));
assert_eq!(policy.backoff_before(4), Duration::from_millis(10));
}
#[test]
fn a_transient_failure_is_retried_until_it_succeeds() {
let inner = MockTransport::new(vec![
Scripted::Fail(timeout()),
Scripted::Fail(timeout()),
Scripted::Answer("No match for EXAMPLE.COM".into()),
]);
let transport = RetryTransport::new(
inner.clone(),
RetryPolicy::fixed(3, Duration::from_millis(1)),
);
let response = transport.fetch(&query()).unwrap();
assert_eq!(response.kind(), ResponseKind::WhoisText);
assert_eq!(inner.call_count(), 3);
}
#[test]
fn attempts_are_capped() {
let inner = MockTransport::new(vec![
Scripted::Fail(timeout()),
Scripted::Fail(timeout()),
Scripted::Fail(timeout()),
Scripted::Answer("too late".into()),
]);
let transport = RetryTransport::new(
inner.clone(),
RetryPolicy::fixed(2, Duration::from_millis(1)),
);
assert!(transport.fetch(&query()).is_err());
assert_eq!(inner.call_count(), 2);
}
#[test]
fn a_permanent_failure_is_not_retried() {
let inner = MockTransport::new(vec![
Scripted::Fail(Error::UnsupportedTld {
tld: crate::domain::Tld::parse("nope").unwrap(),
}),
Scripted::Answer("unreachable".into()),
]);
let transport = RetryTransport::new(inner.clone(), RetryPolicy::DEFAULT);
assert!(transport.fetch(&query()).is_err());
assert_eq!(
inner.call_count(),
1,
"a permanent error must not be retried"
);
}
#[test]
fn rate_limiting_is_retried_but_a_block_is_not() {
for (reason, expected_calls) in [(Refusal::RateLimited, 2), (Refusal::Blocked, 1)] {
let inner = MockTransport::new(vec![
Scripted::Fail(Error::Refused {
server: "whois.example".into(),
reason,
}),
Scripted::Answer("second".into()),
]);
let transport = RetryTransport::new(
inner.clone(),
RetryPolicy::fixed(2, Duration::from_millis(1)),
);
let _ = transport.fetch(&query());
assert_eq!(inner.call_count(), expected_calls, "for {reason:?}");
}
}
#[test]
fn name_shows_the_wrapping() {
let transport = RetryTransport::new(MockTransport::new(vec![]), RetryPolicy::NONE);
assert_eq!(transport.name(), "retry(mock)");
}
}