Skip to main content

rama_http/layer/required_header/
response.rs

1//! Set required headers on the response, if they are missing.
2//!
3//! For now this only sets `Server` and `Date` heades.
4
5use crate::{
6    HeaderValue, Request, Response,
7    header::{self, DATE, RAMA_ID_HEADER_VALUE, SERVER},
8    headers::{Date, HeaderMapExt},
9};
10use rama_core::{Layer, Service};
11use rama_utils::{macros::define_inner_service_accessors, time::now_system_time};
12use std::fmt;
13
14/// Layer that applies [`AddRequiredResponseHeaders`] which adds a request header.
15///
16/// See [`AddRequiredResponseHeaders`] for more details.
17#[derive(Debug, Clone, Default)]
18pub struct AddRequiredResponseHeadersLayer {
19    overwrite: bool,
20    server_header_value: Option<HeaderValue>,
21}
22
23impl AddRequiredResponseHeadersLayer {
24    /// Create a new [`AddRequiredResponseHeadersLayer`].
25    #[must_use]
26    pub const fn new() -> Self {
27        Self {
28            overwrite: false,
29            server_header_value: None,
30        }
31    }
32
33    rama_utils::macros::generate_set_and_with! {
34        /// Set whether to overwrite the existing headers.
35        /// If set to `true`, the headers will be overwritten.
36        ///
37        /// Default is `false`.
38        pub fn overwrite(mut self, overwrite: bool) -> Self {
39            self.overwrite = overwrite;
40            self
41        }
42    }
43
44    rama_utils::macros::generate_set_and_with! {
45        /// Define the custom [`SERVER`] header value.
46        ///
47        /// By default a versioned `rama` value is used.
48        pub fn server_header_value(mut self, value: Option<HeaderValue>) -> Self {
49            self.server_header_value = value;
50            self
51        }
52    }
53}
54
55impl<S> Layer<S> for AddRequiredResponseHeadersLayer {
56    type Service = AddRequiredResponseHeaders<S>;
57
58    fn layer(&self, inner: S) -> Self::Service {
59        AddRequiredResponseHeaders {
60            inner,
61            overwrite: self.overwrite,
62            server_header_value: self.server_header_value.clone(),
63        }
64    }
65
66    fn into_layer(self, inner: S) -> Self::Service {
67        AddRequiredResponseHeaders {
68            inner,
69            overwrite: self.overwrite,
70            server_header_value: self.server_header_value,
71        }
72    }
73}
74
75/// Middleware that sets a header on the request.
76#[derive(Clone)]
77pub struct AddRequiredResponseHeaders<S> {
78    inner: S,
79    overwrite: bool,
80    server_header_value: Option<HeaderValue>,
81}
82
83impl<S> AddRequiredResponseHeaders<S> {
84    /// Create a new [`AddRequiredResponseHeaders`].
85    pub const fn new(inner: S) -> Self {
86        Self {
87            inner,
88            overwrite: false,
89            server_header_value: None,
90        }
91    }
92
93    rama_utils::macros::generate_set_and_with! {
94        /// Set whether to overwrite the existing headers.
95        /// If set to `true`, the headers will be overwritten.
96        ///
97        /// Default is `false`.
98        pub fn overwrite(mut self, overwrite: bool) -> Self {
99            self.overwrite = overwrite;
100            self
101        }
102    }
103
104    rama_utils::macros::generate_set_and_with! {
105        /// Define the custom [`SERVER`] header value.
106        ///
107        /// By default a versioned `rama` value is used.
108        pub fn server_header_value(mut self, value: Option<HeaderValue>) -> Self {
109            self.server_header_value = value;
110            self
111        }
112    }
113
114    define_inner_service_accessors!();
115}
116
117impl<S> fmt::Debug for AddRequiredResponseHeaders<S>
118where
119    S: fmt::Debug,
120{
121    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
122        f.debug_struct("AddRequiredResponseHeaders")
123            .field("inner", &self.inner)
124            .field("server_header_value", &self.server_header_value)
125            .finish()
126    }
127}
128
129impl<ReqBody, ResBody, S> Service<Request<ReqBody>> for AddRequiredResponseHeaders<S>
130where
131    ReqBody: Send + 'static,
132    ResBody: Send + 'static,
133    S: Service<Request<ReqBody>, Output = Response<ResBody>>,
134{
135    type Output = S::Output;
136    type Error = S::Error;
137
138    async fn serve(&self, req: Request<ReqBody>) -> Result<Self::Output, Self::Error> {
139        let mut resp = self.inner.serve(req).await?;
140
141        if self.overwrite {
142            resp.headers_mut().insert(
143                SERVER,
144                self.server_header_value
145                    .as_ref()
146                    .unwrap_or(&RAMA_ID_HEADER_VALUE)
147                    .clone(),
148            );
149        } else if let header::Entry::Vacant(header) = resp.headers_mut().entry(SERVER) {
150            header.insert(
151                self.server_header_value
152                    .as_ref()
153                    .unwrap_or(&RAMA_ID_HEADER_VALUE)
154                    .clone(),
155            );
156        }
157
158        if self.overwrite || !resp.headers().contains_key(DATE) {
159            resp.headers_mut()
160                .typed_insert(Date::from(now_system_time()));
161        }
162
163        Ok(resp)
164    }
165}
166
167#[cfg(test)]
168mod tests {
169    use super::*;
170    use crate::Body;
171    use rama_core::{Layer, service::service_fn};
172    use std::convert::Infallible;
173
174    #[tokio::test]
175    async fn add_required_response_headers() {
176        let svc = AddRequiredResponseHeadersLayer::default().into_layer(service_fn(
177            async |req: Request| {
178                assert!(!req.headers().contains_key(SERVER));
179                assert!(!req.headers().contains_key(DATE));
180                Ok::<_, Infallible>(Response::new(Body::empty()))
181            },
182        ));
183
184        let req = Request::new(Body::empty());
185        let resp = svc.serve(req).await.unwrap();
186
187        assert_eq!(
188            resp.headers().get(SERVER).unwrap(),
189            RAMA_ID_HEADER_VALUE.as_ref()
190        );
191        assert!(resp.headers().contains_key(DATE));
192    }
193
194    #[tokio::test]
195    async fn add_required_response_headers_custom_server() {
196        let svc = AddRequiredResponseHeadersLayer::default()
197            .with_server_header_value(HeaderValue::from_static("foo"))
198            .into_layer(service_fn(async |req: Request| {
199                assert!(!req.headers().contains_key(SERVER));
200                assert!(!req.headers().contains_key(DATE));
201                Ok::<_, Infallible>(Response::new(Body::empty()))
202            }));
203
204        let req = Request::new(Body::empty());
205        let resp = svc.serve(req).await.unwrap();
206
207        assert_eq!(
208            resp.headers().get(SERVER).and_then(|v| v.to_str().ok()),
209            Some("foo")
210        );
211        assert!(resp.headers().contains_key(DATE));
212    }
213
214    #[tokio::test]
215    async fn add_required_response_headers_overwrite() {
216        let svc = AddRequiredResponseHeadersLayer::new()
217            .with_overwrite(true)
218            .into_layer(service_fn(async |req: Request| {
219                assert!(!req.headers().contains_key(SERVER));
220                assert!(!req.headers().contains_key(DATE));
221                Ok::<_, Infallible>(
222                    Response::builder()
223                        .header(SERVER, "foo")
224                        .header(DATE, "bar")
225                        .body(Body::empty())
226                        .unwrap(),
227                )
228            }));
229
230        let req = Request::new(Body::empty());
231        let resp = svc.serve(req).await.unwrap();
232
233        assert_eq!(
234            resp.headers().get(SERVER).unwrap(),
235            RAMA_ID_HEADER_VALUE.to_str().unwrap()
236        );
237        assert_ne!(resp.headers().get(DATE).unwrap(), "bar");
238    }
239
240    #[tokio::test]
241    async fn add_required_response_headers_overwrite_custom_ua() {
242        let svc = AddRequiredResponseHeadersLayer::new()
243            .with_overwrite(true)
244            .with_server_header_value(HeaderValue::from_static("foo"))
245            .into_layer(service_fn(async |req: Request| {
246                assert!(!req.headers().contains_key(SERVER));
247                assert!(!req.headers().contains_key(DATE));
248                Ok::<_, Infallible>(
249                    Response::builder()
250                        .header(SERVER, "foo")
251                        .header(DATE, "bar")
252                        .body(Body::empty())
253                        .unwrap(),
254                )
255            }));
256
257        let req = Request::new(Body::empty());
258        let resp = svc.serve(req).await.unwrap();
259
260        assert_eq!(
261            resp.headers().get(SERVER).and_then(|v| v.to_str().ok()),
262            Some("foo")
263        );
264        assert_ne!(resp.headers().get(DATE).unwrap(), "bar");
265    }
266}