rama_http/layer/required_header/
request.rs1use crate::{
7 HeaderValue, Request, Response,
8 header::{self, HOST, RAMA_ID_HEADER_VALUE, USER_AGENT},
9 headers::HeaderMapExt,
10};
11use rama_core::{
12 Layer, Service,
13 error::{BoxError, ErrorContext},
14 telemetry::tracing,
15};
16use rama_http_headers::Host;
17use rama_net::{AuthorityInputExt, ProtocolInputExt};
18use rama_utils::macros::define_inner_service_accessors;
19use std::fmt;
20
21#[derive(Debug, Clone, Default)]
25pub struct AddRequiredRequestHeadersLayer {
26 overwrite: bool,
27 user_agent_header_value: Option<HeaderValue>,
28}
29
30impl AddRequiredRequestHeadersLayer {
31 #[must_use]
33 pub const fn new() -> Self {
34 Self {
35 overwrite: false,
36 user_agent_header_value: None,
37 }
38 }
39
40 rama_utils::macros::generate_set_and_with! {
41 pub fn overwrite(mut self, overwrite: bool) -> Self {
46 self.overwrite = overwrite;
47 self
48 }
49 }
50
51 rama_utils::macros::generate_set_and_with! {
52 pub fn user_agent_header_value(mut self, value: Option<HeaderValue>) -> Self {
56 self.user_agent_header_value = value;
57 self
58 }
59 }
60}
61
62impl<S> Layer<S> for AddRequiredRequestHeadersLayer {
63 type Service = AddRequiredRequestHeaders<S>;
64
65 fn layer(&self, inner: S) -> Self::Service {
66 AddRequiredRequestHeaders {
67 inner,
68 overwrite: self.overwrite,
69 user_agent_header_value: self.user_agent_header_value.clone(),
70 }
71 }
72
73 fn into_layer(self, inner: S) -> Self::Service {
74 AddRequiredRequestHeaders {
75 inner,
76 overwrite: self.overwrite,
77 user_agent_header_value: self.user_agent_header_value,
78 }
79 }
80}
81
82#[derive(Clone)]
84pub struct AddRequiredRequestHeaders<S> {
85 inner: S,
86 overwrite: bool,
87 user_agent_header_value: Option<HeaderValue>,
88}
89
90impl<S> AddRequiredRequestHeaders<S> {
91 pub const fn new(inner: S) -> Self {
93 Self {
94 inner,
95 overwrite: false,
96 user_agent_header_value: None,
97 }
98 }
99
100 rama_utils::macros::generate_set_and_with! {
101 pub fn overwrite(mut self, overwrite: bool) -> Self {
106 self.overwrite = overwrite;
107 self
108 }
109 }
110
111 rama_utils::macros::generate_set_and_with! {
112 pub fn user_agent_header_value(mut self, value: Option<HeaderValue>) -> Self {
116 self.user_agent_header_value = value;
117 self
118 }
119 }
120
121 define_inner_service_accessors!();
122}
123
124impl<S> fmt::Debug for AddRequiredRequestHeaders<S>
125where
126 S: fmt::Debug,
127{
128 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
129 f.debug_struct("AddRequiredRequestHeaders")
130 .field("inner", &self.inner)
131 .field("user_agent_header_value", &self.user_agent_header_value)
132 .finish()
133 }
134}
135
136impl<ReqBody, ResBody, S> Service<Request<ReqBody>> for AddRequiredRequestHeaders<S>
137where
138 ReqBody: Send + 'static,
139 ResBody: Send + 'static,
140 S: Service<Request<ReqBody>, Output = Response<ResBody>, Error: Into<BoxError>>,
141{
142 type Output = S::Output;
143 type Error = BoxError;
144
145 async fn serve(&self, mut req: Request<ReqBody>) -> Result<Self::Output, Self::Error> {
146 if self.overwrite || !req.headers().contains_key(HOST) {
147 let authority = req
148 .authority()
149 .context("AddRequiredRequestHeaders: resolve authority")?;
150 let protocol = req.protocol();
151 let authority = authority.without_default_port_for(protocol);
152 tracing::trace!(
153 server.address = %authority,
154 "add missing authority as host header",
155 );
156 req.headers_mut().typed_insert(Host::from(authority));
157 }
158
159 if self.overwrite {
160 req.headers_mut().insert(
161 USER_AGENT,
162 self.user_agent_header_value
163 .as_ref()
164 .unwrap_or(&RAMA_ID_HEADER_VALUE)
165 .clone(),
166 );
167 } else if let header::Entry::Vacant(header) = req.headers_mut().entry(USER_AGENT) {
168 header.insert(
169 self.user_agent_header_value
170 .as_ref()
171 .unwrap_or(&RAMA_ID_HEADER_VALUE)
172 .clone(),
173 );
174 }
175
176 self.inner.serve(req).await.into_box_error()
177 }
178}
179
180#[cfg(test)]
181mod test {
182 use super::*;
183 use crate::{Body, Request};
184 use rama_core::service::service_fn;
185 use rama_core::{Layer, Service};
186 use std::convert::Infallible;
187
188 #[tokio::test]
189 async fn add_required_request_headers() {
190 let svc = AddRequiredRequestHeadersLayer::default().into_layer(service_fn(
191 async |req: Request| {
192 assert!(req.headers().contains_key(HOST));
193 assert!(req.headers().contains_key(USER_AGENT));
194 Ok::<_, Infallible>(rama_http_types::Response::new(Body::empty()))
195 },
196 ));
197
198 let req = Request::builder()
199 .uri("http://www.example.com/")
200 .body(Body::empty())
201 .unwrap();
202 let resp = svc.serve(req).await.unwrap();
203
204 assert!(!resp.headers().contains_key(HOST));
205 assert!(!resp.headers().contains_key(USER_AGENT));
206 }
207
208 #[tokio::test]
209 async fn add_required_request_headers_custom_ua() {
210 let svc = AddRequiredRequestHeadersLayer::default()
211 .with_user_agent_header_value(HeaderValue::from_static("foo"))
212 .into_layer(service_fn(async |req: Request| {
213 assert!(req.headers().contains_key(HOST));
214 assert_eq!(
215 req.headers().get(USER_AGENT).and_then(|v| v.to_str().ok()),
216 Some("foo")
217 );
218 Ok::<_, Infallible>(rama_http_types::Response::new(Body::empty()))
219 }));
220
221 let req = Request::builder()
222 .uri("http://www.example.com/")
223 .body(Body::empty())
224 .unwrap();
225 let resp = svc.serve(req).await.unwrap();
226
227 assert!(!resp.headers().contains_key(HOST));
228 assert!(!resp.headers().contains_key(USER_AGENT));
229 }
230
231 #[tokio::test]
232 async fn add_required_request_headers_overwrite() {
233 let svc = AddRequiredRequestHeadersLayer::new()
234 .with_overwrite(true)
235 .into_layer(service_fn(async |req: Request| {
236 assert_eq!(req.headers().get(HOST).unwrap(), "127.0.0.1");
237 assert_eq!(
238 req.headers().get(USER_AGENT).unwrap(),
239 RAMA_ID_HEADER_VALUE.to_str().unwrap()
240 );
241 Ok::<_, Infallible>(rama_http_types::Response::new(Body::empty()))
242 }));
243
244 let req = Request::builder()
245 .uri("http://127.0.0.1/")
246 .header(HOST, "example.com")
247 .header(USER_AGENT, "test")
248 .body(Body::empty())
249 .unwrap();
250
251 let resp = svc.serve(req).await.unwrap();
252
253 assert!(!resp.headers().contains_key(HOST));
254 assert!(!resp.headers().contains_key(USER_AGENT));
255 }
256
257 #[tokio::test]
258 async fn add_required_request_headers_overwrite_custom_ua() {
259 let svc = AddRequiredRequestHeadersLayer::new()
260 .with_overwrite(true)
261 .with_user_agent_header_value(HeaderValue::from_static("foo"))
262 .into_layer(service_fn(async |req: Request| {
263 assert_eq!(req.headers().get(HOST).unwrap(), "127.0.0.1");
264 assert_eq!(
265 req.headers().get(USER_AGENT).and_then(|v| v.to_str().ok()),
266 Some("foo")
267 );
268 Ok::<_, Infallible>(rama_http_types::Response::new(Body::empty()))
269 }));
270
271 let req = Request::builder()
272 .uri("http://127.0.0.1/")
273 .header(HOST, "example.com")
274 .header(USER_AGENT, "test")
275 .body(Body::empty())
276 .unwrap();
277
278 let resp = svc.serve(req).await.unwrap();
279
280 assert!(!resp.headers().contains_key(HOST));
281 assert!(!resp.headers().contains_key(USER_AGENT));
282 }
283}