rama_http/layer/remove_header/
request.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 RemoveRequestHeaderLayer {
41 mode: RemoveRequestHeaderMode,
42}
43
44#[derive(Debug, Clone)]
45enum RemoveRequestHeaderMode {
46 Prefix(SmolStr),
47 Exact(HeaderName),
48 Hop,
49 Sensitive,
50}
51
52impl RemoveRequestHeaderLayer {
53 pub fn prefix(prefix: impl Into<SmolStr>) -> Self {
57 Self {
58 mode: RemoveRequestHeaderMode::Prefix(prefix.into()),
59 }
60 }
61
62 pub fn exact(header: HeaderName) -> Self {
66 Self {
67 mode: RemoveRequestHeaderMode::Exact(header),
68 }
69 }
70
71 #[must_use]
75 pub fn hop_by_hop() -> Self {
76 Self {
77 mode: RemoveRequestHeaderMode::Hop,
78 }
79 }
80
81 #[must_use]
85 pub fn sensitive() -> Self {
86 Self {
87 mode: RemoveRequestHeaderMode::Sensitive,
88 }
89 }
90}
91
92impl<S> Layer<S> for RemoveRequestHeaderLayer {
93 type Service = RemoveRequestHeader<S>;
94
95 fn layer(&self, inner: S) -> Self::Service {
96 RemoveRequestHeader {
97 inner,
98 mode: self.mode.clone(),
99 }
100 }
101
102 fn into_layer(self, inner: S) -> Self::Service {
103 RemoveRequestHeader {
104 inner,
105 mode: self.mode,
106 }
107 }
108}
109
110#[derive(Debug, Clone)]
112pub struct RemoveRequestHeader<S> {
113 inner: S,
114 mode: RemoveRequestHeaderMode,
115}
116
117impl<S> RemoveRequestHeader<S> {
118 pub fn prefix(prefix: impl Into<SmolStr>, inner: S) -> Self {
122 RemoveRequestHeaderLayer::prefix(prefix.into()).into_layer(inner)
123 }
124
125 pub fn exact(header: HeaderName, inner: S) -> Self {
129 RemoveRequestHeaderLayer::exact(header).into_layer(inner)
130 }
131
132 pub fn hop_by_hop(inner: S) -> Self {
136 RemoveRequestHeaderLayer::hop_by_hop().into_layer(inner)
137 }
138
139 pub fn sensitive(inner: S) -> Self {
143 RemoveRequestHeaderLayer::sensitive().into_layer(inner)
144 }
145
146 define_inner_service_accessors!();
147}
148
149impl<ReqBody, ResBody, S> Service<Request<ReqBody>> for RemoveRequestHeader<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 fn serve(
159 &self,
160 mut req: Request<ReqBody>,
161 ) -> impl Future<Output = Result<Self::Output, Self::Error>> + Send + '_ {
162 match &self.mode {
163 RemoveRequestHeaderMode::Hop => {
164 super::remove_hop_by_hop_request_headers(req.headers_mut())
165 }
166 RemoveRequestHeaderMode::Sensitive => {
167 super::remove_sensitive_request_headers(req.headers_mut())
168 }
169 RemoveRequestHeaderMode::Prefix(prefix) => {
170 super::remove_headers_by_prefix(req.headers_mut(), prefix)
171 }
172 RemoveRequestHeaderMode::Exact(header) => {
173 super::remove_headers_by_exact_name(req.headers_mut(), header)
174 }
175 }
176 self.inner.serve(req)
177 }
178}
179
180#[cfg(test)]
181mod test {
182 use super::*;
183 use crate::{Body, Response};
184 use rama_core::{Layer, Service, service::service_fn};
185 use std::convert::Infallible;
186
187 #[tokio::test]
188 async fn remove_request_header_prefix() {
189 let svc = RemoveRequestHeaderLayer::prefix("x-foo").into_layer(service_fn(
190 async |req: Request| {
191 assert!(req.headers().get("x-foo-bar").is_none());
192 assert_eq!(
193 req.headers().get("foo").map(|v| v.to_str().unwrap()),
194 Some("bar")
195 );
196 Ok::<_, Infallible>(Response::new(Body::empty()))
197 },
198 ));
199 let req = Request::builder()
200 .header("x-foo-bar", "baz")
201 .header("foo", "bar")
202 .body(Body::empty())
203 .unwrap();
204 _ = svc.serve(req).await.unwrap();
205 }
206
207 #[tokio::test]
208 async fn remove_request_header_exact() {
209 let svc = RemoveRequestHeaderLayer::exact(HeaderName::from_static("x-foo")).into_layer(
210 service_fn(async |req: Request| {
211 assert!(req.headers().get("x-foo").is_none());
212 assert_eq!(
213 req.headers().get("x-foo-bar").map(|v| v.to_str().unwrap()),
214 Some("baz")
215 );
216 Ok::<_, Infallible>(Response::new(Body::empty()))
217 }),
218 );
219 let req = Request::builder()
220 .header("x-foo", "baz")
221 .header("x-foo-bar", "baz")
222 .body(Body::empty())
223 .unwrap();
224 _ = svc.serve(req).await.unwrap();
225 }
226
227 #[tokio::test]
228 async fn remove_request_header_hop_by_hop() {
229 let svc =
230 RemoveRequestHeaderLayer::hop_by_hop().into_layer(service_fn(async |req: Request| {
231 assert!(req.headers().get("connection").is_none());
232 assert_eq!(
233 req.headers().get("foo").map(|v| v.to_str().unwrap()),
234 Some("bar")
235 );
236 Ok::<_, Infallible>(Response::new(Body::empty()))
237 }));
238 let req = Request::builder()
239 .header("connection", "close")
240 .header("foo", "bar")
241 .body(Body::empty())
242 .unwrap();
243 _ = svc.serve(req).await.unwrap();
244 }
245
246 #[tokio::test]
247 async fn remove_request_header_hop_by_hop_with_connection_list() {
248 let svc =
249 RemoveRequestHeaderLayer::hop_by_hop().into_layer(service_fn(async |req: Request| {
250 assert!(req.headers().get("connection").is_none());
251 assert!(req.headers().get("x-foo").is_none());
252 assert!(req.headers().get("x-bar").is_none());
253 assert!(req.headers().get("connection").is_none());
254 assert_eq!(
255 req.headers().get("foo").map(|v| v.to_str().unwrap()),
256 Some("bar")
257 );
258 Ok::<_, Infallible>(Response::new(Body::empty()))
259 }));
260 let req = Request::builder()
261 .header("connection", "x-foo, x-bar")
262 .header("x-foo", "1")
263 .header("x-real-ip", "1.2.3.4")
264 .header("foo", "bar")
265 .body(Body::empty())
266 .unwrap();
267 _ = svc.serve(req).await.unwrap();
268 }
269}