1use 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#[derive(Debug, Clone)]
29pub struct Retry<P, S> {
30 policy: P,
31 inner: S,
32}
33
34impl<P, S> Retry<P, S> {
37 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)]
49pub 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 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 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 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}