axum_security/headers/
service.rs1use 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#[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}