rama_http/layer/
sensitive_headers.rs1use crate::{HeaderName, Request, Response, header};
42use rama_core::{Layer, Service};
43use rama_utils::macros::define_inner_service_accessors;
44use std::sync::Arc;
45
46#[derive(Clone, Debug)]
54pub struct SetSensitiveHeadersLayer {
55 headers: Arc<[HeaderName]>,
56}
57
58impl SetSensitiveHeadersLayer {
59 pub fn new<I>(headers: I) -> Self
61 where
62 I: IntoIterator<Item = HeaderName>,
63 {
64 let headers = headers.into_iter().collect::<Vec<_>>();
65 Self::from_shared(headers.into())
66 }
67
68 pub fn from_shared(headers: Arc<[HeaderName]>) -> Self {
70 Self { headers }
71 }
72}
73
74impl<S> Layer<S> for SetSensitiveHeadersLayer {
75 type Service = SetSensitiveHeaders<S>;
76
77 fn layer(&self, inner: S) -> Self::Service {
78 SetSensitiveRequestHeaders::from_shared(
79 SetSensitiveResponseHeaders::from_shared(inner, self.headers.clone()),
80 self.headers.clone(),
81 )
82 }
83
84 fn into_layer(self, inner: S) -> Self::Service {
85 SetSensitiveRequestHeaders::from_shared(
86 SetSensitiveResponseHeaders::from_shared(inner, self.headers.clone()),
87 self.headers,
88 )
89 }
90}
91
92pub type SetSensitiveHeaders<S> = SetSensitiveRequestHeaders<SetSensitiveResponseHeaders<S>>;
98
99#[derive(Clone, Debug)]
107pub struct SetSensitiveRequestHeadersLayer {
108 headers: Arc<[HeaderName]>,
109}
110
111impl SetSensitiveRequestHeadersLayer {
112 pub fn new<I>(headers: I) -> Self
114 where
115 I: IntoIterator<Item = HeaderName>,
116 {
117 let headers = headers.into_iter().collect::<Vec<_>>();
118 Self::from_shared(headers.into())
119 }
120
121 pub fn from_shared(headers: Arc<[HeaderName]>) -> Self {
123 Self { headers }
124 }
125}
126
127impl<S> Layer<S> for SetSensitiveRequestHeadersLayer {
128 type Service = SetSensitiveRequestHeaders<S>;
129
130 fn layer(&self, inner: S) -> Self::Service {
131 SetSensitiveRequestHeaders {
132 inner,
133 headers: self.headers.clone(),
134 }
135 }
136
137 fn into_layer(self, inner: S) -> Self::Service {
138 SetSensitiveRequestHeaders {
139 inner,
140 headers: self.headers,
141 }
142 }
143}
144
145#[derive(Clone, Debug)]
151pub struct SetSensitiveRequestHeaders<S> {
152 inner: S,
153 headers: Arc<[HeaderName]>,
154}
155
156impl<S> SetSensitiveRequestHeaders<S> {
157 pub fn new<I>(inner: S, headers: I) -> Self
159 where
160 I: IntoIterator<Item = HeaderName>,
161 {
162 let headers = headers.into_iter().collect::<Vec<_>>();
163 Self::from_shared(inner, headers.into())
164 }
165
166 pub fn from_shared(inner: S, headers: Arc<[HeaderName]>) -> Self {
168 Self { inner, headers }
169 }
170
171 define_inner_service_accessors!();
172}
173
174impl<ReqBody, ResBody, S> Service<Request<ReqBody>> for SetSensitiveRequestHeaders<S>
175where
176 S: Service<Request<ReqBody>, Output = Response<ResBody>>,
177 ReqBody: Send + 'static,
178 ResBody: Send + 'static,
179{
180 type Output = S::Output;
181 type Error = S::Error;
182
183 async fn serve(&self, mut req: Request<ReqBody>) -> Result<Self::Output, Self::Error> {
184 let headers = req.headers_mut();
185 for header in &*self.headers {
186 if let header::Entry::Occupied(mut entry) = headers.entry(header) {
187 for value in entry.iter_mut() {
188 value.set_sensitive(true);
189 }
190 }
191 }
192
193 self.inner.serve(req).await
194 }
195}
196
197#[derive(Clone, Debug)]
205pub struct SetSensitiveResponseHeadersLayer {
206 headers: Arc<[HeaderName]>,
207}
208
209impl SetSensitiveResponseHeadersLayer {
210 pub fn new<I>(headers: I) -> Self
212 where
213 I: IntoIterator<Item = HeaderName>,
214 {
215 let headers = headers.into_iter().collect::<Vec<_>>();
216 Self::from_shared(headers.into())
217 }
218
219 pub fn from_shared(headers: Arc<[HeaderName]>) -> Self {
221 Self { headers }
222 }
223}
224
225impl<S> Layer<S> for SetSensitiveResponseHeadersLayer {
226 type Service = SetSensitiveResponseHeaders<S>;
227
228 fn layer(&self, inner: S) -> Self::Service {
229 SetSensitiveResponseHeaders {
230 inner,
231 headers: self.headers.clone(),
232 }
233 }
234
235 fn into_layer(self, inner: S) -> Self::Service {
236 SetSensitiveResponseHeaders {
237 inner,
238 headers: self.headers,
239 }
240 }
241}
242
243#[derive(Clone, Debug)]
249pub struct SetSensitiveResponseHeaders<S> {
250 inner: S,
251 headers: Arc<[HeaderName]>,
252}
253
254impl<S> SetSensitiveResponseHeaders<S> {
255 pub fn new<I>(inner: S, headers: I) -> Self
257 where
258 I: IntoIterator<Item = HeaderName>,
259 {
260 let headers = headers.into_iter().collect::<Vec<_>>();
261 Self::from_shared(inner, headers.into())
262 }
263
264 pub fn from_shared(inner: S, headers: Arc<[HeaderName]>) -> Self {
266 Self { inner, headers }
267 }
268
269 define_inner_service_accessors!();
270}
271
272impl<ReqBody, ResBody, S> Service<Request<ReqBody>> for SetSensitiveResponseHeaders<S>
273where
274 S: Service<Request<ReqBody>, Output = Response<ResBody>>,
275 ReqBody: Send + 'static,
276 ResBody: Send + 'static,
277{
278 type Output = S::Output;
279 type Error = S::Error;
280
281 async fn serve(&self, req: Request<ReqBody>) -> Result<Self::Output, Self::Error> {
282 let mut res = self.inner.serve(req).await?;
283
284 let headers = res.headers_mut();
285 for header in self.headers.iter() {
286 if let header::Entry::Occupied(mut entry) = headers.entry(header) {
287 for value in entry.iter_mut() {
288 value.set_sensitive(true);
289 }
290 }
291 }
292
293 Ok(res)
294 }
295}
296
297#[cfg(test)]
298mod tests {
299 use super::*;
300 use crate::{HeaderValue, Request, Response, header};
301 use rama_core::service::service_fn;
302
303 #[tokio::test]
304 async fn multiple_value_header() {
305 async fn response_set_cookie(req: Request<()>) -> Result<Response<()>, ()> {
306 let mut iter = req.headers().get_all(header::COOKIE).iter().peekable();
307
308 assert!(iter.peek().is_some());
309
310 for value in iter {
311 assert!(value.is_sensitive())
312 }
313
314 let mut resp = Response::new(());
315 resp.headers_mut()
316 .append(header::CONTENT_TYPE, HeaderValue::from_static("text/html"));
317 resp.headers_mut()
318 .append(header::SET_COOKIE, HeaderValue::from_static("cookie-1"));
319 resp.headers_mut()
320 .append(header::SET_COOKIE, HeaderValue::from_static("cookie-2"));
321 resp.headers_mut()
322 .append(header::SET_COOKIE, HeaderValue::from_static("cookie-3"));
323 Ok(resp)
324 }
325
326 let service = (
327 SetSensitiveRequestHeadersLayer::new(vec![header::COOKIE]),
328 SetSensitiveResponseHeadersLayer::new(vec![header::SET_COOKIE]),
329 )
330 .into_layer(service_fn(response_set_cookie));
331
332 let mut req = Request::new(());
333 req.headers_mut()
334 .append(header::COOKIE, HeaderValue::from_static("cookie+1"));
335 req.headers_mut()
336 .append(header::COOKIE, HeaderValue::from_static("cookie+2"));
337
338 let resp = service.serve(req).await.unwrap();
339
340 assert!(
341 !resp
342 .headers()
343 .get(header::CONTENT_TYPE)
344 .unwrap()
345 .is_sensitive()
346 );
347
348 let mut iter = resp.headers().get_all(header::SET_COOKIE).iter().peekable();
349
350 assert!(iter.peek().is_some());
351
352 for value in iter {
353 assert!(value.is_sensitive())
354 }
355 }
356}