Skip to main content

rama_http/layer/retry/
mod.rs

1//! Middleware for retrying "failed" requests.
2
3use crate::{Request, StreamingBody, body::util::BodyExt};
4use rama_core::Service;
5use rama_core::error::BoxError;
6use rama_core::extensions::ExtensionsRef;
7use rama_utils::macros::define_inner_service_accessors;
8
9mod layer;
10mod policy;
11
12mod body;
13#[doc(inline)]
14pub use body::RetryBody;
15
16pub mod managed;
17pub use managed::ManagedPolicy;
18
19#[cfg(test)]
20mod tests;
21
22pub use self::layer::RetryLayer;
23pub use self::policy::{Policy, PolicyResult};
24
25/// Configure retrying requests of "failed" responses.
26///
27/// A [`Policy`] classifies what is a "failed" response.
28#[derive(Debug, Clone)]
29pub struct Retry<P, S> {
30    policy: P,
31    inner: S,
32}
33
34// ===== impl Retry =====
35
36impl<P, S> Retry<P, S> {
37    /// Retry the inner service depending on this [`Policy`].
38    pub const fn new(policy: P, service: S) -> Self {
39        Self {
40            policy,
41            inner: service,
42        }
43    }
44
45    define_inner_service_accessors!();
46}
47
48#[derive(Debug)]
49/// Error type for [`Retry`]
50pub struct RetryError {
51    kind: RetryErrorKind,
52    inner: Option<BoxError>,
53}
54
55#[derive(Debug)]
56enum RetryErrorKind {
57    BodyConsume,
58    Service,
59}
60
61impl std::fmt::Display for RetryError {
62    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
63        match &self.inner {
64            Some(inner) => write!(f, "{}: {}", self.kind, inner),
65            None => write!(f, "{}", self.kind),
66        }
67    }
68}
69
70impl std::fmt::Display for RetryErrorKind {
71    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
72        match self {
73            Self::BodyConsume => write!(f, "failed to consume body"),
74            Self::Service => write!(f, "service error"),
75        }
76    }
77}
78
79impl std::error::Error for RetryError {
80    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
81        self.inner.as_ref().and_then(|e| e.source())
82    }
83}
84
85impl<P, S, Body> Service<Request<Body>> for Retry<P, S>
86where
87    P: Policy<S::Output, S::Error>,
88    S: Service<Request<RetryBody>, Error: Into<BoxError>>,
89    Body: StreamingBody<Data: Send + 'static, Error: Into<BoxError>> + Send + 'static,
90{
91    type Output = S::Output;
92    type Error = RetryError;
93
94    async fn serve(&self, request: Request<Body>) -> Result<Self::Output, Self::Error> {
95        // consume body so we can clone the request if desired
96        let (parts, body) = request.into_parts();
97        let body = body.collect().await.map_err(|e| RetryError {
98            kind: RetryErrorKind::BodyConsume,
99            inner: Some(e.into()),
100        })?;
101        let body = RetryBody::new(body.to_bytes());
102        let mut request = Request::from_parts(parts, body);
103
104        let mut cloned = self.policy.clone_input(&request);
105
106        let parent_ext = request.extensions().clone();
107        loop {
108            // Fork extensions so we don't leak extensions from failed attempts
109            request.set_extensions(parent_ext.fork());
110
111            let resp = self.inner.serve(request).await;
112            match cloned.take() {
113                Some(cloned_req) => {
114                    let cloned_req = match self.policy.retry(cloned_req, resp).await {
115                        PolicyResult::Abort(result) => {
116                            return result.map_err(|e| RetryError {
117                                kind: RetryErrorKind::Service,
118                                inner: Some(e.into()),
119                            });
120                        }
121                        PolicyResult::Retry { req } => req,
122                    };
123
124                    cloned = self.policy.clone_input(&cloned_req);
125                    request = cloned_req;
126                }
127                // no clone was made, so no possibility to retry
128                None => {
129                    return resp.map_err(|e| RetryError {
130                        kind: RetryErrorKind::Service,
131                        inner: Some(e.into()),
132                    });
133                }
134            }
135        }
136    }
137}
138
139#[cfg(test)]
140mod test {
141    use super::*;
142    use crate::{
143        BodyExtractExt, Response, StatusCode, layer::retry::managed::DoNotRetry,
144        service::web::response::IntoResponse,
145    };
146    use rama_core::{
147        Layer,
148        error::BoxErrorExt,
149        extensions::{Extension, Extensions, ExtensionsRef},
150        service::service_fn,
151    };
152    use rama_utils::{backoff::ExponentialBackoff, rng::HasherRng};
153    use std::{sync::atomic::AtomicUsize, time::Duration};
154
155    #[tokio::test]
156    async fn test_service_with_managed_retry() {
157        let backoff = ExponentialBackoff::new(
158            Duration::from_millis(1),
159            Duration::from_millis(5),
160            0.1,
161            HasherRng::default,
162        )
163        .unwrap();
164
165        #[derive(Debug, Extension)]
166        struct State {
167            retry_counter: AtomicUsize,
168        }
169
170        async fn retry<Body, E>(
171            req: Request<Body>,
172            result: Result<Response, E>,
173        ) -> (Request<Body>, Result<Response, E>, bool) {
174            if req.extensions().contains::<DoNotRetry>() {
175                panic!("unexpected retry: should be disabled");
176            }
177
178            if let Ok(ref res) = result {
179                if res.status().is_server_error() {
180                    req.extensions()
181                        .get_ref::<State>()
182                        .unwrap()
183                        .retry_counter
184                        .fetch_add(1, std::sync::atomic::Ordering::AcqRel);
185                    (req, result, true)
186                } else {
187                    (req, result, false)
188                }
189            } else {
190                req.extensions()
191                    .get_ref::<State>()
192                    .unwrap()
193                    .retry_counter
194                    .fetch_add(1, std::sync::atomic::Ordering::SeqCst);
195                (req, result, true)
196            }
197        }
198
199        let retry_policy = ManagedPolicy::new(retry).with_backoff(backoff);
200
201        let service = RetryLayer::new(retry_policy).into_layer(service_fn(
202            async |req: Request<RetryBody>| {
203                let txt = req.try_into_string().await.unwrap();
204                match txt.as_str() {
205                    "internal" => Ok(StatusCode::INTERNAL_SERVER_ERROR.into_response()),
206                    "error" => Err(BoxError::from_static_str("custom error")),
207                    _ => Ok(txt.into_response()),
208                }
209            },
210        ));
211
212        fn request(s: &'static str) -> Request {
213            Request::builder().body(s.into()).unwrap()
214        }
215
216        fn extensions() -> Extensions {
217            let extensions = Extensions::new();
218            extensions.insert(State {
219                retry_counter: AtomicUsize::new(0),
220            });
221            extensions
222        }
223
224        fn do_not_retry_extensions() -> Extensions {
225            let extensions = extensions();
226            extensions.insert(DoNotRetry::default());
227            extensions
228        }
229
230        async fn assert_serve_ok<E: std::fmt::Debug>(
231            msg: &'static str,
232            input: &'static str,
233            output: &'static str,
234            extensions: Extensions,
235            retried: bool,
236            service: &impl Service<Request, Output = Response, Error = E>,
237        ) {
238            let state = extensions.get_arc::<State>().unwrap();
239
240            let request = request(input);
241            request.extensions().extend(&extensions);
242
243            let fut = service.serve(request);
244            let res = fut.await.unwrap();
245
246            let body = res.try_into_string().await.unwrap();
247            assert_eq!(body, output, "{msg}");
248            if retried {
249                assert!(
250                    state
251                        .retry_counter
252                        .load(std::sync::atomic::Ordering::Acquire)
253                        > 0,
254                    "{msg}"
255                );
256            } else {
257                assert_eq!(
258                    state
259                        .retry_counter
260                        .load(std::sync::atomic::Ordering::Acquire),
261                    0,
262                    "{msg}"
263                );
264            }
265        }
266
267        async fn assert_serve_err<E: std::fmt::Debug>(
268            msg: &'static str,
269            input: &'static str,
270            extensions: Extensions,
271            retried: bool,
272            service: &impl Service<Request, Output = Response, Error = E>,
273        ) {
274            let state = extensions.get_arc::<State>().unwrap();
275
276            let request = request(input);
277            request.extensions().extend(&extensions);
278
279            let fut = service.serve(request);
280            let res = fut.await;
281
282            assert!(res.is_err(), "{msg}");
283            if retried {
284                assert!(
285                    state
286                        .retry_counter
287                        .load(std::sync::atomic::Ordering::Acquire)
288                        > 0,
289                    "{msg}"
290                );
291            } else {
292                assert_eq!(
293                    state
294                        .retry_counter
295                        .load(std::sync::atomic::Ordering::Acquire),
296                    0,
297                    "{msg}"
298                )
299            }
300        }
301
302        assert_serve_ok(
303            "ok response should be aborted as response without retry",
304            "hello",
305            "hello",
306            extensions(),
307            false,
308            &service,
309        )
310        .await;
311        assert_serve_ok(
312            "internal will trigger 500 with a retry",
313            "internal",
314            "",
315            extensions(),
316            true,
317            &service,
318        )
319        .await;
320        assert_serve_err(
321            "error will trigger an actual non-http error with a retry",
322            "error",
323            extensions(),
324            true,
325            &service,
326        )
327        .await;
328
329        assert_serve_ok(
330            "normally internal will trigger a 500 with retry, but using DoNotRetry will disable retrying",
331            "internal",
332            "",
333            do_not_retry_extensions(),
334            false,
335            &service,
336        ).await;
337    }
338}