Skip to main content

ntex_service/
apply.rs

1use std::{cell::Cell, fmt, marker};
2
3use crate::ctx::{Ctx, WaitersRef};
4use crate::{IntoService, IntoServiceFactory, Service, ServiceFactory};
5use crate::{ServiceCaller, ServiceChain, ServiceChainFactory};
6
7/// Applies an asynchronous middleware function to a service.
8///
9/// The function receives an input request and an [`ApplyCtx`] that can call the
10/// wrapped service.
11pub fn apply_fn<S, St, Req, F, In, Out, Err>(
12    service: impl IntoService<S, St, Req>,
13    f: F,
14) -> ServiceChain<Apply<S, St, Req, F, In, Out, Err>, St, In>
15where
16    S: Service<St, Req>,
17    F: AsyncFn(In, &ApplyCtx<'_, S, St, Req>) -> Result<Out, Err>,
18    Err: From<S::Error>,
19{
20    crate::service(Apply::new(service.into_service(), f))
21}
22
23/// Applies an asynchronous middleware function to every service from a factory.
24pub fn apply_fn_factory<Sf, St, Req, F, In, Out, Err>(
25    service: impl IntoServiceFactory<Sf, St, Req>,
26    f: F,
27) -> ServiceChainFactory<ApplyFactory<F, Sf, St, Req, In, Out, Err>, St, In>
28where
29    Sf: ServiceFactory<St, Req>,
30    F: AsyncFn(In, &ApplyCtx<'_, Sf::Service, St, Req>) -> Result<Out, Err> + Clone,
31    Err: From<Sf::Error>,
32{
33    crate::factory(ApplyFactory::new(service.into_factory(), f))
34}
35
36#[derive(Debug)]
37/// Context passed to middleware functions created by [`apply_fn`].
38pub struct ApplyCtx<'a, S, St, Req> {
39    idx: u32,
40    waiters: &'a WaitersRef,
41    service: &'a S,
42    st: &'a St,
43    ready: &'a Cell<bool>,
44    entered: &'a Cell<u32>,
45    r: marker::PhantomData<Req>,
46}
47
48impl<S: Service<St, Req>, St, Req> ApplyCtx<'_, S, St, Req> {
49    /// Returns the pipeline state.
50    #[inline]
51    pub fn st(&self) -> &St {
52        self.st
53    }
54
55    /// Waits for the wrapped service to become ready, then calls it.
56    #[inline]
57    pub async fn call(&self, req: Req) -> Result<S::Res, S::Error> {
58        self.enter(req, self.ready.get()).await
59    }
60
61    async fn enter(&self, req: Req, skip_ready: bool) -> Result<S::Res, S::Error> {
62        let ctx = Ctx::<S, St>::new(self.idx, self.waiters, self.st);
63        if !skip_ready {
64            ctx.ready(self.service).await?;
65        }
66        // entering the service invalidates all readiness observed so far
67        self.ready.set(false);
68        self.entered.set(self.entered.get().wrapping_add(1));
69        ctx.call_nowait(self.service, req).await
70    }
71}
72
73impl<S: Service<St, Req>, St, Req> ServiceCaller<Req, S::Res, S::Error>
74    for ApplyCtx<'_, S, St, Req>
75{
76    #[inline]
77    async fn call_service(&self, req: Req) -> Result<S::Res, S::Error> {
78        self.enter(req, false).await
79    }
80}
81
82/// Service produced by [`apply_fn`].
83pub struct Apply<S, St, Req, F, In, Out, Err> {
84    svc: S,
85    f: F,
86    ready: Cell<bool>,
87    entered: Cell<u32>,
88    r: marker::PhantomData<fn(St, Req) -> (In, Out, Err)>,
89}
90
91impl<S, St, Req, F, In, Out, Err> Apply<S, St, Req, F, In, Out, Err>
92where
93    F: AsyncFn(In, &ApplyCtx<'_, S, St, Req>) -> Result<Out, Err>,
94{
95    pub(crate) fn new(svc: S, f: F) -> Self {
96        Apply {
97            f,
98            svc,
99            ready: Cell::new(false),
100            entered: Cell::new(0),
101            r: marker::PhantomData,
102        }
103    }
104}
105
106impl<S, St, Req, F, In, Out, Err> Clone for Apply<S, St, Req, F, In, Out, Err>
107where
108    S: Clone,
109    F: Clone,
110{
111    fn clone(&self) -> Self {
112        Apply {
113            svc: self.svc.clone(),
114            f: self.f.clone(),
115            ready: Cell::new(false),
116            entered: Cell::new(0),
117            r: marker::PhantomData,
118        }
119    }
120}
121
122impl<S, St, Req, F, In, Out, Err> fmt::Debug for Apply<S, St, Req, F, In, Out, Err>
123where
124    S: fmt::Debug,
125{
126    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
127        f.debug_struct("Apply")
128            .field("svc", &self.svc)
129            .field("map", &std::any::type_name::<F>())
130            .finish()
131    }
132}
133
134impl<S, St, Req, F, In, Out, Err> Service<St, In> for Apply<S, St, Req, F, In, Out, Err>
135where
136    S: Service<St, Req>,
137    F: AsyncFn(In, &ApplyCtx<'_, S, St, Req>) -> Result<Out, Err>,
138    Err: From<S::Error>,
139{
140    type Res = Out;
141    type Error = Err;
142
143    #[inline]
144    async fn ready(&self, ctx: Ctx<'_, Self, St>) -> Result<(), Err> {
145        let entered = self.entered.get();
146        let result = ctx.ready(&self.svc).await.map_err(From::from);
147        if result.is_err() {
148            self.ready.set(false);
149        } else if self.entered.get() == entered {
150            // readiness is valid only if no call entered the service during the check
151            self.ready.set(true);
152        }
153        result
154    }
155
156    #[inline]
157    async fn call(&self, req: In, ctx: Ctx<'_, Self, St>) -> Result<Out, Err> {
158        let (idx, waiters, st) = ctx.inner();
159
160        let ctx = ApplyCtx {
161            idx,
162            waiters,
163            st,
164            ready: &self.ready,
165            entered: &self.entered,
166            service: &self.svc,
167            r: marker::PhantomData,
168        };
169        (self.f)(req, &ctx).await
170    }
171
172    crate::forward_shutdown!(St, svc);
173}
174
175/// Service factory produced by [`apply_fn_factory`].
176pub struct ApplyFactory<F, Sf, St, Req, In, Out, Err>
177where
178    F: AsyncFn(In, &ApplyCtx<'_, Sf::Service, St, Req>) -> Result<Out, Err> + Clone,
179    Sf: ServiceFactory<St, Req>,
180{
181    f: F,
182    sf: Sf,
183    r: marker::PhantomData<fn(St, Req) -> (In, Out)>,
184}
185
186impl<F, Sf, St, Req, In, Out, Err> ApplyFactory<F, Sf, St, Req, In, Out, Err>
187where
188    F: AsyncFn(In, &ApplyCtx<'_, Sf::Service, St, Req>) -> Result<Out, Err> + Clone,
189    Sf: ServiceFactory<St, Req>,
190{
191    /// Creates a new `ApplyFactory`.
192    pub(crate) fn new(sf: Sf, f: F) -> Self
193    where
194        Sf: ServiceFactory<St, Req>,
195        Err: From<Sf::Error>,
196    {
197        Self {
198            f,
199            sf,
200            r: marker::PhantomData,
201        }
202    }
203}
204
205impl<F, Sf, St, Req, In, Out, Err> Clone for ApplyFactory<F, Sf, St, Req, In, Out, Err>
206where
207    F: AsyncFn(In, &ApplyCtx<'_, Sf::Service, St, Req>) -> Result<Out, Err> + Clone,
208    Sf: ServiceFactory<St, Req> + Clone,
209{
210    fn clone(&self) -> Self {
211        Self {
212            f: self.f.clone(),
213            sf: self.sf.clone(),
214            r: marker::PhantomData,
215        }
216    }
217}
218
219impl<F, Sf, St, Req, In, Out, Err> fmt::Debug for ApplyFactory<F, Sf, St, Req, In, Out, Err>
220where
221    F: AsyncFn(In, &ApplyCtx<'_, Sf::Service, St, Req>) -> Result<Out, Err> + Clone,
222    Sf: ServiceFactory<St, Req> + fmt::Debug,
223{
224    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
225        f.debug_struct("ApplyFactory")
226            .field("factory", &self.sf)
227            .field("map", &std::any::type_name::<F>())
228            .finish()
229    }
230}
231
232impl<F, Sf, St, Req, In, Out, Err> ServiceFactory<St, In>
233    for ApplyFactory<F, Sf, St, Req, In, Out, Err>
234where
235    F: AsyncFn(In, &ApplyCtx<'_, Sf::Service, St, Req>) -> Result<Out, Err> + Clone,
236    Sf: ServiceFactory<St, Req>,
237    Err: From<Sf::Error>,
238{
239    type Res = Out;
240    type Error = Err;
241
242    type Service = Apply<Sf::Service, St, Req, F, In, Out, Err>;
243    type InitError = Sf::InitError;
244
245    #[inline]
246    async fn create(&self, st: &St) -> Result<Self::Service, Self::InitError> {
247        self.sf.create(st).await.map(|svc| Apply {
248            svc,
249            f: self.f.clone(),
250            r: marker::PhantomData,
251            ready: Cell::new(false),
252            entered: Cell::new(0),
253        })
254    }
255}
256
257#[cfg(test)]
258mod tests {
259    use std::{cell::Cell, rc::Rc};
260
261    use super::*;
262    use crate::{factory, fn_factory, service};
263
264    #[derive(Debug, Default, Clone)]
265    struct Srv(Rc<Cell<usize>>);
266
267    impl Service<(), ()> for Srv {
268        type Res = ();
269        type Error = ();
270
271        async fn call(&self, _r: (), _: Ctx<'_, Self>) -> Result<(), ()> {
272            Ok(())
273        }
274
275        async fn ready(&self, _: Ctx<'_, Self>) -> Result<(), ()> {
276            self.0.set(self.0.get() + 1);
277            Ok(())
278        }
279
280        async fn shutdown(&self, _: Ctx<'_, Self, ()>) {
281            self.0.set(self.0.get() + 1);
282        }
283    }
284
285    #[derive(Debug, PartialEq, Eq)]
286    struct Err;
287
288    impl From<()> for Err {
289        fn from(_e: ()) -> Self {
290            Err
291        }
292    }
293
294    #[ntex::test]
295    async fn test_call() {
296        let cnt_sht = Rc::new(Cell::new(0));
297        let srv = service(
298            apply_fn(Srv(cnt_sht.clone()), async move |req: &'static str, svc| {
299                svc.call(()).await.unwrap();
300                Ok((req, ()))
301            })
302            .clone(),
303        )
304        .pipeline(());
305
306        assert_eq!(srv.ready().await, Ok::<_, Err>(()));
307
308        srv.shutdown().await;
309        assert_eq!(cnt_sht.get(), 2);
310
311        let res = srv.call("srv").await;
312        assert!(res.is_ok());
313        assert_eq!(res.unwrap(), ("srv", ()));
314    }
315
316    #[ntex::test]
317    async fn test_call_svc() {
318        let cnt_sht = Rc::new(Cell::new(0));
319        let srv = service(Srv(cnt_sht.clone()))
320            .apply_fn(async move |req: &'static str, svc| {
321                svc.st();
322                svc.call(()).await.unwrap();
323                Ok((req, ()))
324            })
325            .clone();
326        let s = format!("{srv:?}");
327        assert!(s.contains("Apply"), "{}", s);
328
329        let srv = srv.pipeline(());
330        assert_eq!(srv.ready().await, Ok::<_, Err>(()));
331
332        srv.shutdown().await;
333        assert_eq!(cnt_sht.get(), 2);
334
335        let res = srv.call("srv").await;
336        assert!(res.is_ok());
337        assert_eq!(res.unwrap(), ("srv", ()));
338        let _ = format!("{srv:?}");
339    }
340
341    #[ntex::test]
342    async fn test_create() {
343        let new_srv = factory(apply_fn_factory(
344            fn_factory(|(): &()| async { Ok::<_, ()>(Srv::default()) }),
345            async move |req: &'static str, srv| {
346                srv.call(()).await.unwrap();
347                Ok((req, ()))
348            },
349        ));
350
351        let srv = new_srv.pipeline(()).await.unwrap();
352
353        assert_eq!(srv.ready().await, Ok::<_, Err>(()));
354
355        let res = srv.call("srv").await;
356        assert!(res.is_ok());
357        assert_eq!(res.unwrap(), ("srv", ()));
358        assert_eq!(Err, Err::from(()));
359    }
360
361    #[ntex::test]
362    async fn test_create_chain() {
363        let new_srv = factory(fn_factory(|(): &()| async { Ok::<_, ()>(Srv::default()) }))
364            .apply_fn(async move |req: &'static str, srv| {
365                srv.call(()).await.unwrap();
366                Ok((req, ()))
367            })
368            .clone();
369
370        let srv = new_srv.pipeline(()).await.unwrap();
371
372        assert_eq!(srv.ready().await, Ok::<_, Err>(()));
373
374        let res = srv.call("srv").await;
375        assert!(res.is_ok());
376        assert_eq!(res.unwrap(), ("srv", ()));
377        let _ = format!("{new_srv:?}");
378    }
379
380    mod ready_flag {
381        use std::rc::Rc;
382
383        use crate::util::tests::{Req, Single, State, concurrent_entry};
384        use crate::{Pipeline, apply::Apply};
385
386        #[ntex::test]
387        async fn concurrent_entry_after_await() {
388            let st = Rc::new(State::default());
389            let pl = Pipeline::new(
390                (),
391                Apply::new(Single(st.clone()), async |(rx1, rx2): Req, svc| {
392                    let _ = rx1.await;
393                    svc.call(rx2).await
394                }),
395            );
396            concurrent_entry(&pl, &st).await;
397        }
398    }
399}