1#[cfg(feature = "server")]
15use axum::http::StatusCode;
16#[cfg(feature = "server")]
17use governor::{DefaultKeyedRateLimiter, Quota, RateLimiter};
18#[cfg(feature = "server")]
19use std::num::NonZeroU32;
20#[cfg(feature = "server")]
21use std::sync::LazyLock;
22#[cfg(feature = "server")]
23use std::time::Duration;
24
25#[cfg(feature = "server")]
26fn env_or(key: &str, default: u32) -> NonZeroU32 {
27 let val = std::env::var(key)
28 .ok()
29 .and_then(|s| s.parse::<u32>().ok())
30 .unwrap_or(default);
31 NonZeroU32::new(val.max(1)).expect("val.max(1) 保证非零,NonZeroU32::new 不可能失败")
33}
34
35#[cfg(feature = "server")]
36static STRICT_LIMITER: LazyLock<DefaultKeyedRateLimiter<String>> = LazyLock::new(|| {
37 RateLimiter::keyed(
38 Quota::per_second(env_or("RATE_LIMIT_STRICT_PER_SEC", 1))
39 .allow_burst(env_or("RATE_LIMIT_STRICT_BURST", 5)),
40 )
41});
42
43#[cfg(feature = "server")]
44static UPLOAD_LIMITER: LazyLock<DefaultKeyedRateLimiter<String>> = LazyLock::new(|| {
45 RateLimiter::keyed(
46 Quota::per_second(env_or("RATE_LIMIT_UPLOAD_PER_SEC", 2))
47 .allow_burst(env_or("RATE_LIMIT_UPLOAD_BURST", 15)),
48 )
49});
50
51#[cfg(feature = "server")]
52static IMAGE_LIMITER: LazyLock<DefaultKeyedRateLimiter<String>> = LazyLock::new(|| {
53 RateLimiter::keyed(
54 Quota::per_second(env_or("RATE_LIMIT_IMAGE_PER_SEC", 10))
55 .allow_burst(env_or("RATE_LIMIT_IMAGE_BURST", 50)),
56 )
57});
58
59#[cfg(feature = "server")]
60static COMMENT_LIMITER: LazyLock<DefaultKeyedRateLimiter<String>> = LazyLock::new(|| {
61 RateLimiter::keyed(
62 Quota::per_second(env_or("RATE_LIMIT_COMMENT_PER_SEC", 1))
63 .allow_burst(env_or("RATE_LIMIT_COMMENT_BURST", 5)),
64 )
65});
66
67#[cfg(feature = "server")]
68static CODE_EXEC_LIMITER: LazyLock<DefaultKeyedRateLimiter<String>> = LazyLock::new(|| {
71 RateLimiter::keyed(
72 Quota::per_second(env_or("RATE_LIMIT_CODE_EXEC_PER_SEC", 1))
73 .allow_burst(env_or("RATE_LIMIT_CODE_EXEC_BURST", 3)),
74 )
75});
76
77#[cfg(feature = "server")]
78static CODE_EXEC_DAILY_LIMITER: LazyLock<DefaultKeyedRateLimiter<String>> = LazyLock::new(|| {
84 RateLimiter::keyed(
85 Quota::with_period(Duration::from_secs(86_400))
86 .expect("with_period 仅在 Duration 为 0 时返回 None;86_400s 必然 Some")
87 .allow_burst(env_or("RATE_LIMIT_CODE_EXEC_DAILY", 50)),
88 )
89});
90
91#[cfg(feature = "server")]
92static UNKNOWN_BUCKET_LIMITER: LazyLock<DefaultKeyedRateLimiter<String>> = LazyLock::new(|| {
99 RateLimiter::keyed(
100 Quota::per_second(env_or("RATE_LIMIT_UNKNOWN_PER_SEC", 30))
101 .allow_burst(env_or("RATE_LIMIT_UNKNOWN_BURST", 100)),
102 )
103});
104
105#[cfg(feature = "server")]
113fn limiter_gc_interval() -> Duration {
114 let secs = std::env::var("RATE_LIMIT_GC_INTERVAL_SECS")
115 .ok()
116 .and_then(|s| s.parse().ok())
117 .unwrap_or(300);
118 Duration::from_secs(secs.max(1))
119}
120
121#[cfg(feature = "server")]
129fn ensure_limiter_gc() {
130 static SPAWNED: std::sync::Once = std::sync::Once::new();
131 SPAWNED.call_once(|| {
132 if tokio::runtime::Handle::try_current().is_err() {
136 return;
137 }
138 let interval = limiter_gc_interval();
139 tokio::spawn(async move {
140 loop {
141 tokio::time::sleep(interval).await;
142 STRICT_LIMITER.retain_recent();
143 UPLOAD_LIMITER.retain_recent();
144 IMAGE_LIMITER.retain_recent();
145 COMMENT_LIMITER.retain_recent();
146 CODE_EXEC_LIMITER.retain_recent();
147 CODE_EXEC_DAILY_LIMITER.retain_recent();
148 UNKNOWN_BUCKET_LIMITER.retain_recent();
149 }
150 });
151 });
152}
153
154#[cfg(feature = "server")]
155pub fn check_comment_limit(ip: &str) -> Result<(), String> {
157 ensure_limiter_gc();
158 COMMENT_LIMITER
159 .check_key(&ip.to_string())
160 .map(|_| ())
161 .map_err(|_| "评论过于频繁,请稍后再试".to_string())
162}
163
164#[cfg(feature = "server")]
165pub fn check_image_limit(ip: &str) -> Result<(), StatusCode> {
167 ensure_limiter_gc();
168 IMAGE_LIMITER
169 .check_key(&ip.to_string())
170 .map(|_| ())
171 .map_err(|_| StatusCode::TOO_MANY_REQUESTS)
172}
173
174#[cfg(feature = "server")]
175fn trusted_proxy_count() -> usize {
176 std::env::var("TRUSTED_PROXY_COUNT")
177 .ok()
178 .and_then(|s| s.parse().ok())
179 .unwrap_or(0)
180}
181
182#[cfg(feature = "server")]
183fn is_valid_ip(ip: &str) -> bool {
184 ip.parse::<std::net::IpAddr>().is_ok()
185}
186
187#[cfg(feature = "server")]
203fn ip_from_x_forwarded_for(value: &str, trusted_proxy_count: usize) -> Option<String> {
204 let parts: Vec<&str> = value
207 .split(',')
208 .map(str::trim)
209 .filter(|s| !s.is_empty())
210 .collect();
211
212 if trusted_proxy_count == 0 || parts.len() <= trusted_proxy_count {
213 return None;
214 }
215
216 let idx = parts.len() - 1 - trusted_proxy_count;
218 let ip = parts[idx].to_string();
219 if is_valid_ip(&ip) {
220 Some(ip)
221 } else {
222 None
223 }
224}
225
226#[cfg(feature = "server")]
227fn ip_from_x_real_ip(value: &str) -> Option<String> {
228 let ip = value.trim().to_string();
229 if is_valid_ip(&ip) {
230 Some(ip)
231 } else {
232 None
233 }
234}
235
236#[cfg(feature = "server")]
237fn get_client_ip_internal(
238 headers: &http::HeaderMap,
239 trusted: usize,
240 peer: Option<std::net::SocketAddr>,
241) -> String {
242 if trusted > 0 {
243 if let Some(value) = headers.get("x-forwarded-for").and_then(|v| v.to_str().ok()) {
244 if let Some(ip) = ip_from_x_forwarded_for(value, trusted) {
245 return ip;
246 }
247 }
248
249 if let Some(ip) = headers
250 .get("x-real-ip")
251 .and_then(|v| v.to_str().ok())
252 .and_then(ip_from_x_real_ip)
253 {
254 return ip;
255 }
256 }
257
258 if let Some(addr) = peer {
259 return addr.ip().to_string();
260 }
261
262 tracing::warn!(
265 "无法获取客户端真实 IP(未配置 TRUSTED_PROXY_COUNT 且无法读取 TCP 对端地址),\
266 限流将按 'unknown' 键聚合"
267 );
268 "unknown".to_string()
269}
270
271#[cfg(feature = "server")]
272pub fn get_client_ip_with_peer(
277 headers: &http::HeaderMap,
278 peer: Option<std::net::SocketAddr>,
279) -> String {
280 get_client_ip_internal(headers, trusted_proxy_count(), peer)
281}
282
283#[cfg(feature = "server")]
284pub fn get_client_ip(headers: &http::HeaderMap) -> String {
289 get_client_ip_internal(headers, trusted_proxy_count(), None)
290}
291
292#[cfg(feature = "server")]
293pub fn check_strict_limit(ip: &str) -> Result<(), String> {
300 ensure_limiter_gc();
301 if ip == "unknown" {
302 UNKNOWN_BUCKET_LIMITER
303 .check_key(&ip.to_string())
304 .map(|_| ())
305 .map_err(|_| "服务繁忙,请稍后再试".to_string())
306 } else {
307 STRICT_LIMITER
308 .check_key(&ip.to_string())
309 .map(|_| ())
310 .map_err(|_| "请求过于频繁,请稍后再试".to_string())
311 }
312}
313
314#[cfg(feature = "server")]
315pub fn check_upload_limit(ip: &str) -> Result<(), String> {
317 ensure_limiter_gc();
318 UPLOAD_LIMITER
319 .check_key(&ip.to_string())
320 .map(|_| ())
321 .map_err(|_| "上传过于频繁,请稍后再试".to_string())
322}
323
324#[cfg(feature = "server")]
325pub fn check_code_exec_limit(ip: &str) -> Result<(), String> {
329 ensure_limiter_gc();
330 CODE_EXEC_LIMITER
331 .check_key(&ip.to_string())
332 .map_err(|_| "请求过于频繁,请稍后再试".to_string())?;
333 CODE_EXEC_DAILY_LIMITER
334 .check_key(&ip.to_string())
335 .map_err(|_| "今日运行次数已达上限,请明天再试".to_string())?;
336 Ok(())
337}
338
339#[cfg(all(test, feature = "server"))]
340mod tests {
341 use super::*;
342 use http::HeaderMap;
343 use std::net::{IpAddr, Ipv4Addr, SocketAddr};
344
345 #[test]
346 fn get_client_ip_from_x_forwarded_for_with_one_trusted_proxy() {
347 let mut headers = HeaderMap::new();
348 headers.insert("x-forwarded-for", "1.2.3.4, 5.6.7.8".parse().unwrap());
349 assert_eq!(
350 get_client_ip_with_trusted_and_peer(&headers, 1, None),
351 "1.2.3.4"
352 );
353 }
354
355 #[test]
356 fn get_client_ip_from_x_forwarded_for_with_two_trusted_proxies() {
357 let mut headers = HeaderMap::new();
358 headers.insert(
359 "x-forwarded-for",
360 "1.2.3.4, 5.6.7.8, 9.10.11.12".parse().unwrap(),
361 );
362 assert_eq!(
363 get_client_ip_with_trusted_and_peer(&headers, 2, None),
364 "1.2.3.4"
365 );
366 }
367
368 #[test]
369 fn get_client_ip_ignores_x_forwarded_for_when_no_trusted_proxies() {
370 let mut headers = HeaderMap::new();
371 headers.insert("x-forwarded-for", "1.2.3.4, 5.6.7.8".parse().unwrap());
372 assert_eq!(
373 get_client_ip_with_trusted_and_peer(&headers, 0, None),
374 "unknown"
375 );
376 }
377
378 #[test]
379 fn get_client_ip_falls_back_to_peer_when_no_trusted_proxies() {
380 let mut headers = HeaderMap::new();
381 headers.insert("x-forwarded-for", "1.2.3.4, 5.6.7.8".parse().unwrap());
382 let peer = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)), 12345);
383 assert_eq!(
384 get_client_ip_with_trusted_and_peer(&headers, 0, Some(peer)),
385 "127.0.0.1"
386 );
387 }
388
389 #[test]
390 fn get_client_ip_from_x_real_ip_when_trusted() {
391 let mut headers = HeaderMap::new();
392 headers.insert("x-real-ip", "9.8.7.6".parse().unwrap());
393 assert_eq!(
394 get_client_ip_with_trusted_and_peer(&headers, 1, None),
395 "9.8.7.6"
396 );
397 }
398
399 #[test]
400 fn get_client_ip_x_real_ip_ignored_when_not_trusted() {
401 let mut headers = HeaderMap::new();
402 headers.insert("x-real-ip", "9.8.7.6".parse().unwrap());
403 let peer = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(192, 168, 1, 1)), 12345);
404 assert_eq!(
405 get_client_ip_with_trusted_and_peer(&headers, 0, Some(peer)),
406 "192.168.1.1"
407 );
408 }
409
410 #[test]
411 fn get_client_ip_x_forwarded_for_takes_priority_over_x_real_ip() {
412 let mut headers = HeaderMap::new();
413 headers.insert("x-forwarded-for", "1.1.1.1, 2.2.2.2".parse().unwrap());
414 headers.insert("x-real-ip", "3.3.3.3".parse().unwrap());
415 assert_eq!(
416 get_client_ip_with_trusted_and_peer(&headers, 1, None),
417 "1.1.1.1"
418 );
419 }
420
421 #[test]
422 fn get_client_ip_no_headers_returns_unknown() {
423 let headers = HeaderMap::new();
424 assert_eq!(
425 get_client_ip_with_trusted_and_peer(&headers, 1, None),
426 "unknown"
427 );
428 }
429
430 #[test]
431 fn get_client_ip_ignores_short_x_forwarded_for_list() {
432 let mut headers = HeaderMap::new();
433 headers.insert("x-forwarded-for", "1.2.3.4".parse().unwrap());
434 assert_eq!(
435 get_client_ip_with_trusted_and_peer(&headers, 2, None),
436 "unknown"
437 );
438 }
439
440 #[test]
441 fn get_client_ip_ignores_x_forwarded_for_equal_to_proxy_count() {
442 let mut headers = HeaderMap::new();
443 headers.insert("x-forwarded-for", "1.2.3.4, 5.6.7.8".parse().unwrap());
444 assert_eq!(
445 get_client_ip_with_trusted_and_peer(&headers, 2, None),
446 "unknown"
447 );
448 }
449
450 #[test]
451 fn get_client_ip_ignores_empty_x_forwarded_for_entries() {
452 let mut headers = HeaderMap::new();
453 headers.insert(
454 "x-forwarded-for",
455 " , 1.2.3.4 , 5.6.7.8 , ".parse().unwrap(),
456 );
457 assert_eq!(
458 get_client_ip_with_trusted_and_peer(&headers, 1, None),
459 "1.2.3.4"
460 );
461 }
462
463 #[test]
464 fn get_client_ip_rejects_invalid_x_forwarded_for_value() {
465 let mut headers = HeaderMap::new();
466 headers.insert("x-forwarded-for", "not-an-ip, 5.6.7.8".parse().unwrap());
467 assert_eq!(
468 get_client_ip_with_trusted_and_peer(&headers, 1, None),
469 "unknown"
470 );
471 }
472
473 #[test]
474 fn get_client_ip_rejects_invalid_x_real_ip_value() {
475 let mut headers = HeaderMap::new();
476 headers.insert("x-real-ip", "not-an-ip".parse().unwrap());
477 assert_eq!(
478 get_client_ip_with_trusted_and_peer(&headers, 1, None),
479 "unknown"
480 );
481 }
482
483 #[test]
484 fn get_client_ip_prefers_xff_over_peer() {
485 let mut headers = HeaderMap::new();
486 headers.insert("x-forwarded-for", "1.2.3.4, 5.6.7.8".parse().unwrap());
487 let peer = SocketAddr::new(IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)), 12345);
488 assert_eq!(
489 get_client_ip_with_trusted_and_peer(&headers, 1, Some(peer)),
490 "1.2.3.4"
491 );
492 }
493
494 #[test]
495 fn get_client_ip_with_env_trusted_proxy_count_zero() {
496 let original = std::env::var("TRUSTED_PROXY_COUNT").ok();
497 std::env::set_var("TRUSTED_PROXY_COUNT", "0");
498
499 let mut headers = HeaderMap::new();
500 headers.insert("x-forwarded-for", "1.2.3.4, 5.6.7.8".parse().unwrap());
501 assert_eq!(get_client_ip(&headers), "unknown");
502
503 match original {
504 Some(value) => std::env::set_var("TRUSTED_PROXY_COUNT", value),
505 None => std::env::remove_var("TRUSTED_PROXY_COUNT"),
506 }
507 }
508
509 #[test]
510 #[serial_test::serial]
511 fn check_strict_unknown_ip_uses_lenient_bucket() {
512 for _ in 0..20 {
515 assert!(
516 super::check_strict_limit("unknown").is_ok(),
517 "unknown bucket should allow small bursts, not hit strict 1 req/s limit"
518 );
519 }
520 }
521
522 #[test]
523 #[serial_test::serial]
524 fn check_strict_real_ip_uses_strict_bucket() {
525 let unique_ip = "198.51.100.42";
528 let mut allowed = 0;
529 let mut blocked = false;
530 for _ in 0..50 {
531 match super::check_strict_limit(unique_ip) {
532 Ok(()) => allowed += 1,
533 Err(_) => blocked = true,
534 }
535 if blocked {
536 break;
537 }
538 }
539 assert!(
540 blocked,
541 "strict bucket should eventually block real IP burst"
542 );
543 assert!(
544 allowed <= 6,
545 "strict burst is 5, allowed should be <= 6, got {allowed}"
546 );
547 }
548
549 fn get_client_ip_with_trusted_and_peer(
551 headers: &HeaderMap,
552 trusted: usize,
553 peer: Option<SocketAddr>,
554 ) -> String {
555 get_client_ip_internal(headers, trusted, peer)
556 }
557}