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>]>,
}
fn should_warn_dropped_origins(kept: usize, rejected: usize) -> bool {
kept > 0 && rejected > 0
}
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,
);
if should_warn_dropped_origins(normalized.len(), rejected.len()) {
tracing::warn!(
rejected = ?rejected,
kept = normalized.len(),
"CsrfConfig::require_origin: {} supplied origin(s) failed to \
canonicalize and were dropped. Requests from them will be \
rejected as if they were never allow-listed.",
rejected.len(),
);
}
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;