use std::sync::Arc;
use salvo::http::HeaderValue;
use salvo::http::header::{
CONTENT_SECURITY_POLICY, HeaderName, REFERRER_POLICY, STRICT_TRANSPORT_SECURITY, X_CONTENT_TYPE_OPTIONS,
X_FRAME_OPTIONS,
};
use salvo::prelude::*;
use serde::Deserialize;
use crate::setup::{GenericServerState, StartupVariant};
#[derive(Clone, Debug, Deserialize)]
#[serde(default)]
pub struct SecurityHeadersOptions {
pub enabled: bool,
pub hsts: Option<String>,
pub x_content_type_options: Option<String>,
pub x_frame_options: Option<String>,
pub referrer_policy: Option<String>,
pub content_security_policy: Option<String>,
pub permissions_policy: Option<String>,
pub cross_origin_opener_policy: Option<String>,
}
impl Default for SecurityHeadersOptions {
fn default() -> Self {
Self {
enabled: true,
hsts: Some("max-age=31536000; includeSubDomains".to_string()),
x_content_type_options: Some("nosniff".to_string()),
x_frame_options: Some("SAMEORIGIN".to_string()),
referrer_policy: Some("strict-origin-when-cross-origin".to_string()),
content_security_policy: None,
permissions_policy: None,
cross_origin_opener_policy: None,
}
}
}
pub struct SecurityHeaders {
always: Arc<Vec<(HeaderName, HeaderValue)>>,
hsts: Option<(HeaderName, HeaderValue)>,
}
fn parse(name: HeaderName, value: &Option<String>) -> Option<(HeaderName, HeaderValue)> {
let raw = value.as_deref()?;
match HeaderValue::from_str(raw) {
Ok(v) => Some((name, v)),
Err(e) => {
tracing::warn!(error = %e, header = %name, raw = %raw, "ignoring invalid security header value");
None
}
}
}
impl SecurityHeaders {
pub fn new(options: &SecurityHeadersOptions) -> Self {
let mut always = Vec::new();
always.extend(parse(X_CONTENT_TYPE_OPTIONS, &options.x_content_type_options));
always.extend(parse(X_FRAME_OPTIONS, &options.x_frame_options));
always.extend(parse(REFERRER_POLICY, &options.referrer_policy));
always.extend(parse(CONTENT_SECURITY_POLICY, &options.content_security_policy));
always.extend(parse(
HeaderName::from_static("permissions-policy"),
&options.permissions_policy,
));
always.extend(parse(
HeaderName::from_static("cross-origin-opener-policy"),
&options.cross_origin_opener_policy,
));
let hsts = parse(STRICT_TRANSPORT_SECURITY, &options.hsts);
Self {
always: Arc::new(always),
hsts,
}
}
fn apply(&self, hsts_allowed: bool, headers: &mut salvo::http::HeaderMap) {
for (name, value) in self.always.iter() {
if !headers.contains_key(name) {
headers.insert(name.clone(), value.clone());
}
}
if hsts_allowed
&& let Some((name, value)) = &self.hsts
&& !headers.contains_key(name)
{
headers.insert(name.clone(), value.clone());
}
}
}
#[salvo::async_trait]
impl Handler for SecurityHeaders {
async fn handle(&self, req: &mut Request, depot: &mut Depot, res: &mut Response, ctrl: &mut FlowCtrl) {
ctrl.call_next(req, depot, res).await;
let hsts_allowed = depot
.obtain::<GenericServerState>()
.map(|s| {
!matches!(
s.startup_variant,
StartupVariant::HttpLocalhost | StartupVariant::UnsafeHttp
)
})
.unwrap_or(true);
self.apply(hsts_allowed, res.headers_mut());
}
}
#[cfg(test)]
mod tests {
use super::*;
use salvo::http::HeaderMap;
fn build(opts: SecurityHeadersOptions) -> SecurityHeaders {
SecurityHeaders::new(&opts)
}
#[test]
fn defaults_emit_the_four_baseline_headers() {
let mut headers = HeaderMap::new();
build(SecurityHeadersOptions::default()).apply(true, &mut headers);
assert_eq!(headers["x-content-type-options"], "nosniff");
assert_eq!(headers["x-frame-options"], "SAMEORIGIN");
assert_eq!(headers["referrer-policy"], "strict-origin-when-cross-origin");
assert_eq!(
headers["strict-transport-security"],
"max-age=31536000; includeSubDomains"
);
assert!(!headers.contains_key("content-security-policy"));
assert!(!headers.contains_key("permissions-policy"));
}
#[test]
fn skips_hsts_over_plain_http() {
let mut headers = HeaderMap::new();
build(SecurityHeadersOptions::default()).apply(false, &mut headers);
assert!(!headers.contains_key("strict-transport-security"));
assert_eq!(headers["x-content-type-options"], "nosniff");
}
#[test]
fn null_disables_specific_header() {
let opts = SecurityHeadersOptions {
x_frame_options: None,
..Default::default()
};
let mut headers = HeaderMap::new();
build(opts).apply(true, &mut headers);
assert!(!headers.contains_key("x-frame-options"));
assert_eq!(headers["x-content-type-options"], "nosniff");
}
#[test]
fn does_not_overwrite_existing_headers() {
let mut headers = HeaderMap::new();
headers.insert("x-frame-options", "DENY".parse().unwrap());
build(SecurityHeadersOptions::default()).apply(true, &mut headers);
assert_eq!(headers["x-frame-options"], "DENY");
}
#[test]
fn optional_headers_appear_when_set() {
let opts = SecurityHeadersOptions {
content_security_policy: Some("default-src 'self'".to_string()),
permissions_policy: Some("camera=()".to_string()),
cross_origin_opener_policy: Some("same-origin".to_string()),
..Default::default()
};
let mut headers = HeaderMap::new();
build(opts).apply(true, &mut headers);
assert_eq!(headers["content-security-policy"], "default-src 'self'");
assert_eq!(headers["permissions-policy"], "camera=()");
assert_eq!(headers["cross-origin-opener-policy"], "same-origin");
}
#[test]
fn invalid_value_is_dropped_not_panicked() {
let opts = SecurityHeadersOptions {
x_frame_options: Some("bad\nvalue".to_string()),
..Default::default()
};
let mut headers = HeaderMap::new();
build(opts).apply(true, &mut headers);
assert!(!headers.contains_key("x-frame-options"));
assert_eq!(headers["x-content-type-options"], "nosniff");
}
}