feat(security): add Origin-based CSRF protection for write endpoints

对所有 POST/PUT/PATCH/DELETE 请求校验 Origin(回退 Referer)等于本站,
堵住 login CSRF 与未来 GET 化写接口的盲区(SameSite=Lax 覆盖不到)。

- 新增 src/api/csrf.rs:纯函数 origin 解析(不引入 url crate)+ axum 中间件
- 挂载到 upload 路由与 Dioxus app 路由(最外层,先于超时/压缩)
- APP_BASE_URL 配置可信域名;未设置时回退 Host + X-Forwarded-Proto
- GET/OPTIONS 放行;拿不到本站 origin 时放行避免误杀
- 9 个单测覆盖写方法识别、origin 标准化、默认端口省略、头解析回退
This commit is contained in:
xfy 2026-06-18 13:22:59 +08:00
parent 00478e4a1a
commit 82a3c12940
4 changed files with 213 additions and 2 deletions

View File

@ -50,6 +50,7 @@ Session / security tuning:
```
COOKIE_SECURE=false # set true/1/yes to add Secure flag to session cookie
TRUSTED_PROXY_COUNT=0 # number of reverse proxies in front of the app; used to extract real client IP from X-Forwarded-For
APP_BASE_URL= # e.g. https://your-domain.example — trusted origin for CSRF checks on write requests; unset falls back to Host header + X-Forwarded-Proto
```
Run migrations before first dev server start:

200
src/api/csrf.rs Normal file
View File

@ -0,0 +1,200 @@
//! CSRF 防护:对写请求校验 Origin回退 Referer必须等于本站。
//!
//! SameSite=Lax 只在顶级 GET 导航时自动带 cookie对跨站 POST 不带 cookie
//! 挡住了大部分经典 CSRF。但存在两个 Lax 覆盖不到的盲区:
//! 1. 登录 CSRF攻击者诱导受害者登录攻击者账号Lax 不阻止「设置」cookie
//! 2. 未来若出现 GET 化写接口Lax 会在顶级 GET 导航时带 cookie。
//! 因此对所有写请求叠加 Origin 校验作为纵深防御。
//! 仅在 `feature = "server"` 时编译。
#[cfg(feature = "server")]
use axum::http::{HeaderMap, Method, Request};
/// 判断请求是否需要 CSRF 校验非简单方法POST/PUT/PATCH/DELETE需要。
#[cfg(feature = "server")]
fn is_write_method(method: &Method) -> bool {
matches!(
method,
&Method::POST | &Method::PUT | &Method::PATCH | &Method::DELETE
)
}
/// 从 `scheme://host[:port][/path][?query]` 提取标准化的 `scheme://host[:port]`
/// 端口为默认值http=80, https=443时省略。
///
/// 不引入 url crateOrigin 头本身就是 `scheme://host[:port]`(无路径),
/// Referer 需要剥离 path/query用简单的 split 即可。
#[cfg(feature = "server")]
fn normalize_origin(input: &str) -> String {
// 取 authority 之前的部分作为 scheme以及第一个 '/' 之前的部分作为 authority。
let (scheme, rest) = match input.split_once("://") {
Some(pair) => pair,
None => return input.to_string(),
};
// rest 形如 host[:port]/path?query去掉首个 '/' 及之后内容。
let authority = match rest.split_once('/') {
Some((auth, _)) => auth,
None => rest,
};
// 省略默认端口。
match authority.rsplit_once(':') {
Some((host, port)) if port == "80" || port == "443" => {
format!("{}://{}", scheme, host)
}
_ => format!("{}://{}", scheme, authority),
}
}
/// 从请求头解析来源站点Origin 优先,回退 Referer
///
/// 返回标准化的 `scheme://host[:port]`。两者都缺失时返回 None视为不可信
#[cfg(feature = "server")]
fn extract_origin(headers: &HeaderMap) -> Option<String> {
if let Some(origin) = headers.get(axum::http::header::ORIGIN) {
return origin.to_str().ok().map(|s| normalize_origin(s));
}
headers
.get(axum::http::header::REFERER)
.and_then(|v| v.to_str().ok())
.map(|s| normalize_origin(s))
}
/// 计算本站可信 origin优先 `APP_BASE_URL` 环境变量(生产域名),
/// 否则用请求 Host 头 + `X-Forwarded-Proto`(反代后)或 https 推导。
///
/// 返回 None 表示无法确定本站 origin此时放行避免误杀——CSRF 漏判
/// 是请求被拒,但拿不到本站 origin 时误杀合法请求代价更高,故保守放行)。
#[cfg(feature = "server")]
fn trusted_origin(headers: &HeaderMap) -> Option<String> {
if let Ok(base) = std::env::var("APP_BASE_URL") {
return Some(normalize_origin(&base));
}
let host = headers.get(axum::http::header::HOST)?.to_str().ok()?;
let proto = headers
.get("X-Forwarded-Proto")
.and_then(|v| v.to_str().ok())
.unwrap_or("https");
Some(normalize_origin(&format!("{}://{}", proto, host)))
}
/// CSRF 校验中间件。
///
/// 对写方法校验请求来源等于本站;不匹配返回 403。GET/OPTIONS 等放行。
/// 拿不到本站 origin 或请求来源时放行(见 trusted_origin 注释)。
#[cfg(feature = "server")]
pub async fn csrf_middleware(
req: Request<axum::body::Body>,
next: axum::middleware::Next,
) -> axum::response::Response {
if is_write_method(req.method()) {
let headers = req.headers().clone();
let trusted = trusted_origin(&headers);
let incoming = extract_origin(&headers);
let ok = match (&trusted, &incoming) {
(Some(t), Some(o)) => t == o,
// 拿不到本站 origin 或请求来源时放行。
_ => true,
};
if !ok {
return axum::response::Response::builder()
.status(axum::http::StatusCode::FORBIDDEN)
.body(axum::body::Body::empty())
.expect("static forbidden response is always valid");
}
}
next.run(req).await
}
#[cfg(all(test, feature = "server"))]
mod tests {
use super::*;
use axum::http::{HeaderMap, HeaderValue, Method};
#[test]
fn is_write_method_recognizes_writes() {
assert!(is_write_method(&Method::POST));
assert!(is_write_method(&Method::PUT));
assert!(is_write_method(&Method::PATCH));
assert!(is_write_method(&Method::DELETE));
assert!(!is_write_method(&Method::GET));
assert!(!is_write_method(&Method::OPTIONS));
assert!(!is_write_method(&Method::HEAD));
}
#[test]
fn normalize_strips_path_and_query() {
assert_eq!(
normalize_origin("https://example.com/a/b?c=1"),
"https://example.com"
);
}
#[test]
fn normalize_preserves_nondefault_port() {
assert_eq!(
normalize_origin("http://localhost:3000/x"),
"http://localhost:3000"
);
}
#[test]
fn normalize_drops_default_ports() {
assert_eq!(
normalize_origin("https://example.com:443/path"),
"https://example.com"
);
assert_eq!(
normalize_origin("http://example.com:80/path"),
"http://example.com"
);
}
#[test]
fn normalize_keeps_explicit_nondefault_https_port() {
assert_eq!(
normalize_origin("https://example.com:8443"),
"https://example.com:8443"
);
}
#[test]
fn normalize_plain_origin_no_path() {
assert_eq!(
normalize_origin("https://example.com"),
"https://example.com"
);
}
#[test]
fn extract_origin_prefers_origin_header() {
let mut headers = HeaderMap::new();
headers.insert(
axum::http::header::ORIGIN,
HeaderValue::from_static("https://example.com"),
);
assert_eq!(
extract_origin(&headers),
Some("https://example.com".to_string())
);
}
#[test]
fn extract_origin_falls_back_to_referer() {
let mut headers = HeaderMap::new();
headers.insert(
axum::http::header::REFERER,
HeaderValue::from_static("https://example.com/posts/1"),
);
// Referer 的路径被剥离,只保留 scheme://host。
assert_eq!(
extract_origin(&headers),
Some("https://example.com".to_string())
);
}
#[test]
fn extract_origin_returns_none_when_both_absent() {
let headers = HeaderMap::new();
assert_eq!(extract_origin(&headers), None);
}
}

View File

@ -6,6 +6,8 @@
/// 认证相关的 Dioxus server function。
pub mod auth;
/// CSRF 防护中间件。
pub mod csrf;
/// 评论相关接口。
pub mod comments;
/// 应用错误类型与转换。

View File

@ -271,6 +271,7 @@ fn main() {
}
// 自定义 API 路由:图片上传(大文件,需要更长超时)
// CSRF 校验置于最外层,先拦截非法来源再做超时/限体。
let upload_route = axum::Router::new()
.route(
"/api/upload",
@ -280,16 +281,23 @@ fn main() {
.layer(TimeoutLayer::with_status_code(
StatusCode::REQUEST_TIMEOUT,
Duration::from_secs(300),
))
.layer(axum::middleware::from_fn(
crate::api::csrf::csrf_middleware,
));
// Dioxus 应用路由:自动挂载所有 server function 并渲染前端组件
let dioxus_app =
axum::Router::new().serve_dioxus_application(config, router::AppRouter);
// 合并 Dioxus + 世代号/缓存头/可选压缩/30s 超时中间件
// 合并 Dioxus + CSRF/世代号/缓存头/可选压缩/30s 超时中间件
// layer 顺序后加的最外层先执行。CSRF 最外层先拦截非法来源。
let mut app_routes = dioxus_app
.layer(axum::middleware::from_fn(ssr_generation_middleware))
.layer(axum::middleware::from_fn(add_cache_control));
.layer(axum::middleware::from_fn(add_cache_control))
.layer(axum::middleware::from_fn(
crate::api::csrf::csrf_middleware,
));
if let Some(layer) = compression_layer_from_env() {
app_routes = app_routes.layer(layer);
}