use crate::call::Call;
use crate::pipeline::{Middleware, Next};
use crate::response::Response;
use async_trait::async_trait;
use http::header::{
HeaderName, CONTENT_SECURITY_POLICY, REFERRER_POLICY, STRICT_TRANSPORT_SECURITY,
X_CONTENT_TYPE_OPTIONS, X_FRAME_OPTIONS,
};
use http::HeaderValue;
static PERMISSIONS_POLICY: HeaderName = HeaderName::from_static("permissions-policy");
static CROSS_ORIGIN_RESOURCE_POLICY: HeaderName =
HeaderName::from_static("cross-origin-resource-policy");
fn header_value(v: &str) -> Option<HeaderValue> {
HeaderValue::from_str(v).ok()
}
#[derive(Debug, Clone)]
pub struct SecurityHeaders {
content_type_options: Option<HeaderValue>,
frame_options: Option<HeaderValue>,
referrer_policy: Option<HeaderValue>,
hsts: Option<HeaderValue>,
csp: Option<HeaderValue>,
permissions_policy: Option<HeaderValue>,
cross_origin_resource_policy: Option<HeaderValue>,
}
impl Default for SecurityHeaders {
fn default() -> Self {
Self {
content_type_options: header_value("nosniff"),
frame_options: header_value("DENY"),
referrer_policy: header_value("no-referrer"),
hsts: header_value("max-age=31536000"),
csp: None,
permissions_policy: header_value(
"camera=(), microphone=(), geolocation=(), payment=(), usb=(), interest-cohort=()",
),
cross_origin_resource_policy: header_value("same-origin"),
}
}
}
impl SecurityHeaders {
pub fn new() -> Self {
Self::default()
}
pub fn content_type_options(mut self, v: Option<&str>) -> Self {
self.content_type_options = v.and_then(header_value);
self
}
pub fn frame_options(mut self, v: Option<&str>) -> Self {
self.frame_options = v.and_then(header_value);
self
}
pub fn referrer_policy(mut self, v: Option<&str>) -> Self {
self.referrer_policy = v.and_then(header_value);
self
}
pub fn hsts(mut self, v: Option<&str>) -> Self {
self.hsts = v.and_then(header_value);
self
}
pub fn content_security_policy(mut self, v: Option<&str>) -> Self {
self.csp = v.and_then(header_value);
self
}
pub fn permissions_policy(mut self, v: Option<&str>) -> Self {
self.permissions_policy = v.and_then(header_value);
self
}
pub fn cross_origin_resource_policy(mut self, v: Option<&str>) -> Self {
self.cross_origin_resource_policy = v.and_then(header_value);
self
}
pub(crate) fn apply_to(&self, headers: &mut http::HeaderMap, over_tls: bool) {
let mut set = |name: &HeaderName, value: &Option<HeaderValue>| {
let Some(v) = value else { return };
if headers.contains_key(name) {
return;
}
headers.insert(name, v.clone());
};
set(&X_CONTENT_TYPE_OPTIONS, &self.content_type_options);
set(&X_FRAME_OPTIONS, &self.frame_options);
set(&REFERRER_POLICY, &self.referrer_policy);
set(&CONTENT_SECURITY_POLICY, &self.csp);
set(&PERMISSIONS_POLICY, &self.permissions_policy);
set(
&CROSS_ORIGIN_RESOURCE_POLICY,
&self.cross_origin_resource_policy,
);
if over_tls {
set(&STRICT_TRANSPORT_SECURITY, &self.hsts);
}
}
pub(crate) fn into_middleware(self, tls_enabled: bool) -> SecurityHeadersMiddleware {
SecurityHeadersMiddleware {
cfg: self,
tls_enabled,
}
}
}
pub(crate) struct SecurityHeadersMiddleware {
cfg: SecurityHeaders,
tls_enabled: bool,
}
#[async_trait]
impl Middleware for SecurityHeadersMiddleware {
async fn handle(&self, call: Call, next: Next) -> Response {
let mut res = next.run(call).await;
self.cfg.apply_to(&mut res.headers, self.tls_enabled);
res
}
}