use std::collections::HashMap; use std::sync::Arc; use std::time::{Duration, Instant}; use tokio::sync::RwLock; use crate::error::AppError; pub struct RateLimiter { requests: Arc>>>, max_requests: usize, window_secs: u64, } impl RateLimiter { pub fn new(max_requests: usize, window_secs: u64) -> Self { Self { requests: Arc::new(RwLock::new(HashMap::new())), max_requests, window_secs, } } pub async fn check_rate_limit(&self, client_ip: &str) -> Result<(), AppError> { let now = Instant::now(); let window = Duration::from_secs(self.window_secs); let mut requests = self.requests.write().await; let client_requests = requests.entry(client_ip.to_string()).or_insert_with(Vec::new); client_requests.retain(|&t| now.duration_since(t) < window); if client_requests.len() >= self.max_requests { return Err(AppError::TooManyRequests( format!("请求过于频繁,请{}秒后再试", self.window_secs) )); } client_requests.push(now); Ok(()) } } pub fn extract_client_ip_from_header(headers: &actix_web::http::header::HeaderMap) -> String { headers .get("X-Forwarded-For") .and_then(|v| v.to_str().ok()) .map(|s| s.split(',').next().unwrap_or(s).trim().to_string()) .or_else(|| { headers .get("X-Real-IP") .and_then(|v| v.to_str().ok()) .map(|s| s.to_string()) }) .unwrap_or_else(|| "unknown".to_string()) } pub fn create_login_rate_limiter() -> Arc { Arc::new(RateLimiter::new(5, 60)) } pub static LOGIN_RATE_LIMITER: std::sync::LazyLock> = std::sync::LazyLock::new(create_login_rate_limiter); #[cfg(test)] mod tests { use super::*; #[tokio::test] async fn test_rate_limiter_allows_requests_under_limit() { let limiter = Arc::new(RateLimiter::new(5, 60)); for _ in 0..5 { let result = limiter.check_rate_limit("192.168.1.1").await; assert!(result.is_ok()); } } #[tokio::test] async fn test_rate_limiter_blocks_excessive_requests() { let limiter = Arc::new(RateLimiter::new(3, 60)); for _ in 0..3 { assert!(limiter.check_rate_limit("192.168.1.2").await.is_ok()); } let result = limiter.check_rate_limit("192.168.1.2").await; assert!(result.is_err()); } #[tokio::test] async fn test_rate_limiter_independent_per_client() { let limiter = Arc::new(RateLimiter::new(2, 60)); assert!(limiter.check_rate_limit("192.168.1.100").await.is_ok()); assert!(limiter.check_rate_limit("192.168.1.100").await.is_ok()); assert!(limiter.check_rate_limit("192.168.1.100").await.is_err()); assert!(limiter.check_rate_limit("192.168.1.101").await.is_ok()); assert!(limiter.check_rate_limit("192.168.1.101").await.is_ok()); assert!(limiter.check_rate_limit("192.168.1.101").await.is_err()); } #[test] fn test_extract_client_ip_from_x_forwarded_for() { use actix_web::http::header::{HeaderMap, HeaderName, HeaderValue}; let mut headers = HeaderMap::new(); headers.insert( HeaderName::from_static("x-forwarded-for"), HeaderValue::from_static("203.0.113.195, 70.41.3.18"), ); let ip = extract_client_ip_from_header(&headers); assert_eq!(ip, "203.0.113.195"); } #[test] fn test_extract_client_ip_from_x_real_ip() { use actix_web::http::header::{HeaderMap, HeaderName, HeaderValue}; let mut headers = HeaderMap::new(); headers.insert( HeaderName::from_static("x-real-ip"), HeaderValue::from_static("203.0.113.195"), ); let ip = extract_client_ip_from_header(&headers); assert_eq!(ip, "203.0.113.195"); } #[test] fn test_extract_client_ip_prefers_x_forwarded_for() { use actix_web::http::header::{HeaderMap, HeaderName, HeaderValue}; let mut headers = HeaderMap::new(); headers.insert( HeaderName::from_static("x-forwarded-for"), HeaderValue::from_static("203.0.113.195"), ); headers.insert( HeaderName::from_static("x-real-ip"), HeaderValue::from_static("198.51.100.178"), ); let ip = extract_client_ip_from_header(&headers); assert_eq!(ip, "203.0.113.195"); } #[test] fn test_extract_client_ip_fallback_to_unknown() { use actix_web::http::header::HeaderMap; let headers = HeaderMap::new(); let ip = extract_client_ip_from_header(&headers); assert_eq!(ip, "unknown"); } }