fix(admin): sign CSRF tokens and escape them in scripts (#2071)

This commit is contained in:
Hampus
2026-08-29 14:54:51 +02:00
committed by GitHub
parent 0496b2f530
commit 79064c3399
5 changed files with 97 additions and 64 deletions
+59 -61
View File
@@ -2,24 +2,43 @@
use axum::{
body::{Body, to_bytes},
extract::Request,
extract::{Request, State},
http::{HeaderValue, Method, StatusCode, header},
middleware::Next,
response::{IntoResponse, Response},
};
use rand::RngExt;
use crate::middleware::auth::AuthContext;
use crate::session::{create_csrf_token, verify_csrf_token};
use crate::state::AppState;
const CSRF_COOKIE_NAME: &str = "csrf_token";
pub const CSRF_FORM_FIELD: &str = "_csrf";
const CSRF_HEADER_NAME: &str = "x-csrf-token";
const TOKEN_LENGTH: usize = 32;
const HOST_CSRF_COOKIE_NAME: &str = "__Host-csrf_token";
const MAX_CSRF_FORM_BYTES: usize = 8 * 1024 * 1024;
const IGNORED_PATH_SUFFIXES: &[&str] = &["/oauth2_callback", "/auth/start"];
pub async fn csrf_protection(mut request: Request, next: Next) -> Response {
let existing_token = extract_csrf_cookie(&request);
let token = existing_token.unwrap_or_else(generate_csrf_token);
pub async fn csrf_protection(
State(state): State<AppState>,
mut request: Request,
next: Next,
) -> Response {
let config = state.config();
let secret = config.secret_key_base.clone();
let admin_endpoint = config.admin_endpoint.clone();
let is_production = config.is_production();
let user_id = request
.extensions()
.get::<AuthContext>()
.map(|ctx| ctx.session.user_id.clone())
.unwrap_or_default();
let token = extract_csrf_cookie(&request)
.filter(|cookie| verify_csrf_token(cookie, &user_id, &secret))
.unwrap_or_else(|| create_csrf_token(&user_id, &secret));
request.extensions_mut().insert(CsrfToken(token.clone()));
if matches!(
@@ -31,6 +50,9 @@ pub async fn csrf_protection(mut request: Request, next: Next) -> Response {
.iter()
.any(|suffix| path.ends_with(suffix));
if !is_ignored {
if !is_same_site_request(&request, &admin_endpoint) {
return StatusCode::FORBIDDEN.into_response();
}
let header_token = extract_csrf_header(&request);
let query_token = extract_csrf_from_query(&request);
let mut submitted = query_token.or(header_token);
@@ -43,21 +65,24 @@ pub async fn csrf_protection(mut request: Request, next: Next) -> Response {
request = restored_request;
submitted = body_token;
}
match submitted {
Some(ref submitted_token) if submitted_token == &token => {}
_ => {
return StatusCode::FORBIDDEN.into_response();
}
let accepted = submitted.as_deref().is_some_and(|submitted_token| {
submitted_token == token && verify_csrf_token(submitted_token, &user_id, &secret)
});
if !accepted {
return StatusCode::FORBIDDEN.into_response();
}
}
}
let mut response = next.run(request).await;
let cookie_value = format!(
"{}={}; Path=/; SameSite=Lax; HttpOnly",
CSRF_COOKIE_NAME, token
);
let cookie_name = if is_production {
HOST_CSRF_COOKIE_NAME
} else {
CSRF_COOKIE_NAME
};
let secure = if is_production { "; Secure" } else { "" };
let cookie_value = format!("{cookie_name}={token}; Path=/; SameSite=Lax; HttpOnly{secure}");
if let Ok(value) = HeaderValue::from_str(&cookie_value) {
response.headers_mut().append(header::SET_COOKIE, value);
}
@@ -69,7 +94,10 @@ fn extract_csrf_cookie(request: &Request) -> Option<String> {
let cookie_header = request.headers().get(header::COOKIE)?.to_str().ok()?;
for pair in cookie_header.split(';') {
let pair = pair.trim();
if let Some(value) = pair.strip_prefix("csrf_token=") {
let value = pair
.strip_prefix("__Host-csrf_token=")
.or_else(|| pair.strip_prefix("csrf_token="));
if let Some(value) = value {
let trimmed = value.trim();
if !trimmed.is_empty() {
return Some(trimmed.to_owned());
@@ -127,18 +155,22 @@ async fn extract_csrf_from_form_body(
Ok((request, token))
}
fn generate_csrf_token() -> String {
let mut rng = rand::rng();
let bytes: [u8; TOKEN_LENGTH] = rng.random();
hex_encode(&bytes)
}
fn hex_encode(bytes: &[u8]) -> String {
let mut s = String::with_capacity(bytes.len() * 2);
for byte in bytes {
s.push_str(&format!("{byte:02x}"));
fn is_same_site_request(request: &Request, admin_endpoint: &str) -> bool {
if let Some(site) = request
.headers()
.get("sec-fetch-site")
.and_then(|value| value.to_str().ok())
{
return matches!(site, "same-origin" | "same-site" | "none");
}
match request
.headers()
.get(header::ORIGIN)
.and_then(|value| value.to_str().ok())
{
Some(origin) => origin == admin_endpoint,
None => true,
}
s
}
#[derive(Clone, Debug)]
@@ -156,40 +188,6 @@ pub fn get_csrf_token(request: &Request) -> String {
mod tests {
use super::*;
#[test]
fn generate_csrf_token_correct_length() {
let token = generate_csrf_token();
assert_eq!(
token.len(),
TOKEN_LENGTH * 2,
"token must be {} hex chars",
TOKEN_LENGTH * 2
);
}
#[test]
fn generate_csrf_token_is_valid_hex() {
let token = generate_csrf_token();
assert!(
token.chars().all(|c| c.is_ascii_hexdigit()),
"token must contain only hex chars: {token}"
);
}
#[test]
fn generate_csrf_token_is_unique() {
let a = generate_csrf_token();
let b = generate_csrf_token();
assert_ne!(a, b, "consecutive tokens must differ");
}
#[test]
fn hex_encode_produces_correct_output() {
assert_eq!(hex_encode(&[0x00, 0xff, 0x0a]), "00ff0a");
assert_eq!(hex_encode(&[]), "");
assert_eq!(hex_encode(&[0xde, 0xad]), "dead");
}
#[test]
fn oauth2_callback_is_exempt() {
let exempt = IGNORED_PATH_SUFFIXES
+1 -1
View File
@@ -348,7 +348,7 @@ fn oauth_callback_page(config: &AdminConfig, code: Option<&str>, state: Option<&
)
}
fn json_string(value: &str) -> String {
pub(crate) fn json_string(value: &str) -> String {
serde_json::to_string(value)
.expect("JSON string serialization cannot fail")
.replace('<', "\\u003c")
+4 -1
View File
@@ -73,7 +73,10 @@ pub fn build_router(config: AdminConfig) -> Router {
.route("/", get(dashboard))
.route("/dashboard", get(dashboard))
.layer(from_fn(middleware::htmx::flash_redirect_to_toast))
.layer(from_fn(middleware::csrf::csrf_protection))
.layer(from_fn_with_state(
state.clone(),
middleware::csrf::csrf_protection,
))
.layer(from_fn(middleware::self_hosted::self_hosted_override))
.layer(from_fn_with_state(
state.clone(),
+31
View File
@@ -3,6 +3,7 @@
use base64::{Engine as _, engine::general_purpose::URL_SAFE_NO_PAD};
use hmac::{Hmac, KeyInit, Mac};
use serde::{Deserialize, Serialize};
use rand::RngExt;
use sha2::Sha256;
use std::time::{SystemTime, UNIX_EPOCH};
@@ -78,6 +79,36 @@ fn verify_signature<'a>(signed_data: &'a str, secret_key: &str) -> Option<&'a st
if diff == 0 { Some(data) } else { None }
}
pub fn create_csrf_token(user_id: &str, secret_key: &str) -> String {
let mut rng = rand::rng();
let nonce: [u8; 16] = rng.random();
let payload = URL_SAFE_NO_PAD.encode(format!(
"{}:{}",
URL_SAFE_NO_PAD.encode(user_id.as_bytes()),
URL_SAFE_NO_PAD.encode(nonce)
));
sign_data(&payload, secret_key)
}
pub fn verify_csrf_token(token: &str, user_id: &str, secret_key: &str) -> bool {
let Some(data) = verify_signature(token, secret_key) else {
return false;
};
let Ok(decoded) = URL_SAFE_NO_PAD.decode(data.as_bytes()) else {
return false;
};
let Ok(payload) = String::from_utf8(decoded) else {
return false;
};
let Some((encoded_uid, _nonce)) = payload.split_once(':') else {
return false;
};
match URL_SAFE_NO_PAD.decode(encoded_uid.as_bytes()) {
Ok(uid_bytes) => uid_bytes == user_id.as_bytes(),
Err(_) => false,
}
}
pub const LEGACY_SESSION_COOKIE_NAME: &str = "session";
pub const SESSION_COOKIE_NAME: &str = "admin_session";
pub const SESSION_MAX_AGE: i64 = MAX_AGE_SECONDS as i64;
@@ -7,6 +7,7 @@ use super::media::user_avatar_url;
use super::nsfw_indicators::{attachment_nsfw_badge, channel_nsfw_state_badge};
use super::user_display::format_user_display;
use crate::config::AdminConfig;
use crate::routes::auth::json_string;
pub struct Attachment {
pub id: String,
@@ -309,7 +310,7 @@ pub fn message_list(
}
pub fn message_deletion_script(csrf_token: &str) -> Markup {
let csrf = serde_json::to_string(csrf_token).unwrap_or_else(|_| "\"\"".into());
let csrf = json_string(csrf_token);
let script = r#"(function() {
var csrf = __CSRF__;
function bp() {