use crate::client::retry::{RequestError, RetryError};
use crate::client::{HttpError, HttpErrorKind};
use crate::{Error, PutMultipartOptions};
use async_trait::async_trait;
use http::StatusCode;
use std::error::Error as StdError;
use std::fmt::Debug;
use std::sync::Arc;
use std::time::Duration;
#[cfg(not(all(target_arch = "wasm32", target_os = "unknown")))]
use std::time::Instant;
#[cfg(all(target_arch = "wasm32", target_os = "unknown"))]
use web_time::Instant;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum RetryFailure {
Status(StatusCode),
Transport(HttpErrorKind),
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub struct RetryContext {
pub failure: RetryFailure,
pub attempt: usize,
pub elapsed: Duration,
}
#[async_trait]
pub trait RetryPolicy: Debug + Send + Sync + 'static {
async fn retry(&self, context: RetryContext) -> bool;
}
#[derive(Clone)]
struct MultipartRetryPolicy(Arc<dyn RetryPolicy>);
impl PutMultipartOptions {
#[must_use]
pub fn with_retry_policy(mut self, policy: Arc<dyn RetryPolicy>) -> Self {
self.extensions.insert(MultipartRetryPolicy(policy));
self
}
pub(crate) fn retry_policy(&self) -> Option<Arc<dyn RetryPolicy>> {
self.extensions
.get::<MultipartRetryPolicy>()
.map(|policy| Arc::clone(&policy.0))
}
}
pub(crate) struct MultipartRetry {
policy: Option<Arc<dyn RetryPolicy>>,
attempt: usize,
start: Instant,
}
impl MultipartRetry {
pub(crate) fn new(policy: Option<Arc<dyn RetryPolicy>>) -> Self {
Self {
policy,
attempt: 0,
start: Instant::now(),
}
}
pub(crate) async fn should_retry(&mut self, error: &Error) -> bool {
let Some(policy) = self.policy.as_ref() else {
return false;
};
let Some(failure) = classify_http_failure(error) else {
return false;
};
self.attempt += 1;
policy
.retry(RetryContext {
failure,
attempt: self.attempt,
elapsed: self.start.elapsed(),
})
.await
}
}
fn classify_http_failure(error: &Error) -> Option<RetryFailure> {
let mut current: &(dyn StdError + 'static) = error;
loop {
if let Some(error) = current.downcast_ref::<RetryError>() {
return classify_request_error(error.inner());
}
if let Some(error) = current.downcast_ref::<HttpError>() {
return Some(RetryFailure::Transport(error.kind()));
}
current = current.source()?;
}
}
fn classify_request_error(error: &RequestError) -> Option<RetryFailure> {
match error {
RequestError::Status { status, .. } | RequestError::Response { status, .. } => {
Some(RetryFailure::Status(*status))
}
RequestError::Http(error) => Some(RetryFailure::Transport(error.kind())),
RequestError::BareRedirect => None,
}
}
#[cfg(test)]
mod tests {
use super::*;
use parking_lot::Mutex;
#[derive(Debug, Default)]
struct RecordingPolicy(Mutex<Vec<RetryContext>>);
#[async_trait]
impl RetryPolicy for RecordingPolicy {
async fn retry(&self, context: RetryContext) -> bool {
self.0.lock().push(context);
true
}
}
#[derive(Debug, thiserror::Error)]
#[error("test error")]
struct TestError;
#[tokio::test]
async fn retries_http_errors() {
let policy = Arc::new(RecordingPolicy::default());
let mut retry = MultipartRetry::new(Some(Arc::clone(&policy) as Arc<dyn RetryPolicy>));
let error = Error::Generic {
store: "test",
source: Box::new(HttpError::new(HttpErrorKind::Timeout, TestError)),
};
assert!(retry.should_retry(&error).await);
let contexts = policy.0.lock();
assert_eq!(contexts.len(), 1);
assert_eq!(contexts[0].attempt, 1);
assert_eq!(
contexts[0].failure,
RetryFailure::Transport(HttpErrorKind::Timeout)
);
}
#[tokio::test]
async fn does_not_retry_non_http_errors() {
let policy = Arc::new(RecordingPolicy::default());
let mut retry = MultipartRetry::new(Some(Arc::clone(&policy) as Arc<dyn RetryPolicy>));
let error = Error::Generic {
store: "test",
source: Box::new(TestError),
};
assert!(!retry.should_retry(&error).await);
assert!(policy.0.lock().is_empty());
}
#[test]
fn classifies_error_responses_with_success_status() {
let error = RequestError::Response {
status: StatusCode::OK,
body: "InternalError".into(),
};
assert_eq!(
classify_request_error(&error),
Some(RetryFailure::Status(StatusCode::OK))
);
}
}