1use 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#[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 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#[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#[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}