use std::future::Future;
use std::io::ErrorKind;
use std::path::Path;
use std::time::Duration;
use super::UdsSecurityError;
use super::rpc::UdsRpcError;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ConnectRetry {
pub attempts: u32,
pub initial_backoff: Duration,
pub max_backoff: Duration,
}
impl ConnectRetry {
#[must_use]
pub const fn per_request() -> Self {
Self {
attempts: 3,
initial_backoff: Duration::from_millis(20),
max_backoff: Duration::from_millis(100),
}
}
#[must_use]
pub const fn startup() -> Self {
Self {
attempts: 8,
initial_backoff: Duration::from_millis(100),
max_backoff: Duration::from_millis(500),
}
}
#[must_use]
pub const fn single_attempt() -> Self {
Self {
attempts: 1,
initial_backoff: Duration::ZERO,
max_backoff: Duration::ZERO,
}
}
#[must_use]
pub fn backoff_after(&self, attempt: u32) -> Duration {
let shift = attempt.saturating_sub(1).min(31);
self.initial_backoff
.checked_mul(1u32 << shift)
.unwrap_or(self.max_backoff)
.min(self.max_backoff)
}
#[must_use]
pub fn backoff_floor(&self) -> Duration {
(1..self.attempts.max(1))
.map(|attempt| self.backoff_after(attempt))
.sum()
}
}
impl Default for ConnectRetry {
fn default() -> Self {
Self::per_request()
}
}
pub(super) async fn with_connect_retry<T, A, AFut, S, SFut>(
path: &Path,
policy: ConnectRetry,
sleep: S,
mut attempt: A,
) -> Result<T, UdsRpcError>
where
A: FnMut(u32) -> AFut,
AFut: Future<Output = Result<T, UdsRpcError>>,
S: Fn(Duration) -> SFut,
SFut: Future<Output = ()>,
{
let total = policy.attempts.max(1);
let mut n: u32 = 1;
let mut last_failure: Option<String> = None;
loop {
let err = match attempt(n).await {
Ok(value) => {
if let Some(last) = last_failure {
tracing::warn!(
socket = %path.display(),
attempts = n,
last_error = %last,
"unix socket dial succeeded after a bounded retry"
);
}
return Ok(value);
}
Err(err) => err,
};
if !is_transient(&err) {
return Err(err);
}
if n >= total {
if n == 1 {
return Err(err);
}
tracing::error!(
socket = %path.display(),
attempts = n,
error = %err,
"unix socket dial failed after every bounded retry"
);
return Err(UdsRpcError::ConnectRetriesExhausted {
path: path.to_path_buf(),
attempts: n,
source: Box::new(err),
});
}
let delay = policy.backoff_after(n);
tracing::debug!(
socket = %path.display(),
attempt = n,
attempts = total,
delay_ms = delay.as_millis() as u64,
error = %err,
"unix socket dial failed; retrying after backoff"
);
last_failure = Some(err.to_string());
sleep(delay).await;
n += 1;
}
}
fn is_transient(err: &UdsRpcError) -> bool {
match err {
UdsRpcError::Dial { source, .. } => dial_is_transient(source),
UdsRpcError::Write { source, .. } => source.kind() == ErrorKind::NotConnected,
_ => false,
}
}
fn dial_is_transient(source: &UdsSecurityError) -> bool {
match source {
UdsSecurityError::Connect { source, .. } => matches!(
source.kind(),
ErrorKind::ConnectionRefused
| ErrorKind::NotConnected
| ErrorKind::WouldBlock
| ErrorKind::Interrupted
),
UdsSecurityError::StatForConnect { source, .. } => source.kind() == ErrorKind::NotFound,
_ => false,
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::cell::RefCell;
use std::path::PathBuf;
fn sock() -> PathBuf {
PathBuf::from("/tmp/trusty-retry-test.sock")
}
fn transient() -> UdsRpcError {
UdsRpcError::Dial {
path: sock(),
source: UdsSecurityError::Connect {
path: sock(),
source: std::io::Error::from(ErrorKind::ConnectionRefused),
},
}
}
fn permanent() -> UdsRpcError {
UdsRpcError::Dial {
path: sock(),
source: UdsSecurityError::UntrustedSocket {
path: sock(),
reason: "socket is mode 0755, not 0600".to_string(),
},
}
}
#[test]
fn backoff_doubles_and_stops_at_the_ceiling() {
let policy = ConnectRetry {
attempts: 6,
initial_backoff: Duration::from_millis(20),
max_backoff: Duration::from_millis(100),
};
let schedule: Vec<u64> = (1..=5)
.map(|n| policy.backoff_after(n).as_millis() as u64)
.collect();
assert_eq!(schedule, vec![20, 40, 80, 100, 100]);
}
#[test]
fn backoff_floor_sums_every_delay_the_policy_will_sleep() {
assert_eq!(
ConnectRetry::per_request().backoff_floor(),
Duration::from_millis(60)
);
assert_eq!(
ConnectRetry::single_attempt().backoff_floor(),
Duration::ZERO
);
}
#[test]
fn transient_classification_covers_the_three_dial_errnos() {
for kind in [
ErrorKind::ConnectionRefused,
ErrorKind::NotConnected,
ErrorKind::WouldBlock,
ErrorKind::Interrupted,
] {
let err = UdsRpcError::Dial {
path: sock(),
source: UdsSecurityError::Connect {
path: sock(),
source: std::io::Error::from(kind),
},
};
assert!(is_transient(&err), "{kind:?} should be retried");
}
let absent = UdsRpcError::Dial {
path: sock(),
source: UdsSecurityError::StatForConnect {
path: sock(),
source: std::io::Error::from(ErrorKind::NotFound),
},
};
assert!(is_transient(&absent), "an absent socket should be retried");
assert!(!is_transient(&UdsRpcError::NoResponse { path: sock() }));
assert!(!is_transient(&UdsRpcError::Write {
path: sock(),
source: std::io::Error::from(ErrorKind::BrokenPipe),
}));
assert!(is_transient(&UdsRpcError::Write {
path: sock(),
source: std::io::Error::from(ErrorKind::NotConnected),
}));
}
#[test]
fn a_failed_half_close_is_never_retried() {
for kind in [
ErrorKind::NotConnected,
ErrorKind::BrokenPipe,
ErrorKind::ConnectionReset,
] {
let err = UdsRpcError::HalfClose {
path: sock(),
source: std::io::Error::from(kind),
};
assert!(
!is_transient(&err),
"the frame is already on the wire; {kind:?} must not be retried"
);
}
}
#[tokio::test]
async fn a_late_non_transient_failure_is_reported_verbatim() {
let seen: RefCell<Vec<u32>> = RefCell::new(Vec::new());
let err = with_connect_retry::<(), _, _, _, _>(
&sock(),
ConnectRetry::per_request(),
|_d| async {},
|n| {
seen.borrow_mut().push(n);
async move {
if n == 1 {
Err(transient())
} else {
Err(permanent())
}
}
},
)
.await
.expect_err("the second attempt fails permanently");
assert_eq!(seen.into_inner(), vec![1, 2]);
assert!(
matches!(err, UdsRpcError::Dial { .. }),
"expected the permanent error verbatim, got {err:?}"
);
}
#[tokio::test]
async fn retry_stops_after_exactly_the_policys_attempt_count() {
let slept: RefCell<Vec<Duration>> = RefCell::new(Vec::new());
let seen: RefCell<Vec<u32>> = RefCell::new(Vec::new());
let policy = ConnectRetry::per_request();
let err = with_connect_retry::<(), _, _, _, _>(
&sock(),
policy,
|d| {
slept.borrow_mut().push(d);
async {}
},
|n| {
seen.borrow_mut().push(n);
async { Err(transient()) }
},
)
.await
.expect_err("every attempt fails");
assert_eq!(seen.into_inner(), vec![1, 2, 3]);
let slept = slept.into_inner();
assert_eq!(
slept,
vec![Duration::from_millis(20), Duration::from_millis(40)]
);
assert_eq!(
slept.iter().sum::<Duration>(),
policy.backoff_floor(),
"the driver must sleep the policy's whole floor"
);
assert!(
matches!(&err, UdsRpcError::ConnectRetriesExhausted { attempts, .. } if *attempts == 3),
"got {err:?}"
);
}
#[tokio::test]
async fn retry_returns_the_first_success_without_further_attempts() {
let seen: RefCell<Vec<u32>> = RefCell::new(Vec::new());
let got = with_connect_retry(
&sock(),
ConnectRetry::per_request(),
|_d| async {},
|n| {
seen.borrow_mut().push(n);
async move { if n < 2 { Err(transient()) } else { Ok(n) } }
},
)
.await
.expect("the second attempt succeeds");
assert_eq!(got, 2);
assert_eq!(seen.into_inner(), vec![1, 2]);
}
#[tokio::test]
async fn a_non_transient_failure_is_not_retried() {
let seen: RefCell<Vec<u32>> = RefCell::new(Vec::new());
let err = with_connect_retry::<(), _, _, _, _>(
&sock(),
ConnectRetry::startup(),
|_d| async { panic!("a permanent refusal must not sleep") },
|n| {
seen.borrow_mut().push(n);
async { Err(permanent()) }
},
)
.await
.expect_err("a security refusal is terminal");
assert_eq!(seen.into_inner(), vec![1]);
assert!(matches!(err, UdsRpcError::Dial { .. }), "got {err:?}");
}
#[tokio::test]
async fn a_single_attempt_policy_returns_the_underlying_error_unwrapped() {
let err = with_connect_retry::<(), _, _, _, _>(
&sock(),
ConnectRetry::single_attempt(),
|_d| async { panic!("a single-attempt policy must not sleep") },
|_n| async { Err(transient()) },
)
.await
.expect_err("one attempt, one error");
assert!(matches!(err, UdsRpcError::Dial { .. }), "got {err:?}");
}
}