1use ntex_service::{Ctx, Middleware, Service};
2
3pub 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;
7
8 fn clone_request(&self, req: &Req) -> Option<Req>;
12}
13
14#[derive(Clone, Debug)]
15pub struct Retry<P> {
19 policy: P,
20}
21
22#[derive(Clone, Debug)]
23pub struct RetryService<P, S> {
25 policy: P,
26 service: S,
27}
28
29impl<P> Retry<P> {
30 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 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 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)]
83pub struct DefaultRetryPolicy(u16);
87
88impl DefaultRetryPolicy {
89 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 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}