use axum::{middleware::Next, response::IntoResponse};
use tower_http::cors::{Any, CorsLayer};
pub fn cors_layer() -> CorsLayer {
tracing::warn!(
"Using permissive CORS settings (allows all origins). \
This is suitable for development only. \
For production, configure cors_origins in server config and use cors_layer_restricted()."
);
CorsLayer::new()
.allow_origin(Any)
.allow_methods(Any)
.allow_headers(Any)
.expose_headers(Any)
}
pub fn cors_layer_restricted(allowed_origins: &[String]) -> CorsLayer {
let origins: Vec<_> = allowed_origins.iter().filter_map(|origin| origin.parse().ok()).collect();
CorsLayer::new()
.allow_origin(origins)
.allow_methods([
axum::http::Method::GET,
axum::http::Method::POST,
axum::http::Method::OPTIONS,
])
.allow_headers([
axum::http::header::CONTENT_TYPE,
axum::http::header::AUTHORIZATION,
])
}
pub async fn security_headers_middleware(
req: axum::extract::Request,
next: Next,
) -> impl IntoResponse {
let mut response = next.run(req).await;
let headers = response.headers_mut();
headers
.entry("X-Content-Type-Options")
.or_insert_with(|| "nosniff".parse().expect("valid header value"));
headers
.entry("X-Frame-Options")
.or_insert_with(|| "DENY".parse().expect("valid header value"));
headers.entry("Strict-Transport-Security").or_insert_with(|| {
"max-age=31536000; includeSubDomains".parse().expect("valid header value")
});
headers
.entry("Referrer-Policy")
.or_insert_with(|| "strict-origin-when-cross-origin".parse().expect("valid header value"));
headers.entry("Content-Security-Policy").or_insert_with(|| {
"default-src 'self'; script-src 'self'; style-src 'self'"
.parse()
.expect("valid header value")
});
headers
.entry("X-XSS-Protection")
.or_insert_with(|| "0".parse().expect("valid header value"));
response
}
#[cfg(test)]
mod security_headers_tests {
#![allow(clippy::unwrap_used)]
use axum::{
Router,
body::Body,
http::{Request, header},
middleware,
response::IntoResponse,
routing::get,
};
use tower::ServiceExt as _;
use super::security_headers_middleware;
#[tokio::test]
async fn sets_all_security_headers() {
async fn ok() -> &'static str {
"ok"
}
let app = Router::new()
.route("/", get(ok))
.layer(middleware::from_fn(security_headers_middleware));
let resp = app
.oneshot(Request::builder().uri("/").body(Body::empty()).unwrap())
.await
.unwrap();
let h = resp.headers();
assert_eq!(h.get("X-Content-Type-Options").unwrap(), "nosniff");
assert_eq!(h.get("X-Frame-Options").unwrap(), "DENY");
assert!(h.contains_key("Strict-Transport-Security"));
assert!(h.contains_key("Referrer-Policy"));
assert!(h.contains_key("Content-Security-Policy"));
assert_eq!(h.get("X-XSS-Protection").unwrap(), "0");
}
#[tokio::test]
async fn preserves_handler_set_csp() {
async fn csp_handler() -> impl IntoResponse {
(
[(header::CONTENT_SECURITY_POLICY, "script-src https://cdn.example.com")],
"html",
)
}
let app = Router::new()
.route("/", get(csp_handler))
.layer(middleware::from_fn(security_headers_middleware));
let resp = app
.oneshot(Request::builder().uri("/").body(Body::empty()).unwrap())
.await
.unwrap();
assert_eq!(
resp.headers().get("Content-Security-Policy").unwrap(),
"script-src https://cdn.example.com"
);
assert_eq!(resp.headers().get("X-Frame-Options").unwrap(), "DENY");
}
}