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