Skip to main content

ntex/web/
stack.rs

1//! Middleware stack types used by applications, scopes, and resources.
2use std::marker::PhantomData;
3
4use crate::error::Failure;
5use crate::service::{Ctx, Middleware, Service, ServiceFactory};
6use crate::web::{State, WebError, WebRequest, WebResponse, WebResponseError};
7
8/// Stack of middlewares.
9#[derive(Debug, Clone)]
10pub struct WebStack<St, Inner, Outer> {
11    inner: Inner,
12    outer: Outer,
13    err: PhantomData<St>,
14}
15
16impl<St, Inner, Outer> WebStack<St, Inner, Outer> {
17    /// Create a stack that applies `inner` first and then wraps it with `outer`.
18    pub fn new(inner: Inner, outer: Outer) -> Self {
19        WebStack {
20            inner,
21            outer,
22            err: PhantomData,
23        }
24    }
25}
26
27impl<S, St, Inner, Outer> Middleware<S, St> for WebStack<St, Inner, Outer>
28where
29    St: State,
30    Inner: Middleware<S, St>,
31    Outer: Middleware<Inner::Service, St>,
32{
33    type Service = WebMiddleware<Outer::Service, St>;
34
35    fn create(&self, st: &St, service: S) -> Self::Service {
36        WebMiddleware {
37            svc: self.outer.create(st, self.inner.create(st, service)),
38            err: PhantomData,
39        }
40    }
41}
42
43/// Service produced by [`WebStack`] layers.
44///
45/// Wraps a middleware service and converts its errors into [`WebError`].
46#[derive(Debug)]
47pub struct WebMiddleware<S, St> {
48    svc: S,
49    err: PhantomData<St>,
50}
51
52impl<S, St> Clone for WebMiddleware<S, St>
53where
54    S: Clone,
55{
56    fn clone(&self) -> Self {
57        Self {
58            svc: self.svc.clone(),
59            err: PhantomData,
60        }
61    }
62}
63
64impl<S, St, In> Service<St, WebRequest<In>> for WebMiddleware<S, St>
65where
66    S: Service<St, WebRequest<In>, Res = WebResponse>,
67    S::Error: WebResponseError<St, St::Error>,
68    St: State,
69{
70    type Res = WebResponse;
71    type Error = WebError<St, St::Error>;
72
73    #[inline]
74    async fn call(
75        &self,
76        req: WebRequest<In>,
77        ctx: Ctx<'_, Self, St>,
78    ) -> Result<Self::Res, Self::Error> {
79        ctx.call(&self.svc, req).await.map_err(WebError::from_err)
80    }
81
82    crate::forward_ready!(St, svc, WebError::from_err);
83    crate::forward_shutdown!(St, svc);
84}
85
86/// Identity request filter.
87///
88/// The default filter of applications, scopes, and resources. It passes
89/// requests through unchanged.
90#[derive(derive_more::Debug)]
91#[debug("Filter")]
92pub struct Filter<St, In>(PhantomData<(St, In)>);
93
94impl<St, In> Filter<St, In> {
95    pub(super) fn new() -> Self {
96        Filter(PhantomData)
97    }
98}
99
100impl<St: State, In> ServiceFactory<St, WebRequest<In>> for Filter<St, In> {
101    type Res = WebRequest<In>;
102    type Error = WebError<St, St::Error>;
103
104    type Service = Filter<St, In>;
105    type InitError = Failure;
106
107    async fn create(&self, _: &St) -> Result<Self::Service, Self::InitError> {
108        Ok(Filter(PhantomData))
109    }
110}
111
112impl<St: State, In> Service<St, WebRequest<In>> for Filter<St, In> {
113    type Res = WebRequest<In>;
114    type Error = WebError<St, St::Error>;
115
116    async fn call(
117        &self,
118        req: WebRequest<In>,
119        _: Ctx<'_, Self, St>,
120    ) -> Result<Self::Res, Self::Error> {
121        Ok(req)
122    }
123}
124
125#[cfg(test)]
126mod tests {
127    use std::io;
128
129    use super::*;
130    use crate::http::StatusCode;
131    use crate::service::{Identity, Pipeline, fn_service};
132    use crate::web::{HttpResponse, test::TestRequest};
133
134    #[crate::rt_test]
135    async fn test_web_middleware() {
136        let svc = fn_service(async |req: WebRequest<()>| {
137            if req.path() == "/err" {
138                Err(io::Error::new(io::ErrorKind::NotFound, "not found"))
139            } else {
140                Ok(req.into_response(HttpResponse::Ok().build()))
141            }
142        });
143        let mw = WebStack::<(), _, _>::new(Identity, Identity).create(&(), svc);
144        let srv = Pipeline::new((), mw.clone());
145
146        let res = srv
147            .call(TestRequest::default().to_srv_request())
148            .await
149            .unwrap();
150        assert_eq!(res.status(), StatusCode::OK);
151
152        let err = srv
153            .call(TestRequest::with_uri("/err").to_srv_request())
154            .await
155            .unwrap_err();
156        assert_eq!(err.to_string(), "not found");
157        let res = WebResponseError::error_response(&err, &());
158        assert_eq!(res.status(), StatusCode::NOT_FOUND);
159    }
160}