Skip to main content

tower_http/set_header/response/
multiple_headers.rs

1//! Set multiple headers on the response.
2//!
3//! See the root [`crate::set_header::response`] module for full documentation and usage examples.
4//!
5use 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
18/// Layer that applies [`SetMultipleResponseHeader`] which adds multiple response headers.
19///
20/// See [`SetMultipleResponseHeader`] for more details.
21pub 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    /// Create a new [`SetMultipleResponseHeadersLayer`] that overrides any existing values for the same header.
43    ///
44    /// If any previous value exists for the same header, it is removed and replaced with the new matching header value.
45    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    /// Create a new [`SetMultipleResponseHeadersLayer`] that appends header values.
55    ///
56    /// The new header is always added, preserving any existing values. If previous values exist, the header will have multiple values.
57    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    /// Create a new [`SetMultipleResponseHeadersLayer`] that only inserts if the header is not already present.
67    ///
68    /// If a previous value exists for the header, the new value is not inserted.
69    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    /// Internal constructor for a new [`SetMultipleResponseHeadersLayer`] from a list of headers.
79    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
95/// Middleware that sets multiple headers on the response.
96pub 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    /// Create a new [`SetMultipleResponseHeader`] that overrides any existing values for the same header.
115    ///
116    /// If a previous value exists for the same header, it is removed and replaced with the new header value.
117    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    /// Create a new [`SetMultipleResponseHeader`] that appends header values.
127    ///
128    /// The new header is always added, preserving any existing values. If previous values exist, the header will have multiple values.
129    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    /// Create a new [`SetMultipleResponseHeader`] that only inserts if the header is not already present.
139    ///
140    /// If a previous value exists for the header, the new value is not inserted.
141    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    /// Internal constructor for a new [`SetMultipleResponseHeader`] from an inner service and a list of headers.
151    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    /// Call the inner service and apply all configured headers to the response.
185    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    /// Response future for [`SetMultipleResponseHeader`].
195    #[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    /// Polls the inner future and applies all configured headers to the response before returning it.
210    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        // Check that the header value is correct by making a header value from meta.make
320        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        // Multiple different headers in the same vec
380        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        // No headers should be set
404        assert_eq!(res.headers().len(), 0);
405    }
406
407    #[tokio::test]
408    async fn test_layer_with_static_and_closure_headers_fixed() {
409        // Wrap the static value
410        let static_meta = (header::CONTENT_TYPE, HeaderValue::from_static("text/html")).into();
411
412        // Wrap the closure
413        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}