use std::collections::BTreeMap;
use std::sync::Arc;
use anyhow::{bail, Context, Result};
use axum::{
extract::{Request, State},
http::{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)>,
}
#[derive(serde::Deserialize)]
struct WireRule {
path: String,
headers: BTreeMap<String, 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));
}
rules.push(RouteHeaderRule {
path: rule.path,
headers,
});
}
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, 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());
}
}
}
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 mut resp = next.run(req).await;
table.apply(&path, 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, &mut headers);
headers
}
#[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("/", &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}");
}
}