rama_http/layer/remove_header/
response.rs1use 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)]
37pub 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 pub fn prefix(prefix: impl Into<SmolStr>) -> Self {
57 Self {
58 mode: RemoveResponseHeaderMode::Prefix(prefix.into()),
59 }
60 }
61
62 pub fn exact(header: HeaderName) -> Self {
66 Self {
67 mode: RemoveResponseHeaderMode::Exact(header),
68 }
69 }
70
71 #[must_use]
75 pub fn hop_by_hop() -> Self {
76 Self {
77 mode: RemoveResponseHeaderMode::Hop,
78 }
79 }
80
81 #[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#[derive(Debug, Clone)]
112pub struct RemoveResponseHeader<S> {
113 inner: S,
114 mode: RemoveResponseHeaderMode,
115}
116
117impl<S> RemoveResponseHeader<S> {
118 pub fn prefix(prefix: impl Into<SmolStr>, inner: S) -> Self {
122 RemoveResponseHeaderLayer::prefix(prefix.into()).into_layer(inner)
123 }
124
125 pub fn exact(header: HeaderName, inner: S) -> Self {
129 RemoveResponseHeaderLayer::exact(header).into_layer(inner)
130 }
131
132 pub fn hop_by_hop(inner: S) -> Self {
136 RemoveResponseHeaderLayer::hop_by_hop().into_layer(inner)
137 }
138
139 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}