Skip to main content

typesafe_system_one/
retry.rs

1//! Retry policy: defaults, delay calculation, retry-after parsing, and stop
2//! conditions.
3//!
4//! The default matches both official SDKs: 2 retries, 500ms initial backoff
5//! doubling to 5s with 25% jitter, retries on 408/429/5xx, honored
6//! `retry-after-ms`/`Retry-After` capped at 60s, connection and timeout
7//! retries, and a 30s total budget.
8
9use std::collections::BTreeSet;
10use std::time::{Duration, SystemTime};
11
12use crate::errors::{ApiError, Error};
13
14/// The `retry-after-ms` response header.
15pub(crate) const RETRY_AFTER_MS_HEADER: &str = "retry-after-ms";
16/// The `Retry-After` response header.
17pub(crate) const RETRY_AFTER_HEADER: &str = "retry-after";
18
19/// Configuration for retry behavior.
20///
21/// All fields are public. The [`Default`] value matches the SPEC and both
22/// official SDKs.
23///
24/// ```
25/// use std::time::Duration;
26/// use typesafe_system_one::{Client, RetryPolicy};
27///
28/// let client = Client::builder()
29///     .api_key("key")
30///     .retry_policy(RetryPolicy {
31///         max_retries: 3,
32///         total_budget: Some(Duration::from_secs(10)),
33///         ..Default::default()
34///     })
35///     .build()
36///     .unwrap();
37/// ```
38#[derive(Debug, Clone, PartialEq)]
39pub struct RetryPolicy {
40    /// Maximum retries after the initial attempt; `0` disables retries.
41    pub max_retries: u32,
42    /// First backoff delay, doubled each retry up to [`Self::backoff_max`];
43    /// zero disables backoff.
44    pub backoff_initial: Duration,
45    /// Maximum backoff delay; zero disables backoff.
46    pub backoff_max: Duration,
47    /// Fraction of each backoff delay randomly subtracted; between 0 and 1.
48    pub backoff_jitter: f64,
49    /// HTTP status codes that are retried.
50    pub http_statuses: BTreeSet<u16>,
51    /// Whether to honor `retry-after-ms` and `Retry-After` headers.
52    pub respect_retry_after: bool,
53    /// Maximum honored server retry delay; a longer delay falls back to backoff.
54    pub max_retry_after: Duration,
55    /// Whether to retry connection errors (no HTTP response at all).
56    pub retry_connection_errors: bool,
57    /// Whether to retry attempts that exceeded the per-attempt timeout.
58    pub retry_timeouts: bool,
59    /// Retry budget per call, measured from the first attempt; `None` or
60    /// `Some(Duration::ZERO)` disables it.
61    ///
62    /// A retry is not started when the elapsed time plus its delay would
63    /// reach or exceed the budget; the call then returns the **last real
64    /// error** rather than an artificial timeout, so the caller sees the
65    /// actual 529. An attempt that has started runs to its own per-attempt
66    /// timeout, so a call can overrun the budget by up to one timeout (the
67    /// same semantics as the official Python SDK). For a hard bound, wrap
68    /// the call in `tokio::time::timeout`.
69    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    /// A policy that never retries (`max_retries = 0`).
95    pub fn none() -> Self {
96        Self {
97            max_retries: 0,
98            ..Default::default()
99        }
100    }
101
102    /// Validates the policy, returning `Err(Error::Config)` on a bad setting.
103    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    /// Whether this error should be retried under this policy.
113    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    /// Whether an HTTP status is retried under this policy.
123    pub fn is_retryable_status(&self, status: u16) -> bool {
124        self.http_statuses.contains(&status)
125    }
126
127    /// Parses the retry delay from response headers, if present and valid.
128    ///
129    /// Prefers `retry-after-ms` (float milliseconds, ≥ 0). Otherwise
130    /// `Retry-After` as float seconds ≥ 0, or as an HTTP date
131    /// (delay = max(0, date − now)). Negative or unparseable values mean
132    /// "absent".
133    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                    // An unrepresentably large value is treated as absent.
146                    return seconds_to_duration(seconds);
147                }
148                // Negative or non-finite: absent.
149                return None;
150            }
151            // Try an HTTP date.
152            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    /// The delay for retry `retry_number` (0-based), following the last error.
166    ///
167    /// If a retry-after is respected (enabled, parseable, and ≤
168    /// [`Self::max_retry_after`]), it is used exactly. Otherwise capped
169    /// exponential backoff: `min(initial * 2^n, max) * (1 - rand[0,1) * jitter)`.
170    /// If initial or max is zero, the backoff delay is zero.
171    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    /// The pure backoff portion of the delay for retry `retry_number`
189    /// (0-based), given a jitter random value in `[0, 1)`.
190    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        // Saturate instead of panicking if a huge configured backoff_max
201        // rounds past Duration::MAX.
202        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        // `try_` because the header is server-controlled: an overflowing
214        // value must be "absent", not a panic.
215        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)); // capped
287        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(); // jitter 0.25
293                                        // random ~ 1: delay = base * 0.75 (allow float rounding).
294        let ms = p.backoff_delay(0, 0.999999).as_secs_f64() * 1000.0;
295        assert!((374.0..=376.0).contains(&ms), "{ms}");
296        // random = 0: delay = base.
297        assert_eq!(p.backoff_delay(0, 0.0), Duration::from_millis(500));
298        // Halfway: 500 * 0.875 = 437.5ms.
299        let ms = p.backoff_delay(0, 0.5).as_secs_f64() * 1000.0;
300        assert!((436.0..=439.0).contains(&ms), "{ms}");
301        // Never above the exponential base or below base * (1 - jitter).
302        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        // A date in the past clamps to zero.
392        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        // Exactly at the cap is honored.
422        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}