use std::collections::BTreeMap;
use std::sync::Arc;
use anyhow::{bail, Context, Result};
use axum::{
extract::{Request, State},
http::{
header::{ACCESS_CONTROL_ALLOW_ORIGIN, ORIGIN, VARY},
HeaderMap, HeaderName, HeaderValue,
},
middleware::Next,
response::Response,
};
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct RouteHeaderTable {
rules: Vec<RouteHeaderRule>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
struct RouteHeaderRule {
path: String,
headers: Vec<(HeaderName, HeaderValue)>,
cors_origins: Vec<String>,
}
#[derive(serde::Deserialize)]
struct WireRule {
path: String,
headers: BTreeMap<String, String>,
#[serde(default)]
cors_origins: Vec<String>,
}
impl RouteHeaderTable {
pub fn parse(raw: &str) -> Result<Self> {
let trimmed = raw.trim();
if trimmed.is_empty() {
return Ok(Self::default());
}
let wire: Vec<WireRule> = serde_json::from_str(trimmed).with_context(|| {
format!(
"route header table is not a JSON array of {{\"path\": \"…\", \"headers\": \
{{…}}}} rules — got {}. This value is produced by \
`DomainConfig::route_headers_json` from `.yah/domains/<name>.toml`, so a \
malformed one means the deploy shipped a broken table",
preview(trimmed),
)
})?;
let mut rules = Vec::with_capacity(wire.len());
for rule in wire {
if rule.path.is_empty() {
bail!("route header rule has an empty `path` — a rule that matches nothing (or everything, depending on who reads it) is not a policy");
}
let mut headers = Vec::with_capacity(rule.headers.len());
for (name, value) in rule.headers {
let header_name = HeaderName::try_from(name.as_str()).with_context(|| {
format!("route {} declares {name:?}, which is not a valid HTTP header name", rule.path)
})?;
let header_value = HeaderValue::try_from(value.as_str()).with_context(|| {
format!(
"route {} declares {name}={value:?}, which is not a valid HTTP header value",
rule.path
)
})?;
headers.push((header_name, header_value));
}
if let Some(bad) = rule
.cors_origins
.iter()
.find(|o| o.as_str() == "*" || o.as_str() == "null")
{
bail!(
"route {} lists {bad:?} in cors_origins, which would grant every site (or \
every sandboxed page) rather than the ones named",
rule.path
);
}
rules.push(RouteHeaderRule {
path: rule.path,
headers,
cors_origins: rule.cors_origins,
});
}
Ok(Self { rules })
}
pub fn is_empty(&self) -> bool {
self.rules.is_empty()
}
pub fn len(&self) -> usize {
self.rules.len()
}
pub fn apply(&self, path: &str, origin: Option<&HeaderValue>, headers: &mut HeaderMap) {
let Some(rule) = self
.rules
.iter()
.find(|rule| matches_route_pattern(&rule.path, path))
else {
return;
};
for (name, value) in &rule.headers {
headers.insert(name.clone(), value.clone());
}
apply_cors(&rule.cors_origins, origin, headers);
}
}
fn apply_cors(cors_origins: &[String], origin: Option<&HeaderValue>, headers: &mut HeaderMap) {
if cors_origins.is_empty() {
return;
}
headers.remove(ACCESS_CONTROL_ALLOW_ORIGIN);
if let Some(origin) =
origin.filter(|o| cors_origins.iter().any(|listed| listed.as_bytes() == o.as_bytes()))
{
headers.insert(ACCESS_CONTROL_ALLOW_ORIGIN, origin.clone());
}
let varies = headers
.get_all(VARY)
.iter()
.filter_map(|v| v.to_str().ok())
.flat_map(|v| v.split(','))
.map(str::trim)
.any(|t| t == "*" || t.eq_ignore_ascii_case("origin"));
if !varies {
headers.append(VARY, HeaderValue::from_static("Origin"));
}
}
fn matches_route_pattern(pattern: &str, path: &str) -> bool {
let Some(head) = pattern.strip_suffix('*') else {
return path == pattern;
};
let prefix = head.trim_end_matches('/');
if prefix.is_empty() {
return true;
}
path == prefix || path.starts_with(&format!("{prefix}/"))
}
pub async fn apply_route_headers(
State(table): State<Arc<RouteHeaderTable>>,
req: Request,
next: Next,
) -> Response {
let path = req.uri().path().to_owned();
let origin = req.headers().get(ORIGIN).cloned();
let mut resp = next.run(req).await;
table.apply(&path, origin.as_ref(), resp.headers_mut());
resp
}
fn preview(raw: &str) -> String {
let cut = raw.char_indices().nth(120).map(|(i, _)| i);
match cut {
Some(i) => format!("{:?}…", &raw[..i]),
None => format!("{raw:?}"),
}
}
#[cfg(test)]
mod tests {
use super::*;
const ISOLATION: &str = r#"[
{"path":"/app/*","headers":{"Cross-Origin-Opener-Policy":"same-origin","Cross-Origin-Embedder-Policy":"require-corp"}},
{"path":"/*","headers":{"X-Tier":"marketing"}}
]"#;
fn applied(table: &RouteHeaderTable, path: &str) -> HeaderMap {
let mut headers = HeaderMap::new();
table.apply(path, None, &mut headers);
headers
}
const CORS: &str = r#"[
{"path":"/app/*","headers":{"Cross-Origin-Resource-Policy":"cross-origin"},
"cors_origins":["https://noisetable.com","https://staging.noisetable.com"]},
{"path":"/*","headers":{"X-Tier":"marketing"}}
]"#;
fn from_origin(table: &RouteHeaderTable, path: &str, origin: Option<&str>) -> HeaderMap {
let mut headers = HeaderMap::new();
let origin = origin.map(|o| HeaderValue::from_str(o).unwrap());
table.apply(path, origin.as_ref(), &mut headers);
headers
}
#[test]
fn cors_echoes_each_listed_origin_and_always_varies() {
let table = RouteHeaderTable::parse(CORS).unwrap();
for origin in ["https://noisetable.com", "https://staging.noisetable.com"] {
let headers = from_origin(&table, "/app/x.wasm", Some(origin));
assert_eq!(headers.get(ACCESS_CONTROL_ALLOW_ORIGIN).unwrap(), origin);
assert_eq!(headers.get(VARY).unwrap(), "Origin");
assert_eq!(
headers.get("cross-origin-resource-policy").unwrap(),
"cross-origin"
);
}
for origin in [Some("https://evil.example"), Some("https://noisetable.com/"), None] {
let headers = from_origin(&table, "/app/x.wasm", origin);
assert!(headers.get(ACCESS_CONTROL_ALLOW_ORIGIN).is_none(), "{origin:?}");
assert_eq!(headers.get(VARY).unwrap(), "Origin", "{origin:?}");
}
let other = from_origin(&table, "/", Some("https://noisetable.com"));
assert!(other.get(ACCESS_CONTROL_ALLOW_ORIGIN).is_none());
assert!(other.get(VARY).is_none(), "a route with no list is untouched");
}
#[test]
fn cors_overrides_a_handler_grant_and_merges_vary() {
let table = RouteHeaderTable::parse(CORS).unwrap();
let mut headers = HeaderMap::new();
headers.insert(ACCESS_CONTROL_ALLOW_ORIGIN, HeaderValue::from_static("*"));
headers.insert(VARY, HeaderValue::from_static("Accept-Encoding"));
table.apply("/app/", None, &mut headers);
assert!(headers.get(ACCESS_CONTROL_ALLOW_ORIGIN).is_none());
let vary: Vec<_> = headers.get_all(VARY).iter().collect();
assert_eq!(vary, ["Accept-Encoding", "Origin"]);
}
#[test]
fn cors_wildcards_are_refused() {
for bad in ["*", "null"] {
let raw = format!(r#"[{{"path":"/*","headers":{{}},"cors_origins":["{bad}"]}}]"#);
assert!(RouteHeaderTable::parse(&raw).is_err(), "{bad}");
}
}
#[test]
fn unset_and_empty_read_as_no_table() {
assert!(RouteHeaderTable::parse("").unwrap().is_empty());
assert!(RouteHeaderTable::parse(" ").unwrap().is_empty());
assert!(RouteHeaderTable::parse("[]").unwrap().is_empty());
}
#[test]
fn prefix_rule_covers_the_bare_prefix_and_everything_under_it() {
let table = RouteHeaderTable::parse(ISOLATION).unwrap();
for path in ["/app", "/app/", "/app/index.html", "/app/pkg/bundle.wasm"] {
let headers = applied(&table, path);
assert_eq!(
headers.get("cross-origin-opener-policy").unwrap(),
"same-origin",
"{path}"
);
assert_eq!(
headers.get("cross-origin-embedder-policy").unwrap(),
"require-corp",
"{path}"
);
}
}
#[test]
fn matching_is_segment_aware() {
let table = RouteHeaderTable::parse(ISOLATION).unwrap();
let headers = applied(&table, "/apple");
assert!(headers.get("cross-origin-opener-policy").is_none());
assert_eq!(headers.get("x-tier").unwrap(), "marketing");
}
#[test]
fn first_match_wins_with_no_merging() {
let table = RouteHeaderTable::parse(ISOLATION).unwrap();
let headers = applied(&table, "/app/");
assert!(
headers.get("x-tier").is_none(),
"the catch-all must not merge into the app route"
);
}
#[test]
fn a_matching_rule_with_no_headers_still_consumes_the_match() {
let table = RouteHeaderTable::parse(
r#"[{"path":"/app/*","headers":{}},{"path":"/*","headers":{"X-Tier":"marketing"}}]"#,
)
.unwrap();
assert!(applied(&table, "/app/x").is_empty());
assert_eq!(applied(&table, "/other").get("x-tier").unwrap(), "marketing");
}
#[test]
fn a_pattern_without_a_star_is_an_exact_match() {
let table =
RouteHeaderTable::parse(r#"[{"path":"/exact","headers":{"X-One":"1"}}]"#).unwrap();
assert_eq!(applied(&table, "/exact").get("x-one").unwrap(), "1");
assert!(applied(&table, "/exact/deeper").is_empty());
}
#[test]
fn declared_headers_overwrite_a_header_the_handler_already_set() {
let table =
RouteHeaderTable::parse(r#"[{"path":"/*","headers":{"Cache-Control":"no-store"}}]"#)
.unwrap();
let mut headers = HeaderMap::new();
headers.insert("cache-control", HeaderValue::from_static("public"));
table.apply("/", None, &mut headers);
assert_eq!(headers.get("cache-control").unwrap(), "no-store");
assert_eq!(headers.get_all("cache-control").iter().count(), 1);
}
#[test]
fn malformed_tables_are_errors_not_empty_tables() {
for raw in [
"{not json",
r#"{"path":"/*","headers":{}}"#, r#"[{"path":"/*"}]"#, r#"[{"headers":{"X":"1"}}]"#, r#"[{"path":"/*","headers":{"X":1}}]"#, r#"[{"path":"","headers":{"X":"1"}}]"#, r#"[{"path":"/*","headers":{"Bad Name":"1"}}]"#,
"[{\"path\":\"/*\",\"headers\":{\"X-Nl\":\"a\\nb\"}}]",
] {
assert!(
RouteHeaderTable::parse(raw).is_err(),
"expected a hard error for {raw}"
);
}
}
#[test]
fn the_parse_error_names_the_offending_input() {
let err = RouteHeaderTable::parse(r#"[{"path":"/app/*","headers":{"Bad Name":"1"}}]"#)
.unwrap_err();
let rendered = format!("{err:#}");
assert!(rendered.contains("/app/*"), "{rendered}");
assert!(rendered.contains("Bad Name"), "{rendered}");
}
}