Skip to main content

ocre/
protect.rs

1//! Middleware every Ocre app runs, in this order (outermost first):
2//!
3//! 1. Security headers on every response.
4//! 2. Host authorization for the hosts listed in `ALLOWED_HOSTS`.
5//! 3. CORS for the origins listed in `ALLOWED_ORIGINS`.
6//! 4. Cross-origin request protection (CSRF).
7//! 5. The session cookie.
8//!
9//! CSRF protection checks where a request comes from instead of embedding
10//! tokens in forms: browsers send `Sec-Fetch-Site` (or at least `Origin`) with
11//! every unsafe request, and a request another site triggers is refused with
12//! 403. Requests without either header do not come from a browser page, so
13//! they cannot carry a victim's cookies by accident. This is the check Go 1.25
14//! ships as `http.CrossOriginProtection`; session cookies are also
15//! `SameSite=Lax`.
16
17use axum::{
18    Router,
19    body::Body,
20    extract::Request,
21    http::{HeaderMap, HeaderName, HeaderValue, Method, StatusCode, header},
22    middleware::{self, Next},
23    response::{IntoResponse, Response},
24};
25use tower_http::cors::{AllowOrigin, CorsLayer};
26
27use crate::session::{Keys, Session};
28
29/// Name of the Worker variable listing extra origins that may call the app from a browser.
30///
31/// Comma-separated; spaces and a trailing `/` are ignored and invalid entries
32/// skipped. Listed origins get CORS headers (methods GET, POST, PUT, PATCH,
33/// DELETE; headers `Content-Type`, `Authorization`, `Accept`; credentials
34/// allowed) and pass the CSRF check. When the variable is unset or empty,
35/// [`serve`](crate::serve) adds no CORS layer and only same-origin browser
36/// requests may change data. Read once per request by `serve`; no binding
37/// call.
38///
39/// ```ts
40/// // worker.env in cloudflare.config.ts
41/// ALLOWED_ORIGINS: bindings.text("https://app.example.com, https://admin.example.com"),
42/// ```
43///
44/// # Examples
45///
46/// ```no_run
47/// use axum::extract::State;
48/// use ocre::{ALLOWED_ORIGINS, Ctx, Result};
49///
50/// async fn origins(State(ctx): State<Ctx>) -> Result<String> {
51///     Ok(ctx.env().var(ALLOWED_ORIGINS).map(|v| v.to_string()).unwrap_or_default())
52/// }
53/// # let _ = origins;
54pub const ALLOWED_ORIGINS: &str = "ALLOWED_ORIGINS";
55
56/// Name of the Worker variable listing the host names the app answers to (Rails' `config.hosts`).
57///
58/// Comma-separated; an entry starting with `.` also allows every subdomain
59/// (`.example.com` allows `example.com` and `www.example.com`). When set,
60/// requests for any other `Host` get a plain-text `403 Forbidden` before any
61/// handler or session code runs; `localhost`, `127.0.0.1` and `[::1]` are
62/// always allowed so `ocre dev` keeps working (Cloudflare only routes your
63/// own host names to the Worker, so these never reach it in production).
64/// Unset or empty: every host is allowed.
65///
66/// On Workers, DNS rebinding cannot reach the app, but the same Worker also
67/// answers on `<name>.<account>.workers.dev` and preview URLs: list your
68/// custom domain to keep search engines and users on it. Read once per
69/// request by [`serve`](crate::serve); no binding call.
70///
71/// ```ts
72/// // worker.env in cloudflare.config.ts
73/// ALLOWED_HOSTS: bindings.text("example.com, .example.com"),
74/// ```
75///
76/// # Examples
77///
78/// ```
79/// assert_eq!(ocre::ALLOWED_HOSTS, "ALLOWED_HOSTS");
80/// ```
81pub const ALLOWED_HOSTS: &str = "ALLOWED_HOSTS";
82
83/// Per-request settings read from the Worker environment.
84pub(crate) struct Config {
85    pub keys: Result<Keys, String>,
86    pub allowed_origins: Vec<HeaderValue>,
87    pub allowed_hosts: Vec<String>,
88}
89
90/// `"https://a.example, https://b.example"` -> the two origins. Invalid
91/// entries are skipped.
92pub(crate) fn parse_origins(value: Option<String>) -> Vec<HeaderValue> {
93    value
94        .unwrap_or_default()
95        .split(',')
96        .map(|origin| origin.trim().trim_end_matches('/'))
97        .filter(|origin| !origin.is_empty())
98        .filter_map(|origin| HeaderValue::from_str(origin).ok())
99        .collect()
100}
101
102/// `"Example.com, .example.com"` -> `["example.com", ".example.com"]`.
103pub(crate) fn parse_hosts(value: Option<String>) -> Vec<String> {
104    let value = value.unwrap_or_default();
105    value.split(',').map(|host| host.trim().to_ascii_lowercase()).filter(|host| !host.is_empty()).collect()
106}
107
108/// Whether a request for `host` (with or without a port) may proceed.
109pub(crate) fn host_allowed(host: Option<&str>, allowed: &[String]) -> bool {
110    if allowed.is_empty() {
111        return true;
112    }
113    let Some(host) = host else { return false };
114    let host = host.to_ascii_lowercase();
115    // `[::1]:8787` -> `[::1]`; `example.com:443` -> `example.com`.
116    let name = match host.rsplit_once(':') {
117        Some((name, port)) if !name.is_empty() && port.bytes().all(|b| b.is_ascii_digit()) => name,
118        _ => host.as_str(),
119    };
120    if matches!(name, "localhost" | "127.0.0.1" | "[::1]") {
121        return true;
122    }
123    allowed.iter().any(|entry| match entry.strip_prefix('.') {
124        Some(domain) => name == domain || name.strip_suffix(domain).is_some_and(|sub| sub.ends_with('.')),
125        None => name == entry,
126    })
127}
128
129/// Wraps the application router with Ocre's middleware.
130pub(crate) fn wrap(router: Router, config: Config) -> Router {
131    let Config { keys, allowed_origins, allowed_hosts } = config;
132    let trusted = allowed_origins.clone();
133    let mut router = router
134        .layer(middleware::from_fn(move |req: Request, next: Next| {
135            let keys = keys.clone();
136            async move { session(keys, req, next).await }
137        }))
138        .layer(middleware::from_fn(move |req: Request, next: Next| {
139            let refused = cross_origin(req.method(), req.headers(), &trusted);
140            async move {
141                match refused {
142                    Some(reason) => (StatusCode::FORBIDDEN, reason).into_response(),
143                    None => next.run(req).await,
144                }
145            }
146        }));
147    if !allowed_origins.is_empty() {
148        router = router.layer(
149            CorsLayer::new()
150                .allow_origin(AllowOrigin::list(allowed_origins))
151                .allow_methods([Method::GET, Method::POST, Method::PUT, Method::PATCH, Method::DELETE])
152                .allow_headers([header::CONTENT_TYPE, header::AUTHORIZATION, header::ACCEPT])
153                .allow_credentials(true),
154        );
155    }
156    if !allowed_hosts.is_empty() {
157        router = router.layer(middleware::from_fn(move |req: Request, next: Next| {
158            let header = req.headers().get(header::HOST).and_then(|host| host.to_str().ok());
159            let allowed = host_allowed(req.uri().host().or(header), &allowed_hosts);
160            async move {
161                if allowed {
162                    next.run(req).await
163                } else {
164                    (StatusCode::FORBIDDEN, "Forbidden: blocked host. Add it to ALLOWED_HOSTS to allow it.")
165                        .into_response()
166                }
167            }
168        }));
169    }
170    router.layer(middleware::from_fn(security_headers))
171}
172
173async fn session(keys: Result<Keys, String>, mut req: Request, next: Next) -> Response {
174    let secure = req.uri().scheme_str() == Some("https");
175    let cookies = crate::cookies::Cookies::from_headers(req.headers(), keys.clone(), secure);
176    let session = Session::from_headers(req.headers(), keys, secure);
177    req.extensions_mut().insert(session.clone());
178    req.extensions_mut().insert(cookies.clone());
179    let mut response = next.run(req).await;
180    for cookie in cookies.set_cookies() {
181        response.headers_mut().append(header::SET_COOKIE, cookie);
182    }
183    match session.set_cookie() {
184        Ok(Some(cookie)) => {
185            response.headers_mut().append(header::SET_COOKIE, cookie);
186            response
187        }
188        Ok(None) => response,
189        #[cfg(feature = "html")]
190        Err(err) => err.into_response(),
191        #[cfg(not(feature = "html"))]
192        Err(err) => crate::ApiError::from(err).into_response(),
193    }
194}
195
196/// Whether the request is a WebSocket handshake (`Upgrade: websocket`).
197pub(crate) fn is_websocket_upgrade(headers: &HeaderMap) -> bool {
198    headers.get(header::UPGRADE).is_some_and(|value| value.as_bytes().eq_ignore_ascii_case(b"websocket"))
199}
200
201/// Why a request is refused, or `None` when it may proceed. WebSocket
202/// handshakes are GETs, but browsers send cookies with them and let any site
203/// open them (cross-site WebSocket hijacking), so they are checked like forms.
204pub(crate) fn cross_origin(method: &Method, headers: &HeaderMap, trusted: &[HeaderValue]) -> Option<&'static str> {
205    if matches!(*method, Method::GET | Method::HEAD | Method::OPTIONS) && !is_websocket_upgrade(headers) {
206        return None;
207    }
208    let origin = headers.get(header::ORIGIN);
209    if origin.is_some_and(|origin| trusted.contains(origin)) {
210        return None;
211    }
212    match headers.get("sec-fetch-site").map(HeaderValue::as_bytes) {
213        Some(b"same-origin" | b"none") => return None,
214        Some(_) => return Some("Forbidden: cross-site request. Add the origin to ALLOWED_ORIGINS to allow it."),
215        None => {}
216    }
217    // Older browsers: compare Origin with Host.
218    let (Some(origin), Some(host)) = (origin, headers.get(header::HOST)) else {
219        return None;
220    };
221    let origin_host = origin.to_str().ok().and_then(|origin| origin.split_once("://")).map(|(_, host)| host);
222    if origin_host.is_some_and(|origin_host| origin_host.as_bytes() == host.as_bytes()) {
223        None
224    } else {
225        Some("Forbidden: cross-origin request. Add the origin to ALLOWED_ORIGINS to allow it.")
226    }
227}
228
229/// Rails' default headers, plus HSTS on HTTPS. A handler that sets one of
230/// these headers keeps its own value.
231async fn security_headers(req: Request<Body>, next: Next) -> Response {
232    let https = req.uri().scheme_str() == Some("https");
233    let mut response = next.run(req).await;
234    let headers = response.headers_mut();
235    let defaults: [(HeaderName, &str); 5] = [
236        (header::X_CONTENT_TYPE_OPTIONS, "nosniff"),
237        (header::X_FRAME_OPTIONS, "SAMEORIGIN"),
238        (header::REFERRER_POLICY, "strict-origin-when-cross-origin"),
239        (header::X_XSS_PROTECTION, "0"),
240        (HeaderName::from_static("x-permitted-cross-domain-policies"), "none"),
241    ];
242    for (name, value) in defaults {
243        headers.entry(name).or_insert(HeaderValue::from_static(value));
244    }
245    if https {
246        headers.entry(header::STRICT_TRANSPORT_SECURITY).or_insert(HeaderValue::from_static("max-age=63072000"));
247    }
248    response
249}
250
251#[cfg(test)]
252#[path = "../tests/protect.rs"]
253mod tests;