Skip to main content

axum_security/headers/
service.rs

1use std::{
2    collections::HashSet,
3    pin::Pin,
4    sync::Arc,
5    task::{Context, Poll, ready},
6};
7
8use axum::extract::Request;
9use http::Response;
10use pin_project_lite::pin_project;
11use tower::{Layer, Service};
12
13use crate::headers::{SecurityHeader, SecurityHeaders};
14
15impl<S> Layer<S> for SecurityHeaders {
16    type Service = SecurityHeadersLayer<S>;
17
18    fn layer(&self, inner: S) -> Self::Service {
19        SecurityHeadersLayer {
20            inner,
21            headers: self.headers.clone(),
22        }
23    }
24}
25
26/// The [`Service`] created by [`SecurityHeaders`]. You don't need to construct this directly.
27#[derive(Clone)]
28pub struct SecurityHeadersLayer<S> {
29    inner: S,
30    headers: Arc<HashSet<SecurityHeader>>,
31}
32
33impl<IB, OB, S> Service<Request<IB>> for SecurityHeadersLayer<S>
34where
35    S: Service<Request<IB>, Response = Response<OB>>,
36{
37    type Response = Response<OB>;
38
39    type Error = S::Error;
40
41    type Future = InsertHeaders<S::Future>;
42
43    fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
44        self.inner.poll_ready(cx)
45    }
46
47    fn call(&mut self, req: Request<IB>) -> Self::Future {
48        InsertHeaders {
49            future: self.inner.call(req),
50            header: self.headers.clone(),
51        }
52    }
53}
54
55pin_project! {
56    pub struct InsertHeaders<F> {
57        #[pin]
58        future: F,
59        header: Arc<HashSet<SecurityHeader>>
60    }
61}
62
63impl<F, B, E> Future for InsertHeaders<F>
64where
65    F: Future<Output = Result<Response<B>, E>>,
66{
67    type Output = Result<Response<B>, E>;
68
69    fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
70        let this = self.project();
71        let res = ready!(this.future.poll(cx));
72
73        Poll::Ready(res.map(|mut res| {
74            let headers = res.headers_mut();
75
76            for header in this.header.iter() {
77                headers.insert(header.name.clone(), header.value.clone());
78            }
79            res
80        }))
81    }
82}
83
84#[cfg(test)]
85mod headers_service {
86    use std::error::Error;
87
88    use axum::{Router, body::Body};
89    use http::Request;
90    use tower::ServiceExt;
91
92    use crate::headers::{
93        CROSS_ORIGIN_OPENER_POLICY, CrossOriginOpenerPolicy, SecurityHeaders, X_XSS_PROTECTION,
94        XssProtection,
95    };
96
97    #[tokio::test]
98    async fn test() -> Result<(), Box<dyn Error>> {
99        let headers = SecurityHeaders::new().add(XssProtection::ZERO);
100        let router = Router::<()>::new().layer(headers);
101
102        let res = router
103            .oneshot(Request::get("/").body(Body::empty())?)
104            .await
105            .unwrap();
106
107        let header = &res.headers()[X_XSS_PROTECTION];
108        assert!(header == XssProtection::ZERO.header_value);
109        Ok(())
110    }
111
112    #[tokio::test]
113    async fn test_multiple() -> Result<(), Box<dyn Error>> {
114        let headers = SecurityHeaders::new()
115            .add(XssProtection::ZERO)
116            .add(CrossOriginOpenerPolicy::SAME_ORIGIN);
117
118        let router = Router::<()>::new().layer(headers);
119
120        let res = router
121            .oneshot(Request::get("/").body(Body::empty())?)
122            .await
123            .unwrap();
124
125        let header = &res.headers()[X_XSS_PROTECTION];
126        assert!(header == XssProtection::ZERO.header_value);
127
128        let header = &res.headers()[CROSS_ORIGIN_OPENER_POLICY];
129        assert!(header == CrossOriginOpenerPolicy::SAME_ORIGIN.header_value);
130
131        Ok(())
132    }
133}