Skip to main content

ntex_util/services/
retry.rs

1use ntex_service::{Ctx, Middleware, Service};
2
3/// Trait defines retry policy
4pub trait Policy<S: Service<St, Req>, St, Req>: Sized + Clone {
5    async fn retry(&mut self, req: &Req, res: &Result<S::Res, S::Error>) -> bool;
6
7    fn clone_request(&self, req: &Req) -> Option<Req>;
8}
9
10#[derive(Clone, Debug)]
11/// Retry middleware
12///
13/// Retry middleware allows to retry service call
14pub struct Retry<P> {
15    policy: P,
16}
17
18#[derive(Clone, Debug)]
19/// Retry service
20///
21/// Retry service allows to retry service call
22pub struct RetryService<P, S> {
23    policy: P,
24    service: S,
25}
26
27impl<P> Retry<P> {
28    /// Create retry middleware
29    pub fn new(policy: P) -> Self {
30        Retry { policy }
31    }
32}
33
34impl<P: Clone, S, St> Middleware<S, St> for Retry<P> {
35    type Service = RetryService<P, S>;
36
37    fn create(&self, _: &St, service: S) -> Self::Service {
38        RetryService {
39            service,
40            policy: self.policy.clone(),
41        }
42    }
43}
44
45impl<P, S> RetryService<P, S> {
46    /// Create retry service
47    pub fn new(policy: P, service: S) -> Self {
48        RetryService { policy, service }
49    }
50}
51
52impl<P, S, St, Req> Service<St, Req> for RetryService<P, S>
53where
54    P: Policy<S, St, Req>,
55    S: Service<St, Req>,
56{
57    type Res = S::Res;
58    type Error = S::Error;
59
60    async fn call(&self, mut req: Req, ctx: Ctx<'_, Self, St>) -> Result<S::Res, S::Error> {
61        let mut policy = self.policy.clone();
62        let mut cloned = policy.clone_request(&req);
63
64        loop {
65            let result = ctx.call(&self.service, req).await;
66
67            cloned = if let Some(r) = cloned.take() {
68                if policy.retry(&r, &result).await {
69                    req = r;
70                    policy.clone_request(&req)
71                } else {
72                    return result;
73                }
74            } else {
75                return result;
76            }
77        }
78    }
79
80    ntex_service::forward_ready!(St, service);
81    ntex_service::forward_shutdown!(St, service);
82}
83
84#[derive(Copy, Clone, Debug)]
85/// Default retry policy
86///
87/// This policy retries on any error. By default retry count is 3
88pub struct DefaultRetryPolicy(u16);
89
90impl DefaultRetryPolicy {
91    /// Create default retry policy
92    pub fn new(retry: u16) -> Self {
93        DefaultRetryPolicy(retry)
94    }
95}
96
97impl Default for DefaultRetryPolicy {
98    fn default() -> Self {
99        DefaultRetryPolicy::new(3)
100    }
101}
102
103impl<S, St, Req> Policy<S, St, Req> for DefaultRetryPolicy
104where
105    S: Service<St, Req>,
106    Req: Clone,
107{
108    async fn retry(&mut self, _: &Req, res: &Result<S::Res, S::Error>) -> bool {
109        if res.is_err() {
110            if self.0 == 0 {
111                false
112            } else {
113                self.0 -= 1;
114                true
115            }
116        } else {
117            false
118        }
119    }
120
121    fn clone_request(&self, req: &Req) -> Option<Req> {
122        Some(req.clone())
123    }
124}
125
126#[cfg(test)]
127mod tests {
128    #![allow(clippy::unused_async_trait_impl)]
129    use std::{cell::Cell, rc::Rc};
130
131    use ntex_service::{Pipeline, apply, fn_factory};
132
133    use super::*;
134
135    #[derive(Clone, Debug, PartialEq)]
136    struct TestService(Rc<Cell<usize>>);
137
138    impl Service<(), ()> for TestService {
139        type Res = ();
140        type Error = ();
141
142        async fn call(&self, _r: (), _: Ctx<'_, Self>) -> Result<(), ()> {
143            let cnt = self.0.get();
144            if cnt == 0 {
145                Ok(())
146            } else {
147                self.0.set(cnt - 1);
148                Err(())
149            }
150        }
151    }
152
153    #[ntex::test]
154    async fn test_retry() {
155        let cnt = Rc::new(Cell::new(5));
156        let svc = Pipeline::new(
157            (),
158            RetryService::new(DefaultRetryPolicy::default(), TestService(cnt.clone())).clone(),
159        );
160        assert_eq!(svc.call(()).await, Err(()));
161        assert_eq!(svc.ready().await, Ok(()));
162        svc.shutdown().await;
163        assert_eq!(cnt.get(), 1);
164
165        let factory = apply(
166            Retry::new(DefaultRetryPolicy::new(3)).clone(),
167            fn_factory(|(): &()| async { Ok::<_, ()>(TestService(Rc::new(Cell::new(2)))) }),
168        );
169        let srv = factory.pipeline(()).await.unwrap();
170        assert_eq!(srv.call(()).await, Ok(()));
171
172        let factory = apply(
173            Retry::new(DefaultRetryPolicy::new(3)).clone(),
174            fn_factory(|(): &()| async { Ok::<_, ()>(TestService(Rc::new(Cell::new(2)))) }),
175        );
176        let srv = factory.pipeline(()).await.unwrap();
177        assert_eq!(srv.call(()).await, Ok(()));
178    }
179}