Skip to main content

rest/
security.rs

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/// Adds conservative browser security headers without replacing handler-provided values.
11#[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}