1#[cfg(feature = "server")]
11use axum::http::{HeaderMap, Method, Request};
12
13#[cfg(feature = "server")]
15fn is_write_method(method: &Method) -> bool {
16 matches!(
17 method,
18 &Method::POST | &Method::PUT | &Method::PATCH | &Method::DELETE
19 )
20}
21
22#[cfg(feature = "server")]
32pub fn warn_if_app_base_url_unset() {
33 if app_base_url_is_set() {
34 return;
35 }
36 tracing::warn!(
37 "APP_BASE_URL 未设置。CSRF 校验将回退到请求 Host 头推导本站 origin,\
38 反向代理后若 Host 头可被客户端影响存在绕过风险。\
39 生产环境应显式设置为站点完整 origin,如 https://your-domain.example。"
40 );
41}
42
43#[cfg(feature = "server")]
45fn app_base_url_is_set() -> bool {
46 std::env::var("APP_BASE_URL")
47 .ok()
48 .map(|v| !v.trim().is_empty())
49 .unwrap_or(false)
50}
51
52#[cfg(feature = "server")]
58fn normalize_origin(input: &str) -> String {
59 let (scheme, rest) = match input.split_once("://") {
61 Some(pair) => pair,
62 None => return input.to_string(),
63 };
64 let authority = match rest.split_once('/') {
66 Some((auth, _)) => auth,
67 None => rest,
68 };
69 match authority.rsplit_once(':') {
71 Some((host, port)) if port == "80" || port == "443" => {
72 format!("{}://{}", scheme, host)
73 }
74 _ => format!("{}://{}", scheme, authority),
75 }
76}
77
78#[cfg(feature = "server")]
82fn extract_origin(headers: &HeaderMap) -> Option<String> {
83 if let Some(origin) = headers.get(axum::http::header::ORIGIN) {
84 return origin.to_str().ok().map(normalize_origin);
85 }
86 headers
87 .get(axum::http::header::REFERER)
88 .and_then(|v| v.to_str().ok())
89 .map(normalize_origin)
90}
91
92#[cfg(feature = "server")]
98fn trusted_origin(headers: &HeaderMap) -> Option<String> {
99 if let Ok(base) = std::env::var("APP_BASE_URL") {
100 return Some(normalize_origin(&base));
101 }
102 let host = headers.get(axum::http::header::HOST)?.to_str().ok()?;
103 let proto = headers
104 .get("X-Forwarded-Proto")
105 .and_then(|v| v.to_str().ok())
106 .unwrap_or("https");
107 Some(normalize_origin(&format!("{}://{}", proto, host)))
108}
109
110#[cfg(feature = "server")]
115pub async fn csrf_middleware(
116 req: Request<axum::body::Body>,
117 next: axum::middleware::Next,
118) -> axum::response::Response {
119 if is_write_method(req.method()) {
120 let headers = req.headers().clone();
121 let trusted = trusted_origin(&headers);
122 let incoming = extract_origin(&headers);
123 let ok = match (&trusted, &incoming) {
124 (Some(t), Some(o)) => t == o,
125 _ => true,
127 };
128 if !ok {
129 return axum::response::Response::builder()
130 .status(axum::http::StatusCode::FORBIDDEN)
131 .body(axum::body::Body::empty())
132 .expect("static forbidden response is always valid");
133 }
134 }
135 next.run(req).await
136}
137
138#[cfg(all(test, feature = "server"))]
139mod tests {
140 use super::*;
141 use axum::http::{HeaderMap, HeaderValue, Method};
142
143 #[test]
144 fn is_write_method_recognizes_writes() {
145 assert!(is_write_method(&Method::POST));
146 assert!(is_write_method(&Method::PUT));
147 assert!(is_write_method(&Method::PATCH));
148 assert!(is_write_method(&Method::DELETE));
149 assert!(!is_write_method(&Method::GET));
150 assert!(!is_write_method(&Method::OPTIONS));
151 assert!(!is_write_method(&Method::HEAD));
152 }
153
154 #[test]
155 fn normalize_strips_path_and_query() {
156 assert_eq!(
157 normalize_origin("https://example.com/a/b?c=1"),
158 "https://example.com"
159 );
160 }
161
162 #[test]
163 fn normalize_preserves_nondefault_port() {
164 assert_eq!(
165 normalize_origin("http://localhost:3000/x"),
166 "http://localhost:3000"
167 );
168 }
169
170 #[test]
171 fn normalize_drops_default_ports() {
172 assert_eq!(
173 normalize_origin("https://example.com:443/path"),
174 "https://example.com"
175 );
176 assert_eq!(
177 normalize_origin("http://example.com:80/path"),
178 "http://example.com"
179 );
180 }
181
182 #[test]
183 fn normalize_keeps_explicit_nondefault_https_port() {
184 assert_eq!(
185 normalize_origin("https://example.com:8443"),
186 "https://example.com:8443"
187 );
188 }
189
190 #[test]
191 fn normalize_plain_origin_no_path() {
192 assert_eq!(
193 normalize_origin("https://example.com"),
194 "https://example.com"
195 );
196 }
197
198 #[test]
199 fn extract_origin_prefers_origin_header() {
200 let mut headers = HeaderMap::new();
201 headers.insert(
202 axum::http::header::ORIGIN,
203 HeaderValue::from_static("https://example.com"),
204 );
205 assert_eq!(
206 extract_origin(&headers),
207 Some("https://example.com".to_string())
208 );
209 }
210
211 #[test]
212 fn extract_origin_falls_back_to_referer() {
213 let mut headers = HeaderMap::new();
214 headers.insert(
215 axum::http::header::REFERER,
216 HeaderValue::from_static("https://example.com/posts/1"),
217 );
218 assert_eq!(
220 extract_origin(&headers),
221 Some("https://example.com".to_string())
222 );
223 }
224
225 #[test]
226 fn extract_origin_returns_none_when_both_absent() {
227 let headers = HeaderMap::new();
228 assert_eq!(extract_origin(&headers), None);
229 }
230
231 #[test]
236 #[serial_test::serial]
237 fn app_base_url_is_set_false_when_unset() {
238 let original = std::env::var("APP_BASE_URL").ok();
239 std::env::remove_var("APP_BASE_URL");
240 assert!(!app_base_url_is_set());
241 restore_env("APP_BASE_URL", original);
242 }
243
244 #[test]
245 #[serial_test::serial]
246 fn app_base_url_is_set_false_when_empty() {
247 let original = std::env::var("APP_BASE_URL").ok();
248 std::env::set_var("APP_BASE_URL", "");
249 assert!(!app_base_url_is_set());
250 restore_env("APP_BASE_URL", original);
251 }
252
253 #[test]
254 #[serial_test::serial]
255 fn app_base_url_is_set_false_when_whitespace_only() {
256 let original = std::env::var("APP_BASE_URL").ok();
257 std::env::set_var("APP_BASE_URL", " \t ");
258 assert!(!app_base_url_is_set());
259 restore_env("APP_BASE_URL", original);
260 }
261
262 #[test]
263 #[serial_test::serial]
264 fn app_base_url_is_set_true_when_set() {
265 let original = std::env::var("APP_BASE_URL").ok();
266 std::env::set_var("APP_BASE_URL", "https://example.com");
267 assert!(app_base_url_is_set());
268 restore_env("APP_BASE_URL", original);
269 }
270
271 #[test]
272 #[serial_test::serial]
273 fn app_base_url_is_set_trims_surrounding_whitespace() {
274 let original = std::env::var("APP_BASE_URL").ok();
275 std::env::set_var("APP_BASE_URL", " https://example.com ");
276 assert!(app_base_url_is_set());
277 restore_env("APP_BASE_URL", original);
278 }
279
280 fn restore_env(key: &str, original: Option<String>) {
282 match original {
283 Some(value) => std::env::set_var(key, value),
284 None => std::env::remove_var(key),
285 }
286 }
287}