mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
feat: add security middleware, rate limiting, auth, encryption, conversation management
Borrowed patterns from ZeroClaw: - server/middleware/security: request body size limit + timeout layer - server/middleware/rate_limit: sliding window rate limiter per client IP - server/middleware/idempotency: idempotency key store for duplicate prevention - server/auth/pairing: pairing code auth with brute force protection - core/sanitizer: credential sanitizer with 9 builtin patterns - core/router/hint_router: message prefix hint routing ([reasoning], [fast]) - processor/conversation_manager: conversation history trimming - processor/conversation_summarizer: LLM-based conversation summarization - credential/encryption: ChaCha20-Poly1305 AEAD encryption for API keys
This commit is contained in:
@@ -20,6 +20,7 @@ serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
tokio.workspace = true
|
||||
futures.workspace = true
|
||||
hex.workspace = true
|
||||
axum.workspace = true
|
||||
tower.workspace = true
|
||||
tower-http.workspace = true
|
||||
@@ -35,6 +36,8 @@ subtle.workspace = true
|
||||
async-stream.workspace = true
|
||||
urlencoding.workspace = true
|
||||
parking_lot.workspace = true
|
||||
rand.workspace = true
|
||||
sha2.workspace = true
|
||||
tokio-util.workspace = true
|
||||
dirs.workspace = true
|
||||
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
//! 认证模块
|
||||
|
||||
pub mod pairing;
|
||||
@@ -0,0 +1,309 @@
|
||||
//! 配对认证系统
|
||||
//!
|
||||
//! 提供一次性配对码认证流程:
|
||||
//! 1. 启动时生成配对码
|
||||
//! 2. 客户端通过配对码获取 bearer token
|
||||
//! 3. 后续请求使用 bearer token 认证
|
||||
//! 4. 暴力破解保护
|
||||
|
||||
use parking_lot::Mutex;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use sha2::{Digest, Sha256};
|
||||
use std::collections::HashSet;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
const MAX_FAILED_ATTEMPTS: u32 = 5;
|
||||
const LOCKOUT_DURATION_SECS: u64 = 300; // 5 分钟
|
||||
|
||||
/// 配对认证配置
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct PairingConfig {
|
||||
/// 是否启用配对认证
|
||||
#[serde(default)]
|
||||
pub enabled: bool,
|
||||
}
|
||||
|
||||
impl Default for PairingConfig {
|
||||
fn default() -> Self {
|
||||
Self { enabled: false }
|
||||
}
|
||||
}
|
||||
|
||||
/// 失败尝试状态
|
||||
#[derive(Debug)]
|
||||
struct FailureState {
|
||||
count: u32,
|
||||
window_start: Option<Instant>,
|
||||
blocked_until: Option<Instant>,
|
||||
}
|
||||
|
||||
impl Default for FailureState {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
count: 0,
|
||||
window_start: None,
|
||||
blocked_until: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 配对结果
|
||||
#[derive(Debug)]
|
||||
pub enum PairingResult {
|
||||
/// 配对成功,返回 token
|
||||
Success { token: String },
|
||||
/// 配对码错误
|
||||
InvalidCode,
|
||||
/// 被锁定
|
||||
Locked { retry_after_secs: u64 },
|
||||
/// 配对未启用
|
||||
Disabled,
|
||||
}
|
||||
|
||||
/// 认证结果
|
||||
#[derive(Debug, PartialEq)]
|
||||
pub enum AuthResult {
|
||||
/// 认证成功
|
||||
Authenticated,
|
||||
/// 未认证
|
||||
Unauthenticated,
|
||||
/// 配对未启用(允许通过)
|
||||
Disabled,
|
||||
}
|
||||
|
||||
/// 配对认证守卫
|
||||
pub struct PairingGuard {
|
||||
config: PairingConfig,
|
||||
/// 当前配对码
|
||||
pairing_code: Mutex<Option<String>>,
|
||||
/// 已配对的 token(存储 SHA-256 哈希)
|
||||
paired_tokens: Mutex<HashSet<String>>,
|
||||
/// 失败尝试追踪
|
||||
failed_attempts: Mutex<FailureState>,
|
||||
}
|
||||
|
||||
impl PairingGuard {
|
||||
pub fn new(config: PairingConfig) -> Self {
|
||||
let code = if config.enabled {
|
||||
Some(Self::generate_pairing_code())
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
if let Some(ref code) = code {
|
||||
tracing::info!("========================================");
|
||||
tracing::info!("配对码: {}", code);
|
||||
tracing::info!("========================================");
|
||||
}
|
||||
|
||||
Self {
|
||||
config,
|
||||
pairing_code: Mutex::new(code),
|
||||
paired_tokens: Mutex::new(HashSet::new()),
|
||||
failed_attempts: Mutex::new(FailureState::default()),
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建带指定配对码的守卫(用于测试)
|
||||
#[cfg(test)]
|
||||
fn with_code(config: PairingConfig, code: String) -> Self {
|
||||
Self {
|
||||
config,
|
||||
pairing_code: Mutex::new(Some(code)),
|
||||
paired_tokens: Mutex::new(HashSet::new()),
|
||||
failed_attempts: Mutex::new(FailureState::default()),
|
||||
}
|
||||
}
|
||||
|
||||
/// 生成 6 位配对码
|
||||
fn generate_pairing_code() -> String {
|
||||
use rand::Rng;
|
||||
let mut rng = rand::thread_rng();
|
||||
format!("{:06}", rng.gen_range(0..1_000_000))
|
||||
}
|
||||
|
||||
/// 生成 bearer token
|
||||
fn generate_token() -> String {
|
||||
use rand::Rng;
|
||||
let mut rng = rand::thread_rng();
|
||||
let bytes: Vec<u8> = (0..32).map(|_| rng.gen()).collect();
|
||||
hex::encode(bytes)
|
||||
}
|
||||
|
||||
/// 计算 token 的 SHA-256 哈希
|
||||
fn hash_token(token: &str) -> String {
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(token.as_bytes());
|
||||
hex::encode(hasher.finalize())
|
||||
}
|
||||
|
||||
/// 尝试配对
|
||||
pub fn pair(&self, code: &str) -> PairingResult {
|
||||
if !self.config.enabled {
|
||||
return PairingResult::Disabled;
|
||||
}
|
||||
|
||||
// 检查是否被锁定
|
||||
{
|
||||
let state = self.failed_attempts.lock();
|
||||
if let Some(blocked_until) = state.blocked_until {
|
||||
if Instant::now() < blocked_until {
|
||||
let remaining = blocked_until.duration_since(Instant::now()).as_secs();
|
||||
return PairingResult::Locked {
|
||||
retry_after_secs: remaining + 1,
|
||||
};
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 验证配对码
|
||||
let valid = {
|
||||
let pairing_code = self.pairing_code.lock();
|
||||
pairing_code.as_deref() == Some(code)
|
||||
};
|
||||
|
||||
if valid {
|
||||
// 重置失败计数
|
||||
{
|
||||
let mut state = self.failed_attempts.lock();
|
||||
*state = FailureState::default();
|
||||
}
|
||||
|
||||
// 生成 token
|
||||
let token = Self::generate_token();
|
||||
let hash = Self::hash_token(&token);
|
||||
self.paired_tokens.lock().insert(hash);
|
||||
|
||||
PairingResult::Success { token }
|
||||
} else {
|
||||
// 记录失败
|
||||
let mut state = self.failed_attempts.lock();
|
||||
let now = Instant::now();
|
||||
|
||||
match state.window_start {
|
||||
Some(start)
|
||||
if now.duration_since(start) < Duration::from_secs(LOCKOUT_DURATION_SECS) =>
|
||||
{
|
||||
state.count += 1;
|
||||
}
|
||||
_ => {
|
||||
state.count = 1;
|
||||
state.window_start = Some(now);
|
||||
}
|
||||
}
|
||||
|
||||
if state.count >= MAX_FAILED_ATTEMPTS {
|
||||
state.blocked_until = Some(now + Duration::from_secs(LOCKOUT_DURATION_SECS));
|
||||
tracing::warn!(
|
||||
"配对认证:暴力破解保护触发,锁定 {} 秒",
|
||||
LOCKOUT_DURATION_SECS
|
||||
);
|
||||
}
|
||||
|
||||
PairingResult::InvalidCode
|
||||
}
|
||||
}
|
||||
|
||||
/// 验证 bearer token
|
||||
pub fn authenticate(&self, token: &str) -> AuthResult {
|
||||
if !self.config.enabled {
|
||||
return AuthResult::Disabled;
|
||||
}
|
||||
|
||||
let hash = Self::hash_token(token);
|
||||
let tokens = self.paired_tokens.lock();
|
||||
|
||||
if tokens.contains(&hash) {
|
||||
AuthResult::Authenticated
|
||||
} else {
|
||||
AuthResult::Unauthenticated
|
||||
}
|
||||
}
|
||||
|
||||
/// 是否启用
|
||||
pub fn is_enabled(&self) -> bool {
|
||||
self.config.enabled
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn enabled_config() -> PairingConfig {
|
||||
PairingConfig { enabled: true }
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_disabled_pairing() {
|
||||
let guard = PairingGuard::new(PairingConfig::default());
|
||||
assert!(!guard.is_enabled());
|
||||
assert!(matches!(guard.pair("anything"), PairingResult::Disabled));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_successful_pairing() {
|
||||
let guard = PairingGuard::with_code(enabled_config(), "123456".to_string());
|
||||
|
||||
match guard.pair("123456") {
|
||||
PairingResult::Success { token } => {
|
||||
assert_eq!(token.len(), 64); // 32 bytes hex
|
||||
assert!(token.chars().all(|c| c.is_ascii_hexdigit()));
|
||||
}
|
||||
other => panic!("期望 Success,得到 {:?}", other),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_invalid_code() {
|
||||
let guard = PairingGuard::with_code(enabled_config(), "123456".to_string());
|
||||
assert!(matches!(guard.pair("000000"), PairingResult::InvalidCode));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_authentication() {
|
||||
let guard = PairingGuard::with_code(enabled_config(), "123456".to_string());
|
||||
|
||||
let token = match guard.pair("123456") {
|
||||
PairingResult::Success { token } => token,
|
||||
_ => panic!("配对应成功"),
|
||||
};
|
||||
|
||||
assert_eq!(guard.authenticate(&token), AuthResult::Authenticated);
|
||||
assert_eq!(guard.authenticate("bad_token"), AuthResult::Unauthenticated);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_brute_force_protection() {
|
||||
let guard = PairingGuard::with_code(enabled_config(), "123456".to_string());
|
||||
|
||||
// 5 次失败触发锁定
|
||||
for _ in 0..MAX_FAILED_ATTEMPTS {
|
||||
assert!(matches!(guard.pair("000000"), PairingResult::InvalidCode));
|
||||
}
|
||||
|
||||
// 第 6 次应被锁定
|
||||
match guard.pair("000000") {
|
||||
PairingResult::Locked { retry_after_secs } => {
|
||||
assert!(retry_after_secs > 0);
|
||||
assert!(retry_after_secs <= LOCKOUT_DURATION_SECS + 1);
|
||||
}
|
||||
other => panic!("期望 Locked,得到 {:?}", other),
|
||||
}
|
||||
|
||||
// 即使用正确码也应被锁定
|
||||
assert!(matches!(guard.pair("123456"), PairingResult::Locked { .. }));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_disabled_auth_allows_all() {
|
||||
let guard = PairingGuard::new(PairingConfig::default());
|
||||
assert_eq!(guard.authenticate("any_token"), AuthResult::Disabled);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_default_config() {
|
||||
let config = PairingConfig::default();
|
||||
assert!(!config.enabled);
|
||||
}
|
||||
}
|
||||
@@ -1,6 +1,8 @@
|
||||
//! HTTP API 服务器
|
||||
|
||||
pub mod auth;
|
||||
pub mod client_detector;
|
||||
pub mod middleware;
|
||||
|
||||
use axum::{
|
||||
extract::{DefaultBodyLimit, Path, State},
|
||||
@@ -43,6 +45,7 @@ use std::path::PathBuf;
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::{oneshot, RwLock};
|
||||
use tower_http::cors::CorsLayer;
|
||||
use tower_http::timeout::TimeoutLayer;
|
||||
|
||||
/// 记录请求统计到遥测系统
|
||||
pub fn record_request_telemetry(
|
||||
@@ -1053,6 +1056,10 @@ async fn run_server(
|
||||
.merge(batch_api_routes)
|
||||
.layer(cors_layer)
|
||||
.layer(DefaultBodyLimit::max(body_limit))
|
||||
.layer(TimeoutLayer::with_status_code(
|
||||
StatusCode::REQUEST_TIMEOUT,
|
||||
std::time::Duration::from_secs(300),
|
||||
))
|
||||
.with_state(state);
|
||||
|
||||
let addr: std::net::SocketAddr = format!("{host}:{port}")
|
||||
|
||||
@@ -0,0 +1,270 @@
|
||||
//! 幂等性中间件
|
||||
//!
|
||||
//! 通过 Idempotency-Key header 防止重复请求
|
||||
|
||||
use parking_lot::Mutex;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
/// 幂等性配置
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct IdempotencyConfig {
|
||||
/// 是否启用
|
||||
#[serde(default)]
|
||||
pub enabled: bool,
|
||||
/// 缓存 TTL(秒)
|
||||
#[serde(default = "default_ttl_secs")]
|
||||
pub ttl_secs: u64,
|
||||
/// Header 名称
|
||||
#[serde(default = "default_header_name")]
|
||||
pub header_name: String,
|
||||
}
|
||||
|
||||
fn default_ttl_secs() -> u64 {
|
||||
86400 // 24 小时
|
||||
}
|
||||
fn default_header_name() -> String {
|
||||
"Idempotency-Key".to_string()
|
||||
}
|
||||
|
||||
impl Default for IdempotencyConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
enabled: false,
|
||||
ttl_secs: default_ttl_secs(),
|
||||
header_name: default_header_name(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 幂等性检查结果
|
||||
#[derive(Debug, PartialEq)]
|
||||
pub enum IdempotencyCheck {
|
||||
/// 新请求,可以处理
|
||||
New,
|
||||
/// 正在处理中(返回 409 Conflict)
|
||||
InProgress,
|
||||
/// 已完成,有缓存响应
|
||||
Completed { status: u16, body: String },
|
||||
}
|
||||
|
||||
/// 请求状态
|
||||
#[derive(Debug, Clone)]
|
||||
enum RequestState {
|
||||
/// 正在处理
|
||||
InProgress { started_at: Instant },
|
||||
/// 已完成
|
||||
Completed {
|
||||
status: u16,
|
||||
body: String,
|
||||
completed_at: Instant,
|
||||
},
|
||||
}
|
||||
|
||||
/// 幂等性存储
|
||||
pub struct IdempotencyStore {
|
||||
config: IdempotencyConfig,
|
||||
entries: Mutex<HashMap<String, RequestState>>,
|
||||
}
|
||||
|
||||
impl IdempotencyStore {
|
||||
pub fn new(config: IdempotencyConfig) -> Self {
|
||||
Self {
|
||||
config,
|
||||
entries: Mutex::new(HashMap::new()),
|
||||
}
|
||||
}
|
||||
|
||||
/// 检查幂等性键
|
||||
pub fn check(&self, key: &str) -> IdempotencyCheck {
|
||||
if !self.config.enabled {
|
||||
return IdempotencyCheck::New;
|
||||
}
|
||||
|
||||
let mut entries = self.entries.lock();
|
||||
let ttl = Duration::from_secs(self.config.ttl_secs);
|
||||
let now = Instant::now();
|
||||
|
||||
match entries.get(key) {
|
||||
Some(RequestState::InProgress { started_at }) => {
|
||||
// 如果处理超过 TTL,视为过期
|
||||
if now.duration_since(*started_at) > ttl {
|
||||
entries.insert(
|
||||
key.to_string(),
|
||||
RequestState::InProgress { started_at: now },
|
||||
);
|
||||
IdempotencyCheck::New
|
||||
} else {
|
||||
IdempotencyCheck::InProgress
|
||||
}
|
||||
}
|
||||
Some(RequestState::Completed {
|
||||
status,
|
||||
body,
|
||||
completed_at,
|
||||
}) => {
|
||||
if now.duration_since(*completed_at) > ttl {
|
||||
entries.insert(
|
||||
key.to_string(),
|
||||
RequestState::InProgress { started_at: now },
|
||||
);
|
||||
IdempotencyCheck::New
|
||||
} else {
|
||||
IdempotencyCheck::Completed {
|
||||
status: *status,
|
||||
body: body.clone(),
|
||||
}
|
||||
}
|
||||
}
|
||||
None => {
|
||||
entries.insert(
|
||||
key.to_string(),
|
||||
RequestState::InProgress { started_at: now },
|
||||
);
|
||||
IdempotencyCheck::New
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 标记请求完成
|
||||
pub fn complete(&self, key: &str, status: u16, body: String) {
|
||||
if !self.config.enabled {
|
||||
return;
|
||||
}
|
||||
let mut entries = self.entries.lock();
|
||||
entries.insert(
|
||||
key.to_string(),
|
||||
RequestState::Completed {
|
||||
status,
|
||||
body,
|
||||
completed_at: Instant::now(),
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
/// 移除键(请求失败时调用,允许重试)
|
||||
pub fn remove(&self, key: &str) {
|
||||
let mut entries = self.entries.lock();
|
||||
entries.remove(key);
|
||||
}
|
||||
|
||||
/// 清理过期条目
|
||||
pub fn cleanup(&self) {
|
||||
let ttl = Duration::from_secs(self.config.ttl_secs);
|
||||
let now = Instant::now();
|
||||
let mut entries = self.entries.lock();
|
||||
entries.retain(|_, state| match state {
|
||||
RequestState::InProgress { started_at } => now.duration_since(*started_at) < ttl,
|
||||
RequestState::Completed { completed_at, .. } => now.duration_since(*completed_at) < ttl,
|
||||
});
|
||||
}
|
||||
|
||||
/// 获取当前条目数
|
||||
pub fn len(&self) -> usize {
|
||||
self.entries.lock().len()
|
||||
}
|
||||
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.entries.lock().is_empty()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::thread;
|
||||
|
||||
fn enabled_config(ttl_secs: u64) -> IdempotencyConfig {
|
||||
IdempotencyConfig {
|
||||
enabled: true,
|
||||
ttl_secs,
|
||||
header_name: "Idempotency-Key".to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_disabled_always_new() {
|
||||
let store = IdempotencyStore::new(IdempotencyConfig::default());
|
||||
assert_eq!(store.check("key1"), IdempotencyCheck::New);
|
||||
assert_eq!(store.check("key1"), IdempotencyCheck::New);
|
||||
assert!(store.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_new_request() {
|
||||
let store = IdempotencyStore::new(enabled_config(60));
|
||||
assert_eq!(store.check("key1"), IdempotencyCheck::New);
|
||||
assert_eq!(store.len(), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_in_progress_request() {
|
||||
let store = IdempotencyStore::new(enabled_config(60));
|
||||
assert_eq!(store.check("key1"), IdempotencyCheck::New);
|
||||
// 同一 key 再次检查应返回 InProgress
|
||||
assert_eq!(store.check("key1"), IdempotencyCheck::InProgress);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_completed_request() {
|
||||
let store = IdempotencyStore::new(enabled_config(60));
|
||||
assert_eq!(store.check("key1"), IdempotencyCheck::New);
|
||||
|
||||
store.complete("key1", 200, "ok".to_string());
|
||||
|
||||
assert_eq!(
|
||||
store.check("key1"),
|
||||
IdempotencyCheck::Completed {
|
||||
status: 200,
|
||||
body: "ok".to_string(),
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_expired_entry() {
|
||||
let store = IdempotencyStore::new(enabled_config(1)); // 1 秒 TTL
|
||||
|
||||
assert_eq!(store.check("key1"), IdempotencyCheck::New);
|
||||
store.complete("key1", 200, "ok".to_string());
|
||||
|
||||
// 等待过期
|
||||
thread::sleep(Duration::from_millis(1100));
|
||||
|
||||
// 过期后应视为新请求
|
||||
assert_eq!(store.check("key1"), IdempotencyCheck::New);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_cleanup() {
|
||||
let store = IdempotencyStore::new(enabled_config(1));
|
||||
assert_eq!(store.check("key1"), IdempotencyCheck::New);
|
||||
assert_eq!(store.check("key2"), IdempotencyCheck::New);
|
||||
store.complete("key1", 200, "ok".to_string());
|
||||
|
||||
thread::sleep(Duration::from_millis(1100));
|
||||
|
||||
store.cleanup();
|
||||
assert!(store.is_empty(), "清理后应无过期条目");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_remove_allows_retry() {
|
||||
let store = IdempotencyStore::new(enabled_config(60));
|
||||
assert_eq!(store.check("key1"), IdempotencyCheck::New);
|
||||
assert_eq!(store.check("key1"), IdempotencyCheck::InProgress);
|
||||
|
||||
// 移除后应可重试
|
||||
store.remove("key1");
|
||||
assert_eq!(store.check("key1"), IdempotencyCheck::New);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_default_config() {
|
||||
let config = IdempotencyConfig::default();
|
||||
assert!(!config.enabled);
|
||||
assert_eq!(config.ttl_secs, 86400);
|
||||
assert_eq!(config.header_name, "Idempotency-Key");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,5 @@
|
||||
//! 服务器中间件模块
|
||||
|
||||
pub mod idempotency;
|
||||
pub mod rate_limit;
|
||||
pub mod security;
|
||||
@@ -0,0 +1,238 @@
|
||||
//! 滑动窗口速率限制中间件
|
||||
//!
|
||||
//! 基于客户端 IP 的请求速率限制,防止 API 滥用
|
||||
|
||||
use parking_lot::Mutex;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::{
|
||||
collections::HashMap,
|
||||
time::{Duration, Instant},
|
||||
};
|
||||
|
||||
/// 速率限制配置
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct RateLimitConfig {
|
||||
/// 是否启用
|
||||
#[serde(default = "default_enabled")]
|
||||
pub enabled: bool,
|
||||
/// 窗口内最大请求数
|
||||
#[serde(default = "default_requests_per_minute")]
|
||||
pub requests_per_minute: u32,
|
||||
/// 窗口大小(秒)
|
||||
#[serde(default = "default_window_secs")]
|
||||
pub window_secs: u64,
|
||||
}
|
||||
|
||||
fn default_enabled() -> bool {
|
||||
false
|
||||
}
|
||||
fn default_requests_per_minute() -> u32 {
|
||||
60
|
||||
}
|
||||
fn default_window_secs() -> u64 {
|
||||
60
|
||||
}
|
||||
|
||||
impl Default for RateLimitConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
enabled: false,
|
||||
requests_per_minute: 60,
|
||||
window_secs: 60,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 滑动窗口速率限制器
|
||||
pub struct SlidingWindowRateLimiter {
|
||||
config: RateLimitConfig,
|
||||
/// 客户端 IP -> 请求时间戳列表
|
||||
requests: Mutex<HashMap<String, Vec<Instant>>>,
|
||||
}
|
||||
|
||||
impl SlidingWindowRateLimiter {
|
||||
pub fn new(config: RateLimitConfig) -> Self {
|
||||
Self {
|
||||
config,
|
||||
requests: Mutex::new(HashMap::new()),
|
||||
}
|
||||
}
|
||||
|
||||
/// 检查是否允许请求
|
||||
pub fn check_rate_limit(&self, client_id: &str) -> RateLimitResult {
|
||||
if !self.config.enabled {
|
||||
return RateLimitResult::Allowed;
|
||||
}
|
||||
|
||||
let now = Instant::now();
|
||||
let window = Duration::from_secs(self.config.window_secs);
|
||||
let mut requests = self.requests.lock();
|
||||
|
||||
let timestamps = requests.entry(client_id.to_string()).or_default();
|
||||
|
||||
// 清理窗口外的请求
|
||||
timestamps.retain(|t| now.duration_since(*t) < window);
|
||||
|
||||
if timestamps.len() >= self.config.requests_per_minute as usize {
|
||||
// 计算最早请求到窗口结束的剩余时间
|
||||
let oldest = timestamps.first().copied();
|
||||
let retry_after = oldest
|
||||
.map(|t| window.saturating_sub(now.duration_since(t)))
|
||||
.unwrap_or(window);
|
||||
RateLimitResult::Limited { retry_after }
|
||||
} else {
|
||||
timestamps.push(now);
|
||||
RateLimitResult::Allowed
|
||||
}
|
||||
}
|
||||
|
||||
/// 清理过期条目(应定期调用)
|
||||
pub fn cleanup(&self) {
|
||||
let now = Instant::now();
|
||||
let window = Duration::from_secs(self.config.window_secs);
|
||||
let mut requests = self.requests.lock();
|
||||
|
||||
requests.retain(|_, timestamps| {
|
||||
timestamps.retain(|t| now.duration_since(*t) < window);
|
||||
!timestamps.is_empty()
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
/// 速率限制检查结果
|
||||
#[derive(Debug)]
|
||||
pub enum RateLimitResult {
|
||||
/// 允许
|
||||
Allowed,
|
||||
/// 被限制
|
||||
Limited {
|
||||
/// 建议重试等待时间
|
||||
retry_after: Duration,
|
||||
},
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::thread;
|
||||
|
||||
#[test]
|
||||
fn test_disabled_allows_all() {
|
||||
let limiter = SlidingWindowRateLimiter::new(RateLimitConfig {
|
||||
enabled: false,
|
||||
requests_per_minute: 1,
|
||||
window_secs: 60,
|
||||
});
|
||||
|
||||
// 即使超过限制,禁用时也应全部允许
|
||||
for _ in 0..100 {
|
||||
assert!(matches!(
|
||||
limiter.check_rate_limit("client1"),
|
||||
RateLimitResult::Allowed
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_within_limit() {
|
||||
let limiter = SlidingWindowRateLimiter::new(RateLimitConfig {
|
||||
enabled: true,
|
||||
requests_per_minute: 5,
|
||||
window_secs: 60,
|
||||
});
|
||||
|
||||
for _ in 0..5 {
|
||||
assert!(matches!(
|
||||
limiter.check_rate_limit("client1"),
|
||||
RateLimitResult::Allowed
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_exceeds_limit() {
|
||||
let limiter = SlidingWindowRateLimiter::new(RateLimitConfig {
|
||||
enabled: true,
|
||||
requests_per_minute: 3,
|
||||
window_secs: 60,
|
||||
});
|
||||
|
||||
// 前 3 个请求应允许
|
||||
for _ in 0..3 {
|
||||
assert!(matches!(
|
||||
limiter.check_rate_limit("client1"),
|
||||
RateLimitResult::Allowed
|
||||
));
|
||||
}
|
||||
|
||||
// 第 4 个应被限制
|
||||
match limiter.check_rate_limit("client1") {
|
||||
RateLimitResult::Limited { retry_after } => {
|
||||
assert!(retry_after.as_secs() <= 60);
|
||||
}
|
||||
RateLimitResult::Allowed => panic!("应该被限制"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_window_expiry() {
|
||||
let limiter = SlidingWindowRateLimiter::new(RateLimitConfig {
|
||||
enabled: true,
|
||||
requests_per_minute: 2,
|
||||
window_secs: 1, // 1 秒窗口,方便测试过期
|
||||
});
|
||||
|
||||
// 用完配额
|
||||
assert!(matches!(
|
||||
limiter.check_rate_limit("client1"),
|
||||
RateLimitResult::Allowed
|
||||
));
|
||||
assert!(matches!(
|
||||
limiter.check_rate_limit("client1"),
|
||||
RateLimitResult::Allowed
|
||||
));
|
||||
assert!(matches!(
|
||||
limiter.check_rate_limit("client1"),
|
||||
RateLimitResult::Limited { .. }
|
||||
));
|
||||
|
||||
// 等待窗口过期
|
||||
thread::sleep(Duration::from_millis(1100));
|
||||
|
||||
// 窗口过期后应重新允许
|
||||
assert!(matches!(
|
||||
limiter.check_rate_limit("client1"),
|
||||
RateLimitResult::Allowed
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_cleanup() {
|
||||
let limiter = SlidingWindowRateLimiter::new(RateLimitConfig {
|
||||
enabled: true,
|
||||
requests_per_minute: 10,
|
||||
window_secs: 1,
|
||||
});
|
||||
|
||||
// 添加一些请求
|
||||
limiter.check_rate_limit("client1");
|
||||
limiter.check_rate_limit("client2");
|
||||
|
||||
// 等待窗口过期
|
||||
thread::sleep(Duration::from_millis(1100));
|
||||
|
||||
// 清理应移除过期条目
|
||||
limiter.cleanup();
|
||||
|
||||
let requests = limiter.requests.lock();
|
||||
assert!(requests.is_empty(), "清理后应无过期条目");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_default_config() {
|
||||
let config = RateLimitConfig::default();
|
||||
assert!(!config.enabled);
|
||||
assert_eq!(config.requests_per_minute, 60);
|
||||
assert_eq!(config.window_secs, 60);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,62 @@
|
||||
//! 安全中间件
|
||||
//!
|
||||
//! 提供请求体大小限制和请求超时控制
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::time::Duration;
|
||||
|
||||
/// 安全中间件配置
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct SecurityMiddlewareConfig {
|
||||
/// 最大请求体大小(字节),默认 10MB
|
||||
#[serde(default = "default_max_body_size")]
|
||||
pub max_body_size: usize,
|
||||
/// 请求超时(秒),默认 300 秒(LLM 请求可能很长)
|
||||
#[serde(default = "default_request_timeout_secs")]
|
||||
pub request_timeout_secs: u64,
|
||||
}
|
||||
|
||||
fn default_max_body_size() -> usize {
|
||||
10 * 1024 * 1024 // 10MB
|
||||
}
|
||||
|
||||
fn default_request_timeout_secs() -> u64 {
|
||||
300 // 5 分钟
|
||||
}
|
||||
|
||||
impl Default for SecurityMiddlewareConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
max_body_size: default_max_body_size(),
|
||||
request_timeout_secs: default_request_timeout_secs(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl SecurityMiddlewareConfig {
|
||||
/// 获取请求超时 Duration
|
||||
pub fn request_timeout(&self) -> Duration {
|
||||
Duration::from_secs(self.request_timeout_secs)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_default_config() {
|
||||
let config = SecurityMiddlewareConfig::default();
|
||||
assert_eq!(config.max_body_size, 10 * 1024 * 1024);
|
||||
assert_eq!(config.request_timeout_secs, 300);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_request_timeout() {
|
||||
let config = SecurityMiddlewareConfig {
|
||||
max_body_size: 1024,
|
||||
request_timeout_secs: 60,
|
||||
};
|
||||
assert_eq!(config.request_timeout(), Duration::from_secs(60));
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user