Skip to main content

rama_http/layer/required_header/
request.rs

1//! Set required headers on the request, if they are missing.
2//!
3//! For now this only sets `Host` header on http/1.1,
4//! as well as always a User-Agent for all versions.
5
6use 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/// Layer that applies [`AddRequiredRequestHeaders`] which adds a request header.
22///
23/// See [`AddRequiredRequestHeaders`] for more details.
24#[derive(Debug, Clone, Default)]
25pub struct AddRequiredRequestHeadersLayer {
26    overwrite: bool,
27    user_agent_header_value: Option<HeaderValue>,
28}
29
30impl AddRequiredRequestHeadersLayer {
31    /// Create a new [`AddRequiredRequestHeadersLayer`].
32    #[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        /// Set whether to overwrite the existing headers.
42        /// If set to `true`, the headers will be overwritten.
43        ///
44        /// Default is `false`.
45        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        /// Define the custom [`USER_AGENT`] header value to be used.
53        ///
54        /// By default a versioned `rama` value is used.
55        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/// Middleware that sets a header on the request.
83#[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    /// Create a new [`AddRequiredRequestHeaders`].
92    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        /// Set whether to overwrite the existing headers.
102        /// If set to `true`, the headers will be overwritten.
103        ///
104        /// Default is `false`.
105        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        /// Define the custom [`USER_AGENT`] header value.
113        ///
114        /// By default a versioned `rama` value is used.
115        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}