use axess_rng::{SecureRng, SystemRng};
use axum::{
body::Body,
http::{HeaderValue, Request, Response, StatusCode, header},
response::IntoResponse,
};
use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD};
use hmac::Mac;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::task::{Context, Poll};
use subtle::ConstantTimeEq;
use tower::{Layer, Service};
use url::Url;
use crate::session::layer::SessionHandle;
pub const DEFAULT_CSRF_COOKIE: &str = "axess.csrf";
pub const DEFAULT_CSRF_HEADER: &str = "x-csrf-token";
pub const DEFAULT_CSRF_FORM_FIELD: &str = "_csrf";
pub const DEFAULT_CSRF_FORM_BODY_LIMIT: usize = 64 * 1024;
const TOKEN_NONCE_BYTES: usize = 32;
use crate::cookies::MAX_COOKIE_VALUE_BYTES;
#[derive(Clone)]
pub struct CsrfConfig {
signing_key: Arc<[u8; 32]>,
cookie_name: Arc<str>,
header_name: Arc<str>,
form_field_name: Arc<str>,
form_body_limit: usize,
secure: bool,
same_site: tower_cookies::cookie::SameSite,
path: Arc<str>,
allowed_origins: Arc<[Arc<str>]>,
}
impl CsrfConfig {
pub fn new(signing_key: [u8; 32]) -> Self {
Self {
signing_key: Arc::new(signing_key),
cookie_name: DEFAULT_CSRF_COOKIE.into(),
header_name: DEFAULT_CSRF_HEADER.into(),
form_field_name: DEFAULT_CSRF_FORM_FIELD.into(),
form_body_limit: DEFAULT_CSRF_FORM_BODY_LIMIT,
secure: true,
same_site: tower_cookies::cookie::SameSite::Lax,
path: "/".into(),
allowed_origins: Arc::from(Vec::new()),
}
}
pub fn cookie_name(mut self, name: impl Into<Arc<str>>) -> Self {
self.cookie_name = name.into();
self
}
pub fn header_name(mut self, name: impl Into<Arc<str>>) -> Self {
self.header_name = name.into();
self
}
pub fn form_field_name(mut self, name: impl Into<Arc<str>>) -> Self {
self.form_field_name = name.into();
self
}
pub fn form_body_limit(mut self, limit: usize) -> Self {
self.form_body_limit = limit;
self
}
pub fn secure(mut self, secure: bool) -> Self {
self.secure = secure;
self
}
pub fn same_site(mut self, same_site: tower_cookies::cookie::SameSite) -> Self {
self.same_site = same_site;
self
}
pub fn require_origin(mut self, origins: impl IntoIterator<Item = impl AsRef<str>>) -> Self {
let mut normalized: Vec<Arc<str>> = Vec::new();
let mut supplied_any = false;
let mut rejected: Vec<String> = Vec::new();
for raw in origins {
supplied_any = true;
let raw = raw.as_ref();
match normalize_origin(raw) {
Some(n) => normalized.push(Arc::from(n)),
None => rejected.push(raw.to_owned()),
}
}
assert!(
!(supplied_any && normalized.is_empty()),
"CsrfConfig::require_origin: all {} supplied origin(s) failed to canonicalize \
(missing scheme, opaque-origin scheme, or unparseable URL): {:?}. \
Leaving the list empty would silently disable the Origin/Referer gate.",
rejected.len(),
rejected,
);
self.allowed_origins = normalized.into();
self
}
}
#[derive(Clone, Debug)]
pub struct CsrfToken(pub String);
impl CsrfToken {
pub fn as_str(&self) -> &str {
&self.0
}
}
#[derive(Clone)]
pub struct CsrfLayer {
config: CsrfConfig,
}
impl CsrfLayer {
pub fn new(config: CsrfConfig) -> Self {
Self { config }
}
}
impl<S> Layer<S> for CsrfLayer {
type Service = CsrfService<S>;
fn layer(&self, inner: S) -> Self::Service {
CsrfService {
inner,
config: self.config.clone(),
}
}
}
#[derive(Clone)]
pub struct CsrfService<S> {
inner: S,
config: CsrfConfig,
}
impl<S> Service<Request<Body>> for CsrfService<S>
where
S: Service<Request<Body>, Response = Response<Body>> + Clone + Send + 'static,
S::Future: Send + 'static,
S::Error: Send + 'static,
{
type Response = S::Response;
type Error = S::Error;
type Future = Pin<Box<dyn Future<Output = Result<Self::Response, Self::Error>> + Send>>;
fn poll_ready(&mut self, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.inner.poll_ready(cx)
}
fn call(&mut self, mut req: Request<Body>) -> Self::Future {
let config = self.config.clone();
let mut inner = self.inner.clone();
std::mem::swap(&mut inner, &mut self.inner);
Box::pin(async move {
let cookie_token = extract_cookie_token(&req, &config.cookie_name);
let method = req.method().clone();
let session_handle = req.extensions().get::<SessionHandle>().cloned();
let session_id = match &session_handle {
Some(handle) => Some(handle.0.read().await.id),
None => None,
};
if is_state_changing(&method) {
let Some(session_id) = session_id else {
tracing::warn!(
method = %method,
path = %req.uri().path(),
"csrf: no session id on state-changing request \
(CsrfLayer must be layered inside the session layer)"
);
return Ok((StatusCode::FORBIDDEN, "CSRF validation failed").into_response());
};
if !config.allowed_origins.is_empty()
&& !origin_matches_allowed(&req, &config.allowed_origins)
{
tracing::warn!(
method = %method,
path = %req.uri().path(),
"csrf: Origin/Referer validation failed"
);
return Ok((StatusCode::FORBIDDEN, "CSRF validation failed").into_response());
}
let (extracted, req_after) = match extract_presented_token(req, &config).await {
Ok(pair) => pair,
Err(reason) => {
tracing::warn!(
method = %method,
reason = %reason,
"csrf: request rejected during token extraction"
);
return Ok(
(StatusCode::FORBIDDEN, "CSRF validation failed").into_response()
);
}
};
req = req_after;
let presented = extracted;
let cookie_present = cookie_token.as_deref();
if !validate_pair(
cookie_present,
presented.as_deref(),
session_id.as_bytes(),
&config.signing_key,
) {
tracing::warn!(
method = %method,
path = %req.uri().path(),
cookie_present = cookie_present.is_some(),
header_or_form_present = presented.is_some(),
"csrf: token validation failed"
);
return Ok((StatusCode::FORBIDDEN, "CSRF validation failed").into_response());
}
}
let provisional_mint = match (&cookie_token, session_id.as_ref()) {
(Some(existing), _) if !existing.is_empty() => None,
(_, Some(sid_at_entry)) => {
Some(mint_token(sid_at_entry.as_bytes(), &config.signing_key))
}
(_, None) => None,
};
let extension_token = provisional_mint
.clone()
.or_else(|| cookie_token.clone())
.unwrap_or_default();
req.extensions_mut().insert(CsrfToken(extension_token));
let mut response = inner.call(req).await?;
let session_id_after = match &session_handle {
Some(handle) => Some(handle.0.read().await.id),
None => None,
};
let effective_cookie = provisional_mint
.as_deref()
.or(cookie_token.as_deref())
.filter(|t| !t.is_empty());
let token_to_set = match (effective_cookie, session_id_after.as_ref()) {
(Some(existing), Some(sid_after))
if validate_token(existing, sid_after.as_bytes(), &config.signing_key) =>
{
provisional_mint
}
(_, Some(sid_after)) => {
Some(mint_token(sid_after.as_bytes(), &config.signing_key))
}
(_, None) => None,
};
if let Some(new_token) = token_to_set {
let cookie = build_cookie(&config, &new_token);
if let Ok(hv) = HeaderValue::from_str(&cookie) {
response.headers_mut().append(header::SET_COOKIE, hv);
}
}
Ok(response)
})
}
}
fn is_state_changing(method: &axum::http::Method) -> bool {
matches!(
*method,
axum::http::Method::POST
| axum::http::Method::PUT
| axum::http::Method::PATCH
| axum::http::Method::DELETE
)
}
fn extract_cookie_token(req: &Request<Body>, cookie_name: &str) -> Option<String> {
crate::cookies::extract_named_cookie(req.headers(), cookie_name, MAX_COOKIE_VALUE_BYTES)
}
async fn extract_presented_token(
req: Request<Body>,
config: &CsrfConfig,
) -> Result<(Option<String>, Request<Body>), &'static str> {
if let Some(value) = req.headers().get(config.header_name.as_ref())
&& let Ok(s) = value.to_str()
&& !s.is_empty()
{
return Ok((Some(s.to_string()), req));
}
let is_urlencoded = req
.headers()
.get(header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.map(|s| {
let base = s.split(';').next().unwrap_or("").trim();
base.eq_ignore_ascii_case("application/x-www-form-urlencoded")
})
.unwrap_or(false);
if !is_urlencoded {
return Ok((None, req));
}
if let Some(declared) = req
.headers()
.get(header::CONTENT_LENGTH)
.and_then(|v| v.to_str().ok())
.and_then(|s| s.parse::<usize>().ok())
&& declared > config.form_body_limit
{
return Err("form body exceeds csrf buffer cap");
}
let (parts, body) = req.into_parts();
let bytes = match axum::body::to_bytes(body, config.form_body_limit).await {
Ok(b) => b,
Err(_) => return Err("form body exceeds csrf buffer cap"),
};
let field = form_urlencoded::parse(&bytes)
.find(|(k, _)| k.as_ref() == config.form_field_name.as_ref())
.map(|(_, v)| v.into_owned())
.filter(|s| !s.is_empty());
let restored = Request::from_parts(parts, Body::from(bytes));
Ok((field, restored))
}
fn normalize_origin(raw: &str) -> Option<String> {
let parsed = Url::parse(raw).ok()?;
match parsed.origin() {
url::Origin::Tuple(..) => Some(parsed.origin().ascii_serialization()),
url::Origin::Opaque(_) => None,
}
}
fn origin_matches_allowed(req: &Request<Body>, allowed: &[Arc<str>]) -> bool {
if let Some(header) = req.headers().get(header::ORIGIN) {
let Ok(s) = header.to_str() else { return false };
let Some(origin) = normalize_origin(s) else {
return false;
};
return allowed.iter().any(|a| a.as_ref() == origin);
}
let Some(referer) = req.headers().get(header::REFERER) else {
return false;
};
let Ok(referer_str) = referer.to_str() else {
return false;
};
let Some(origin) = normalize_origin(referer_str) else {
return false;
};
allowed.iter().any(|a| a.as_ref() == origin)
}
fn mint_token(session_id: &[u8], signing_key: &[u8; 32]) -> String {
let mut nonce = [0u8; TOKEN_NONCE_BYTES];
SystemRng.fill_bytes(&mut nonce);
let tag = compute_tag(&nonce, session_id, signing_key);
let mut combined = Vec::with_capacity(TOKEN_NONCE_BYTES + tag.len());
combined.extend_from_slice(&nonce);
combined.extend_from_slice(&tag);
URL_SAFE_NO_PAD.encode(&combined)
}
fn compute_tag(nonce: &[u8], session_id: &[u8], signing_key: &[u8; 32]) -> [u8; 32] {
let mut mac = crate::hmac::new_signer(signing_key);
mac.update(nonce);
mac.update(session_id);
mac.finalize().into_bytes().into()
}
fn validate_token(token: &str, session_id: &[u8], signing_key: &[u8; 32]) -> bool {
let bytes = match URL_SAFE_NO_PAD.decode(token) {
Ok(b) => b,
Err(_) => return false,
};
if bytes.len() != TOKEN_NONCE_BYTES + 32 {
return false;
}
let (nonce, tag) = bytes.split_at(TOKEN_NONCE_BYTES);
let expected = compute_tag(nonce, session_id, signing_key);
expected.ct_eq(tag).into()
}
fn validate_pair(
cookie_token: Option<&str>,
presented: Option<&str>,
session_id: &[u8],
signing_key: &[u8; 32],
) -> bool {
let (Some(c), Some(p)) = (cookie_token, presented) else {
return false;
};
if c.is_empty() || p.is_empty() {
return false;
}
bool::from(c.as_bytes().ct_eq(p.as_bytes())) && validate_token(c, session_id, signing_key)
}
fn build_cookie(config: &CsrfConfig, token: &str) -> String {
use tower_cookies::Cookie;
let mut cookie = Cookie::new(config.cookie_name.as_ref().to_string(), token.to_string());
cookie.set_http_only(false);
cookie.set_secure(config.secure);
cookie.set_same_site(config.same_site);
cookie.set_path(config.path.as_ref().to_string());
cookie.to_string()
}
#[cfg(test)]
mod tests {
use super::*;
const SID: &[u8] = b"session-a";
fn session_handle(seed: u64) -> SessionHandle {
use crate::session::data::SessionData;
use crate::session::id::SessionId;
use crate::session::layer::SessionInner;
use crate::testing::mock_random::MockRng;
use std::sync::Arc;
use tokio::sync::RwLock;
let rng = MockRng::new(seed);
let inner = SessionInner {
id: SessionId::new(&rng),
data: SessionData::default(),
modified: false,
regenerate: false,
pre_cycle_id: None,
pending_fingerprint: None,
max_custom_bytes: 64 * 1024,
};
SessionHandle(Arc::new(RwLock::new(inner)))
}
#[test]
fn token_round_trip_validates() {
let key = [7u8; 32];
let token = mint_token(SID, &key);
assert!(validate_token(&token, SID, &key));
}
#[test]
fn token_with_wrong_key_rejected() {
let key = [7u8; 32];
let other_key = [9u8; 32];
let token = mint_token(SID, &key);
assert!(!validate_token(&token, SID, &other_key));
}
#[test]
fn truncated_token_rejected() {
let key = [7u8; 32];
let token = mint_token(SID, &key);
let truncated = &token[..token.len() - 4];
assert!(!validate_token(truncated, SID, &key));
}
#[test]
fn empty_token_rejected() {
let key = [7u8; 32];
assert!(!validate_token("", SID, &key));
}
#[test]
fn token_bound_to_session_rejects_other_session() {
let key = [7u8; 32];
let token = mint_token(b"session-A".as_slice(), &key);
assert!(
validate_token(&token, b"session-A".as_slice(), &key),
"token must validate under the session it was minted for"
);
assert!(
!validate_token(&token, b"session-B".as_slice(), &key),
"token minted under session A must NOT validate under session B"
);
assert!(
validate_pair(Some(&token), Some(&token), b"session-A".as_slice(), &key),
"cookie==header token must validate under its own session"
);
assert!(
!validate_pair(Some(&token), Some(&token), b"session-B".as_slice(), &key),
"cookie==header token must NOT validate under a different session"
);
}
#[test]
fn validate_pair_requires_both_match_and_signature() {
let key = [7u8; 32];
let valid = mint_token(SID, &key);
assert!(validate_pair(Some(&valid), Some(&valid), SID, &key));
let other = mint_token(SID, &key);
assert!(!validate_pair(Some(&valid), Some(&other), SID, &key));
assert!(!validate_pair(None, Some(&valid), SID, &key));
assert!(!validate_pair(Some(&valid), None, SID, &key));
let forged =
"AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA";
assert!(!validate_pair(Some(forged), Some(forged), SID, &key));
}
#[test]
fn is_state_changing_only_unsafe_verbs() {
assert!(!is_state_changing(&axum::http::Method::GET));
assert!(!is_state_changing(&axum::http::Method::HEAD));
assert!(!is_state_changing(&axum::http::Method::OPTIONS));
assert!(is_state_changing(&axum::http::Method::POST));
assert!(is_state_changing(&axum::http::Method::PUT));
assert!(is_state_changing(&axum::http::Method::PATCH));
assert!(is_state_changing(&axum::http::Method::DELETE));
}
#[test]
fn validate_pair_rejects_empty_strings() {
let key = [7u8; 32];
assert!(!validate_pair(Some(""), Some(""), SID, &key));
let valid = mint_token(SID, &key);
assert!(!validate_pair(Some(""), Some(&valid), SID, &key));
assert!(!validate_pair(Some(&valid), Some(""), SID, &key));
}
#[test]
fn validate_token_rejects_non_base64() {
let key = [7u8; 32];
assert!(!validate_token("not-valid-base64!!!", SID, &key));
}
#[test]
fn validate_token_rejects_wrong_length_payload() {
let key = [7u8; 32];
let short = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(b"too_short");
assert!(!validate_token(&short, SID, &key));
}
#[test]
fn extract_cookie_token_parses_correctly() {
use axum::http::Request;
let req = Request::builder()
.header("cookie", "other=abc; axess.csrf=my_token; third=xyz")
.body(Body::empty())
.unwrap();
assert_eq!(
extract_cookie_token(&req, "axess.csrf"),
Some("my_token".to_string())
);
}
#[test]
fn extract_cookie_token_missing_returns_none() {
use axum::http::Request;
let req = Request::builder()
.header("cookie", "other=abc")
.body(Body::empty())
.unwrap();
assert_eq!(extract_cookie_token(&req, "axess.csrf"), None);
}
#[test]
fn extract_cookie_token_no_cookie_header_returns_none() {
use axum::http::Request;
let req = Request::builder().body(Body::empty()).unwrap();
assert_eq!(extract_cookie_token(&req, "axess.csrf"), None);
}
#[test]
fn extract_cookie_token_rejects_oversize_value() {
use axum::http::Request;
let oversize = "x".repeat(MAX_COOKIE_VALUE_BYTES + 1);
let header = format!("axess.csrf={oversize}");
let req = Request::builder()
.header("cookie", header)
.body(Body::empty())
.unwrap();
assert_eq!(extract_cookie_token(&req, "axess.csrf"), None);
}
#[test]
fn extract_cookie_token_accepts_value_at_cap() {
use axum::http::Request;
let at_cap = "x".repeat(MAX_COOKIE_VALUE_BYTES);
let header = format!("axess.csrf={at_cap}");
let req = Request::builder()
.header("cookie", header)
.body(Body::empty())
.unwrap();
assert_eq!(
extract_cookie_token(&req, "axess.csrf").map(|v| v.len()),
Some(MAX_COOKIE_VALUE_BYTES)
);
}
#[test]
fn csrf_token_as_str_returns_inner_value() {
let t = CsrfToken("abc.defg.hij".to_string());
assert_eq!(t.as_str(), "abc.defg.hij");
let empty = CsrfToken(String::new());
assert_eq!(empty.as_str(), "");
}
#[tokio::test]
async fn extract_presented_token_returns_header_value() {
use axum::http::Request;
let key = [9u8; 32];
let config = CsrfConfig::new(key);
let req = Request::builder()
.header(config.header_name.as_ref(), "presented-csrf-value")
.body(Body::empty())
.unwrap();
let (extracted, _req) = extract_presented_token(req, &config).await.unwrap();
assert_eq!(
extracted,
Some("presented-csrf-value".to_string()),
"must return the exact header value, not None / empty / 'xyzzy'"
);
let req = Request::builder().body(Body::empty()).unwrap();
let (extracted, _req) = extract_presented_token(req, &config).await.unwrap();
assert!(
extracted.is_none(),
"missing header must return None, not Some(...)"
);
let req = Request::builder()
.header(config.header_name.as_ref(), "")
.body(Body::empty())
.unwrap();
let (extracted, _req) = extract_presented_token(req, &config).await.unwrap();
assert!(extracted.is_none(), "empty header must return None");
}
#[tokio::test]
async fn extract_presented_token_reads_form_field_and_restores_body() {
use axum::http::Request;
use http_body_util::BodyExt;
let config = CsrfConfig::new([9u8; 32]);
let body_str = "username=alice&_csrf=form-token-value&password=x";
let req = Request::builder()
.method(axum::http::Method::POST)
.header(
axum::http::header::CONTENT_TYPE,
"application/x-www-form-urlencoded",
)
.body(Body::from(body_str))
.unwrap();
let (extracted, restored) = extract_presented_token(req, &config).await.unwrap();
assert_eq!(
extracted,
Some("form-token-value".to_string()),
"must extract the _csrf field from the urlencoded body"
);
let bytes = restored.into_body().collect().await.unwrap().to_bytes();
assert_eq!(
&bytes[..],
body_str.as_bytes(),
"body must be restored verbatim so handlers can still parse the form"
);
}
#[tokio::test]
async fn extract_presented_token_skips_multipart_body() {
use axum::http::Request;
let config = CsrfConfig::new([9u8; 32]);
let req = Request::builder()
.method(axum::http::Method::POST)
.header(
axum::http::header::CONTENT_TYPE,
"multipart/form-data; boundary=xxx",
)
.body(Body::from(b"unused".as_slice()))
.unwrap();
let (extracted, _req) = extract_presented_token(req, &config).await.unwrap();
assert!(
extracted.is_none(),
"multipart Content-Type must skip form-field extraction"
);
}
#[tokio::test]
async fn extract_presented_token_rejects_oversized_body() {
use axum::http::Request;
let config = CsrfConfig::new([9u8; 32]).form_body_limit(64);
let big_body = "a=".to_string() + &"x".repeat(500);
let req = Request::builder()
.method(axum::http::Method::POST)
.header(
axum::http::header::CONTENT_TYPE,
"application/x-www-form-urlencoded",
)
.header(axum::http::header::CONTENT_LENGTH, big_body.len())
.body(Body::from(big_body))
.unwrap();
let err = extract_presented_token(req, &config).await.unwrap_err();
assert!(
err.contains("cap"),
"oversized body must return a cap-related error, got {err:?}"
);
}
#[test]
fn origin_matches_allowed_exact_origin() {
use axum::http::Request;
let allowed: Vec<Arc<str>> = vec!["https://app.example.com".into()];
let req = Request::builder()
.header(axum::http::header::ORIGIN, "https://app.example.com")
.body(Body::empty())
.unwrap();
assert!(origin_matches_allowed(&req, &allowed));
let req = Request::builder()
.header(axum::http::header::ORIGIN, "http://app.example.com")
.body(Body::empty())
.unwrap();
assert!(!origin_matches_allowed(&req, &allowed));
let req = Request::builder()
.header(axum::http::header::ORIGIN, "https://evil.example.com")
.body(Body::empty())
.unwrap();
assert!(!origin_matches_allowed(&req, &allowed));
}
#[test]
fn origin_matches_allowed_rejects_null_origin() {
use axum::http::Request;
let allowed: Vec<Arc<str>> = vec!["https://app.example.com".into()];
let req = Request::builder()
.header(axum::http::header::ORIGIN, "null")
.body(Body::empty())
.unwrap();
assert!(
!origin_matches_allowed(&req, &allowed),
"'Origin: null' must be treated as failure"
);
}
#[test]
fn origin_matches_allowed_falls_back_to_referer() {
use axum::http::Request;
let allowed: Vec<Arc<str>> = vec!["https://app.example.com".into()];
let req = Request::builder()
.header(
axum::http::header::REFERER,
"https://app.example.com/login?next=/dashboard",
)
.body(Body::empty())
.unwrap();
assert!(
origin_matches_allowed(&req, &allowed),
"Referer with matching origin component must match when Origin is absent"
);
let req = Request::builder()
.header(axum::http::header::REFERER, "https://evil.example.com/pwn")
.body(Body::empty())
.unwrap();
assert!(!origin_matches_allowed(&req, &allowed));
}
#[test]
fn origin_matches_allowed_fails_when_both_headers_absent() {
use axum::http::Request;
let allowed: Vec<Arc<str>> = vec!["https://app.example.com".into()];
let req = Request::builder().body(Body::empty()).unwrap();
assert!(!origin_matches_allowed(&req, &allowed));
}
#[test]
fn origin_matches_allowed_is_case_insensitive() {
use axum::http::Request;
let allowed: Vec<Arc<str>> = vec!["https://app.example.com".into()];
let req = Request::builder()
.header(axum::http::header::ORIGIN, "HTTPS://APP.EXAMPLE.COM")
.body(Body::empty())
.unwrap();
assert!(
origin_matches_allowed(&req, &allowed),
"scheme+host comparison must be case-insensitive"
);
}
#[test]
fn origin_matches_allowed_strips_userinfo_from_referer() {
use axum::http::Request;
let allowed: Vec<Arc<str>> = vec!["https://app.example.com".into()];
let req = Request::builder()
.header(
axum::http::header::REFERER,
"https://attacker@app.example.com/",
)
.body(Body::empty())
.unwrap();
assert!(
origin_matches_allowed(&req, &allowed),
"userinfo must be stripped before origin comparison"
);
}
#[test]
fn require_origin_normalizes_default_port() {
let config = CsrfConfig::new([0u8; 32]).require_origin(["https://app.example.com:443"]);
assert_eq!(
config.allowed_origins.as_ref(),
&[Arc::<str>::from("https://app.example.com")]
);
}
#[test]
#[should_panic(expected = "all 2 supplied origin(s) failed to canonicalize")]
fn require_origin_panics_when_all_inputs_are_malformed() {
let _ = CsrfConfig::new([0u8; 32]).require_origin(["not-a-url", "javascript:alert(1)"]);
}
#[test]
fn require_origin_with_no_inputs_stays_disabled() {
let empty: [&str; 0] = [];
let config = CsrfConfig::new([0u8; 32]).require_origin(empty);
assert!(config.allowed_origins.is_empty());
}
#[tokio::test]
async fn csrf_service_end_to_end_drives_call_path() {
use axum::http::{Method, Request};
use std::convert::Infallible;
use tower::{Layer, ServiceExt, service_fn};
let echo_body = service_fn(|req: Request<Body>| {
tracing::trace!(method = %req.method(), uri = %req.uri(), "EchoBody call");
async move {
Ok::<_, Infallible>(Response::builder().status(200).body(Body::empty()).unwrap())
}
});
let key = [13u8; 32];
let config = CsrfConfig::new(key);
let service = CsrfLayer::new(config.clone()).layer(echo_body);
let handle = session_handle(1);
let req = Request::builder()
.method(Method::GET)
.uri("/safe")
.extension(handle.clone())
.body(Body::empty())
.unwrap();
let resp = service.clone().oneshot(req).await.unwrap();
assert_eq!(resp.status(), 200, "safe verb must pass through");
let set_cookie = resp
.headers()
.get(header::SET_COOKIE)
.expect("GET without cookie must mint a fresh CSRF cookie")
.to_str()
.unwrap()
.to_string();
assert!(
set_cookie.starts_with(&format!("{}=", config.cookie_name)),
"minted cookie must be named {}",
config.cookie_name
);
let token = set_cookie
.split('=')
.nth(1)
.unwrap()
.split(';')
.next()
.unwrap()
.to_string();
assert!(!token.is_empty(), "minted cookie value must not be empty");
let req = Request::builder()
.method(Method::GET)
.uri("/safe")
.header("cookie", format!("{}={}", config.cookie_name, token))
.extension(handle.clone())
.body(Body::empty())
.unwrap();
let resp = service.clone().oneshot(req).await.unwrap();
assert_eq!(resp.status(), 200);
assert!(
resp.headers().get(header::SET_COOKIE).is_none(),
"existing non-empty cookie must NOT trigger a fresh mint \
; pins `match guard !existing.is_empty() -> false` and `delete !`"
);
let req = Request::builder()
.method(Method::GET)
.uri("/safe")
.header("cookie", format!("{}=", config.cookie_name))
.extension(handle.clone())
.body(Body::empty())
.unwrap();
let resp = service.clone().oneshot(req).await.unwrap();
assert_eq!(resp.status(), 200);
let minted = resp
.headers()
.get(header::SET_COOKIE)
.expect(
"empty cookie value must trigger a fresh mint \
(otherwise the client never gets a token)",
)
.to_str()
.unwrap();
assert!(
minted.starts_with(&format!("{}=", config.cookie_name)),
"minted cookie must be named {}",
config.cookie_name
);
let minted_value = minted.split('=').nth(1).unwrap().split(';').next().unwrap();
assert!(
!minted_value.is_empty(),
"minted cookie value must not itself be empty"
);
let req = Request::builder()
.method(Method::POST)
.uri("/state-changing")
.header("cookie", format!("{}={}", config.cookie_name, token))
.header(config.header_name.as_ref(), &token)
.extension(handle.clone())
.body(Body::empty())
.unwrap();
let resp = service.clone().oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
200,
"POST with valid cookie+header must reach inner service \
; pins `delete !` on the `if !validate_pair(...)` guard at line 201"
);
let req = Request::builder()
.method(Method::POST)
.uri("/state-changing")
.extension(handle.clone())
.body(Body::empty())
.unwrap();
let resp = service.clone().oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::FORBIDDEN,
"state-changing request without tokens must be rejected as 403"
);
}
#[tokio::test]
async fn csrf_service_rejects_token_replayed_under_different_session() {
use axum::http::{Method, Request};
use std::convert::Infallible;
use tower::{Layer, ServiceExt, service_fn};
let echo_body = service_fn(|_req: Request<Body>| async move {
Ok::<_, Infallible>(Response::builder().status(200).body(Body::empty()).unwrap())
});
let key = [21u8; 32];
let config = CsrfConfig::new(key);
let service = CsrfLayer::new(config.clone()).layer(echo_body);
let handle_a = session_handle(1);
let handle_b = session_handle(2);
let req = Request::builder()
.method(Method::GET)
.uri("/safe")
.extension(handle_a.clone())
.body(Body::empty())
.unwrap();
let resp = service.clone().oneshot(req).await.unwrap();
let set_cookie = resp
.headers()
.get(header::SET_COOKIE)
.expect("GET under session A must mint a token")
.to_str()
.unwrap()
.to_string();
let token = set_cookie
.split('=')
.nth(1)
.unwrap()
.split(';')
.next()
.unwrap()
.to_string();
assert!(!token.is_empty());
let req = Request::builder()
.method(Method::POST)
.uri("/state-changing")
.header("cookie", format!("{}={}", config.cookie_name, token))
.header(config.header_name.as_ref(), &token)
.extension(handle_b.clone())
.body(Body::empty())
.unwrap();
let resp = service.clone().oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::FORBIDDEN,
"a token minted under session A must be rejected under session B"
);
let req = Request::builder()
.method(Method::POST)
.uri("/state-changing")
.header("cookie", format!("{}={}", config.cookie_name, token))
.header(config.header_name.as_ref(), &token)
.extension(handle_a.clone())
.body(Body::empty())
.unwrap();
let resp = service.clone().oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
200,
"the A-minted token must still be accepted under session A"
);
}
#[tokio::test]
async fn csrf_service_fails_closed_without_session_handle() {
use axum::http::{Method, Request};
use std::convert::Infallible;
use tower::{Layer, ServiceExt, service_fn};
let echo_body = service_fn(|_req: Request<Body>| async move {
Ok::<_, Infallible>(Response::builder().status(200).body(Body::empty()).unwrap())
});
let key = [23u8; 32];
let config = CsrfConfig::new(key);
let service = CsrfLayer::new(config.clone()).layer(echo_body);
let token = mint_token(b"some-session", &key);
let req = Request::builder()
.method(Method::POST)
.uri("/state-changing")
.header("cookie", format!("{}={}", config.cookie_name, token))
.header(config.header_name.as_ref(), &token)
.body(Body::empty())
.unwrap();
let resp = service.clone().oneshot(req).await.unwrap();
assert_eq!(
resp.status(),
StatusCode::FORBIDDEN,
"no session handle on a state-changing request must fail closed (403)"
);
}
#[tokio::test]
async fn csrf_service_remints_stale_cookie_on_safe_verb() {
use axum::http::{Method, Request};
use std::convert::Infallible;
use tower::{Layer, ServiceExt, service_fn};
let echo_body = service_fn(|_req: Request<Body>| async move {
Ok::<_, Infallible>(Response::builder().status(200).body(Body::empty()).unwrap())
});
let key = [29u8; 32];
let config = CsrfConfig::new(key);
let service = CsrfLayer::new(config.clone()).layer(echo_body);
let handle_a = session_handle(1);
let stale_token = mint_token(b"some-other-session", &key);
let req = Request::builder()
.method(Method::GET)
.uri("/safe")
.header("cookie", format!("{}={}", config.cookie_name, stale_token))
.extension(handle_a.clone())
.body(Body::empty())
.unwrap();
let resp = service.clone().oneshot(req).await.unwrap();
assert_eq!(resp.status(), 200, "safe verb must pass through");
let set_cookie = resp
.headers()
.get(header::SET_COOKIE)
.expect("stale cookie on safe verb must trigger a fresh mint")
.to_str()
.unwrap()
.to_string();
let refreshed = set_cookie
.split('=')
.nth(1)
.unwrap()
.split(';')
.next()
.unwrap()
.to_string();
assert!(!refreshed.is_empty(), "refreshed cookie must not be empty");
assert_ne!(
refreshed, stale_token,
"refreshed cookie must differ from the stale one"
);
let sid_a = handle_a.0.read().await.id;
assert!(
validate_token(&refreshed, sid_a.as_bytes(), &key),
"refreshed cookie must be HMAC-bound to the current session id"
);
}
#[tokio::test]
async fn csrf_service_remints_when_handler_rotates_session() {
use axum::http::{Method, Request};
use std::convert::Infallible;
use tower::{Layer, ServiceExt, service_fn};
let rotating_body = service_fn(|req: Request<Body>| async move {
if let Some(handle) = req.extensions().get::<SessionHandle>() {
let mut guard = handle.0.write().await;
guard.rotate_id();
guard.modified = true;
}
Ok::<_, Infallible>(Response::builder().status(200).body(Body::empty()).unwrap())
});
let key = [31u8; 32];
let config = CsrfConfig::new(key);
let service = CsrfLayer::new(config.clone()).layer(rotating_body);
let handle = session_handle(1);
let sid_before = handle.0.read().await.id;
let token_before = mint_token(sid_before.as_bytes(), &key);
let req = Request::builder()
.method(Method::GET)
.uri("/rotate")
.header("cookie", format!("{}={}", config.cookie_name, token_before))
.extension(handle.clone())
.body(Body::empty())
.unwrap();
let resp = service.clone().oneshot(req).await.unwrap();
assert_eq!(resp.status(), 200);
let sid_after = handle.0.read().await.id;
assert_ne!(
sid_before, sid_after,
"test precondition: inner service rotated the session id"
);
let set_cookie = resp
.headers()
.get(header::SET_COOKIE)
.expect("session rotation mid-request must trigger a fresh Set-Cookie")
.to_str()
.unwrap()
.to_string();
let refreshed = set_cookie
.split('=')
.nth(1)
.unwrap()
.split(';')
.next()
.unwrap()
.to_string();
assert_ne!(
refreshed, token_before,
"refreshed cookie must differ from the pre-rotation token"
);
assert!(
validate_token(&refreshed, sid_after.as_bytes(), &key),
"refreshed cookie must bind to the post-rotation session id"
);
}
#[tokio::test]
async fn csrf_token_for_fresh_guest_survives_session_finalize() {
use crate::session::layer::SessionLayer;
use crate::session::store::MemorySessionStore;
use axum::Router;
use axum::http::Method;
use axum::routing::{get, post};
use tower::ServiceExt;
let session_layer =
SessionLayer::new(MemorySessionStore::new(), [42u8; 32]).with_secure(false);
let csrf_layer = CsrfLayer::new(CsrfConfig::new([7u8; 32]).secure(false));
let app = Router::new()
.route("/safe", get(|| async { "ok" }))
.route("/change", post(|| async { "changed" }))
.layer(csrf_layer) .layer(session_layer);
let resp = app
.clone()
.oneshot(
Request::builder()
.method(Method::GET)
.uri("/safe")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let cookie = |name: &str| {
resp.headers()
.get_all(header::SET_COOKIE)
.iter()
.filter_map(|hv| hv.to_str().ok())
.find_map(|c| {
let (k, v) = c.split(';').next()?.split_once('=')?;
(k.trim() == name).then(|| v.trim().to_string())
})
};
let sid = cookie("axess.sid").expect("GET must set the session cookie");
let csrf = cookie(DEFAULT_CSRF_COOKIE).expect("GET must mint the CSRF cookie");
let resp = app
.oneshot(
Request::builder()
.method(Method::POST)
.uri("/change")
.header(
"cookie",
format!("axess.sid={sid}; {DEFAULT_CSRF_COOKIE}={csrf}"),
)
.header(DEFAULT_CSRF_HEADER, &csrf)
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(
resp.status(),
StatusCode::OK,
"a CSRF token minted for a fresh guest must validate on its next \
state-changing request; a 403 here means finalize handed the \
response cookie a different session id than the token was bound to"
);
}
}