rama_http/layer/required_header/
response.rs1use 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#[derive(Debug, Clone, Default)]
18pub struct AddRequiredResponseHeadersLayer {
19 overwrite: bool,
20 server_header_value: Option<HeaderValue>,
21}
22
23impl AddRequiredResponseHeadersLayer {
24 #[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 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 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#[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 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 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 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}