1use actix_web::{
2 body::MessageBody,
3 dev::{Service, ServiceRequest, ServiceResponse, Transform},
4 http::header::{HeaderMap, HeaderName, HeaderValue},
5 Error,
6};
7use futures::future::{ok, LocalBoxFuture, Ready};
8use std::task::{Context, Poll};
9
10#[derive(Clone)]
12pub struct SecurityHeaders {
13 headers: HeaderMap,
14}
15
16impl Default for SecurityHeaders {
17 fn default() -> Self {
18 let mut headers = HeaderMap::new();
19 headers.insert(
20 HeaderName::from_static("content-security-policy"),
21 HeaderValue::from_static("default-src 'self'; frame-ancestors 'none'"),
22 );
23 headers.insert(
24 HeaderName::from_static("x-content-type-options"),
25 HeaderValue::from_static("nosniff"),
26 );
27 headers.insert(
28 HeaderName::from_static("x-frame-options"),
29 HeaderValue::from_static("DENY"),
30 );
31 headers.insert(
32 HeaderName::from_static("referrer-policy"),
33 HeaderValue::from_static("no-referrer"),
34 );
35 headers.insert(
36 HeaderName::from_static("permissions-policy"),
37 HeaderValue::from_static("camera=(), geolocation=(), microphone=()"),
38 );
39 Self { headers }
40 }
41}
42
43impl SecurityHeaders {
44 pub fn new() -> Self {
45 Self::default()
46 }
47
48 pub fn with_header(mut self, name: HeaderName, value: HeaderValue) -> Self {
49 self.headers.insert(name, value);
50 self
51 }
52}
53
54impl<S, B> Transform<S, ServiceRequest> for SecurityHeaders
55where
56 S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
57 S::Future: 'static,
58 B: MessageBody + 'static,
59{
60 type Response = ServiceResponse<B>;
61 type Error = Error;
62 type Transform = SecurityHeadersMiddleware<S>;
63 type InitError = ();
64 type Future = Ready<Result<Self::Transform, Self::InitError>>;
65
66 fn new_transform(&self, service: S) -> Self::Future {
67 ok(SecurityHeadersMiddleware {
68 service,
69 headers: self.headers.clone(),
70 })
71 }
72}
73
74pub struct SecurityHeadersMiddleware<S> {
75 service: S,
76 headers: HeaderMap,
77}
78
79impl<S, B> Service<ServiceRequest> for SecurityHeadersMiddleware<S>
80where
81 S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
82 S::Future: 'static,
83 B: MessageBody + 'static,
84{
85 type Response = ServiceResponse<B>;
86 type Error = Error;
87 type Future = LocalBoxFuture<'static, Result<Self::Response, Self::Error>>;
88
89 fn poll_ready(&self, context: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
90 self.service.poll_ready(context)
91 }
92
93 fn call(&self, request: ServiceRequest) -> Self::Future {
94 let headers = self.headers.clone();
95 let future = self.service.call(request);
96 Box::pin(async move {
97 let mut response = future.await?;
98 for (name, value) in &headers {
99 if !response.headers().contains_key(name) {
100 response.headers_mut().insert(name.clone(), value.clone());
101 }
102 }
103 Ok(response)
104 })
105 }
106}
107
108#[cfg(test)]
109mod tests {
110 use super::*;
111 use actix_web::{test, web, App, HttpResponse};
112
113 #[actix_rt::test]
114 async fn sets_defaults_but_preserves_handler_headers() {
115 let app = test::init_service(App::new().wrap(SecurityHeaders::new()).route(
116 "/",
117 web::get().to(|| async {
118 HttpResponse::Ok()
119 .insert_header(("x-frame-options", "SAMEORIGIN"))
120 .finish()
121 }),
122 ))
123 .await;
124
125 let response =
126 test::call_service(&app, test::TestRequest::get().uri("/").to_request()).await;
127
128 assert_eq!(
129 response.headers().get("x-frame-options").unwrap(),
130 "SAMEORIGIN"
131 );
132 assert_eq!(
133 response.headers().get("x-content-type-options").unwrap(),
134 "nosniff"
135 );
136 assert!(response.headers().contains_key("content-security-policy"));
137 }
138}