tower_http/set_header/response/
multiple_headers.rs1use http::{Request, Response};
6use pin_project_lite::pin_project;
7use std::{
8 fmt,
9 future::Future,
10 pin::Pin,
11 task::{ready, Context, Poll},
12};
13use tower_layer::Layer;
14use tower_service::Service;
15
16use crate::set_header::{HeaderInsertionConfig, HeaderMetadata, InsertHeaderMode};
17
18pub struct SetMultipleResponseHeadersLayer<M> {
22 headers: Vec<HeaderInsertionConfig<M>>,
23}
24
25impl<M> Clone for SetMultipleResponseHeadersLayer<M> {
26 fn clone(&self) -> Self {
27 Self {
28 headers: self.headers.clone(),
29 }
30 }
31}
32
33impl<M> fmt::Debug for SetMultipleResponseHeadersLayer<M> {
34 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
35 f.debug_struct("SetMultipleResponseHeadersLayer")
36 .field("headers", &self.headers)
37 .finish()
38 }
39}
40
41impl<M> SetMultipleResponseHeadersLayer<M> {
42 pub fn overriding(metadata: Vec<HeaderMetadata<M>>) -> Self {
46 let headers: Vec<HeaderInsertionConfig<M>> = metadata
47 .into_iter()
48 .map(|m| m.build_config(InsertHeaderMode::Override))
49 .collect();
50
51 Self::new(headers)
52 }
53
54 pub fn appending(metadata: Vec<HeaderMetadata<M>>) -> Self {
58 let headers: Vec<HeaderInsertionConfig<M>> = metadata
59 .into_iter()
60 .map(|m| m.build_config(InsertHeaderMode::Append))
61 .collect();
62
63 Self::new(headers)
64 }
65
66 pub fn if_not_present(metadata: Vec<HeaderMetadata<M>>) -> Self {
70 let headers: Vec<HeaderInsertionConfig<M>> = metadata
71 .into_iter()
72 .map(|m| m.build_config(InsertHeaderMode::IfNotPresent))
73 .collect();
74
75 Self::new(headers)
76 }
77
78 fn new(headers: Vec<HeaderInsertionConfig<M>>) -> Self {
80 Self { headers }
81 }
82}
83
84impl<S, M> Layer<S> for SetMultipleResponseHeadersLayer<M> {
85 type Service = SetMultipleResponseHeader<S, M>;
86
87 fn layer(&self, inner: S) -> Self::Service {
88 SetMultipleResponseHeader {
89 inner,
90 headers: self.headers.clone(),
91 }
92 }
93}
94
95pub struct SetMultipleResponseHeader<S, M> {
97 inner: S,
98 headers: Vec<HeaderInsertionConfig<M>>,
99}
100
101impl<S, M> Clone for SetMultipleResponseHeader<S, M>
102where
103 S: Clone,
104{
105 fn clone(&self) -> Self {
106 Self {
107 inner: self.inner.clone(),
108 headers: self.headers.clone(),
109 }
110 }
111}
112
113impl<S, M> SetMultipleResponseHeader<S, M> {
114 pub fn overriding(inner: S, metadata: Vec<HeaderMetadata<M>>) -> Self {
118 let headers: Vec<HeaderInsertionConfig<M>> = metadata
119 .into_iter()
120 .map(|m| m.build_config(InsertHeaderMode::Override))
121 .collect();
122
123 Self::new(inner, headers)
124 }
125
126 pub fn appending(inner: S, metadata: Vec<HeaderMetadata<M>>) -> Self {
130 let headers: Vec<HeaderInsertionConfig<M>> = metadata
131 .into_iter()
132 .map(|m| m.build_config(InsertHeaderMode::Append))
133 .collect();
134
135 Self::new(inner, headers)
136 }
137
138 pub fn if_not_present(inner: S, metadata: Vec<HeaderMetadata<M>>) -> Self {
142 let headers: Vec<HeaderInsertionConfig<M>> = metadata
143 .into_iter()
144 .map(|m| m.build_config(InsertHeaderMode::IfNotPresent))
145 .collect();
146
147 Self::new(inner, headers)
148 }
149
150 fn new(inner: S, headers: Vec<HeaderInsertionConfig<M>>) -> Self {
152 Self { inner, headers }
153 }
154
155 define_inner_service_accessors!();
156}
157
158impl<S, M> fmt::Debug for SetMultipleResponseHeader<S, M>
159where
160 S: fmt::Debug,
161{
162 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
163 f.debug_struct("SetMultipleResponseHeader")
164 .field("inner", &self.inner)
165 .field("headers", &self.headers)
166 .finish()
167 }
168}
169
170impl<ReqBody, ResBody, S> Service<Request<ReqBody>>
171 for SetMultipleResponseHeader<S, Response<ResBody>>
172where
173 S: Service<Request<ReqBody>, Response = Response<ResBody>>,
174{
175 type Response = S::Response;
176 type Error = S::Error;
177 type Future = ResponseFuture<S::Future, Response<ResBody>>;
178
179 #[inline]
180 fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
181 self.inner.poll_ready(cx)
182 }
183
184 fn call(&mut self, req: Request<ReqBody>) -> Self::Future {
186 ResponseFuture {
187 future: self.inner.call(req),
188 headers: self.headers.clone(),
189 }
190 }
191}
192
193pin_project! {
194 #[derive(Debug)]
196 pub struct ResponseFuture<F, M> {
197 #[pin]
198 future: F,
199 headers: Vec<HeaderInsertionConfig<M>>,
200 }
201}
202
203impl<F, ResBody, E> Future for ResponseFuture<F, Response<ResBody>>
204where
205 F: Future<Output = Result<Response<ResBody>, E>>,
206{
207 type Output = F::Output;
208
209 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
211 let this = self.project();
212 let mut res = ready!(this.future.poll(cx)?);
213
214 for header in this.headers {
215 header
216 .mode
217 .apply(&header.header_name, &mut res, &mut header.make);
218 }
219
220 Poll::Ready(Ok(res))
221 }
222}
223
224#[cfg(test)]
225mod tests {
226 use super::*;
227 use crate::{
228 set_header::{BoxedMakeHeaderValue, MakeHeaderValue as _},
229 test_helpers::Body,
230 };
231 use http::{header, HeaderName, HeaderValue};
232 use std::convert::Infallible;
233 use tower::{service_fn, ServiceExt};
234
235 #[tokio::test]
236 async fn test_override_mode() {
237 let svc = SetMultipleResponseHeader::overriding(
238 service_fn(|_req: Request<Body>| async {
239 let res = Response::builder()
240 .header(header::CONTENT_TYPE, "good-content")
241 .body(Body::empty())
242 .unwrap();
243 Ok::<_, Infallible>(res)
244 }),
245 vec![(header::CONTENT_TYPE, HeaderValue::from_static("text/html")).into()],
246 );
247
248 let res = svc.oneshot(Request::new(Body::empty())).await.unwrap();
249
250 let mut values = res.headers().get_all(header::CONTENT_TYPE).iter();
251 assert_eq!(values.next().unwrap(), "text/html");
252 assert_eq!(values.next(), None);
253 }
254
255 #[tokio::test]
256 async fn test_append_mode() {
257 let svc = SetMultipleResponseHeader::appending(
258 service_fn(|_req: Request<Body>| async {
259 let res = Response::builder()
260 .header(header::CONTENT_TYPE, "good-content")
261 .body(Body::empty())
262 .unwrap();
263 Ok::<_, Infallible>(res)
264 }),
265 vec![(header::CONTENT_TYPE, HeaderValue::from_static("text/html")).into()],
266 );
267
268 let res = svc.oneshot(Request::new(Body::empty())).await.unwrap();
269
270 let mut values = res.headers().get_all(header::CONTENT_TYPE).iter();
271 assert_eq!(values.next().unwrap(), "good-content");
272 assert_eq!(values.next().unwrap(), "text/html");
273 assert_eq!(values.next(), None);
274 }
275
276 #[tokio::test]
277 async fn test_skip_if_present_mode() {
278 let svc = SetMultipleResponseHeader::if_not_present(
279 service_fn(|_req: Request<Body>| async {
280 let res = Response::builder()
281 .header(header::CONTENT_TYPE, "good-content")
282 .body(Body::empty())
283 .unwrap();
284 Ok::<_, Infallible>(res)
285 }),
286 vec![(header::CONTENT_TYPE, HeaderValue::from_static("text/html")).into()],
287 );
288
289 let res = svc.oneshot(Request::new(Body::empty())).await.unwrap();
290
291 let mut values = res.headers().get_all(header::CONTENT_TYPE).iter();
292 assert_eq!(values.next().unwrap(), "good-content");
293 assert_eq!(values.next(), None);
294 }
295
296 #[tokio::test]
297 async fn test_skip_if_present_mode_when_not_present() {
298 let svc = SetMultipleResponseHeader::if_not_present(
299 service_fn(|_req: Request<Body>| async {
300 let res = Response::builder().body(Body::empty()).unwrap();
301 Ok::<_, Infallible>(res)
302 }),
303 vec![(header::CONTENT_TYPE, HeaderValue::from_static("text/html")).into()],
304 );
305
306 let res = svc.oneshot(Request::new(Body::empty())).await.unwrap();
307
308 let mut values = res.headers().get_all(header::CONTENT_TYPE).iter();
309 assert_eq!(values.next().unwrap(), "text/html");
310 assert_eq!(values.next(), None);
311 }
312
313 #[test]
314 fn test_tuple_metadata_impl() {
315 let tuple: (HeaderName, HeaderValue) =
316 (header::CONTENT_TYPE, HeaderValue::from_static("foo"));
317 let meta: HeaderMetadata<HeaderValue> = tuple.into();
318 assert_eq!(meta.header_name, header::CONTENT_TYPE);
319 let mut make = meta.make.clone();
321 assert_eq!(
322 make.make_header_value(&HeaderValue::from_static("foo")),
323 Some(HeaderValue::from_static("foo"))
324 );
325 }
326
327 #[test]
328 fn test_convert_to_header_config_struct_and_tuple() {
329 let meta: HeaderMetadata<HeaderValue> = HeaderMetadata::<HeaderValue> {
330 header_name: header::CONTENT_TYPE,
331 make: BoxedMakeHeaderValue::new(HeaderValue::from_static("bar")),
332 };
333 let rh = meta.build_config(crate::set_header::InsertHeaderMode::Override);
334 assert_eq!(rh.header_name, header::CONTENT_TYPE);
335 let mut make = rh.make.clone();
336 assert_eq!(
337 make.make_header_value(&HeaderValue::from_static("bar")),
338 Some(HeaderValue::from_static("bar"))
339 );
340
341 let tuple: (HeaderName, HeaderValue) =
342 (header::CONTENT_TYPE, HeaderValue::from_static("baz"));
343 let meta: HeaderMetadata<HeaderValue> = tuple.into();
344 let rh2 = meta.build_config(crate::set_header::InsertHeaderMode::Override);
345 assert_eq!(rh2.header_name, header::CONTENT_TYPE);
346 let mut make2 = rh2.make.clone();
347 assert_eq!(
348 make2.make_header_value(&HeaderValue::from_static("baz")),
349 Some(HeaderValue::from_static("baz"))
350 );
351 }
352
353 #[test]
354 fn test_debug_impls() {
355 let meta: HeaderMetadata<HeaderValue> =
356 (header::CONTENT_TYPE, HeaderValue::from_static("bar")).into();
357 let rh = meta
358 .clone()
359 .build_config(crate::set_header::InsertHeaderMode::Override);
360 let layer = SetMultipleResponseHeadersLayer::overriding(vec![meta]);
361 let debug_str = format!("{:?}", layer);
362 assert!(debug_str.contains("SetMultipleResponseHeadersLayer"));
363 let debug_rh = format!("{:?}", rh);
364 assert!(debug_rh.contains("HeaderInsertionConfig"));
365
366 let svc = SetMultipleResponseHeader::overriding(
367 tower::service_fn(|_req: Request<Body>| async {
368 Ok::<_, std::convert::Infallible>(Response::new(Body::empty()))
369 }),
370 vec![(header::CONTENT_TYPE, HeaderValue::from_static("foo")).into()]
371 as Vec<HeaderMetadata<HeaderValue>>,
372 );
373 let debug_svc = format!("{:?}", svc);
374 assert!(debug_svc.contains("SetMultipleResponseHeader"));
375 }
376
377 #[tokio::test]
378 async fn test_layer_construction_and_multiple_headers() {
379 let svc = tower::ServiceBuilder::new()
381 .layer(SetMultipleResponseHeadersLayer::overriding(vec![
382 (header::CONTENT_TYPE, HeaderValue::from_static("text/html")).into(),
383 (header::CACHE_CONTROL, HeaderValue::from_static("no-cache")).into(),
384 ]))
385 .service(service_fn(|_req: Request<Body>| async {
386 Ok::<_, Infallible>(Response::new(Body::empty()))
387 }));
388
389 let res = svc.oneshot(Request::new(Body::empty())).await.unwrap();
390 assert_eq!(res.headers()["content-type"], "text/html");
391 assert_eq!(res.headers()["cache-control"], "no-cache");
392 }
393
394 #[tokio::test]
395 async fn test_layer_with_empty_vec() {
396 let svc = tower::ServiceBuilder::new()
397 .layer(SetMultipleResponseHeadersLayer::<Response<Body>>::overriding(vec![]))
398 .service(service_fn(|_req: Request<Body>| async {
399 Ok::<_, Infallible>(Response::new(Body::empty()))
400 }));
401
402 let res = svc.oneshot(Request::new(Body::empty())).await.unwrap();
403 assert_eq!(res.headers().len(), 0);
405 }
406
407 #[tokio::test]
408 async fn test_layer_with_static_and_closure_headers_fixed() {
409 let static_meta = (header::CONTENT_TYPE, HeaderValue::from_static("text/html")).into();
411
412 let closure_meta = (header::X_FRAME_OPTIONS, |_res: &Response<Body>| {
414 Some(HeaderValue::from_static("DENY"))
415 })
416 .into();
417
418 let svc = tower::ServiceBuilder::new()
419 .layer(SetMultipleResponseHeadersLayer::overriding(vec![
420 static_meta,
421 closure_meta,
422 ]))
423 .service(service_fn(|_req: Request<Body>| async {
424 Ok::<_, Infallible>(Response::new(Body::empty()))
425 }));
426
427 let res = svc.oneshot(Request::new(Body::empty())).await.unwrap();
428 assert_eq!(res.headers()["content-type"], "text/html");
429 assert_eq!(res.headers()["x-frame-options"], "DENY");
430 }
431
432 #[test]
433 fn test_debug_layer_and_service() {
434 let meta: HeaderMetadata<HeaderValue> =
435 (header::CONTENT_TYPE, HeaderValue::from_static("foo")).into();
436 let layer = SetMultipleResponseHeadersLayer::overriding(vec![meta]);
437 let debug_str = format!("{:?}", layer);
438 assert!(debug_str.contains("SetMultipleResponseHeadersLayer"));
439 }
440
441 #[test]
442 fn test_service_clone() {
443 struct NonCloneBody;
444 let svc = tower::ServiceBuilder::new()
445 .layer(SetMultipleResponseHeadersLayer::<Response<NonCloneBody>>::overriding(vec![]))
446 .check_clone()
447 .service(service_fn(|_: Request<NonCloneBody>| async move {
448 Ok::<_, Infallible>(Response::new(NonCloneBody))
449 }));
450
451 fn check_service_and_clone<T: Service<Request<NonCloneBody>> + Clone>(_: T) {}
452 check_service_and_clone(svc);
453 }
454}