1use std::cell::Cell;
2
3use super::{Ctx, Service, ServiceFactory, util};
4
5#[derive(Debug)]
6pub struct AndThen<A, B> {
10 svc1: A,
11 svc2: B,
12 ready: Cell<bool>,
13 entered: Cell<u32>,
14}
15
16impl<A, B> AndThen<A, B> {
17 pub(crate) fn new(svc1: A, svc2: B) -> Self {
19 Self {
20 svc1,
21 svc2,
22 ready: Cell::new(false),
23 entered: Cell::new(0),
24 }
25 }
26}
27
28impl<A: Clone, B: Clone> Clone for AndThen<A, B> {
29 fn clone(&self) -> Self {
30 Self::new(self.svc1.clone(), self.svc2.clone())
31 }
32}
33
34impl<A, B, St, Req> Service<St, Req> for AndThen<A, B>
35where
36 A: Service<St, Req>,
37 B: Service<St, A::Res, Error = A::Error>,
38{
39 type Res = B::Res;
40 type Error = A::Error;
41
42 #[inline]
43 async fn call(&self, req: Req, ctx: Ctx<'_, Self, St>) -> Result<B::Res, A::Error> {
44 let result = ctx.call_nowait(&self.svc1, req).await?;
45
46 if !self.ready.get() {
47 ctx.ready(&self.svc2).await?;
48 }
49 self.ready.set(false);
51 self.entered.set(self.entered.get().wrapping_add(1));
52 ctx.call_nowait(&self.svc2, result).await
53 }
54
55 #[inline]
56 async fn ready(&self, ctx: Ctx<'_, Self, St>) -> Result<(), Self::Error> {
57 let entered = self.entered.get();
58 let res = util::ready(&self.svc1, &self.svc2, ctx).await;
59 if res.is_err() {
60 self.ready.set(false);
61 } else if self.entered.get() == entered {
62 self.ready.set(true);
64 }
65 res
66 }
67
68 #[inline]
69 async fn shutdown(&self, ctx: Ctx<'_, Self, St>) {
70 util::shutdown(&self.svc1, &self.svc2, ctx).await;
71 }
72}
73
74#[derive(Debug, Clone)]
75pub struct AndThenFactory<A, B> {
79 svc1: A,
80 svc2: B,
81}
82
83impl<A, B> AndThenFactory<A, B> {
84 pub fn new(svc1: A, svc2: B) -> Self {
86 Self { svc1, svc2 }
87 }
88}
89
90impl<A, B, St, Req> ServiceFactory<St, Req> for AndThenFactory<A, B>
91where
92 A: ServiceFactory<St, Req>,
93 B: ServiceFactory<St, A::Res, Error = A::Error, InitError = A::InitError>,
94{
95 type Res = B::Res;
96 type Error = A::Error;
97
98 type Service = AndThen<A::Service, B::Service>;
99 type InitError = A::InitError;
100
101 #[inline]
102 async fn create(&self, st: &St) -> Result<Self::Service, Self::InitError> {
103 Ok(AndThen {
104 svc1: self.svc1.create(st).await?,
105 svc2: self.svc2.create(st).await?,
106 ready: Cell::new(false),
107 entered: Cell::new(0),
108 })
109 }
110}
111
112#[cfg(test)]
113mod tests {
114 use std::{cell::Cell, rc::Rc};
115
116 use crate::{Ctx, Service, factory, fn_factory, service};
117
118 #[derive(Debug, Clone)]
119 struct Srv1(Rc<Cell<usize>>, Rc<Cell<usize>>);
120
121 impl Service<(), &'static str> for Srv1 {
122 type Res = &'static str;
123 type Error = ();
124
125 async fn ready(&self, _: Ctx<'_, Self>) -> Result<(), Self::Error> {
126 self.0.set(self.0.get() + 1);
127 Ok(())
128 }
129
130 async fn call(&self, req: &'static str, _: Ctx<'_, Self>) -> Result<Self::Res, ()> {
131 Ok(req)
132 }
133
134 async fn shutdown(&self, _: Ctx<'_, Self, ()>) {
135 self.1.set(self.1.get() + 1);
136 }
137 }
138
139 #[derive(Debug, Clone)]
140 struct Srv2(Rc<Cell<usize>>, Rc<Cell<usize>>);
141
142 impl Service<(), &'static str> for Srv2 {
143 type Res = (&'static str, &'static str);
144 type Error = ();
145
146 async fn ready(&self, _: Ctx<'_, Self>) -> Result<(), Self::Error> {
147 self.0.set(self.0.get() + 1);
148 Ok(())
149 }
150
151 async fn call(&self, req: &'static str, _: Ctx<'_, Self>) -> Result<Self::Res, ()> {
152 Ok((req, "srv2"))
153 }
154
155 async fn shutdown(&self, _: Ctx<'_, Self, ()>) {
156 self.1.set(self.1.get() + 1);
157 }
158 }
159
160 #[ntex::test]
161 async fn test_ready() {
162 let cnt = Rc::new(Cell::new(0));
163 let cnt_sht = Rc::new(Cell::new(0));
164 let srv = service(Rc::new(Srv1(cnt.clone(), cnt_sht.clone())))
165 .clone()
166 .and_then(crate::boxed::service(Srv2(cnt.clone(), cnt_sht.clone())));
167 assert!(format!("{srv:?}").contains("AndThen"));
168
169 let srv = srv.pipeline(());
170 let res = srv.ready().await;
171 assert_eq!(res, Ok(()));
172 assert_eq!(cnt.get(), 2);
173
174 srv.shutdown().await;
175 assert_eq!(cnt_sht.get(), 2);
176 }
177
178 #[ntex::test]
179 async fn test_ready2() {
180 let cnt = Rc::new(Cell::new(0));
181 let srv = Box::new(
182 service(Srv1(cnt.clone(), Rc::new(Cell::new(0))))
183 .and_then(Srv2(cnt.clone(), Rc::new(Cell::new(0)))),
184 )
185 .pipeline(());
186 let res = srv.ready().await;
187 assert_eq!(res, Ok(()));
188 assert_eq!(cnt.get(), 2);
189 }
190
191 #[ntex::test]
192 async fn test_call() {
193 let cnt = Rc::new(Cell::new(0));
194 let cnt_sht = Rc::new(Cell::new(0));
195 let srv = Srv1(cnt.clone(), cnt_sht.clone())
196 .and_then(Srv2(cnt, cnt_sht.clone()))
197 .pipeline(());
198 let res = srv.call("srv1").await;
199 assert!(res.is_ok());
200 assert_eq!(res.unwrap(), ("srv1", "srv2"));
201
202 srv.shutdown().await;
203 assert_eq!(cnt_sht.get(), 2);
204 }
205
206 #[ntex::test]
207 async fn test_factory() {
208 let cnt = Rc::new(Cell::new(0));
209 let cnt2 = cnt.clone();
210 let new_srv = factory(fn_factory(move |(): &()| {
211 let cnt = cnt2.clone();
212 async move { Ok::<_, ()>(Srv1(cnt, Rc::new(Cell::new(0)))) }
213 }))
214 .and_then(fn_factory(move |(): &()| {
215 let cnt = cnt.clone();
216 async move { Ok(Srv2(cnt.clone(), Rc::new(Cell::new(0)))) }
217 }))
218 .clone();
219
220 let srv = new_srv.pipeline(()).await.unwrap();
221 let res = srv.call("srv1").await;
222 assert!(res.is_ok());
223 assert_eq!(res.unwrap(), ("srv1", "srv2"));
224 }
225
226 mod ready_flag {
227 use std::{cell::Cell, cell::RefCell, future::poll_fn, rc::Rc, task::Poll, task::Waker};
228
229 use ntex::channel::oneshot::{self, Receiver};
230
231 use crate::util::tests::{Req, Single, State, concurrent_entry, start, tick};
232 use crate::{Ctx, Pipeline, Service, and_then::AndThen};
233
234 #[derive(Default)]
235 struct Gate {
236 closed: Cell<bool>,
237 waker: RefCell<Option<Waker>>,
238 }
239
240 struct Svc1(Rc<Gate>);
242
243 impl Service<(), Req> for Svc1 {
244 type Res = Receiver<()>;
245 type Error = ();
246
247 async fn ready(&self, _: Ctx<'_, Self>) -> Result<(), ()> {
248 poll_fn(|cx| {
249 if self.0.closed.get() {
250 *self.0.waker.borrow_mut() = Some(cx.waker().clone());
251 Poll::Pending
252 } else {
253 Poll::Ready(Ok(()))
254 }
255 })
256 .await
257 }
258
259 async fn call(&self, (rx1, rx2): Req, _: Ctx<'_, Self>) -> Result<Receiver<()>, ()> {
260 let _ = rx1.await;
261 Ok(rx2)
262 }
263 }
264
265 fn setup() -> (Rc<Gate>, Rc<State>, Pipeline<Req, (), ()>) {
266 let gate = Rc::new(Gate::default());
267 let st = Rc::new(State::default());
268 let pl = Pipeline::new((), AndThen::new(Svc1(gate.clone()), Single(st.clone())));
269 (gate, st, pl)
270 }
271
272 #[ntex::test]
273 async fn cached_svc2_readiness() {
274 let (gate, st, pl) = setup();
275 let (a1, _a2) = start(&pl);
276 tick().await; gate.closed.set(true);
278 let (b1, _b2) = start(&pl);
279 tick().await; let _ = b1.send(());
281 let _ = a1.send(());
282 tick().await; assert_eq!(st.active.get(), 1);
284 gate.closed.set(false);
285 if let Some(w) = gate.waker.borrow_mut().take() {
286 w.wake();
287 }
288 tick().await; assert_eq!(st.max.get(), 1, "svc2 capacity exceeded");
290 }
291
292 #[ntex::test]
293 async fn concurrent_svc2_entry() {
294 let (_gate, st, pl) = setup();
295 concurrent_entry(&pl, &st).await;
296 }
297
298 #[ntex::test]
299 async fn svc2_readiness_reused() {
300 let (_gate, st, pl) = setup();
301 for _ in 0..3 {
302 let (tx1, rx1) = oneshot::channel();
303 let (tx2, rx2) = oneshot::channel();
304 let _ = tx1.send(());
305 let _ = tx2.send(());
306 pl.call((rx1, rx2)).await.unwrap();
307 }
308 assert_eq!(st.checks.get(), 3);
309 }
310 }
311}