Files
asd-backend/src/rate_limiter.rs
Milky0217 e6048ea010 feat(auth): 实现登录限流和性能优化
性能优化:
- 添加数据库索引优化查询性能 (006)
  - weather_data: user_id, date, is_favorite 索引
  - users: openid 索引
  - payment_orders: status 索引
- 新增 refresh_tokens 表支持双 Token 机制 (004)
- 新增 web_login_codes 表支持网页端扫码登录 (005)

安全增强:
- 实现基于 IP 的登录限流 (rate_limiter.rs)
  - 滑动窗口算法: 5次/分钟/IP
  - 自动清理过期记录
  - 429 TooManyRequests 响应

新模块:
- src/rate_limiter.rs: 限流模块
- src/alipay.rs: 支付宝签名模块 (RSA2)
- src/error.rs: 统一错误类型 (含 TooManyRequests)
- src/handlers/meta.rs: 元数据处理器

代码清理:
- 修复 .gitignore 规则,正确跟踪 src/ 和 migrations/
2026-04-24 16:13:49 +08:00

158 lines
4.8 KiB
Rust

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<RwLock<HashMap<String, Vec<Instant>>>>,
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<RateLimiter> {
Arc::new(RateLimiter::new(5, 60))
}
pub static LOGIN_RATE_LIMITER: std::sync::LazyLock<Arc<RateLimiter>> =
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");
}
}