Skip to main content

axum_security/headers/
hsts.rs

1use std::time::Duration;
2
3use axum::http::{HeaderValue, header::STRICT_TRANSPORT_SECURITY};
4use http::HeaderName;
5use tower::Layer;
6
7use crate::{headers::IntoSecurityHeader, utils::headers::InsertHeadersService};
8
9// in seconds
10const PRELOAD_MIN_MAX_AGE: u64 = 365 * 24 * 60 * 60;
11
12/// `Strict-Transport-Security` header.
13///
14/// Build one with [`StrictTransportSecurity::builder`]. The result implements
15/// [`Layer`] and [`IntoSecurityHeader`](super::IntoSecurityHeader).
16///
17/// # Example
18///
19/// ```rust
20/// use axum_security::headers::StrictTransportSecurity;
21///
22/// let hsts = StrictTransportSecurity::builder()
23///     .max_age_years(1)
24///     .include_subdomains()
25///     .preload()
26///     .build();
27/// ```
28#[derive(Clone)]
29pub struct StrictTransportSecurity {
30    header_value: HeaderValue,
31}
32
33impl StrictTransportSecurity {
34    /// Create a builder for this header.
35    pub fn builder() -> HstsBuilder {
36        HstsBuilder {
37            max_age_secs: None,
38            include_subdomains: false,
39            preload: false,
40        }
41    }
42}
43
44/// Builder for [`StrictTransportSecurity`].
45///
46/// Set a max-age with one of the `max_age_*` methods, then call
47/// [`build`](HstsBuilder::build) (panics if no max-age is set) or
48/// [`try_build`](HstsBuilder::try_build).
49pub struct HstsBuilder {
50    /// Max age in seconds.
51    max_age_secs: Option<u64>,
52    include_subdomains: bool,
53    preload: bool,
54}
55
56impl HstsBuilder {
57    pub fn max_age(mut self, duration: Duration) -> Self {
58        self.max_age_secs = Some(duration.as_secs());
59        self
60    }
61
62    pub fn max_age_seconds(mut self, max_age: u64) -> Self {
63        self.max_age_secs = Some(max_age);
64        self
65    }
66
67    /// 24h in a day
68    pub fn max_age_days(self, max_age: u64) -> Self {
69        self.max_age_seconds(max_age * 24 * 60 * 60)
70    }
71
72    /// 365 days in a year
73    pub fn max_age_years(self, max_age: u64) -> Self {
74        self.max_age_days(max_age * 365)
75    }
76
77    pub fn include_subdomains(mut self) -> Self {
78        self.include_subdomains = true;
79        self
80    }
81
82    pub fn preload(mut self) -> Self {
83        self.preload = true;
84        self
85    }
86
87    pub fn try_build(self) -> Result<StrictTransportSecurity, StrictTransportSecurityBuilderError> {
88        let Some(max_age) = self.max_age_secs else {
89            return Err(StrictTransportSecurityBuilderError::NoMaxAge);
90        };
91
92        let mut header = format!("max-age={max_age}");
93
94        if self.include_subdomains {
95            header.push_str("; includeSubDomains");
96        }
97
98        if self.preload {
99            if max_age < PRELOAD_MIN_MAX_AGE {
100                return Err(StrictTransportSecurityBuilderError::InvalidMaxAge);
101            } else if !self.include_subdomains {
102                return Err(StrictTransportSecurityBuilderError::IncludeSubdomainsRequired);
103            }
104
105            header.push_str("; preload");
106        }
107
108        let header_value =
109            HeaderValue::from_str(&header).expect("Hsts header does not contain invalid bytes");
110
111        Ok(StrictTransportSecurity { header_value })
112    }
113
114    pub fn build(self) -> StrictTransportSecurity {
115        self.try_build().unwrap()
116    }
117}
118
119/// Error returned when building a [`StrictTransportSecurity`] header with invalid options.
120#[derive(Debug)]
121pub enum StrictTransportSecurityBuilderError {
122    /// No max-age was set.
123    NoMaxAge,
124    /// `preload` requires a max-age of at least 1 year.
125    InvalidMaxAge,
126    /// `preload` requires `include_subdomains`.
127    IncludeSubdomainsRequired,
128}
129
130impl<S> Layer<S> for StrictTransportSecurity {
131    type Service = InsertHeadersService<S>;
132
133    fn layer(&self, inner: S) -> Self::Service {
134        InsertHeadersService {
135            inner,
136            header_name: STRICT_TRANSPORT_SECURITY,
137            header_value: self.header_value.clone(),
138        }
139    }
140}
141
142impl IntoSecurityHeader for StrictTransportSecurity {
143    fn into_header(self) -> (HeaderName, HeaderValue) {
144        (STRICT_TRANSPORT_SECURITY, self.header_value)
145    }
146}
147
148impl IntoSecurityHeader for HstsBuilder {
149    fn into_header(self) -> (HeaderName, HeaderValue) {
150        self.build().into_header()
151    }
152}
153
154#[cfg(test)]
155mod hsts_tests {
156    use axum::{Router, body::Body, extract::Request, http::header::STRICT_TRANSPORT_SECURITY};
157
158    use crate::headers::{StrictTransportSecurity, StrictTransportSecurityBuilderError};
159    use tower::ServiceExt;
160
161    #[test]
162    fn builder() {
163        let hsts = StrictTransportSecurity::builder().try_build();
164        assert!(matches!(
165            hsts,
166            Err(StrictTransportSecurityBuilderError::NoMaxAge)
167        ));
168
169        let hsts = StrictTransportSecurity::builder()
170            .max_age_days(364)
171            .preload()
172            .try_build();
173        assert!(matches!(
174            hsts,
175            Err(StrictTransportSecurityBuilderError::InvalidMaxAge)
176        ));
177
178        let hsts = StrictTransportSecurity::builder()
179            .max_age_years(1)
180            .preload()
181            .try_build();
182        assert!(matches!(
183            hsts,
184            Err(StrictTransportSecurityBuilderError::IncludeSubdomainsRequired)
185        ));
186    }
187
188    #[test]
189    fn header() {
190        let hsts = StrictTransportSecurity::builder()
191            .max_age_seconds(1)
192            .build();
193        assert!(hsts.header_value == "max-age=1");
194
195        let hsts = StrictTransportSecurity::builder()
196            .max_age_seconds(1)
197            .include_subdomains()
198            .build();
199        assert!(hsts.header_value == "max-age=1; includeSubDomains");
200
201        let hsts = StrictTransportSecurity::builder()
202            .max_age_years(1)
203            .include_subdomains()
204            .preload()
205            .build();
206        assert!(hsts.header_value == "max-age=31536000; includeSubDomains; preload");
207    }
208
209    #[tokio::test]
210    async fn basic() {
211        let hsts = StrictTransportSecurity::builder()
212            .max_age_years(1)
213            .include_subdomains()
214            .preload()
215            .build();
216
217        let router = Router::<()>::new().layer(hsts);
218
219        let res = router
220            .oneshot(Request::get("/").body(Body::empty()).unwrap())
221            .await
222            .unwrap();
223
224        assert_eq!(
225            res.headers()[STRICT_TRANSPORT_SECURITY],
226            "max-age=31536000; includeSubDomains; preload"
227        );
228    }
229}