use std::convert::Infallible;
use axum::extract::FromRequestParts;
use axum::http::HeaderMap;
use axum::http::header::COOKIE;
use axum::http::request::Parts;
pub const CONSENT_COOKIE_NAME: &str = "autumn.consent";
#[cfg(feature = "maud")]
pub const DEFAULT_CSRF_COOKIE_NAME: &str = "autumn-csrf";
#[cfg(feature = "maud")]
pub const DEFAULT_CSRF_FORM_FIELD: &str = "_csrf";
pub const NECESSARY: &str = "necessary";
const MAX_AGE_SECS: u64 = 180 * 24 * 60 * 60;
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct Consent {
categories: Vec<String>,
policy_version: u32,
decided_at: Option<String>,
}
impl Consent {
#[must_use]
pub fn undecided() -> Self {
Self::default()
}
#[must_use]
pub const fn is_decided(&self) -> bool {
self.decided_at.is_some()
}
#[must_use]
pub const fn policy_version(&self) -> u32 {
self.policy_version
}
#[must_use]
pub fn decided_at(&self) -> Option<&str> {
self.decided_at.as_deref()
}
#[must_use]
pub fn categories(&self) -> &[String] {
&self.categories
}
#[must_use]
pub fn allows(&self, category: &str, current_policy_version: u32) -> bool {
if category == NECESSARY {
return true;
}
self.is_decided()
&& self.policy_version == current_policy_version
&& self.categories.iter().any(|c| c == category)
}
#[must_use]
pub const fn needs_prompt(&self, current_policy_version: u32) -> bool {
!self.is_decided() || self.policy_version != current_policy_version
}
fn from_cookie_value(raw: &str) -> Option<Self> {
let decoded = percent_decode(raw)?;
let mut parts = decoded.splitn(3, '|');
let version: u32 = parts.next()?.parse().ok()?;
let decided_at = parts.next()?;
if decided_at.is_empty() || chrono::DateTime::parse_from_rfc3339(decided_at).is_err() {
return None;
}
let categories_field = parts.next().unwrap_or("");
let categories = if categories_field.is_empty() {
Vec::new()
} else {
categories_field
.split(',')
.map(str::to_owned)
.collect::<Vec<_>>()
};
Some(Self {
categories,
policy_version: version,
decided_at: Some(decided_at.to_owned()),
})
}
#[must_use]
pub fn from_headers(headers: &HeaderMap) -> Self {
find_cookie(headers, CONSENT_COOKIE_NAME)
.and_then(|raw| Self::from_cookie_value(&raw))
.unwrap_or_default()
}
}
impl<S> FromRequestParts<S> for Consent
where
S: Send + Sync,
{
type Rejection = Infallible;
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
Ok(Self::from_headers(&parts.headers))
}
}
pub(crate) fn find_cookie(headers: &HeaderMap, name: &str) -> Option<String> {
let mut found = None;
for cookie_header in headers.get_all(COOKIE) {
let Ok(cookie_str) = cookie_header.to_str() else {
continue;
};
for pair in cookie_str.split(';') {
let pair = pair.trim();
let Some((k, v)) = pair.split_once('=') else {
continue;
};
if k.trim() != name {
continue;
}
if found.is_some() {
return None;
}
found = Some(v.trim().to_owned());
}
}
found
}
fn encode_cookie_value(policy_version: u32, decided_at: &str, categories: &[&str]) -> String {
let payload = format!("{policy_version}|{decided_at}|{}", categories.join(","));
percent_encode(&payload)
}
fn percent_encode(value: &str) -> String {
let mut out = String::with_capacity(value.len());
for byte in value.bytes() {
if is_unreserved_byte(byte) {
out.push(char::from(byte));
} else {
push_percent_encoded(&mut out, byte);
}
}
out
}
fn percent_decode(value: &str) -> Option<String> {
let bytes = value.as_bytes();
let mut out = Vec::with_capacity(bytes.len());
let mut i = 0;
while i < bytes.len() {
if bytes[i] == b'%' {
let hi = hex_value(*bytes.get(i + 1)?)?;
let lo = hex_value(*bytes.get(i + 2)?)?;
out.push((hi << 4) | lo);
i += 3;
} else {
out.push(bytes[i]);
i += 1;
}
}
String::from_utf8(out).ok()
}
const fn is_unreserved_byte(byte: u8) -> bool {
byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'_' | b'.')
}
fn push_percent_encoded(output: &mut String, byte: u8) {
const HEX: &[u8; 16] = b"0123456789ABCDEF";
output.push('%');
output.push(char::from(HEX[(byte >> 4) as usize]));
output.push(char::from(HEX[(byte & 0x0f) as usize]));
}
const fn hex_value(byte: u8) -> Option<u8> {
match byte {
b'0'..=b'9' => Some(byte - b'0'),
b'a'..=b'f' => Some(byte - b'a' + 10),
b'A'..=b'F' => Some(byte - b'A' + 10),
_ => None,
}
}
#[must_use]
pub fn accept_all_cookie(categories: &[&str], policy_version: u32) -> String {
build_consent_cookie(categories, policy_version)
}
#[must_use]
pub fn reject_non_essential_cookie(policy_version: u32) -> String {
build_consent_cookie(&[], policy_version)
}
fn build_consent_cookie(categories: &[&str], policy_version: u32) -> String {
let decided_at = chrono::Utc::now().to_rfc3339();
let value = encode_cookie_value(policy_version, &decided_at, categories);
format!(
"{CONSENT_COOKIE_NAME}={value}; Path=/; Max-Age={MAX_AGE_SECS}; HttpOnly; Secure; SameSite=Lax"
)
}
#[must_use]
pub fn expire_consent_cookie() -> String {
format!("{CONSENT_COOKIE_NAME}=; Path=/; Max-Age=0; HttpOnly; Secure; SameSite=Lax")
}
#[must_use]
pub fn safe_redirect_target(path: &str) -> &str {
let is_safe = path.starts_with('/')
&& !path.starts_with("//")
&& !path.contains("://")
&& !path.contains('\\')
&& path.bytes().all(|b| !b.is_ascii_control());
if is_safe { path } else { "/" }
}
#[must_use]
pub fn redirect_target_from_referer(referer: Option<&str>) -> String {
let Some(value) = referer else {
return "/".to_owned();
};
let after_scheme = value.split_once("://").map_or(value, |(_, rest)| rest);
let path = after_scheme.find('/').map_or("/", |i| &after_scheme[i..]);
safe_redirect_target(path).to_owned()
}
#[cfg(feature = "maud")]
mod banner;
#[cfg(feature = "maud")]
pub use banner::{consent_banner_markup, inject_consent_banner};
#[cfg(test)]
mod tests {
use super::*;
use axum::http::HeaderValue;
fn headers_with_cookie(raw: &str) -> HeaderMap {
let mut headers = HeaderMap::new();
headers.insert(COOKIE, HeaderValue::from_str(raw).unwrap());
headers
}
#[test]
fn undecided_is_not_decided() {
let consent = Consent::undecided();
assert!(!consent.is_decided());
assert_eq!(consent.policy_version(), 0);
assert_eq!(consent.decided_at(), None);
assert!(consent.categories().is_empty());
}
#[test]
fn from_headers_with_no_cookie_header_is_undecided() {
let consent = Consent::from_headers(&HeaderMap::new());
assert!(!consent.is_decided());
}
#[test]
fn from_headers_with_unrelated_cookies_is_undecided() {
let headers = headers_with_cookie("autumn.sid=abc123; theme=dark");
let consent = Consent::from_headers(&headers);
assert!(!consent.is_decided());
}
#[test]
fn malformed_pair_before_target_cookie_does_not_hide_it() {
let headers = headers_with_cookie("junk; autumn.sid=abc123");
assert_eq!(find_cookie(&headers, "autumn.sid"), Some("abc123".into()));
}
#[test]
fn duplicate_consent_cookie_is_rejected_as_possible_tossing() {
let headers = headers_with_cookie("autumn.consent=aaa; autumn.consent=bbb");
assert_eq!(find_cookie(&headers, "autumn.consent"), None);
let consent = Consent::from_headers(&headers);
assert!(
!consent.is_decided(),
"ambiguous duplicate cookie must not be trusted"
);
}
#[test]
fn necessary_category_always_allowed_even_when_undecided() {
let consent = Consent::undecided();
assert!(consent.allows(NECESSARY, 1));
}
#[test]
fn non_necessary_category_denied_when_undecided() {
let consent = Consent::undecided();
assert!(!consent.allows("analytics", 1));
assert!(!consent.allows("marketing", 1));
}
#[test]
fn accept_all_cookie_round_trips_and_allows_category() {
let set_cookie = accept_all_cookie(&["analytics", "marketing"], 1);
let raw_value = set_cookie
.split(';')
.next()
.unwrap()
.strip_prefix("autumn.consent=")
.unwrap();
let headers = headers_with_cookie(&format!("autumn.consent={raw_value}"));
let consent = Consent::from_headers(&headers);
assert!(consent.is_decided());
assert_eq!(consent.policy_version(), 1);
assert!(consent.allows("analytics", 1));
assert!(consent.allows("marketing", 1));
assert!(consent.allows(NECESSARY, 1));
assert!(!consent.allows("unrelated-category", 1));
assert!(consent.decided_at().is_some());
}
#[test]
fn accept_all_cookie_has_expected_attributes() {
let cookie = accept_all_cookie(&["analytics"], 1);
assert!(cookie.starts_with("autumn.consent="));
assert!(cookie.contains("Path=/"));
assert!(cookie.contains("HttpOnly"));
assert!(cookie.contains("Secure"));
assert!(cookie.contains("SameSite=Lax"));
assert!(cookie.contains(&format!("Max-Age={MAX_AGE_SECS}")));
}
#[test]
fn reject_non_essential_cookie_round_trips_and_denies_categories() {
let set_cookie = reject_non_essential_cookie(1);
let raw_value = set_cookie
.split(';')
.next()
.unwrap()
.strip_prefix("autumn.consent=")
.unwrap();
let headers = headers_with_cookie(&format!("autumn.consent={raw_value}"));
let consent = Consent::from_headers(&headers);
assert!(
consent.is_decided(),
"reject is still a recorded decision, not undecided"
);
assert!(!consent.allows("analytics", 1));
assert!(consent.allows(NECESSARY, 1), "necessary always allowed");
assert!(consent.categories().is_empty());
}
#[test]
fn expire_consent_cookie_has_zero_max_age() {
let cookie = expire_consent_cookie();
assert!(cookie.starts_with("autumn.consent="));
assert!(cookie.contains("Max-Age=0"));
}
#[test]
fn expiring_a_prior_accept_returns_to_undecided_and_reopens_the_gate() {
let accept = accept_all_cookie(&["analytics"], 1);
let raw_accept = accept
.split(';')
.next()
.unwrap()
.strip_prefix("autumn.consent=")
.unwrap();
let decided = Consent::from_headers(&headers_with_cookie(&format!(
"autumn.consent={raw_accept}"
)));
assert!(decided.allows("analytics", 1));
let expired = expire_consent_cookie();
assert!(
expired.contains("Max-Age=0"),
"confirms the cookie a browser would discard"
);
let withdrawn = Consent::from_headers(&HeaderMap::new());
assert!(!withdrawn.is_decided());
assert!(!withdrawn.allows("analytics", 1));
assert!(withdrawn.needs_prompt(1), "banner must reappear");
}
#[test]
fn needs_prompt_true_when_undecided() {
assert!(Consent::undecided().needs_prompt(1));
}
#[test]
fn needs_prompt_false_when_decided_under_current_version() {
let headers = decided_headers(&["analytics"], 1);
let consent = Consent::from_headers(&headers);
assert!(!consent.needs_prompt(1));
}
#[test]
fn needs_prompt_true_after_policy_version_bump() {
let headers = decided_headers(&["analytics"], 1);
let consent = Consent::from_headers(&headers);
assert!(consent.needs_prompt(2));
}
#[test]
fn allows_denies_category_after_policy_version_bump_even_if_previously_accepted() {
let headers = decided_headers(&["analytics"], 1);
let consent = Consent::from_headers(&headers);
assert!(consent.allows("analytics", 1));
assert!(!consent.allows("analytics", 2));
assert!(consent.allows(NECESSARY, 2));
}
fn decided_headers(categories: &[&str], policy_version: u32) -> HeaderMap {
let set_cookie = accept_all_cookie(categories, policy_version);
let raw_value = set_cookie
.split(';')
.next()
.unwrap()
.strip_prefix("autumn.consent=")
.unwrap();
headers_with_cookie(&format!("autumn.consent={raw_value}"))
}
#[test]
fn malformed_cookie_value_is_undecided_not_a_panic() {
for bogus in [
"not-percent-encoded-but-fine-chars",
"1",
"1|",
"%zz",
"",
"abc|2024-01-01T00:00:00Z|analytics",
] {
let headers = headers_with_cookie(&format!("autumn.consent={bogus}"));
let consent = Consent::from_headers(&headers);
assert!(
!consent.is_decided(),
"bogus value {bogus:?} must decode to undecided, not panic"
);
}
}
#[test]
fn cookie_with_invalid_timestamp_is_undecided() {
let value = percent_encode("1|not-a-timestamp|analytics");
let headers = headers_with_cookie(&format!("autumn.consent={value}"));
let consent = Consent::from_headers(&headers);
assert!(!consent.is_decided());
}
#[test]
fn session_and_csrf_cookies_survive_alongside_an_undecided_consent() {
let headers = headers_with_cookie("autumn.sid=session-value; autumn-csrf=csrf-value");
let consent = Consent::from_headers(&headers);
assert!(!consent.is_decided());
assert!(consent.allows(NECESSARY, 1));
}
#[test]
fn percent_encode_decode_round_trips_reserved_characters() {
let original = "1|2024-01-01T00:00:00+00:00|analytics,marketing";
let encoded = percent_encode(original);
assert!(!encoded.contains('|'), "pipe must be encoded: {encoded}");
assert!(!encoded.contains(':'), "colon must be encoded: {encoded}");
assert_eq!(percent_decode(&encoded).unwrap(), original);
}
#[test]
fn percent_decode_rejects_truncated_escape() {
assert_eq!(percent_decode("%4"), None);
assert_eq!(percent_decode("%"), None);
}
#[test]
fn percent_decode_rejects_invalid_hex() {
assert_eq!(percent_decode("%zz"), None);
}
#[test]
fn safe_redirect_target_allows_plain_relative_path() {
assert_eq!(safe_redirect_target("/blog/post-1"), "/blog/post-1");
assert_eq!(safe_redirect_target("/"), "/");
assert_eq!(safe_redirect_target("/a?b=c"), "/a?b=c");
}
#[test]
fn safe_redirect_target_rejects_scheme_relative_open_redirect() {
assert_eq!(safe_redirect_target("//evil.example.com"), "/");
}
#[test]
fn safe_redirect_target_rejects_absolute_url() {
assert_eq!(safe_redirect_target("https://evil.example.com"), "/");
assert_eq!(safe_redirect_target("javascript://alert(1)"), "/");
}
#[test]
fn safe_redirect_target_rejects_missing_leading_slash() {
assert_eq!(safe_redirect_target("evil.example.com"), "/");
assert_eq!(safe_redirect_target(""), "/");
}
#[test]
fn safe_redirect_target_rejects_embedded_crlf() {
assert_eq!(safe_redirect_target("/x\r\nSet-Cookie: pwned=1"), "/");
}
#[test]
fn safe_redirect_target_rejects_tab_based_scheme_relative_bypass() {
assert_eq!(safe_redirect_target("/\t/evil.example"), "/");
}
#[test]
fn safe_redirect_target_rejects_any_ascii_control_byte() {
assert_eq!(safe_redirect_target("/x\0evil"), "/");
assert_eq!(safe_redirect_target("/x\x0Bevil"), "/");
assert_eq!(safe_redirect_target("/x\x7Fevil"), "/");
}
#[test]
fn safe_redirect_target_rejects_backslash_based_scheme_relative_bypass() {
assert_eq!(safe_redirect_target("/\\evil.example"), "/");
assert_eq!(safe_redirect_target("/\\\\evil.example"), "/");
assert_eq!(safe_redirect_target("/ok/path\\evil.example"), "/");
}
#[test]
fn redirect_target_from_referer_extracts_path_and_query() {
assert_eq!(
redirect_target_from_referer(Some("https://app.example.com/blog/post-1?x=1")),
"/blog/post-1?x=1"
);
}
#[test]
fn redirect_target_from_referer_falls_back_to_root_when_absent() {
assert_eq!(redirect_target_from_referer(None), "/");
}
#[test]
fn redirect_target_from_referer_discards_scheme_and_host_entirely() {
assert_eq!(
redirect_target_from_referer(Some("https://good.example.com//evil.example.com")),
"/"
);
}
#[test]
fn redirect_target_from_referer_falls_back_when_no_path_present() {
assert_eq!(
redirect_target_from_referer(Some("https://example.com")),
"/"
);
}
#[test]
fn redirect_target_from_referer_does_not_silently_accept_a_scheme_truncated_target() {
let result = redirect_target_from_referer(Some(
"https://app.example.com/docs?source=https://vendor.example.com/x",
));
assert_ne!(
result, "/docs?source=https",
"must not silently truncate to (and accept) a corrupted target: {result}"
);
assert_eq!(result, "/");
}
}