use topcoat_core::context::Cx;
use crate::{
Body, IntoPath, Layer, LayerFuture, Method, Next, Path, PathBuf,
error::forbidden,
header,
request::{headers, method, uri},
};
#[derive(Debug)]
pub struct OriginPolicy {
inner: Inner,
}
#[derive(Debug)]
enum Inner {
Verify {
trusted_origins: Vec<String>,
exempt_paths: Vec<PathBuf>,
},
Disabled,
}
impl OriginPolicy {
#[must_use]
pub fn new() -> Self {
Self {
inner: Inner::Verify {
trusted_origins: Vec::new(),
exempt_paths: Vec::new(),
},
}
}
#[must_use]
pub fn trust_origins<I>(mut self, origins: I) -> Self
where
I: IntoIterator,
I::Item: Into<String>,
{
if let Inner::Verify {
trusted_origins, ..
} = &mut self.inner
{
trusted_origins.extend(origins.into_iter().map(Into::into));
}
self
}
#[must_use]
pub fn exempt_paths<I>(mut self, paths: I) -> Self
where
I: IntoIterator,
I::Item: IntoPath,
{
if let Inner::Verify { exempt_paths, .. } = &mut self.inner {
exempt_paths.extend(paths.into_iter().map(|path| path.into_path().into_owned()));
}
self
}
#[must_use]
pub fn dangerous_disable() -> Self {
Self {
inner: Inner::Disabled,
}
}
fn check(&self, cx: &Cx) -> OriginVerdict {
match &self.inner {
Inner::Disabled => OriginVerdict::Allow,
Inner::Verify {
trusted_origins,
exempt_paths,
} => {
let headers = headers(cx);
let header = |name: &str| headers.get(name).and_then(|value| value.to_str().ok());
let upgrades_to_websocket = header(header::UPGRADE.as_str())
.is_some_and(|value| value.trim().eq_ignore_ascii_case("websocket"));
if !upgrades_to_websocket
&& matches!(method(cx), &Method::GET | &Method::HEAD | &Method::OPTIONS)
{
return OriginVerdict::Allow;
}
if exempt_paths.iter().any(|path| path.matches(uri(cx).path())) {
return OriginVerdict::Allow;
}
let origin = header(header::ORIGIN.as_str());
if origin.is_some_and(|origin| {
trusted_origins
.iter()
.any(|trusted| trusted.eq_ignore_ascii_case(origin))
}) {
return OriginVerdict::Allow;
}
if let Some(site) = header("sec-fetch-site") {
return if matches!(site, "same-origin" | "none") {
OriginVerdict::Allow
} else {
OriginVerdict::Deny
};
}
if let Some(origin) = origin {
let origin = origin.split_once("://").map(|(_, host)| host);
let host = header(header::HOST.as_str())
.or_else(|| uri(cx).authority().map(http::uri::Authority::as_str));
return if let Some(host) = host
&& let Some(origin) = origin
&& origin.eq_ignore_ascii_case(host)
{
OriginVerdict::Allow
} else {
OriginVerdict::Deny
};
}
OriginVerdict::Allow
}
}
}
}
impl Default for OriginPolicy {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, PartialEq, Eq)]
enum OriginVerdict {
Allow,
Deny,
}
#[derive(Debug)]
pub struct OriginLayer {
policy: OriginPolicy,
}
impl OriginLayer {
#[must_use]
pub fn new(policy: OriginPolicy) -> Self {
Self { policy }
}
}
impl Layer for OriginLayer {
fn path(&self) -> Option<&Path> {
None
}
fn handle<'a>(&'a self, cx: &'a Cx, body: Body, next: Next<'a>) -> LayerFuture<'a> {
match self.policy.check(cx) {
OriginVerdict::Allow => next.run(cx, body),
OriginVerdict::Deny => Box::pin(async { Err(forbidden().into()) }),
}
}
}
#[cfg(test)]
mod tests {
use http::Request;
use topcoat_core::context::CxTestBuilder;
use super::*;
fn cx_with(method: Method, uri: &str, headers: &[(&str, &str)]) -> Cx {
let mut builder = Request::builder().method(method).uri(uri);
for (name, value) in headers {
builder = builder.header(*name, *value);
}
let (parts, ()) = builder.body(()).expect("request should build").into_parts();
CxTestBuilder::new().request_context(parts).build()
}
fn check(policy: &OriginPolicy, method: Method, headers: &[(&str, &str)]) -> OriginVerdict {
policy.check(&cx_with(method, "/x", headers))
}
#[test]
fn non_browser_requests_pass() {
let policy = OriginPolicy::new();
assert_eq!(check(&policy, Method::POST, &[]), OriginVerdict::Allow);
}
#[test]
fn same_origin_requests_pass() {
let policy = OriginPolicy::new();
for site in ["same-origin", "none"] {
assert_eq!(
check(&policy, Method::POST, &[("sec-fetch-site", site)]),
OriginVerdict::Allow
);
}
}
#[test]
fn cross_origin_safe_methods_pass() {
let policy = OriginPolicy::new();
for method in [Method::GET, Method::HEAD, Method::OPTIONS] {
assert_eq!(
check(&policy, method, &[("sec-fetch-site", "cross-site")]),
OriginVerdict::Allow
);
}
}
#[test]
fn cross_origin_unsafe_methods_are_denied() {
let policy = OriginPolicy::new();
for method in [Method::POST, Method::PUT, Method::DELETE, Method::PATCH] {
assert_eq!(
check(&policy, method, &[("sec-fetch-site", "cross-site")]),
OriginVerdict::Deny
);
}
}
#[test]
fn cross_origin_websocket_upgrades_are_denied() {
let policy = OriginPolicy::new();
assert_eq!(
check(
&policy,
Method::GET,
&[("sec-fetch-site", "cross-site"), ("upgrade", "WebSocket")]
),
OriginVerdict::Deny
);
}
#[test]
fn same_origin_websocket_upgrades_pass() {
let policy = OriginPolicy::new();
assert_eq!(
check(
&policy,
Method::GET,
&[("sec-fetch-site", "same-origin"), ("upgrade", "websocket")]
),
OriginVerdict::Allow
);
}
#[test]
fn matching_origin_and_host_pass() {
let policy = OriginPolicy::new();
assert_eq!(
check(
&policy,
Method::POST,
&[("origin", "https://example.com"), ("host", "EXAMPLE.com")]
),
OriginVerdict::Allow
);
}
#[test]
fn mismatched_origin_and_host_are_denied_for_unsafe_methods() {
let policy = OriginPolicy::new();
let headers = [("origin", "https://evil.example"), ("host", "example.com")];
assert_eq!(check(&policy, Method::POST, &headers), OriginVerdict::Deny);
assert_eq!(check(&policy, Method::GET, &headers), OriginVerdict::Allow);
}
#[test]
fn trusted_origins_pass() {
let policy = OriginPolicy::new().trust_origins(["https://accounts.example.com"]);
assert_eq!(
check(
&policy,
Method::POST,
&[
("origin", "https://ACCOUNTS.example.com"),
("sec-fetch-site", "cross-site"),
]
),
OriginVerdict::Allow
);
assert_eq!(
check(
&policy,
Method::POST,
&[
("origin", "https://other.example.com"),
("sec-fetch-site", "cross-site"),
]
),
OriginVerdict::Deny
);
}
#[test]
fn exempt_paths_pass_unchecked() {
let policy = OriginPolicy::new().exempt_paths(["/ws/{id}"]);
let headers = [("sec-fetch-site", "cross-site"), ("upgrade", "websocket")];
assert_eq!(
policy.check(&cx_with(Method::GET, "/ws/42", &headers)),
OriginVerdict::Allow
);
assert_eq!(
policy.check(&cx_with(Method::GET, "/other", &headers)),
OriginVerdict::Deny
);
}
#[test]
fn disabled_policy_passes_everything() {
let policy = OriginPolicy::dangerous_disable();
assert_eq!(
check(
&policy,
Method::POST,
&[("sec-fetch-site", "cross-site"), ("upgrade", "websocket")]
),
OriginVerdict::Allow
);
}
}