1use std::time::{Duration, SystemTime};
4
5use reqwest::header::HeaderMap;
6
7#[derive(Clone, Debug, PartialEq)]
34pub struct RetryPolicy {
35 pub max_retries: u32,
36 pub backoff_initial: Duration,
37 pub backoff_max: Duration,
38 pub jitter: f64,
39 pub statuses: Vec<u16>,
40 pub respect_retry_after: bool,
41 pub retry_connection_errors: bool,
42 pub budget: Option<Duration>,
43}
44
45impl Default for RetryPolicy {
46 fn default() -> Self {
47 RetryPolicy {
48 max_retries: 2,
49 backoff_initial: Duration::from_millis(500),
50 backoff_max: Duration::from_secs(5),
51 jitter: 0.25,
52 statuses: [408, 429].into_iter().chain(500..=599).collect(),
53 respect_retry_after: true,
54 retry_connection_errors: true,
55 budget: Some(Duration::from_secs(30)),
56 }
57 }
58}
59
60impl RetryPolicy {
61 pub fn disabled() -> Self {
63 RetryPolicy {
64 max_retries: 0,
65 ..RetryPolicy::default()
66 }
67 }
68
69 pub(crate) fn next_delay(
73 &self,
74 attempts: u32,
75 failure: Failure,
76 remaining: Option<Duration>,
77 ) -> Option<Duration> {
78 if attempts > self.max_retries {
79 return None;
80 }
81 let delay = match failure {
82 Failure::Status { status, .. } if !self.statuses.contains(&status) => return None,
83 Failure::Transport if !self.retry_connection_errors => return None,
84 Failure::Status {
85 retry_after: Some(wait),
86 ..
87 } if self.respect_retry_after => wait,
88 _ => self.backoff(attempts),
89 };
90 match remaining {
91 Some(left) if delay >= left => None,
92 _ => Some(delay),
93 }
94 }
95
96 fn backoff(&self, attempts: u32) -> Duration {
97 let doubling = 2u32.saturating_pow(attempts.saturating_sub(1));
98 let exponential = self
99 .backoff_initial
100 .saturating_mul(doubling)
101 .min(self.backoff_max);
102 exponential.mul_f64(1.0 - fastrand::f64() * self.jitter.clamp(0.0, 1.0))
103 }
104}
105
106#[derive(Clone, Copy, Debug)]
108pub(crate) enum Failure {
109 Status {
110 status: u16,
111 retry_after: Option<Duration>,
112 },
113 Transport,
114}
115
116pub(crate) fn retry_after(headers: &HeaderMap) -> Option<Duration> {
118 let header = |name: &str| headers.get(name)?.to_str().ok().map(str::trim);
119 let seconds = |raw: &str, scale: f64| {
120 raw.parse::<f64>()
121 .ok()
122 .filter(|v| v.is_finite() && *v >= 0.0)
123 .map(|v| Duration::from_secs_f64(v * scale))
124 };
125 if let Some(wait) = header("retry-after-ms").and_then(|raw| seconds(raw, 0.001)) {
126 return Some(wait);
127 }
128 let raw = header("retry-after")?;
129 seconds(raw, 1.0).or_else(|| {
130 let at = httpdate::parse_http_date(raw).ok()?;
131 Some(
132 at.duration_since(SystemTime::now())
133 .unwrap_or(Duration::ZERO),
134 )
135 })
136}
137
138#[cfg(test)]
139mod tests {
140 use super::*;
141 use reqwest::header::HeaderValue;
142
143 fn status(status: u16) -> Failure {
144 Failure::Status {
145 status,
146 retry_after: None,
147 }
148 }
149
150 fn no_jitter() -> RetryPolicy {
151 RetryPolicy {
152 jitter: 0.0,
153 ..RetryPolicy::default()
154 }
155 }
156
157 #[test]
158 fn backs_off_exponentially_to_the_maximum() {
159 let policy = RetryPolicy {
160 max_retries: 10,
161 budget: None,
162 ..no_jitter()
163 };
164 let delays: Vec<u128> = (1..=6)
165 .map(|n| policy.next_delay(n, status(529), None).unwrap().as_millis())
166 .collect();
167 assert_eq!(delays, [500, 1000, 2000, 4000, 5000, 5000]);
168 }
169
170 #[test]
171 fn jitter_only_takes_time_off() {
172 let policy = RetryPolicy::default();
173 for _ in 0..100 {
174 let delay = policy.next_delay(1, status(429), None).unwrap();
175 assert!(delay <= Duration::from_millis(500));
176 assert!(delay >= Duration::from_millis(375));
177 }
178 }
179
180 #[test]
181 fn stops_after_max_retries() {
182 let policy = no_jitter();
183 assert!(policy.next_delay(1, status(500), None).is_some());
184 assert!(policy.next_delay(2, status(500), None).is_some());
185 assert!(policy.next_delay(3, status(500), None).is_none());
186 assert!(
187 RetryPolicy::disabled()
188 .next_delay(1, status(500), None)
189 .is_none()
190 );
191 }
192
193 #[test]
194 fn retries_only_listed_statuses() {
195 let policy = no_jitter();
196 for retried in [408, 429, 500, 503, 529] {
197 assert!(
198 policy.next_delay(1, status(retried), None).is_some(),
199 "{retried}"
200 );
201 }
202 for not in [400, 401, 403, 404, 422] {
203 assert!(policy.next_delay(1, status(not), None).is_none(), "{not}");
204 }
205 }
206
207 #[test]
208 fn connection_errors_follow_the_flag() {
209 assert!(
210 no_jitter()
211 .next_delay(1, Failure::Transport, None)
212 .is_some()
213 );
214 let off = RetryPolicy {
215 retry_connection_errors: false,
216 ..no_jitter()
217 };
218 assert!(off.next_delay(1, Failure::Transport, None).is_none());
219 }
220
221 #[test]
222 fn honours_the_server_and_the_budget() {
223 let asked = Failure::Status {
224 status: 429,
225 retry_after: Some(Duration::from_secs(3)),
226 };
227 let policy = no_jitter();
228 assert_eq!(
229 policy.next_delay(1, asked, None),
230 Some(Duration::from_secs(3))
231 );
232 assert_eq!(
233 policy.next_delay(1, asked, Some(Duration::from_secs(10))),
234 Some(Duration::from_secs(3))
235 );
236 assert_eq!(
237 policy.next_delay(1, asked, Some(Duration::from_secs(3))),
238 None
239 );
240
241 let ignoring = RetryPolicy {
242 respect_retry_after: false,
243 ..no_jitter()
244 };
245 assert_eq!(
246 ignoring.next_delay(1, asked, None),
247 Some(Duration::from_millis(500))
248 );
249 }
250
251 #[test]
252 fn parses_retry_after_headers() {
253 let parse = |pairs: &[(&'static str, &str)]| {
254 let mut headers = HeaderMap::new();
255 for (name, value) in pairs {
256 headers.insert(*name, HeaderValue::from_str(value).unwrap());
257 }
258 retry_after(&headers)
259 };
260 assert_eq!(
261 parse(&[("retry-after-ms", "250")]),
262 Some(Duration::from_millis(250))
263 );
264 assert_eq!(parse(&[("retry-after", "2")]), Some(Duration::from_secs(2)));
265 assert_eq!(
266 parse(&[("retry-after", "0.5")]),
267 Some(Duration::from_millis(500))
268 );
269 assert_eq!(
270 parse(&[("retry-after-ms", "100"), ("retry-after", "9")]),
271 Some(Duration::from_millis(100))
272 );
273 assert_eq!(
274 parse(&[("retry-after-ms", "-1"), ("retry-after", "1")]),
275 Some(Duration::from_secs(1))
276 );
277 assert_eq!(parse(&[("retry-after", "soon")]), None);
278 assert_eq!(parse(&[]), None);
279
280 let past = httpdate::fmt_http_date(SystemTime::now() - Duration::from_secs(60));
281 assert_eq!(parse(&[("retry-after", &past)]), Some(Duration::ZERO));
282 let future = httpdate::fmt_http_date(SystemTime::now() + Duration::from_secs(120));
283 let wait = parse(&[("retry-after", &future)]).unwrap();
284 assert!(wait > Duration::from_secs(100) && wait <= Duration::from_secs(120));
285 }
286}