1use std::collections::BTreeSet;
10use std::time::{Duration, SystemTime};
11
12use crate::errors::{ApiError, Error};
13
14pub(crate) const RETRY_AFTER_MS_HEADER: &str = "retry-after-ms";
16pub(crate) const RETRY_AFTER_HEADER: &str = "retry-after";
18
19#[derive(Debug, Clone, PartialEq)]
39pub struct RetryPolicy {
40 pub max_retries: u32,
42 pub backoff_initial: Duration,
45 pub backoff_max: Duration,
47 pub backoff_jitter: f64,
49 pub http_statuses: BTreeSet<u16>,
51 pub respect_retry_after: bool,
53 pub max_retry_after: Duration,
55 pub retry_connection_errors: bool,
57 pub retry_timeouts: bool,
59 pub total_budget: Option<Duration>,
70}
71
72impl Default for RetryPolicy {
73 fn default() -> Self {
74 let mut http_statuses = BTreeSet::new();
75 http_statuses.insert(408);
76 http_statuses.insert(429);
77 http_statuses.extend(500..=599);
78 Self {
79 max_retries: 2,
80 backoff_initial: Duration::from_millis(500),
81 backoff_max: Duration::from_secs(5),
82 backoff_jitter: 0.25,
83 http_statuses,
84 respect_retry_after: true,
85 max_retry_after: Duration::from_secs(60),
86 retry_connection_errors: true,
87 retry_timeouts: true,
88 total_budget: Some(Duration::from_secs(30)),
89 }
90 }
91}
92
93impl RetryPolicy {
94 pub fn none() -> Self {
96 Self {
97 max_retries: 0,
98 ..Default::default()
99 }
100 }
101
102 pub(super) fn validate(&self) -> Result<(), Error> {
104 if !(0.0..=1.0).contains(&self.backoff_jitter) || !self.backoff_jitter.is_finite() {
105 return Err(Error::Config(
106 "retry policy backoff_jitter must be between 0 and 1.".into(),
107 ));
108 }
109 Ok(())
110 }
111
112 pub(super) fn is_retryable(&self, error: &Error) -> bool {
114 match error {
115 Error::Timeout { .. } => self.retry_timeouts,
116 Error::Connection { .. } => self.retry_connection_errors,
117 Error::Api(api) => self.is_retryable_status(api.status),
118 _ => false,
119 }
120 }
121
122 pub fn is_retryable_status(&self, status: u16) -> bool {
124 self.http_statuses.contains(&status)
125 }
126
127 pub(super) fn parse_retry_after(
134 &self,
135 headers: &reqwest::header::HeaderMap,
136 ) -> Option<Duration> {
137 if let Some(raw) = header_str(headers, RETRY_AFTER_MS_HEADER) {
138 if let Some(delay) = parse_ms(raw) {
139 return Some(delay);
140 }
141 }
142 if let Some(raw) = header_str(headers, RETRY_AFTER_HEADER) {
143 if let Ok(seconds) = raw.trim().parse::<f64>() {
144 if seconds.is_finite() && seconds >= 0.0 {
145 return seconds_to_duration(seconds);
147 }
148 return None;
150 }
151 if let Ok(date) = httpdate::parse_http_date(raw.trim()) {
153 if let Ok(now) = SystemTime::now().duration_since(SystemTime::UNIX_EPOCH) {
154 let date = date
155 .duration_since(SystemTime::UNIX_EPOCH)
156 .unwrap_or_default();
157 let delay = date.saturating_sub(now);
158 return Some(delay);
159 }
160 }
161 }
162 None
163 }
164
165 pub(super) fn delay_for_retry(
172 &self,
173 retry_number: u32,
174 api_error: Option<&ApiError>,
175 ) -> Duration {
176 if self.respect_retry_after {
177 if let Some(api_error) = api_error {
178 if let Some(delay) = api_error.retry_after {
179 if delay <= self.max_retry_after {
180 return delay;
181 }
182 }
183 }
184 }
185 self.backoff_delay(retry_number, fastrand::f64())
186 }
187
188 pub fn backoff_delay(&self, retry_number: u32, random: f64) -> Duration {
191 if self.backoff_initial.is_zero() || self.backoff_max.is_zero() {
192 return Duration::ZERO;
193 }
194 let exponential = self
195 .backoff_initial
196 .saturating_mul(1u32 << retry_number.min(31))
197 .min(self.backoff_max);
198 let jittered = 1.0 - random * self.backoff_jitter;
199 let scaled = exponential.as_secs_f64() * jittered;
200 Duration::try_from_secs_f64(scaled.max(0.0)).unwrap_or(Duration::MAX)
203 }
204}
205
206fn header_str<'a>(headers: &'a reqwest::header::HeaderMap, name: &str) -> Option<&'a str> {
207 headers.get(name)?.to_str().ok()
208}
209
210fn parse_ms(raw: &str) -> Option<Duration> {
211 let value: f64 = raw.trim().parse().ok()?;
212 if value.is_finite() && value >= 0.0 {
213 Duration::try_from_secs_f64(value / 1000.0).ok()
216 } else {
217 None
218 }
219}
220
221fn seconds_to_duration(seconds: f64) -> Option<Duration> {
222 Duration::try_from_secs_f64(seconds.max(0.0)).ok()
223}
224
225#[cfg(test)]
226mod tests {
227 use super::*;
228 use reqwest::header::{HeaderMap, HeaderName, HeaderValue};
229
230 fn headers(pairs: &[(&str, &str)]) -> HeaderMap {
231 let mut map = HeaderMap::new();
232 for (name, value) in pairs {
233 map.insert(
234 HeaderName::from_bytes(name.as_bytes()).unwrap(),
235 HeaderValue::from_str(value).unwrap(),
236 );
237 }
238 map
239 }
240
241 fn policy() -> RetryPolicy {
242 RetryPolicy {
243 backoff_jitter: 0.0,
244 ..Default::default()
245 }
246 }
247
248 #[test]
249 fn defaults_match_spec() {
250 let p = RetryPolicy::default();
251 assert_eq!(p.max_retries, 2);
252 assert_eq!(p.backoff_initial, Duration::from_millis(500));
253 assert_eq!(p.backoff_max, Duration::from_secs(5));
254 assert_eq!(p.backoff_jitter, 0.25);
255 assert_eq!(
256 p.http_statuses,
257 [408u16, 429]
258 .into_iter()
259 .chain(500u16..=599)
260 .collect::<BTreeSet<u16>>()
261 );
262 for status in 500..=599 {
263 assert!(p.http_statuses.contains(&status));
264 }
265 assert!(!p.http_statuses.contains(&400));
266 assert!(!p.http_statuses.contains(&600));
267 assert!(p.respect_retry_after);
268 assert_eq!(p.max_retry_after, Duration::from_secs(60));
269 assert!(p.retry_connection_errors);
270 assert!(p.retry_timeouts);
271 assert_eq!(p.total_budget, Some(Duration::from_secs(30)));
272 }
273
274 #[test]
275 fn none_policy_disables_retries() {
276 assert_eq!(RetryPolicy::none().max_retries, 0);
277 }
278
279 #[test]
280 fn backoff_formula_jitter_zero() {
281 let p = policy();
282 assert_eq!(p.backoff_delay(0, 0.0), Duration::from_millis(500));
283 assert_eq!(p.backoff_delay(1, 0.0), Duration::from_millis(1000));
284 assert_eq!(p.backoff_delay(2, 0.0), Duration::from_millis(2000));
285 assert_eq!(p.backoff_delay(3, 0.0), Duration::from_millis(4000));
286 assert_eq!(p.backoff_delay(4, 0.0), Duration::from_millis(5000)); assert_eq!(p.backoff_delay(9, 0.0), Duration::from_millis(5000));
288 }
289
290 #[test]
291 fn backoff_formula_jitter_bounds() {
292 let p = RetryPolicy::default(); let ms = p.backoff_delay(0, 0.999999).as_secs_f64() * 1000.0;
295 assert!((374.0..=376.0).contains(&ms), "{ms}");
296 assert_eq!(p.backoff_delay(0, 0.0), Duration::from_millis(500));
298 let ms = p.backoff_delay(0, 0.5).as_secs_f64() * 1000.0;
300 assert!((436.0..=439.0).contains(&ms), "{ms}");
301 for random in [0.0, 0.25, 0.5, 0.75, 0.9999] {
303 let delay = p.backoff_delay(1, random);
304 assert!(delay <= Duration::from_millis(1000), "{delay:?}");
305 assert!(delay >= Duration::from_millis(749), "{delay:?}");
306 }
307 }
308
309 #[test]
310 fn backoff_zero_initial_or_max_is_zero() {
311 let p = RetryPolicy {
312 backoff_initial: Duration::ZERO,
313 ..policy()
314 };
315 assert_eq!(p.backoff_delay(0, 0.0), Duration::ZERO);
316 let p = RetryPolicy {
317 backoff_max: Duration::ZERO,
318 ..policy()
319 };
320 assert_eq!(p.backoff_delay(3, 0.0), Duration::ZERO);
321 }
322
323 #[test]
324 fn parse_retry_after_ms() {
325 let p = policy();
326 assert_eq!(
327 p.parse_retry_after(&headers(&[(RETRY_AFTER_MS_HEADER, "250")])),
328 Some(Duration::from_millis(250))
329 );
330 assert_eq!(
331 p.parse_retry_after(&headers(&[(RETRY_AFTER_MS_HEADER, "0")])),
332 Some(Duration::ZERO)
333 );
334 assert_eq!(
335 p.parse_retry_after(&headers(&[(RETRY_AFTER_MS_HEADER, "10.5")])),
336 Some(Duration::from_nanos(10_500_000))
337 );
338 assert_eq!(
339 p.parse_retry_after(&headers(&[(RETRY_AFTER_MS_HEADER, "nope")])),
340 None
341 );
342 assert_eq!(
343 p.parse_retry_after(&headers(&[(RETRY_AFTER_MS_HEADER, "-5")])),
344 None
345 );
346 }
347
348 #[test]
349 fn parse_retry_after_seconds() {
350 let p = policy();
351 assert_eq!(
352 p.parse_retry_after(&headers(&[(RETRY_AFTER_HEADER, "3")])),
353 Some(Duration::from_secs(3))
354 );
355 assert_eq!(
356 p.parse_retry_after(&headers(&[(RETRY_AFTER_HEADER, "0")])),
357 Some(Duration::ZERO)
358 );
359 assert_eq!(
360 p.parse_retry_after(&headers(&[(RETRY_AFTER_HEADER, "1.5")])),
361 Some(Duration::from_millis(1500))
362 );
363 assert_eq!(
364 p.parse_retry_after(&headers(&[(RETRY_AFTER_HEADER, "-5")])),
365 None
366 );
367 assert_eq!(
368 p.parse_retry_after(&headers(&[(RETRY_AFTER_HEADER, "soon")])),
369 None
370 );
371 }
372
373 #[test]
374 fn parse_retry_after_prefers_ms() {
375 let p = policy();
376 let h = headers(&[(RETRY_AFTER_MS_HEADER, "250"), (RETRY_AFTER_HEADER, "3")]);
377 assert_eq!(p.parse_retry_after(&h), Some(Duration::from_millis(250)));
378 }
379
380 #[test]
381 fn parse_retry_after_http_date() {
382 let p = policy();
383 let date = SystemTime::now() + Duration::from_secs(5);
384 let date_str = httpdate::fmt_http_date(date);
385 let parsed = p
386 .parse_retry_after(&headers(&[(RETRY_AFTER_HEADER, date_str.as_str())]))
387 .unwrap();
388 assert!(parsed >= Duration::from_secs(4), "{parsed:?}");
389 assert!(parsed <= Duration::from_secs(6), "{parsed:?}");
390
391 let past = httpdate::fmt_http_date(SystemTime::now() - Duration::from_secs(60));
393 assert_eq!(
394 p.parse_retry_after(&headers(&[(RETRY_AFTER_HEADER, past.as_str())])),
395 Some(Duration::ZERO)
396 );
397 }
398
399 #[test]
400 fn delay_uses_retry_after_when_within_cap() {
401 let p = policy();
402 let error = retry_error(
403 429,
404 &[(RETRY_AFTER_MS_HEADER, "1500")],
405 Duration::from_millis(1500),
406 );
407 assert_eq!(
408 p.delay_for_retry(0, Some(&error)),
409 Duration::from_millis(1500)
410 );
411 }
412
413 #[test]
414 fn delay_falls_back_to_backoff_above_cap() {
415 let p = policy();
416 let error = retry_error(429, &[(RETRY_AFTER_HEADER, "61")], Duration::from_secs(61));
417 assert_eq!(
418 p.delay_for_retry(0, Some(&error)),
419 Duration::from_millis(500)
420 );
421 let error = retry_error(429, &[(RETRY_AFTER_HEADER, "60")], Duration::from_secs(60));
423 assert_eq!(p.delay_for_retry(0, Some(&error)), Duration::from_secs(60));
424 }
425
426 #[test]
427 fn delay_ignores_retry_after_when_disabled() {
428 let p = RetryPolicy {
429 respect_retry_after: false,
430 ..policy()
431 };
432 let error = retry_error(
433 429,
434 &[(RETRY_AFTER_MS_HEADER, "1500")],
435 Duration::from_millis(1500),
436 );
437 assert_eq!(
438 p.delay_for_retry(0, Some(&error)),
439 Duration::from_millis(500)
440 );
441 }
442
443 fn retry_error(status: u16, headers: &[(&str, &str)], retry_after: Duration) -> ApiError {
444 ApiError {
445 status,
446 kind: crate::errors::ApiErrorKind::from_status(status),
447 message: "msg".into(),
448 body: None,
449 headers: headers
450 .iter()
451 .map(|(name, value)| {
452 (
453 HeaderName::from_bytes(name.as_bytes()).unwrap(),
454 HeaderValue::from_str(value).unwrap(),
455 )
456 })
457 .collect(),
458 request_id: None,
459 endpoint: "POST https://api.typesafe.ai/v1/systemone".into(),
460 retry_after: Some(retry_after),
461 }
462 }
463}