From 82a3c129409547b45028e7590d81a49343209da4 Mon Sep 17 00:00:00 2001 From: xfy Date: Thu, 18 Jun 2026 13:22:59 +0800 Subject: [PATCH] feat(security): add Origin-based CSRF protection for write endpoints MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 对所有 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 标准化、默认端口省略、头解析回退 --- AGENTS.md | 1 + src/api/csrf.rs | 200 ++++++++++++++++++++++++++++++++++++++++++++++++ src/api/mod.rs | 2 + src/main.rs | 12 ++- 4 files changed, 213 insertions(+), 2 deletions(-) create mode 100644 src/api/csrf.rs diff --git a/AGENTS.md b/AGENTS.md index 6a2ac39..3dc3eb4 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -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: diff --git a/src/api/csrf.rs b/src/api/csrf.rs new file mode 100644 index 0000000..82c210f --- /dev/null +++ b/src/api/csrf.rs @@ -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 crate:Origin 头本身就是 `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 { + 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 { + 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, + 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); + } +} diff --git a/src/api/mod.rs b/src/api/mod.rs index fe97f9c..a2df475 100644 --- a/src/api/mod.rs +++ b/src/api/mod.rs @@ -6,6 +6,8 @@ /// 认证相关的 Dioxus server function。 pub mod auth; +/// CSRF 防护中间件。 +pub mod csrf; /// 评论相关接口。 pub mod comments; /// 应用错误类型与转换。 diff --git a/src/main.rs b/src/main.rs index c440869..fc51f4c 100644 --- a/src/main.rs +++ b/src/main.rs @@ -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); }