Skip to main content

ntex_util/services/
retry.rs

1use ntex_service::{Ctx, Middleware, Service};
2
3/// Determines whether and how a failed service call is retried.
4pub trait Policy<S: Service<St, Req>, St, Req>: Sized + Clone {
5    /// Returns whether the call should be retried.
6    async fn retry(&mut self, req: &Req, res: &Result<S::Res, S::Error>) -> bool;
7
8    /// Clones or reconstructs a request for a possible retry.
9    ///
10    /// Returning `None` prevents retries for this request.
11    fn clone_request(&self, req: &Req) -> Option<Req>;
12}
13
14#[derive(Clone, Debug)]
15/// Middleware that retries service calls according to a [`Policy`].
16///
17/// The policy is cloned for each request.
18pub struct Retry<P> {
19    policy: P,
20}
21
22#[derive(Clone, Debug)]
23/// A service that retries calls according to a [`Policy`].
24pub struct RetryService<P, S> {
25    policy: P,
26    service: S,
27}
28
29impl<P> Retry<P> {
30    /// Creates retry middleware with the specified policy.
31    pub fn new(policy: P) -> Self {
32        Retry { policy }
33    }
34}
35
36impl<P: Clone, S, St> Middleware<S, St> for Retry<P> {
37    type Service = RetryService<P, S>;
38
39    fn create(&self, _: &St, service: S) -> Self::Service {
40        RetryService {
41            service,
42            policy: self.policy.clone(),
43        }
44    }
45}
46
47impl<P, S> RetryService<P, S> {
48    /// Wraps a service with the specified retry policy.
49    pub fn new(policy: P, service: S) -> Self {
50        RetryService { policy, service }
51    }
52}
53
54impl<P, S, St, Req> Service<St, Req> for RetryService<P, S>
55where
56    P: Policy<S, St, Req>,
57    S: Service<St, Req>,
58{
59    type Res = S::Res;
60    type Error = S::Error;
61
62    async fn call(&self, req: Req, ctx: Ctx<'_, Self, St>) -> Result<S::Res, S::Error> {
63        let mut policy = self.policy.clone();
64        let mut cloned = policy.clone_request(&req);
65        // the first call is covered by the outer readiness check
66        let mut result = ctx.call_nowait(&self.service, req).await;
67
68        while let Some(r) = cloned.take() {
69            if !policy.retry(&r, &result).await {
70                break;
71            }
72            cloned = policy.clone_request(&r);
73            result = ctx.call(&self.service, r).await;
74        }
75        result
76    }
77
78    ntex_service::forward_ready!(St, service);
79    ntex_service::forward_shutdown!(St, service);
80}
81
82#[derive(Copy, Clone, Debug)]
83/// A retry policy that retries every service error.
84///
85/// The default policy permits up to three retries after the initial call.
86pub struct DefaultRetryPolicy(u16);
87
88impl DefaultRetryPolicy {
89    /// Creates a policy that permits up to `retry` retries.
90    pub fn new(retry: u16) -> Self {
91        DefaultRetryPolicy(retry)
92    }
93}
94
95impl Default for DefaultRetryPolicy {
96    fn default() -> Self {
97        DefaultRetryPolicy::new(3)
98    }
99}
100
101impl<S, St, Req> Policy<S, St, Req> for DefaultRetryPolicy
102where
103    S: Service<St, Req>,
104    Req: Clone,
105{
106    async fn retry(&mut self, _: &Req, res: &Result<S::Res, S::Error>) -> bool {
107        if res.is_err() {
108            if self.0 == 0 {
109                false
110            } else {
111                self.0 -= 1;
112                true
113            }
114        } else {
115            false
116        }
117    }
118
119    fn clone_request(&self, req: &Req) -> Option<Req> {
120        Some(req.clone())
121    }
122}
123
124#[cfg(test)]
125mod tests {
126    #![allow(clippy::unused_async_trait_impl)]
127    use std::{cell::Cell, rc::Rc};
128
129    use ntex_service::{Pipeline, apply, fn_factory};
130
131    use super::*;
132
133    #[derive(Clone, Debug, PartialEq)]
134    struct TestService(Rc<Cell<usize>>);
135
136    impl Service<(), ()> for TestService {
137        type Res = ();
138        type Error = ();
139
140        async fn call(&self, _r: (), _: Ctx<'_, Self>) -> Result<(), ()> {
141            let cnt = self.0.get();
142            if cnt == 0 {
143                Ok(())
144            } else {
145                self.0.set(cnt - 1);
146                Err(())
147            }
148        }
149    }
150
151    #[ntex::test]
152    async fn test_retry() {
153        let cnt = Rc::new(Cell::new(5));
154        let svc = Pipeline::new(
155            (),
156            RetryService::new(DefaultRetryPolicy::default(), TestService(cnt.clone())).clone(),
157        );
158        assert_eq!(svc.call(()).await, Err(()));
159        assert_eq!(svc.ready().await, Ok(()));
160        svc.shutdown().await;
161        assert_eq!(cnt.get(), 1);
162
163        let factory = apply(
164            Retry::new(DefaultRetryPolicy::new(3)).clone(),
165            fn_factory(|(): &()| async { Ok::<_, ()>(TestService(Rc::new(Cell::new(2)))) }),
166        );
167        let srv = factory.pipeline(()).await.unwrap();
168        assert_eq!(srv.call(()).await, Ok(()));
169
170        let factory = apply(
171            Retry::new(DefaultRetryPolicy::new(3)).clone(),
172            fn_factory(|(): &()| async { Ok::<_, ()>(TestService(Rc::new(Cell::new(2)))) }),
173        );
174        let srv = factory.pipeline(()).await.unwrap();
175        assert_eq!(srv.call(()).await, Ok(()));
176    }
177
178    #[derive(Clone)]
179    struct CloneOnce(Rc<Cell<usize>>);
180
181    impl Policy<TestService, (), ()> for CloneOnce {
182        async fn retry(&mut self, (): &(), res: &Result<(), ()>) -> bool {
183            res.is_err()
184        }
185
186        fn clone_request(&self, (): &()) -> Option<()> {
187            let n = self.0.get();
188            self.0.set(n + 1);
189            if n == 0 { Some(()) } else { None }
190        }
191    }
192
193    #[ntex::test]
194    async fn test_retry_without_clone() {
195        // the retried request cannot be cloned again, its result is returned
196        let cnt = Rc::new(Cell::new(5));
197        let clones = Rc::new(Cell::new(0));
198        let svc = Pipeline::new(
199            (),
200            RetryService::new(CloneOnce(clones.clone()), TestService(cnt.clone())),
201        );
202        assert_eq!(svc.call(()).await, Err(()));
203        assert_eq!(cnt.get(), 3);
204        assert_eq!(clones.get(), 2);
205    }
206}