1use std::collections::{BTreeMap, HashSet};
22use std::convert::Infallible;
23
24use axum::extract::{FromRequestParts, MatchedPath, Request, State};
25use axum::http::header::{
26 CONTENT_SECURITY_POLICY, REFERRER_POLICY, STRICT_TRANSPORT_SECURITY, X_CONTENT_TYPE_OPTIONS,
27 X_FRAME_OPTIONS,
28};
29use axum::http::request::Parts;
30use axum::http::{HeaderName, HeaderValue, Method};
31use axum::middleware::Next;
32use axum::response::{IntoResponse, Response};
33
34use crate::AppState;
35use crate::config::{Config, CspMode, Environment};
36use crate::routing::RouteInfo;
37
38#[derive(Debug, Clone, Default)]
40pub struct Csp {
41 extra: BTreeMap<String, Vec<String>>,
42}
43
44impl Csp {
45 pub fn allow(&mut self, directive: &str, source: &str) -> &mut Self {
48 self.extra
49 .entry(directive.trim().to_ascii_lowercase())
50 .or_default()
51 .push(source.trim().to_owned());
52 self
53 }
54}
55
56#[derive(Debug, Clone)]
58pub struct CspNonce(pub String);
59
60impl<S: Send + Sync> FromRequestParts<S> for CspNonce {
61 type Rejection = Infallible;
62
63 async fn from_request_parts(parts: &mut Parts, _: &S) -> Result<Self, Infallible> {
64 Ok(parts
65 .extensions
66 .get::<CspNonce>()
67 .cloned()
68 .unwrap_or(CspNonce(String::new())))
69 }
70}
71
72const NONCE: &str = "{nonce}";
73
74#[derive(Debug)]
76pub(crate) struct Security {
77 pub mode: CspMode,
78 policy: Option<String>,
80 hsts: bool,
81 csrf_exempt: HashSet<(String, String)>,
83 webhook_paths: HashSet<String>,
85 pub xsrf_cookie: bool,
87 trusted_hosts: Vec<String>,
89}
90
91impl Security {
92 pub fn new(config: &Config, csp: &Csp, routes: &[RouteInfo]) -> Self {
93 let script = match config.csp {
94 CspMode::Strict => format!("'self' 'nonce-{NONCE}'"),
95 _ => "'self' 'unsafe-inline' 'unsafe-eval'".to_owned(),
96 };
97 let mut directives: Vec<(String, String)> = [
98 ("default-src", "'self'".to_owned()),
99 ("script-src", script),
100 ("style-src", "'self' 'unsafe-inline'".to_owned()),
101 ("img-src", "'self' data: blob: https:".to_owned()),
102 ("font-src", "'self' data: https:".to_owned()),
103 ("connect-src", "'self'".to_owned()),
104 ("frame-ancestors", "'self'".to_owned()),
105 ("base-uri", "'self'".to_owned()),
106 ("object-src", "'none'".to_owned()),
107 ]
108 .into_iter()
109 .map(|(name, value)| (name.to_owned(), value))
110 .collect();
111 let mut extra = csp.extra.clone();
112 for (directive, source) in crate::seo::csp_sources(config) {
113 extra
114 .entry(directive.to_owned())
115 .or_default()
116 .push(source.to_owned());
117 }
118 for (name, sources) in &extra {
119 let sources = sources.join(" ");
120 match directives.iter_mut().find(|(n, _)| n == name) {
121 Some((_, value)) => {
122 value.push(' ');
123 value.push_str(&sources);
124 }
125 None => directives.push((name.clone(), format!("'self' {sources}"))),
128 }
129 }
130 let policy = (config.csp != CspMode::Off).then(|| {
131 directives
132 .iter()
133 .map(|(name, value)| format!("{name} {value}"))
134 .collect::<Vec<_>>()
135 .join("; ")
136 });
137 let csrf_exempt = routes
138 .iter()
139 .filter(|route| route.middleware.iter().any(|m| m == "no-csrf"))
140 .flat_map(|route| {
142 route
143 .method
144 .split('|')
145 .map(|method| (method.to_owned(), route.path.clone()))
146 })
147 .collect();
148 let webhook_paths = routes
149 .iter()
150 .filter(|route| route.middleware.iter().any(|m| m.starts_with("webhook:")))
151 .map(|route| route.path.clone())
152 .collect();
153 Self {
154 webhook_paths,
155 mode: config.csp,
156 policy,
157 hsts: config.env == Environment::Production && config.url.starts_with("https://"),
158 csrf_exempt,
159 xsrf_cookie: false,
160 trusted_hosts: trusted_hosts(config),
161 }
162 }
163
164 pub fn allows_host(&self, host: Option<&str>) -> bool {
166 if self.trusted_hosts.is_empty() {
167 return true;
168 }
169 let Some(host) = host.map(str::to_ascii_lowercase) else {
170 return false;
171 };
172 self.trusted_hosts
173 .iter()
174 .any(|allowed| match allowed.strip_prefix("*.") {
175 Some(parent) => host
176 .strip_suffix(parent)
177 .is_some_and(|sub| sub.len() > 1 && sub.ends_with('.')),
178 None => *allowed == host,
179 })
180 }
181
182 pub fn is_webhook(&self, path: Option<&MatchedPath>) -> bool {
184 path.is_some_and(|path| self.webhook_paths.contains(path.as_str()))
185 }
186
187 pub fn skips_csrf(&self, method: &Method, path: Option<&MatchedPath>) -> bool {
189 let Some(path) = path else {
190 return false;
191 };
192 let path = path.as_str().to_owned();
193 self.csrf_exempt
194 .contains(&(method.as_str().to_owned(), path.clone()))
195 || self.csrf_exempt.contains(&("*".to_owned(), path))
196 }
197}
198
199#[derive(Debug, Clone, Copy)]
201pub(crate) struct WantsEtag;
202
203const ETAG_LIMIT: u64 = 2 * 1024 * 1024;
205
206async fn etag(
212 res: Response,
213 method: &Method,
214 if_none_match: Option<HeaderValue>,
215 nonce: &str,
216) -> Response {
217 use axum::body::HttpBody as _;
218 use axum::http::StatusCode;
219 let wanted = res.extensions().get::<WantsEtag>().is_some()
220 && matches!(*method, Method::GET | Method::HEAD)
221 && res.status() == StatusCode::OK
222 && !res.headers().contains_key(axum::http::header::ETAG)
223 && res
224 .body()
225 .size_hint()
226 .exact()
227 .is_some_and(|size| size <= ETAG_LIMIT);
228 if !wanted {
229 return res;
230 }
231 let (mut parts, body) = res.into_parts();
232 let Ok(bytes) = axum::body::to_bytes(body, ETAG_LIMIT as usize).await else {
233 return (StatusCode::INTERNAL_SERVER_ERROR, "could not read the page").into_response();
234 };
235 let hash = crate::webhook::sha256_hex(without(&bytes, nonce.as_bytes()));
236 let tag = format!("\"{}\"", &hash[..32]);
237 let matches = if_none_match
238 .as_ref()
239 .and_then(|v| v.to_str().ok())
240 .is_some_and(|v| {
241 v.split(',')
242 .map(|t| t.trim().trim_start_matches("W/"))
243 .any(|t| t == tag || t == "*")
244 });
245 if let Ok(value) = HeaderValue::from_str(&tag) {
246 parts.headers.insert(axum::http::header::ETAG, value);
247 }
248 if matches {
249 parts.status = StatusCode::NOT_MODIFIED;
250 parts.headers.remove(axum::http::header::CONTENT_LENGTH);
251 parts.headers.remove(axum::http::header::CONTENT_TYPE);
252 return Response::from_parts(parts, axum::body::Body::empty());
253 }
254 Response::from_parts(parts, axum::body::Body::from(bytes))
255}
256
257fn without<'a>(bytes: &'a [u8], needle: &[u8]) -> std::borrow::Cow<'a, [u8]> {
260 if needle.is_empty() || !bytes.windows(needle.len()).any(|w| w == needle) {
261 return std::borrow::Cow::Borrowed(bytes);
262 }
263 let mut out = Vec::with_capacity(bytes.len());
264 let mut rest = bytes;
265 while let Some(at) = rest.windows(needle.len()).position(|w| w == needle) {
266 out.extend_from_slice(&rest[..at]);
267 rest = &rest[at + needle.len()..];
268 }
269 out.extend_from_slice(rest);
270 std::borrow::Cow::Owned(out)
271}
272
273fn trusted_hosts(config: &Config) -> Vec<String> {
274 let mut hosts = config.trusted_hosts.clone();
275 if !hosts.is_empty()
276 && let Ok(url) = config.url.parse::<axum::http::Uri>()
277 && let Some(host) = url.host()
278 {
279 hosts.push(host.to_ascii_lowercase());
280 }
281 hosts
282}
283
284pub(crate) async fn middleware(
285 State(state): State<AppState>,
286 mut req: Request,
287 next: Next,
288) -> Response {
289 if req.uri().path() != "/health"
293 && !state
294 .security
295 .allows_host(crate::domain::host(&req).as_deref())
296 {
297 return (
298 axum::http::StatusCode::BAD_REQUEST,
299 "This host is not served here.",
300 )
301 .into_response();
302 }
303 let client = crate::client_ip::resolve(&req, &state.config.trusted_proxies);
304 req.extensions_mut().insert(client);
305 let security = &state.security;
306 let nonce = crate::crypto::random_token();
307 req.extensions_mut().insert(CspNonce(nonce.clone()));
308 let wants_json = crate::error::wants_json(req.headers());
309 let method = req.method().clone();
310 let if_none_match = req
311 .headers()
312 .get(axum::http::header::IF_NONE_MATCH)
313 .cloned();
314 let mut res = next.run(req).await;
315 res = etag(res, &method, if_none_match, &nonce).await;
316 if wants_json && let Some(page) = res.extensions_mut().remove::<crate::error::ErrorPage>() {
318 res = page.json(state.config.debug);
319 }
320
321 let not_modified = res.status() == axum::http::StatusCode::NOT_MODIFIED;
322 let headers = res.headers_mut();
323 let mut set = |name: HeaderName, value: &str| {
324 if !headers.contains_key(&name)
325 && let Ok(value) = HeaderValue::from_str(value)
326 {
327 headers.insert(name, value);
328 }
329 };
330 set(X_CONTENT_TYPE_OPTIONS, "nosniff");
331 set(REFERRER_POLICY, "strict-origin-when-cross-origin");
332 set(X_FRAME_OPTIONS, "SAMEORIGIN");
333 if security.hsts {
334 set(STRICT_TRANSPORT_SECURITY, "max-age=31536000");
335 }
336 if let Some(policy) = &security.policy
339 && !not_modified
340 {
341 set(CONTENT_SECURITY_POLICY, &policy.replace(NONCE, &nonce));
342 }
343 res
344}
345
346#[cfg(test)]
347mod tests {
348 use std::pin::Pin;
349 use std::task::{Context, Poll};
350
351 use axum::body::{Body, Bytes, HttpBody};
352 use axum::http::StatusCode;
353 use axum::response::IntoResponse;
354
355 use super::*;
356
357 #[test]
359 fn the_nonce_is_left_out_of_the_hash() {
360 assert_eq!(&*without(b"a-NONCE-b-NONCE", b"NONCE"), b"a--b-");
361 assert_eq!(&*without(b"plain", b"NONCE"), b"plain");
362 assert_eq!(&*without(b"plain", b""), b"plain");
363 }
364
365 struct Broken;
367
368 impl HttpBody for Broken {
369 type Data = Bytes;
370 type Error = std::io::Error;
371
372 fn poll_frame(
373 self: Pin<&mut Self>,
374 _: &mut Context<'_>,
375 ) -> Poll<Option<Result<http_body::Frame<Bytes>, std::io::Error>>> {
376 Poll::Ready(Some(Err(std::io::Error::other("disk went away"))))
377 }
378
379 fn size_hint(&self) -> http_body::SizeHint {
380 http_body::SizeHint::with_exact(10)
381 }
382 }
383
384 #[tokio::test]
385 async fn an_etag_page_that_cant_be_read_is_a_500() {
386 let mut res = (StatusCode::OK, Body::new(Broken)).into_response();
387 res.extensions_mut().insert(WantsEtag);
388 let res = etag(res, &Method::GET, None, "").await;
389 assert_eq!(res.status(), StatusCode::INTERNAL_SERVER_ERROR);
390 }
391}