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