Skip to main content

ntex_service/
and_then.rs

1use std::cell::Cell;
2
3use super::{Ctx, Service, ServiceFactory, util};
4
5#[derive(Debug)]
6/// Service produced by the `and_then` combinator.
7///
8/// This is created by the `Service::and_then()` and `ServiceChain::and_then()` methods.
9pub 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    /// Creates a new `AndThen` service.
18    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        // entering svc2 invalidates all svc2 readiness observed so far
50        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            // svc2 readiness is valid only if no call entered svc2 during the check
63            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)]
75/// Service factory produced by the `and_then` combinator.
76///
77/// This is created by the `ServiceChainFactory::and_then()` method.
78pub struct AndThenFactory<A, B> {
79    svc1: A,
80    svc2: B,
81}
82
83impl<A, B> AndThenFactory<A, B> {
84    /// Creates a new `AndThenFactory`.
85    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        /// Readiness is gated, a call waits for its first sender
241        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; // A: ready, waits in svc1
277            gate.closed.set(true);
278            let (b1, _b2) = start(&pl);
279            tick().await; // B: svc2 readiness is cached while svc1 is not ready
280            let _ = b1.send(());
281            let _ = a1.send(());
282            tick().await; // A enters svc2
283            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; // B: readiness completes, B leaves svc1
289            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}