Skip to main content

cognite/
retry.rs

1// This file is adapted from reqwest-retry, which was a bit too opinionated for our use.
2// https://github.com/TrueLayer/reqwest-middleware
3
4use 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
11/// Middleware for retrying requests.
12pub 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    /// Create a new retry middleware instance.
33    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            // Check if the error can be retried.
58            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                    // If the response failed and the error type was transient
66                    // we can safely try to retry the request.
67                    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                    // Jitter so we land between initial * 2 ** attempt * 3/4 and initial * 2 ** attempt * 5/4
72                    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    /// The failure was due to something that might resolve in the future.
86    Transient,
87    /// Unresolvable error.
88    Fatal,
89    /// Unauthorized. This is _maybe_ resolvable, if the last request wasn't also a 401.
90    Unauthorized,
91}
92
93impl Retryable {
94    /// Try to map a `reqwest` response into `Retryable`.
95    ///
96    /// Returns `None` if the response object does not contain any errors.
97    ///
98    /// # Arguments
99    ///
100    /// * `res` - Request response.
101    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                        // We omit checking if error.is_status() since we check that already.
144                        // However, if Response::error_for_status is used the status will still
145                        // remain in the response object.
146                        None
147                    }
148                }
149            },
150        }
151    }
152}