Skip to main content

ntex/web/middleware/
defaultheaders.rs

1//! Middleware for setting default response headers
2use std::rc::Rc;
3
4use crate::http::error::HttpError;
5use crate::http::header::{CONTENT_TYPE, HeaderMap, HeaderName, HeaderValue};
6use crate::service::{Ctx, Middleware, Service};
7use crate::web::{WebRequest, WebResponse};
8
9/// `Middleware` for setting default response headers.
10///
11/// This middleware does not set header if response headers already contains it.
12///
13/// ```rust
14/// use ntex::http;
15/// use ntex::web::{self, middleware, App, HttpResponse};
16///
17/// fn main() {
18///     let app = App::default()
19///         .middleware(middleware::DefaultHeaders::new().header("X-Version", "0.2"))
20///         .service(
21///             web::resource("/test")
22///                 .route(web::get().to(async || { HttpResponse::Ok() }))
23///                 .route(web::method(http::Method::HEAD).to(async || { HttpResponse::MethodNotAllowed() }))
24///         );
25/// }
26/// ```
27#[derive(Clone, Debug)]
28pub struct DefaultHeaders {
29    inner: Rc<Inner>,
30}
31
32#[derive(Debug)]
33struct Inner {
34    ct: bool,
35    headers: HeaderMap,
36}
37
38impl Default for DefaultHeaders {
39    fn default() -> Self {
40        DefaultHeaders {
41            inner: Rc::new(Inner {
42                ct: false,
43                headers: HeaderMap::new(),
44            }),
45        }
46    }
47}
48
49impl DefaultHeaders {
50    #[must_use]
51    /// Construct `DefaultHeaders` middleware.
52    pub fn new() -> DefaultHeaders {
53        DefaultHeaders::default()
54    }
55
56    #[must_use]
57    /// Set a header.
58    pub fn header<K, V>(mut self, key: K, value: V) -> Self
59    where
60        HeaderName: TryFrom<K>,
61        <HeaderName as TryFrom<K>>::Error: Into<HttpError>,
62        HeaderValue: TryFrom<V>,
63        <HeaderValue as TryFrom<V>>::Error: Into<HttpError>,
64    {
65        #[allow(clippy::match_wild_err_arm)]
66        match HeaderName::try_from(key) {
67            Ok(key) => match HeaderValue::try_from(value) {
68                Ok(value) => {
69                    Rc::get_mut(&mut self.inner)
70                        .expect("Multiple copies exist")
71                        .headers
72                        .append(key, value);
73                }
74                Err(_) => panic!("Cannot create header value"),
75            },
76            Err(_) => panic!("Cannot create header name"),
77        }
78        self
79    }
80
81    #[must_use]
82    /// Set *CONTENT-TYPE* header if response does not contain this header.
83    ///
84    /// The header is set to `application/octet-stream`.
85    pub fn content_type(mut self) -> Self {
86        Rc::get_mut(&mut self.inner)
87            .expect("Multiple copies exist")
88            .ct = true;
89        self
90    }
91}
92
93impl<S, St> Middleware<S, St> for DefaultHeaders {
94    type Service = DefaultHeadersMiddleware<S>;
95
96    fn create(&self, _: &St, service: S) -> Self::Service {
97        DefaultHeadersMiddleware {
98            service,
99            inner: self.inner.clone(),
100        }
101    }
102}
103
104#[derive(Debug)]
105pub struct DefaultHeadersMiddleware<S> {
106    service: S,
107    inner: Rc<Inner>,
108}
109
110impl<S, St, In> Service<St, WebRequest<In>> for DefaultHeadersMiddleware<S>
111where
112    S: Service<St, WebRequest<In>, Res = WebResponse>,
113{
114    type Res = WebResponse;
115    type Error = S::Error;
116
117    crate::forward_ready!(St, service);
118    crate::forward_shutdown!(St, service);
119
120    async fn call(&self, r: WebRequest<In>, ctx: Ctx<'_, Self, St>) -> Result<Self::Res, S::Error> {
121        let mut res = ctx.call(&self.service, r).await?;
122
123        // set response headers
124        for (key, value) in &self.inner.headers {
125            if !res.headers().contains_key(key) {
126                res.headers_mut().insert(key.clone(), value.clone());
127            }
128        }
129        // default content-type
130        if self.inner.ct && !res.headers().contains_key(&CONTENT_TYPE) {
131            res.headers_mut().insert(
132                CONTENT_TYPE,
133                HeaderValue::from_static("application/octet-stream"),
134            );
135        }
136        Ok(res)
137    }
138}
139
140#[cfg(test)]
141#[allow(unused_must_use)]
142mod tests {
143    use std::convert::Infallible;
144
145    use super::*;
146    use crate::web::{HttpResponse, test::TestRequest, test::ok_service};
147    use crate::{Pipeline, fn_service, util::lazy};
148
149    #[crate::rt_test]
150    async fn test_default_headers() {
151        let mw = Pipeline::new(
152            (),
153            Middleware::create(
154                &DefaultHeaders::new().header(CONTENT_TYPE, "0001"),
155                &(),
156                ok_service(),
157            ),
158        );
159
160        assert!(lazy(|cx| mw.poll_ready(cx).is_ready()).await);
161        assert!(lazy(|cx| mw.poll_shutdown(cx).is_ready()).await);
162
163        let req = TestRequest::default().to_srv_request();
164        let resp = mw.call(req).await.unwrap();
165        assert_eq!(resp.headers().get(CONTENT_TYPE).unwrap(), "0001");
166
167        let req = TestRequest::default().to_srv_request();
168        let srv = fn_service(async move |req: WebRequest<()>| {
169            Ok::<_, Infallible>(
170                req.into_response(HttpResponse::Ok().header(CONTENT_TYPE, "0002").build()),
171            )
172        });
173        let mw = Pipeline::new(
174            (),
175            Middleware::create(
176                &DefaultHeaders::new().header(CONTENT_TYPE, "0001"),
177                &(),
178                srv,
179            ),
180        );
181        let resp = mw.call(req).await.unwrap();
182        assert_eq!(resp.headers().get(CONTENT_TYPE).unwrap(), "0002");
183    }
184
185    #[crate::rt_test]
186    #[should_panic(expected = "Cannot create header name")]
187    async fn test_invalid_header_name() {
188        DefaultHeaders::new().header("no existing header name", "0001");
189    }
190
191    #[crate::rt_test]
192    #[should_panic(expected = "Cannot create header value")]
193    async fn test_invalid_header_value() {
194        DefaultHeaders::new().header(CONTENT_TYPE, "\n");
195    }
196
197    #[crate::rt_test]
198    async fn test_content_type() {
199        let srv = fn_service(async move |req: WebRequest<()>| {
200            Ok::<_, Infallible>(req.into_response(HttpResponse::Ok().build()))
201        });
202        let mw = Pipeline::new(
203            (),
204            Middleware::create(&DefaultHeaders::new().content_type(), &(), srv),
205        );
206
207        let req = TestRequest::default().to_srv_request();
208        let resp = mw.call(req).await.unwrap();
209        assert_eq!(
210            resp.headers().get(CONTENT_TYPE).unwrap(),
211            "application/octet-stream"
212        );
213    }
214}