ntex/web/middleware/
defaultheaders.rs1use 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#[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 pub fn new() -> DefaultHeaders {
53 DefaultHeaders::default()
54 }
55
56 #[must_use]
57 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 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 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 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}