Skip to main content

rama_http/layer/remove_header/
response.rs

1//! Remove headers from a response.
2//!
3//! # Example
4//!
5//! ```
6//! use rama_http::layer::remove_header::RemoveResponseHeaderLayer;
7//! use rama_http::{Body, Request, Response, header::{self, HeaderValue}};
8//! use rama_core::service::service_fn;
9//! use rama_core::{Service, Layer};
10//! use rama_core::error::BoxError;
11//!
12//! # #[tokio::main]
13//! # async fn main() -> Result<(), BoxError> {
14//! # let http_client = service_fn(async |_: Request| {
15//! #     Ok::<_, std::convert::Infallible>(Response::new(Body::empty()))
16//! # });
17//! #
18//! let mut svc = (
19//!     // Layer that removes all response headers with the prefix `x-foo`.
20//!     RemoveResponseHeaderLayer::prefix("x-foo"),
21//! ).into_layer(http_client);
22//!
23//! let request = Request::new(Body::empty());
24//!
25//! let response = svc.serve(request).await?;
26//! #
27//! # Ok(())
28//! # }
29//! ```
30
31use crate::{HeaderName, Request, Response};
32use rama_core::{Layer, Service};
33use rama_utils::macros::define_inner_service_accessors;
34use rama_utils::str::smol_str::SmolStr;
35
36#[derive(Debug, Clone)]
37/// Layer that applies [`RemoveResponseHeader`] which removes response headers.
38///
39/// See [`RemoveResponseHeader`] for more details.
40pub struct RemoveResponseHeaderLayer {
41    mode: RemoveResponseHeaderMode,
42}
43
44#[derive(Debug, Clone)]
45enum RemoveResponseHeaderMode {
46    Prefix(SmolStr),
47    Exact(HeaderName),
48    Hop,
49    Sensitive,
50}
51
52impl RemoveResponseHeaderLayer {
53    /// Create a new [`RemoveResponseHeaderLayer`].
54    ///
55    /// Removes response headers by prefix.
56    pub fn prefix(prefix: impl Into<SmolStr>) -> Self {
57        Self {
58            mode: RemoveResponseHeaderMode::Prefix(prefix.into()),
59        }
60    }
61
62    /// Create a new [`RemoveResponseHeaderLayer`].
63    ///
64    /// Removes the response header with the exact name.
65    pub fn exact(header: HeaderName) -> Self {
66        Self {
67            mode: RemoveResponseHeaderMode::Exact(header),
68        }
69    }
70
71    /// Create a new [`RemoveResponseHeaderLayer`].
72    ///
73    /// Removes all hop-by-hop request headers as specified in [RFC 9110](https://github.com/plabayo/rama/blob/main/rama-http-core/specifications/rfc9110.txt#section-7.6.1).
74    #[must_use]
75    pub fn hop_by_hop() -> Self {
76        Self {
77            mode: RemoveResponseHeaderMode::Hop,
78        }
79    }
80
81    /// Create a new [`RemoveResponseHeaderLayer`].
82    ///
83    /// Removes all sensitive response headers.
84    #[must_use]
85    pub fn sensitive() -> Self {
86        Self {
87            mode: RemoveResponseHeaderMode::Sensitive,
88        }
89    }
90}
91
92impl<S> Layer<S> for RemoveResponseHeaderLayer {
93    type Service = RemoveResponseHeader<S>;
94
95    fn layer(&self, inner: S) -> Self::Service {
96        RemoveResponseHeader {
97            inner,
98            mode: self.mode.clone(),
99        }
100    }
101
102    fn into_layer(self, inner: S) -> Self::Service {
103        RemoveResponseHeader {
104            inner,
105            mode: self.mode,
106        }
107    }
108}
109
110/// Middleware that removes response headers from a request.
111#[derive(Debug, Clone)]
112pub struct RemoveResponseHeader<S> {
113    inner: S,
114    mode: RemoveResponseHeaderMode,
115}
116
117impl<S> RemoveResponseHeader<S> {
118    /// Create a new [`RemoveResponseHeader`].
119    ///
120    /// Removes response headers by prefix.
121    pub fn prefix(prefix: impl Into<SmolStr>, inner: S) -> Self {
122        RemoveResponseHeaderLayer::prefix(prefix.into()).into_layer(inner)
123    }
124
125    /// Create a new [`RemoveResponseHeader`].
126    ///
127    /// Removes the response header with the exact name.
128    pub fn exact(header: HeaderName, inner: S) -> Self {
129        RemoveResponseHeaderLayer::exact(header).into_layer(inner)
130    }
131
132    /// Create a new [`RemoveResponseHeader`].
133    ///
134    /// Removes all hop-by-hop request headers as specified in [RFC 9110](https://github.com/plabayo/rama/blob/main/rama-http-core/specifications/rfc9110.txt#section-7.6.1).
135    pub fn hop_by_hop(inner: S) -> Self {
136        RemoveResponseHeaderLayer::hop_by_hop().into_layer(inner)
137    }
138
139    /// Create a new [`RemoveResponseHeader`].
140    ///
141    /// Removes all sensitive response headers.
142    pub fn sensitive(inner: S) -> Self {
143        RemoveResponseHeaderLayer::sensitive().into_layer(inner)
144    }
145
146    define_inner_service_accessors!();
147}
148
149impl<ReqBody, ResBody, S> Service<Request<ReqBody>> for RemoveResponseHeader<S>
150where
151    ReqBody: Send + 'static,
152    ResBody: Send + 'static,
153    S: Service<Request<ReqBody>, Output = Response<ResBody>>,
154{
155    type Output = S::Output;
156    type Error = S::Error;
157
158    async fn serve(&self, req: Request<ReqBody>) -> Result<Self::Output, Self::Error> {
159        let mut resp = self.inner.serve(req).await?;
160        match &self.mode {
161            RemoveResponseHeaderMode::Hop => {
162                super::remove_hop_by_hop_response_headers(resp.headers_mut())
163            }
164            RemoveResponseHeaderMode::Sensitive => {
165                super::remove_sensitive_response_headers(resp.headers_mut())
166            }
167            RemoveResponseHeaderMode::Prefix(prefix) => {
168                super::remove_headers_by_prefix(resp.headers_mut(), prefix)
169            }
170            RemoveResponseHeaderMode::Exact(header) => {
171                super::remove_headers_by_exact_name(resp.headers_mut(), header)
172            }
173        }
174        Ok(resp)
175    }
176}
177
178#[cfg(test)]
179mod test {
180    use super::*;
181    use crate::{Body, Response};
182    use rama_core::{Layer, Service, service::service_fn};
183    use std::convert::Infallible;
184
185    #[tokio::test]
186    async fn remove_response_header_prefix() {
187        let svc = RemoveResponseHeaderLayer::prefix("x-foo").into_layer(service_fn(
188            async |_req: Request| {
189                Ok::<_, Infallible>(
190                    Response::builder()
191                        .header("x-foo-bar", "baz")
192                        .header("foo", "bar")
193                        .body(Body::empty())
194                        .unwrap(),
195                )
196            },
197        ));
198        let req = Request::builder().body(Body::empty()).unwrap();
199        let res = svc.serve(req).await.unwrap();
200        assert!(res.headers().get("x-foo-bar").is_none());
201        assert_eq!(
202            res.headers().get("foo").map(|v| v.to_str().unwrap()),
203            Some("bar")
204        );
205    }
206
207    #[tokio::test]
208    async fn remove_response_header_exact() {
209        let svc = RemoveResponseHeaderLayer::exact(HeaderName::from_static("foo")).into_layer(
210            service_fn(async |_req: Request| {
211                Ok::<_, Infallible>(
212                    Response::builder()
213                        .header("x-foo", "baz")
214                        .header("foo", "bar")
215                        .body(Body::empty())
216                        .unwrap(),
217                )
218            }),
219        );
220        let req = Request::builder().body(Body::empty()).unwrap();
221        let res = svc.serve(req).await.unwrap();
222        assert!(res.headers().get("foo").is_none());
223        assert_eq!(
224            res.headers().get("x-foo").map(|v| v.to_str().unwrap()),
225            Some("baz")
226        );
227    }
228
229    #[tokio::test]
230    async fn remove_response_header_hop_by_hop() {
231        let svc = RemoveResponseHeaderLayer::hop_by_hop().into_layer(service_fn(
232            async |_req: Request| {
233                Ok::<_, Infallible>(
234                    Response::builder()
235                        .header("connection", "close")
236                        .header("keep-alive", "timeout=5")
237                        .header("foo", "bar")
238                        .body(Body::empty())
239                        .unwrap(),
240                )
241            },
242        ));
243        let req = Request::builder().body(Body::empty()).unwrap();
244        let res = svc.serve(req).await.unwrap();
245        assert!(res.headers().get("connection").is_none());
246        assert!(res.headers().get("keep-alive").is_none());
247        assert_eq!(
248            res.headers().get("foo").map(|v| v.to_str().unwrap()),
249            Some("bar")
250        );
251    }
252
253    #[tokio::test]
254    async fn remove_response_header_hop_by_hop_with_headers_in_connect() {
255        let svc = RemoveResponseHeaderLayer::hop_by_hop().into_layer(service_fn(
256            async |_req: Request| {
257                Ok::<_, Infallible>(
258                    Response::builder()
259                        .header("connection", "x-foo, x-bar")
260                        .header("keep-alive", "timeout=5")
261                        .header("x-foo", "1")
262                        .header("foo", "bar")
263                        .body(Body::empty())
264                        .unwrap(),
265                )
266            },
267        ));
268        let req = Request::builder().body(Body::empty()).unwrap();
269        let res = svc.serve(req).await.unwrap();
270        assert!(res.headers().get("connection").is_none());
271        assert!(res.headers().get("x-foo").is_none());
272        assert!(res.headers().get("x-bar").is_none());
273        assert!(res.headers().get("keep-alive").is_none());
274        assert_eq!(
275            res.headers().get("foo").map(|v| v.to_str().unwrap()),
276            Some("bar")
277        );
278    }
279}