1use async_trait::async_trait;
5use http::Extensions;
6use rand::{rng, RngExt};
7use reqwest::{Request, Response, StatusCode};
8use reqwest_middleware::{Middleware, Next, Result};
9use std::time::Duration;
10
11pub struct CustomRetryMiddleware {
13 max_retries: u32,
14 max_delay_ms: u64,
15 initial_delay_ms: u64,
16}
17
18#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
19#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
20impl Middleware for CustomRetryMiddleware {
21 async fn handle(
22 &self,
23 req: Request,
24 extensions: &mut Extensions,
25 next: Next<'_>,
26 ) -> Result<Response> {
27 self.execute_with_retry(req, next, extensions).await
28 }
29}
30
31impl CustomRetryMiddleware {
32 pub fn new(max_retries: u32, max_delay_ms: u64, initial_delay_ms: u64) -> Self {
34 Self {
35 max_retries: max_retries.min(10),
36 max_delay_ms,
37 initial_delay_ms,
38 }
39 }
40
41 async fn execute_with_retry<'a>(
42 &'a self,
43 req: Request,
44 next: Next<'a>,
45 ext: &'a mut Extensions,
46 ) -> Result<Response> {
47 let mut n_past_retries = 0;
48 let mut last_req_401 = false;
49 loop {
50 let duplicate_request = match req.try_clone() {
51 Some(x) => x,
52 None => return next.run(req, ext).await,
53 };
54
55 let result = next.clone().run(duplicate_request, ext).await;
56
57 break match Retryable::from_reqwest_response(&result) {
59 Some(retryable)
60 if (retryable == Retryable::Transient
61 || retryable == Retryable::Unauthorized && !last_req_401)
62 && n_past_retries < self.max_retries =>
63 {
64 last_req_401 = retryable == Retryable::Unauthorized;
65 let mut retry_delay = self.initial_delay_ms * 2u64.pow(n_past_retries);
68 if retry_delay > self.max_delay_ms {
69 retry_delay = self.max_delay_ms;
70 }
71 retry_delay = retry_delay / 4 * 3 + rng().random_range(0..=(retry_delay / 2));
73 futures_timer::Delay::new(Duration::from_millis(retry_delay)).await;
74 n_past_retries += 1;
75 continue;
76 }
77 Some(_) | None => result,
78 };
79 }
80 }
81}
82
83#[derive(PartialEq, Eq)]
84pub(crate) enum Retryable {
85 Transient,
87 Fatal,
89 Unauthorized,
91}
92
93impl Retryable {
94 pub fn from_reqwest_response(
102 res: &reqwest_middleware::Result<reqwest::Response>,
103 ) -> Option<Self> {
104 match res {
105 Ok(success) => {
106 let status = success.status();
107 if status.is_success() {
108 None
109 } else if status == StatusCode::UNAUTHORIZED {
110 Some(Retryable::Unauthorized)
111 } else if status.is_server_error()
112 || status == StatusCode::REQUEST_TIMEOUT
113 || status == StatusCode::TOO_MANY_REQUESTS
114 || success
115 .headers()
116 .get("cdf-is-auto-retryable")
117 .and_then(|v| v.to_str().ok())
118 .is_some_and(|v| v == "true")
119 {
120 Some(Retryable::Transient)
121 } else {
122 Some(Retryable::Fatal)
123 }
124 }
125 Err(error) => match error {
126 reqwest_middleware::Error::Middleware(_) => Some(Retryable::Fatal),
127 reqwest_middleware::Error::Reqwest(error) => {
128 #[cfg(not(target_arch = "wasm32"))]
129 let is_connect = error.is_connect();
130 #[cfg(target_arch = "wasm32")]
131 let is_connect = false;
132
133 if error.is_timeout() || is_connect {
134 Some(Retryable::Transient)
135 } else if error.is_body()
136 || error.is_decode()
137 || error.is_builder()
138 || error.is_redirect()
139 || error.is_request()
140 {
141 Some(Retryable::Fatal)
142 } else {
143 None
147 }
148 }
149 },
150 }
151 }
152}