Skip to main content

ntex_service/
middleware.rs

1use std::{fmt, marker::PhantomData, rc::Rc};
2
3use crate::dev::{Apply, ApplyCtx};
4use crate::{IntoServiceFactory, Service, ServiceChainFactory, ServiceFactory};
5
6/// Applies middleware to every service produced by a factory.
7pub fn apply<Sf, St, Req, M>(
8    mw: M,
9    factory: impl IntoServiceFactory<Sf, St, Req>,
10) -> ServiceChainFactory<ApplyMiddleware<M, Sf>, St, Req>
11where
12    Sf: ServiceFactory<St, Req>,
13    M: Middleware<Sf::Service, St>,
14{
15    ServiceChainFactory {
16        factory: ApplyMiddleware::new(mw, factory.into_factory()),
17        _t: PhantomData,
18    }
19}
20
21/// Wraps an inner service during service construction.
22///
23/// Middleware can run before and after the inner service, and can modify
24/// requests, responses, or errors.
25///
26/// For example, timeout middleware:
27///
28/// ```rust
29/// use ntex_service::{Ctx, Service};
30/// use ntex::{time::sleep, util::Either, util::select};
31///
32/// pub struct Timeout<S> {
33///     service: S,
34///     timeout: std::time::Duration,
35/// }
36///
37/// pub enum TimeoutError<E> {
38///    Service(E),
39///    Timeout,
40/// }
41///
42/// impl<S, R> Service<(), R> for Timeout<S>
43/// where
44///     S: Service<(), R>,
45/// {
46///     type Res = S::Res;
47///     type Error = TimeoutError<S::Error>;
48///
49///     async fn ready(&self, ctx: Ctx<'_, Self>) -> Result<(), Self::Error> {
50///         ctx.ready(&self.service).await.map_err(TimeoutError::Service)
51///     }
52///
53///     async fn call(&self, req: R, ctx: Ctx<'_, Self>) -> Result<Self::Res, Self::Error> {
54///         match select(sleep(self.timeout), ctx.call(&self.service, req)).await {
55///             Either::Left(_) => Err(TimeoutError::Timeout),
56///             Either::Right(res) => res.map_err(TimeoutError::Service),
57///         }
58///     }
59/// }
60/// ```
61///
62/// The timeout service is independent of the wrapped service implementation and
63/// can be applied to any compatible service.
64///
65/// A middleware factory for `Timeout` could look like this:
66///
67/// ```rust
68/// # use ntex_service::Middleware;
69/// # pub struct Timeout<S> {
70/// #     service: S,
71/// #     timeout: std::time::Duration,
72/// # }
73/// pub struct TimeoutMiddleware {
74///     timeout: std::time::Duration,
75/// }
76///
77/// impl<S> Middleware<S, ()> for TimeoutMiddleware
78/// {
79///     type Service = Timeout<S>;
80///
81///     fn create(&self, _: &(), service: S) -> Self::Service {
82///         Timeout {
83///             service,
84///             timeout: self.timeout,
85///         }
86///     }
87/// }
88/// ```
89pub trait Middleware<S, St> {
90    /// Service created by this middleware.
91    type Service;
92
93    /// Creates and returns a new middleware service.
94    fn create(&self, st: &St, service: S) -> Self::Service;
95
96    /// Creates a service factory that instantiates a service and applies
97    /// the current middleware to it.
98    ///
99    /// This is equivalent to `apply(self, factory)`.
100    fn apply_to<Sf, Req>(
101        self,
102        factory: Sf,
103    ) -> ServiceChainFactory<ApplyMiddleware<Self, Sf>, St, Req>
104    where
105        Sf: ServiceFactory<St, Req, Service = S>,
106        Self: Sized,
107        Self::Service: Service<St, Req>,
108    {
109        crate::factory(ApplyMiddleware::new(self, factory))
110    }
111}
112
113impl<M, S, St> Middleware<S, St> for Rc<M>
114where
115    M: Middleware<S, St>,
116{
117    type Service = M::Service;
118
119    fn create(&self, st: &St, service: S) -> M::Service {
120        self.as_ref().create(st, service)
121    }
122}
123
124/// A service factory with middleware applied.
125pub struct ApplyMiddleware<M, Sf>(Rc<(M, Sf)>);
126
127impl<M, Sf> ApplyMiddleware<M, Sf> {
128    /// Creates a new `ApplyMiddleware` factory.
129    pub(crate) fn new(mw: M, sf: Sf) -> Self {
130        Self(Rc::new((mw, sf)))
131    }
132}
133
134impl<M, Sf> Clone for ApplyMiddleware<M, Sf> {
135    fn clone(&self) -> Self {
136        Self(self.0.clone())
137    }
138}
139
140impl<M, Sf> fmt::Debug for ApplyMiddleware<M, Sf>
141where
142    M: fmt::Debug,
143    Sf: fmt::Debug,
144{
145    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
146        f.debug_struct("ApplyMiddleware")
147            .field("factory", &self.0.1)
148            .field("middleware", &self.0.0)
149            .finish()
150    }
151}
152
153impl<M, Sf, St, Req> ServiceFactory<St, Req> for ApplyMiddleware<M, Sf>
154where
155    Sf: ServiceFactory<St, Req>,
156    M: Middleware<Sf::Service, St>,
157    M::Service: Service<St, Req>,
158{
159    type Res = <M::Service as Service<St, Req>>::Res;
160    type Error = <M::Service as Service<St, Req>>::Error;
161
162    type Service = M::Service;
163    type InitError = Sf::InitError;
164
165    #[inline]
166    async fn create(&self, st: &St) -> Result<Self::Service, Self::InitError> {
167        Ok(self.0.0.create(st, self.0.1.create(st).await?))
168    }
169}
170
171/// Middleware that returns the wrapped service unchanged.
172#[derive(Debug, Clone, Copy)]
173pub struct Identity;
174
175impl<S, St> Middleware<S, St> for Identity {
176    type Service = S;
177
178    #[inline]
179    fn create(&self, _: &St, service: S) -> Self::Service {
180        service
181    }
182}
183
184/// Two middleware values applied in sequence.
185///
186/// The inner middleware is applied first, then the outer middleware wraps its
187/// service.
188#[derive(Debug, Clone)]
189pub struct Stack<Inner, Outer> {
190    inner: Inner,
191    outer: Outer,
192}
193
194impl<Inner, Outer> Stack<Inner, Outer> {
195    /// Creates a middleware stack.
196    pub fn new(inner: Inner, outer: Outer) -> Self {
197        Stack { inner, outer }
198    }
199}
200
201impl<S, St, Inner, Outer> Middleware<S, St> for Stack<Inner, Outer>
202where
203    Inner: Middleware<S, St>,
204    Outer: Middleware<Inner::Service, St>,
205{
206    type Service = Outer::Service;
207
208    fn create(&self, st: &St, service: S) -> Self::Service {
209        self.outer.create(st, self.inner.create(st, service))
210    }
211}
212
213#[doc(hidden)]
214/// Creates middleware from an asynchronous function.
215///
216/// The function receives an input request and an [`ApplyCtx`] that can call the
217/// wrapped service.
218pub fn fn_layer<F, S, St, Req, In, Out, Err>(f: F) -> FnMiddleware<F, S, St, Req, In, Out, Err>
219where
220    F: AsyncFn(In, &ApplyCtx<'_, S, St, Req>) -> Result<Out, Err> + Clone,
221    S: Service<St, Req>,
222{
223    FnMiddleware { f, r: PhantomData }
224}
225
226/// Middleware backed by an asynchronous function.
227pub struct FnMiddleware<F, S, St, Req, In, Out, Err> {
228    f: F,
229    r: PhantomData<fn(S, St, Req) -> (In, Out, Err)>,
230}
231
232impl<F, S, St, Req, In, Out, Err> Clone for FnMiddleware<F, S, St, Req, In, Out, Err>
233where
234    F: Clone,
235{
236    fn clone(&self) -> Self {
237        FnMiddleware {
238            f: self.f.clone(),
239            r: PhantomData,
240        }
241    }
242}
243
244impl<F, S, St, Req, In, Out, Err> fmt::Debug for FnMiddleware<F, S, St, Req, In, Out, Err> {
245    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
246        f.debug_struct("FnMiddleware")
247            .field("layer", &std::any::type_name::<F>())
248            .finish()
249    }
250}
251
252impl<F, S, St, Req, In, Out, Err> Middleware<S, St> for FnMiddleware<F, S, St, Req, In, Out, Err>
253where
254    S: Service<St, Req>,
255    F: AsyncFn(In, &ApplyCtx<'_, S, St, Req>) -> Result<Out, Err> + Clone,
256    Err: From<S::Error>,
257{
258    type Service = Apply<S, St, Req, F, In, Out, Err>;
259
260    fn create(&self, _: &St, service: S) -> Self::Service {
261        Apply::new(service, self.f.clone())
262    }
263}
264
265#[cfg(test)]
266#[allow(clippy::redundant_clone)]
267mod tests {
268    use std::{cell::Cell, rc::Rc};
269
270    use super::*;
271    use crate::{Ctx, Pipeline, factory, fn_service};
272
273    #[derive(Debug, Clone)]
274    struct Mw(Rc<Cell<usize>>);
275
276    impl<S, St> Middleware<S, St> for Mw {
277        type Service = Srv<S>;
278
279        fn create(&self, _: &St, service: S) -> Self::Service {
280            self.0.set(self.0.get() + 1);
281            Srv(service, self.0.clone())
282        }
283    }
284
285    #[derive(Debug, Clone)]
286    struct Srv<S>(S, Rc<Cell<usize>>);
287
288    impl<S: Service<(), R>, R> Service<(), R> for Srv<S> {
289        type Res = S::Res;
290        type Error = S::Error;
291
292        async fn ready(&self, ctx: Ctx<'_, Self>) -> Result<(), Self::Error> {
293            ctx.ready(&self.0).await
294        }
295
296        async fn call(&self, req: R, ctx: Ctx<'_, Self>) -> Result<S::Res, S::Error> {
297            ctx.call(&self.0, req).await
298        }
299
300        async fn shutdown(&self, _: Ctx<'_, Self, ()>) {
301            self.1.set(self.1.get() + 1);
302        }
303    }
304
305    #[ntex::test]
306    async fn middleware() {
307        let cnt_sht = Rc::new(Cell::new(0));
308        let fac = apply(
309            Rc::new(Mw(cnt_sht.clone()).clone()),
310            fn_service(|i: usize| async move { Ok::<_, ()>(i * 2) }),
311        )
312        .clone();
313
314        let srv = Pipeline::new((), fac.create(&()).await.unwrap().clone());
315        let res = srv.call(10).await;
316        assert!(res.is_ok());
317        assert_eq!(res.unwrap(), 20);
318        let _ = format!("{fac:?} {srv:?}");
319
320        assert_eq!(srv.ready().await, Ok(()));
321        srv.shutdown().await;
322        assert_eq!(cnt_sht.get(), 2);
323
324        let fac = factory(fn_service(|i: usize| async move { Ok::<_, ()>(i * 2) }))
325            .apply(Rc::new(Mw(Rc::new(Cell::new(0))).clone()))
326            .clone();
327
328        let srv = Pipeline::new((), fac.create(&()).await.unwrap().clone());
329        let res = srv.call(10).await;
330        assert!(res.is_ok());
331        assert_eq!(res.unwrap(), 20);
332        let _ = format!("{fac:?} {srv:?}");
333
334        assert_eq!(srv.ready().await, Ok(()));
335    }
336
337    #[ntex::test]
338    async fn middleware_apply() {
339        let cnt_sht = Rc::new(Cell::new(0));
340        let fac = Mw(cnt_sht.clone())
341            .apply_to(factory(async |i: usize| Ok::<_, ()>(i * 2)))
342            .boxed();
343
344        let srv = Pipeline::new((), fac.create(&()).await.unwrap());
345        let res = srv.call(10).await;
346        assert!(res.is_ok());
347        assert_eq!(res.unwrap(), 20);
348        let _ = format!("{fac:?} {srv:?}");
349
350        assert_eq!(srv.ready().await, Ok(()));
351        srv.shutdown().await;
352        assert_eq!(cnt_sht.get(), 2);
353    }
354
355    #[ntex::test]
356    async fn middleware_chain() {
357        let cnt_sht = Rc::new(Cell::new(0));
358        let fac = factory(fn_service(async move |i: usize| Ok::<_, ()>(i * 2)))
359            .apply(Mw(cnt_sht.clone()).clone());
360
361        let srv = Pipeline::new((), fac.create(&()).await.unwrap().clone());
362        let res = srv.call(10).await;
363        assert!(res.is_ok());
364        assert_eq!(res.unwrap(), 20);
365        let _ = format!("{fac:?} {srv:?}");
366
367        assert_eq!(srv.ready().await, Ok(()));
368        srv.shutdown().await;
369        assert_eq!(cnt_sht.get(), 2);
370    }
371
372    #[ntex::test]
373    async fn stack() {
374        let cnt_sht = Rc::new(Cell::new(0));
375        let mw = Stack::new(Identity, Mw(cnt_sht.clone()));
376        let _ = format!("{mw:?}");
377
378        let pl = Pipeline::new(
379            (),
380            Middleware::create(
381                &mw,
382                &(),
383                fn_service(|i: usize| async move { Ok::<_, ()>(i * 2) }),
384            ),
385        );
386        let res = pl.call(10).await;
387        assert!(res.is_ok());
388        assert_eq!(res.unwrap(), 20);
389        assert_eq!(pl.ready().await, Ok(()));
390        pl.shutdown().await;
391        assert_eq!(cnt_sht.get(), 2);
392    }
393
394    #[ntex::test]
395    async fn fn_middleware_service() {
396        let cnt_sht = Rc::new(Cell::new(0));
397        let cnt_sht2 = cnt_sht.clone();
398        let mw = fn_layer(async move |req: &'static str, svc| {
399            cnt_sht2.set(cnt_sht2.get() + 1);
400            let result = svc.call(1).await?;
401            Ok::<_, ()>((req, result))
402        })
403        .clone();
404        let _ = format!("{mw:?}");
405
406        let svc = Pipeline::new(
407            (),
408            mw.create(&(), fn_service(async move |i: usize| Ok::<_, ()>(i * 2))),
409        );
410
411        let res = svc.call("test").await;
412        assert!(res.is_ok());
413        assert_eq!(res.unwrap(), ("test", 2));
414        let _ = format!("{svc:?}");
415
416        assert_eq!(svc.ready().await, Ok(()));
417        svc.shutdown().await;
418        assert_eq!(cnt_sht.get(), 1);
419    }
420}