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:
@@ -0,0 +1,31 @@
|
||||
# ZeroClaw → ProxyCast/Aster-Rust 借鉴计划 - 实施状态
|
||||
|
||||
## 阶段 1:快速胜利 ✅ 完成
|
||||
|
||||
| # | 任务 | 层 | 状态 | 文件 |
|
||||
|---|------|-----|------|------|
|
||||
| 1-A | 错误分类和智能重试 | Aster | ✅ | `core/retry_logic.rs` |
|
||||
| 1-B | 统一 Observer Trait | Aster | ✅ | `observability/` |
|
||||
| 1-C | 请求体大小和超时限制 | ProxyCast | ✅ | `server/middleware/security.rs` |
|
||||
| 1-D | 滑动窗口速率限制 | ProxyCast | ✅ | `server/middleware/rate_limit.rs` |
|
||||
| 1-E | 凭证清理 | ProxyCast | ✅ | `core/sanitizer.rs` |
|
||||
| 1-F | 历史修剪策略 | ProxyCast | ✅ | `processor/conversation_manager.rs` |
|
||||
|
||||
## 阶段 2:核心增强 ✅ 完成
|
||||
|
||||
| # | 任务 | 层 | 状态 | 文件 |
|
||||
|---|------|-----|------|------|
|
||||
| 2-A | 组件监督者模式 | Aster | ✅ | `core/supervisor.rs` |
|
||||
| 2-B | HeartbeatEngine | Aster | ✅ | `heartbeat/` |
|
||||
| 2-C | SecurityPolicy Trait | Aster | ✅ | `security/policy.rs` |
|
||||
| 2-D | 配对认证系统 | ProxyCast | ✅ | `server/auth/pairing.rs` |
|
||||
| 2-E | 幂等性中间件 | ProxyCast | ✅ | `server/middleware/idempotency.rs` |
|
||||
| 2-F | 提示路由系统 | ProxyCast | ✅ | `core/router/hint_router.rs` |
|
||||
|
||||
## 阶段 3:高级功能 ✅ 完成
|
||||
|
||||
| # | 任务 | 层 | 状态 | 文件 |
|
||||
|---|------|-----|------|------|
|
||||
| 3-A | ChaCha20-Poly1305 加密 | ProxyCast | ✅ | `credential/encryption.rs` |
|
||||
| 3-B | 对话摘要功能 | ProxyCast | ✅ | `processor/conversation_summarizer.rs` |
|
||||
| 3-C | 配置热重载增强 | Aster | ✅ | `config/watcher.rs` |
|
||||
Generated
+70
@@ -8,6 +8,16 @@ version = "2.0.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "320119579fcad9c21884f5c4861d16174d0e06250625266f50fe6898340abefa"
|
||||
|
||||
[[package]]
|
||||
name = "aead"
|
||||
version = "0.5.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d122413f284cf2d62fb1b7db97e02edb8cda96d769b16e443a4f6195e35662b0"
|
||||
dependencies = [
|
||||
"crypto-common",
|
||||
"generic-array",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "aes"
|
||||
version = "0.8.4"
|
||||
@@ -1654,6 +1664,30 @@ version = "0.2.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724"
|
||||
|
||||
[[package]]
|
||||
name = "chacha20"
|
||||
version = "0.9.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c3613f74bd2eac03dad61bd53dbe620703d4371614fe0bc3b9f04dd36fe4e818"
|
||||
dependencies = [
|
||||
"cfg-if",
|
||||
"cipher",
|
||||
"cpufeatures",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "chacha20poly1305"
|
||||
version = "0.10.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "10cd79432192d1c0f4e1a0fef9527696cc039165d729fb41b3f4f4f354c2dc35"
|
||||
dependencies = [
|
||||
"aead",
|
||||
"chacha20",
|
||||
"cipher",
|
||||
"poly1305",
|
||||
"zeroize",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "chrono"
|
||||
version = "0.4.43"
|
||||
@@ -1686,6 +1720,7 @@ checksum = "773f3b9af64447d2ce9850330c473515014aa235e6a783b02db81ff39e4a3dad"
|
||||
dependencies = [
|
||||
"crypto-common",
|
||||
"inout",
|
||||
"zeroize",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
@@ -2161,6 +2196,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a"
|
||||
dependencies = [
|
||||
"generic-array",
|
||||
"rand_core 0.6.4",
|
||||
"typenum",
|
||||
]
|
||||
|
||||
@@ -5769,6 +5805,12 @@ version = "1.70.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe"
|
||||
|
||||
[[package]]
|
||||
name = "opaque-debug"
|
||||
version = "0.3.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381"
|
||||
|
||||
[[package]]
|
||||
name = "open"
|
||||
version = "5.3.3"
|
||||
@@ -6423,6 +6465,17 @@ dependencies = [
|
||||
"windows-sys 0.61.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "poly1305"
|
||||
version = "0.8.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "8159bd90725d2df49889a078b54f4f79e87f1f8a8444194cdca81d38f5393abf"
|
||||
dependencies = [
|
||||
"cpufeatures",
|
||||
"opaque-debug",
|
||||
"universal-hash",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "portable-atomic"
|
||||
version = "1.13.1"
|
||||
@@ -6811,14 +6864,18 @@ name = "proxycast-credential"
|
||||
version = "0.68.0"
|
||||
dependencies = [
|
||||
"axum 0.7.9",
|
||||
"base64 0.22.1",
|
||||
"chacha20poly1305",
|
||||
"chrono",
|
||||
"dashmap 5.5.3",
|
||||
"proptest",
|
||||
"proxycast-core",
|
||||
"proxycast-infra",
|
||||
"rand 0.8.5",
|
||||
"reqwest 0.12.28",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"sha2",
|
||||
"tempfile",
|
||||
"tokio",
|
||||
"tracing",
|
||||
@@ -6970,6 +7027,7 @@ dependencies = [
|
||||
"chrono",
|
||||
"dirs 5.0.1",
|
||||
"futures",
|
||||
"hex",
|
||||
"parking_lot",
|
||||
"proptest",
|
||||
"proxycast-agent",
|
||||
@@ -6983,11 +7041,13 @@ dependencies = [
|
||||
"proxycast-server-utils",
|
||||
"proxycast-services",
|
||||
"proxycast-websocket",
|
||||
"rand 0.8.5",
|
||||
"regex",
|
||||
"reqwest 0.12.28",
|
||||
"rusqlite",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"sha2",
|
||||
"subtle",
|
||||
"tokio",
|
||||
"tokio-util",
|
||||
@@ -10305,6 +10365,16 @@ version = "0.1.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "39ec24b3121d976906ece63c9daad25b85969647682eee313cb5779fdd69e14e"
|
||||
|
||||
[[package]]
|
||||
name = "universal-hash"
|
||||
version = "0.5.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "fc1de2c688dc15305988b563c3854064043356019f97a4b46276fe734c4f07ea"
|
||||
dependencies = [
|
||||
"crypto-common",
|
||||
"subtle",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "unsafe-libyaml"
|
||||
version = "0.2.11"
|
||||
|
||||
@@ -55,7 +55,7 @@ tracing-subscriber = "0.3"
|
||||
axum = { version = "0.7", features = ["ws"] }
|
||||
axum-server = { version = "0.7", features = ["tls-rustls"] }
|
||||
tower = "0.5"
|
||||
tower-http = { version = "0.6", features = ["limit", "cors"] }
|
||||
tower-http = { version = "0.6", features = ["limit", "cors", "timeout"] }
|
||||
|
||||
# HTTP 客户端
|
||||
reqwest = { version = "0.12", features = ["json", "stream", "gzip", "brotli", "deflate"] }
|
||||
|
||||
@@ -59,6 +59,9 @@ pub mod event_emit;
|
||||
// 网络工具
|
||||
pub mod network;
|
||||
|
||||
// 凭证清理(敏感信息过滤)
|
||||
pub mod sanitizer;
|
||||
|
||||
// 数据层
|
||||
pub mod content;
|
||||
pub mod database;
|
||||
|
||||
@@ -0,0 +1,285 @@
|
||||
//! 提示路由器
|
||||
//!
|
||||
//! 支持通过消息前缀提示(hint)将请求路由到不同的 Provider 和模型。
|
||||
//!
|
||||
//! 提示格式:`[hint] 消息内容`
|
||||
//! 例如:`[reasoning] 请分析这段代码的复杂度`
|
||||
//! `[fast] 翻译这句话`
|
||||
//! `[code] 实现一个排序算法`
|
||||
|
||||
use crate::ProviderType;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
|
||||
/// 提示路由配置
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct HintRouterConfig {
|
||||
/// 是否启用提示路由
|
||||
#[serde(default)]
|
||||
pub enabled: bool,
|
||||
/// 提示路由规则
|
||||
#[serde(default)]
|
||||
pub routes: Vec<HintRouteEntry>,
|
||||
}
|
||||
|
||||
impl Default for HintRouterConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
enabled: false,
|
||||
routes: Vec::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 单条提示路由配置
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct HintRouteEntry {
|
||||
/// 提示关键词(如 "reasoning", "fast", "code")
|
||||
pub hint: String,
|
||||
/// 目标 Provider
|
||||
pub provider: ProviderType,
|
||||
/// 目标模型
|
||||
pub model: String,
|
||||
}
|
||||
|
||||
/// 已解析的提示路由
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct HintRoute {
|
||||
/// 提示关键词
|
||||
pub hint: String,
|
||||
/// 目标 Provider
|
||||
pub provider: ProviderType,
|
||||
/// 目标模型
|
||||
pub model: String,
|
||||
}
|
||||
|
||||
/// 提示匹配结果
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct HintMatch {
|
||||
/// 匹配到的路由
|
||||
pub route: HintRoute,
|
||||
/// 去除提示前缀后的消息内容
|
||||
pub stripped_message: String,
|
||||
}
|
||||
|
||||
/// 提示路由器
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct HintRouter {
|
||||
enabled: bool,
|
||||
/// hint 关键词 -> 路由(小写匹配)
|
||||
routes: HashMap<String, HintRoute>,
|
||||
}
|
||||
|
||||
impl HintRouter {
|
||||
/// 从配置创建提示路由器
|
||||
pub fn from_config(config: &HintRouterConfig) -> Self {
|
||||
let mut routes = HashMap::new();
|
||||
|
||||
if config.enabled {
|
||||
for entry in &config.routes {
|
||||
let key = entry.hint.to_lowercase();
|
||||
routes.insert(
|
||||
key,
|
||||
HintRoute {
|
||||
hint: entry.hint.clone(),
|
||||
provider: entry.provider,
|
||||
model: entry.model.clone(),
|
||||
},
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Self {
|
||||
enabled: config.enabled,
|
||||
routes,
|
||||
}
|
||||
}
|
||||
|
||||
/// 是否启用
|
||||
pub fn is_enabled(&self) -> bool {
|
||||
self.enabled
|
||||
}
|
||||
|
||||
/// 获取已注册的路由数量
|
||||
pub fn route_count(&self) -> usize {
|
||||
self.routes.len()
|
||||
}
|
||||
|
||||
/// 根据提示关键词查找路由
|
||||
pub fn route_by_hint(&self, hint: &str) -> Option<&HintRoute> {
|
||||
if !self.enabled {
|
||||
return None;
|
||||
}
|
||||
self.routes.get(&hint.to_lowercase())
|
||||
}
|
||||
|
||||
/// 从消息中提取提示并匹配路由
|
||||
///
|
||||
/// 支持格式:`[hint] 消息内容` 或 `[hint]消息内容`
|
||||
/// 提示匹配不区分大小写
|
||||
pub fn match_message(&self, message: &str) -> Option<HintMatch> {
|
||||
if !self.enabled {
|
||||
return None;
|
||||
}
|
||||
|
||||
let trimmed = message.trim_start();
|
||||
if !trimmed.starts_with('[') {
|
||||
return None;
|
||||
}
|
||||
|
||||
let close_bracket = trimmed.find(']')?;
|
||||
let hint = trimmed[1..close_bracket].trim();
|
||||
|
||||
if hint.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
let route = self.routes.get(&hint.to_lowercase())?;
|
||||
|
||||
// 提取去除前缀后的消息
|
||||
let rest = trimmed[close_bracket + 1..].trim_start();
|
||||
|
||||
Some(HintMatch {
|
||||
route: route.clone(),
|
||||
stripped_message: rest.to_string(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for HintRouter {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
enabled: false,
|
||||
routes: HashMap::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn test_config() -> HintRouterConfig {
|
||||
HintRouterConfig {
|
||||
enabled: true,
|
||||
routes: vec![
|
||||
HintRouteEntry {
|
||||
hint: "reasoning".to_string(),
|
||||
provider: ProviderType::Kiro,
|
||||
model: "claude-sonnet-4-5-20250514".to_string(),
|
||||
},
|
||||
HintRouteEntry {
|
||||
hint: "fast".to_string(),
|
||||
provider: ProviderType::Gemini,
|
||||
model: "gemini-2.0-flash".to_string(),
|
||||
},
|
||||
HintRouteEntry {
|
||||
hint: "code".to_string(),
|
||||
provider: ProviderType::Kiro,
|
||||
model: "claude-sonnet-4-5-20250514".to_string(),
|
||||
},
|
||||
],
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_disabled_router() {
|
||||
let config = HintRouterConfig::default();
|
||||
let router = HintRouter::from_config(&config);
|
||||
assert!(!router.is_enabled());
|
||||
assert_eq!(router.route_count(), 0);
|
||||
assert!(router.route_by_hint("reasoning").is_none());
|
||||
assert!(router.match_message("[reasoning] test").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_route_by_hint() {
|
||||
let router = HintRouter::from_config(&test_config());
|
||||
assert!(router.is_enabled());
|
||||
assert_eq!(router.route_count(), 3);
|
||||
|
||||
let route = router.route_by_hint("reasoning").unwrap();
|
||||
assert_eq!(route.provider, ProviderType::Kiro);
|
||||
|
||||
let route = router.route_by_hint("fast").unwrap();
|
||||
assert_eq!(route.provider, ProviderType::Gemini);
|
||||
|
||||
assert!(router.route_by_hint("unknown").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_case_insensitive_hint() {
|
||||
let router = HintRouter::from_config(&test_config());
|
||||
assert!(router.route_by_hint("Reasoning").is_some());
|
||||
assert!(router.route_by_hint("FAST").is_some());
|
||||
assert!(router.route_by_hint("Code").is_some());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_match_message_basic() {
|
||||
let router = HintRouter::from_config(&test_config());
|
||||
|
||||
let m = router.match_message("[reasoning] 请分析这段代码").unwrap();
|
||||
assert_eq!(m.route.hint, "reasoning");
|
||||
assert_eq!(m.route.provider, ProviderType::Kiro);
|
||||
assert_eq!(m.stripped_message, "请分析这段代码");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_match_message_no_space() {
|
||||
let router = HintRouter::from_config(&test_config());
|
||||
|
||||
let m = router.match_message("[fast]翻译这句话").unwrap();
|
||||
assert_eq!(m.route.hint, "fast");
|
||||
assert_eq!(m.stripped_message, "翻译这句话");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_match_message_case_insensitive() {
|
||||
let router = HintRouter::from_config(&test_config());
|
||||
|
||||
let m = router.match_message("[REASONING] test").unwrap();
|
||||
assert_eq!(m.route.hint, "reasoning");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_match_message_leading_whitespace() {
|
||||
let router = HintRouter::from_config(&test_config());
|
||||
|
||||
let m = router.match_message(" [fast] hello").unwrap();
|
||||
assert_eq!(m.route.hint, "fast");
|
||||
assert_eq!(m.stripped_message, "hello");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_match_message_no_hint() {
|
||||
let router = HintRouter::from_config(&test_config());
|
||||
assert!(router.match_message("普通消息").is_none());
|
||||
assert!(router.match_message("").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_match_message_unknown_hint() {
|
||||
let router = HintRouter::from_config(&test_config());
|
||||
assert!(router.match_message("[unknown] test").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_match_message_empty_hint() {
|
||||
let router = HintRouter::from_config(&test_config());
|
||||
assert!(router.match_message("[] test").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_match_message_no_closing_bracket() {
|
||||
let router = HintRouter::from_config(&test_config());
|
||||
assert!(router.match_message("[reasoning test").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_default_config() {
|
||||
let config = HintRouterConfig::default();
|
||||
assert!(!config.enabled);
|
||||
assert!(config.routes.is_empty());
|
||||
}
|
||||
}
|
||||
@@ -10,13 +10,18 @@
|
||||
//!
|
||||
//! 模型映射:
|
||||
//! - 支持模型别名映射(如 `gpt-4` -> `claude-sonnet-4-5-20250514`)
|
||||
//!
|
||||
//! 提示路由:
|
||||
//! - 支持消息前缀提示路由(如 `[reasoning] 请分析...`)
|
||||
|
||||
mod amp_router;
|
||||
mod hint_router;
|
||||
mod mapper;
|
||||
mod provider_router;
|
||||
mod route_registry;
|
||||
mod rules;
|
||||
|
||||
pub use amp_router::AmpRouter;
|
||||
pub use hint_router::{HintMatch, HintRoute, HintRouter, HintRouterConfig, HintRouteEntry};
|
||||
pub use mapper::ModelMapper;
|
||||
pub use rules::Router;
|
||||
|
||||
@@ -0,0 +1,238 @@
|
||||
//! 凭证清理模块
|
||||
//!
|
||||
//! 使用正则表达式从文本中清理敏感信息(API 密钥、token、密码等)
|
||||
|
||||
use regex::Regex;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::sync::OnceLock;
|
||||
|
||||
/// 清理配置
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct SanitizeConfig {
|
||||
/// 是否启用
|
||||
#[serde(default = "default_enabled")]
|
||||
pub enabled: bool,
|
||||
/// 替换文本
|
||||
#[serde(default = "default_replacement")]
|
||||
pub replacement: String,
|
||||
/// 用户自定义正则模式
|
||||
#[serde(default)]
|
||||
pub custom_patterns: Vec<String>,
|
||||
}
|
||||
|
||||
fn default_enabled() -> bool {
|
||||
true
|
||||
}
|
||||
fn default_replacement() -> String {
|
||||
"[REDACTED]".to_string()
|
||||
}
|
||||
|
||||
impl Default for SanitizeConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
enabled: true,
|
||||
replacement: default_replacement(),
|
||||
custom_patterns: Vec::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 凭证清理器
|
||||
pub struct CredentialSanitizer {
|
||||
config: SanitizeConfig,
|
||||
custom_regexes: Vec<Regex>,
|
||||
}
|
||||
|
||||
/// 内置的敏感信息正则模式
|
||||
fn builtin_patterns() -> &'static [Regex] {
|
||||
static PATTERNS: OnceLock<Vec<Regex>> = OnceLock::new();
|
||||
PATTERNS.get_or_init(|| {
|
||||
let patterns = [
|
||||
// OpenAI / Anthropic API 密钥
|
||||
r"sk-[a-zA-Z0-9_-]{20,}",
|
||||
// Anthropic 密钥
|
||||
r"sk-ant-[a-zA-Z0-9_-]{20,}",
|
||||
// AWS Access Key
|
||||
r"AKIA[0-9A-Z]{16}",
|
||||
// Groq 密钥
|
||||
r"gsk_[a-zA-Z0-9]{20,}",
|
||||
// Google API 密钥
|
||||
r"AIza[0-9A-Za-z_-]{35}",
|
||||
// Bearer token
|
||||
r"Bearer\s+[a-zA-Z0-9_\-.]+",
|
||||
// 通用 key=value 模式
|
||||
r"(?i)(api[_-]?key|secret[_-]?key|access[_-]?token|auth[_-]?token|password|passwd|secret)\s*[=:]\s*\S+",
|
||||
// GitHub token
|
||||
r"gh[pousr]_[A-Za-z0-9_]{36,}",
|
||||
// 通用长 hex/base64 token(40+ 字符)
|
||||
r#"(?i)(token|key|secret|credential)\s*[=:]\s*['"]?[a-zA-Z0-9+/=_-]{40,}['"]?"#,
|
||||
];
|
||||
patterns
|
||||
.iter()
|
||||
.filter_map(|p| Regex::new(p).ok())
|
||||
.collect()
|
||||
})
|
||||
}
|
||||
|
||||
impl CredentialSanitizer {
|
||||
/// 创建新的清理器
|
||||
pub fn new(config: SanitizeConfig) -> Self {
|
||||
let custom_regexes = config
|
||||
.custom_patterns
|
||||
.iter()
|
||||
.filter_map(|p| Regex::new(p).ok())
|
||||
.collect();
|
||||
Self {
|
||||
config,
|
||||
custom_regexes,
|
||||
}
|
||||
}
|
||||
|
||||
/// 创建默认清理器
|
||||
pub fn with_defaults() -> Self {
|
||||
Self::new(SanitizeConfig::default())
|
||||
}
|
||||
|
||||
/// 清理文本中的敏感信息
|
||||
pub fn sanitize(&self, text: &str) -> String {
|
||||
if !self.config.enabled {
|
||||
return text.to_string();
|
||||
}
|
||||
|
||||
let mut result = text.to_string();
|
||||
let replacement = &self.config.replacement;
|
||||
|
||||
// 应用内置模式
|
||||
for pattern in builtin_patterns() {
|
||||
result = pattern.replace_all(&result, replacement.as_str()).to_string();
|
||||
}
|
||||
|
||||
// 应用自定义模式
|
||||
for pattern in &self.custom_regexes {
|
||||
result = pattern.replace_all(&result, replacement.as_str()).to_string();
|
||||
}
|
||||
|
||||
result
|
||||
}
|
||||
|
||||
/// 检查文本是否包含敏感信息
|
||||
pub fn contains_sensitive(&self, text: &str) -> bool {
|
||||
for pattern in builtin_patterns() {
|
||||
if pattern.is_match(text) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
for pattern in &self.custom_regexes {
|
||||
if pattern.is_match(text) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
false
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_sanitize_openai_key() {
|
||||
let s = CredentialSanitizer::with_defaults();
|
||||
let input = "my key is sk-abc123def456ghi789jkl012mno";
|
||||
let result = s.sanitize(input);
|
||||
assert!(!result.contains("sk-abc123"));
|
||||
assert!(result.contains("[REDACTED]"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_sanitize_anthropic_key() {
|
||||
let s = CredentialSanitizer::with_defaults();
|
||||
let input = "key: sk-ant-api03-abcdefghijklmnopqrstuvwxyz";
|
||||
let result = s.sanitize(input);
|
||||
assert!(!result.contains("sk-ant-"));
|
||||
assert!(result.contains("[REDACTED]"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_sanitize_aws_key() {
|
||||
let s = CredentialSanitizer::with_defaults();
|
||||
let input = "aws_access_key_id = AKIAIOSFODNN7EXAMPLE";
|
||||
let result = s.sanitize(input);
|
||||
assert!(!result.contains("AKIAIOSFODNN7EXAMPLE"));
|
||||
assert!(result.contains("[REDACTED]"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_sanitize_bearer_token() {
|
||||
let s = CredentialSanitizer::with_defaults();
|
||||
let input = "Authorization: Bearer eyJhbGciOiJIUzI1NiJ9.test";
|
||||
let result = s.sanitize(input);
|
||||
assert!(!result.contains("eyJhbGci"));
|
||||
assert!(result.contains("[REDACTED]"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_sanitize_key_value_pairs() {
|
||||
let s = CredentialSanitizer::with_defaults();
|
||||
let input = "api_key=super_secret_value_123";
|
||||
let result = s.sanitize(input);
|
||||
assert!(!result.contains("super_secret_value_123"));
|
||||
assert!(result.contains("[REDACTED]"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_sanitize_github_token() {
|
||||
let s = CredentialSanitizer::with_defaults();
|
||||
let input = "token: ghp_ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmn";
|
||||
let result = s.sanitize(input);
|
||||
assert!(!result.contains("ghp_"));
|
||||
assert!(result.contains("[REDACTED]"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_disabled_returns_original() {
|
||||
let config = SanitizeConfig {
|
||||
enabled: false,
|
||||
..Default::default()
|
||||
};
|
||||
let s = CredentialSanitizer::new(config);
|
||||
let input = "sk-abc123def456ghi789jkl012mno";
|
||||
assert_eq!(s.sanitize(input), input);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_custom_patterns() {
|
||||
let config = SanitizeConfig {
|
||||
custom_patterns: vec![r"my-custom-\d+".to_string()],
|
||||
..Default::default()
|
||||
};
|
||||
let s = CredentialSanitizer::new(config);
|
||||
let input = "value is my-custom-12345 here";
|
||||
let result = s.sanitize(input);
|
||||
assert!(!result.contains("my-custom-12345"));
|
||||
assert!(result.contains("[REDACTED]"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_contains_sensitive() {
|
||||
let s = CredentialSanitizer::with_defaults();
|
||||
assert!(s.contains_sensitive("sk-abc123def456ghi789jkl012mno"));
|
||||
assert!(s.contains_sensitive("AKIAIOSFODNN7EXAMPLE"));
|
||||
assert!(!s.contains_sensitive("hello world"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_no_false_positives() {
|
||||
let s = CredentialSanitizer::with_defaults();
|
||||
let normal_texts = [
|
||||
"Hello, this is a normal message.",
|
||||
"The temperature is 72 degrees.",
|
||||
"Please check the documentation at docs.rs",
|
||||
"User ID: 12345",
|
||||
"sk-short",
|
||||
];
|
||||
for text in &normal_texts {
|
||||
assert_eq!(s.sanitize(text), *text, "False positive on: {text}");
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -30,6 +30,12 @@ chrono.workspace = true
|
||||
# 并发
|
||||
dashmap.workspace = true
|
||||
|
||||
# 加密
|
||||
chacha20poly1305 = "0.10"
|
||||
base64.workspace = true
|
||||
rand.workspace = true
|
||||
sha2.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
proptest.workspace = true
|
||||
tempfile.workspace = true
|
||||
|
||||
@@ -0,0 +1,273 @@
|
||||
//! ChaCha20-Poly1305 AEAD 加密模块
|
||||
//!
|
||||
//! 提供凭证加密/解密功能:
|
||||
//! - ChaCha20-Poly1305 认证加密(防篡改)
|
||||
//! - 随机 nonce(每次加密生成新的 12 字节 nonce)
|
||||
//! - 密钥派生(SHA-256)
|
||||
//! - 格式:enc2:base64(nonce || ciphertext || tag)
|
||||
|
||||
use base64::{engine::general_purpose::STANDARD as BASE64, Engine};
|
||||
use chacha20poly1305::{
|
||||
aead::{Aead, KeyInit, OsRng},
|
||||
ChaCha20Poly1305, Nonce,
|
||||
};
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
/// 加密前缀标识
|
||||
const ENCRYPTED_PREFIX: &str = "enc2:";
|
||||
|
||||
/// Nonce 长度(12 字节)
|
||||
const NONCE_SIZE: usize = 12;
|
||||
|
||||
/// 加密器
|
||||
pub struct Encryptor {
|
||||
cipher: ChaCha20Poly1305,
|
||||
}
|
||||
|
||||
impl Encryptor {
|
||||
/// 从密码/密钥创建加密器
|
||||
///
|
||||
/// 使用 SHA-256 将任意长度的密钥派生为 256-bit 密钥
|
||||
pub fn new(key: &str) -> Self {
|
||||
let derived_key = Self::derive_key(key);
|
||||
let cipher = ChaCha20Poly1305::new(&derived_key.into());
|
||||
Self { cipher }
|
||||
}
|
||||
|
||||
/// 从原始 32 字节密钥创建加密器
|
||||
pub fn from_raw_key(key: &[u8; 32]) -> Self {
|
||||
let cipher = ChaCha20Poly1305::new(key.into());
|
||||
Self { cipher }
|
||||
}
|
||||
|
||||
/// 使用 SHA-256 派生 256-bit 密钥
|
||||
fn derive_key(password: &str) -> [u8; 32] {
|
||||
let mut hasher = Sha256::new();
|
||||
hasher.update(password.as_bytes());
|
||||
let result = hasher.finalize();
|
||||
let mut key = [0u8; 32];
|
||||
key.copy_from_slice(&result);
|
||||
key
|
||||
}
|
||||
|
||||
/// 加密明文
|
||||
///
|
||||
/// 返回格式:enc2:base64(nonce || ciphertext)
|
||||
pub fn encrypt(&self, plaintext: &str) -> Result<String, EncryptionError> {
|
||||
use chacha20poly1305::aead::AeadCore;
|
||||
|
||||
// 生成随机 nonce
|
||||
let nonce = ChaCha20Poly1305::generate_nonce(&mut OsRng);
|
||||
|
||||
// 加密
|
||||
let ciphertext = self
|
||||
.cipher
|
||||
.encrypt(&nonce, plaintext.as_bytes())
|
||||
.map_err(|_| EncryptionError::EncryptionFailed)?;
|
||||
|
||||
// 组合 nonce + ciphertext
|
||||
let mut combined = Vec::with_capacity(NONCE_SIZE + ciphertext.len());
|
||||
combined.extend_from_slice(&nonce);
|
||||
combined.extend_from_slice(&ciphertext);
|
||||
|
||||
// Base64 编码并添加前缀
|
||||
Ok(format!("{}{}", ENCRYPTED_PREFIX, BASE64.encode(&combined)))
|
||||
}
|
||||
|
||||
/// 解密密文
|
||||
///
|
||||
/// 输入格式:enc2:base64(nonce || ciphertext)
|
||||
pub fn decrypt(&self, encrypted: &str) -> Result<String, EncryptionError> {
|
||||
// 检查前缀
|
||||
let encoded = encrypted
|
||||
.strip_prefix(ENCRYPTED_PREFIX)
|
||||
.ok_or(EncryptionError::InvalidFormat)?;
|
||||
|
||||
// Base64 解码
|
||||
let combined = BASE64
|
||||
.decode(encoded)
|
||||
.map_err(|_| EncryptionError::InvalidBase64)?;
|
||||
|
||||
// 分离 nonce 和 ciphertext
|
||||
if combined.len() < NONCE_SIZE {
|
||||
return Err(EncryptionError::InvalidFormat);
|
||||
}
|
||||
|
||||
let (nonce_bytes, ciphertext) = combined.split_at(NONCE_SIZE);
|
||||
let nonce = Nonce::from_slice(nonce_bytes);
|
||||
|
||||
// 解密
|
||||
let plaintext = self
|
||||
.cipher
|
||||
.decrypt(nonce, ciphertext)
|
||||
.map_err(|_| EncryptionError::DecryptionFailed)?;
|
||||
|
||||
String::from_utf8(plaintext).map_err(|_| EncryptionError::InvalidUtf8)
|
||||
}
|
||||
|
||||
/// 检查文本是否已加密
|
||||
pub fn is_encrypted(text: &str) -> bool {
|
||||
text.starts_with(ENCRYPTED_PREFIX)
|
||||
}
|
||||
|
||||
/// 加密(如果尚未加密)
|
||||
pub fn encrypt_if_needed(&self, text: &str) -> Result<String, EncryptionError> {
|
||||
if Self::is_encrypted(text) {
|
||||
Ok(text.to_string())
|
||||
} else {
|
||||
self.encrypt(text)
|
||||
}
|
||||
}
|
||||
|
||||
/// 解密(如果已加密)
|
||||
pub fn decrypt_if_needed(&self, text: &str) -> Result<String, EncryptionError> {
|
||||
if Self::is_encrypted(text) {
|
||||
self.decrypt(text)
|
||||
} else {
|
||||
Ok(text.to_string())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 加密错误
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum EncryptionError {
|
||||
/// 加密失败
|
||||
EncryptionFailed,
|
||||
/// 解密失败(密钥错误或数据被篡改)
|
||||
DecryptionFailed,
|
||||
/// 无效的格式(缺少 enc2: 前缀)
|
||||
InvalidFormat,
|
||||
/// 无效的 Base64 编码
|
||||
InvalidBase64,
|
||||
/// 无效的 UTF-8
|
||||
InvalidUtf8,
|
||||
}
|
||||
|
||||
impl std::fmt::Display for EncryptionError {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::EncryptionFailed => write!(f, "加密失败"),
|
||||
Self::DecryptionFailed => write!(f, "解密失败:密钥错误或数据被篡改"),
|
||||
Self::InvalidFormat => write!(f, "无效的加密格式"),
|
||||
Self::InvalidBase64 => write!(f, "无效的 Base64 编码"),
|
||||
Self::InvalidUtf8 => write!(f, "无效的 UTF-8 编码"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl std::error::Error for EncryptionError {}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_encrypt_decrypt_roundtrip() {
|
||||
let enc = Encryptor::new("test-password");
|
||||
let plaintext = "sk-abc123-secret-api-key";
|
||||
let encrypted = enc.encrypt(plaintext).unwrap();
|
||||
assert!(encrypted.starts_with(ENCRYPTED_PREFIX));
|
||||
let decrypted = enc.decrypt(&encrypted).unwrap();
|
||||
assert_eq!(decrypted, plaintext);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_different_nonces() {
|
||||
let enc = Encryptor::new("test-password");
|
||||
let plaintext = "same-plaintext";
|
||||
let encrypted1 = enc.encrypt(plaintext).unwrap();
|
||||
let encrypted2 = enc.encrypt(plaintext).unwrap();
|
||||
assert_ne!(encrypted1, encrypted2);
|
||||
// 两者都能正确解密
|
||||
assert_eq!(enc.decrypt(&encrypted1).unwrap(), plaintext);
|
||||
assert_eq!(enc.decrypt(&encrypted2).unwrap(), plaintext);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_wrong_key_fails() {
|
||||
let enc1 = Encryptor::new("correct-password");
|
||||
let enc2 = Encryptor::new("wrong-password");
|
||||
let encrypted = enc1.encrypt("secret").unwrap();
|
||||
assert_eq!(enc2.decrypt(&encrypted), Err(EncryptionError::DecryptionFailed));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_is_encrypted() {
|
||||
assert!(Encryptor::is_encrypted("enc2:abc123"));
|
||||
assert!(!Encryptor::is_encrypted("plain-text"));
|
||||
assert!(!Encryptor::is_encrypted("enc1:old-format"));
|
||||
assert!(!Encryptor::is_encrypted(""));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_encrypt_if_needed_already_encrypted() {
|
||||
let enc = Encryptor::new("key");
|
||||
let already = "enc2:already-encrypted-data";
|
||||
let result = enc.encrypt_if_needed(already).unwrap();
|
||||
assert_eq!(result, already);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_decrypt_if_needed_not_encrypted() {
|
||||
let enc = Encryptor::new("key");
|
||||
let plain = "not-encrypted";
|
||||
let result = enc.decrypt_if_needed(plain).unwrap();
|
||||
assert_eq!(result, plain);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_invalid_format() {
|
||||
let enc = Encryptor::new("key");
|
||||
assert_eq!(enc.decrypt("no-prefix"), Err(EncryptionError::InvalidFormat));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_invalid_base64() {
|
||||
let enc = Encryptor::new("key");
|
||||
assert_eq!(enc.decrypt("enc2:!!!invalid-base64!!!"), Err(EncryptionError::InvalidBase64));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tampered_data() {
|
||||
let enc = Encryptor::new("key");
|
||||
let encrypted = enc.encrypt("secret").unwrap();
|
||||
// 篡改密文中的一个字符
|
||||
let encoded = encrypted.strip_prefix(ENCRYPTED_PREFIX).unwrap();
|
||||
let mut bytes = BASE64.decode(encoded).unwrap();
|
||||
if let Some(last) = bytes.last_mut() {
|
||||
*last ^= 0xFF;
|
||||
}
|
||||
let tampered = format!("{}{}", ENCRYPTED_PREFIX, BASE64.encode(&bytes));
|
||||
assert_eq!(enc.decrypt(&tampered), Err(EncryptionError::DecryptionFailed));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_empty_string() {
|
||||
let enc = Encryptor::new("key");
|
||||
let encrypted = enc.encrypt("").unwrap();
|
||||
assert_eq!(enc.decrypt(&encrypted).unwrap(), "");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_unicode_content() {
|
||||
let enc = Encryptor::new("密钥");
|
||||
let plaintext = "你好世界 🌍 こんにちは";
|
||||
let encrypted = enc.encrypt(plaintext).unwrap();
|
||||
assert_eq!(enc.decrypt(&encrypted).unwrap(), plaintext);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_from_raw_key() {
|
||||
let raw_key: [u8; 32] = [
|
||||
0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08,
|
||||
0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x0f, 0x10,
|
||||
0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18,
|
||||
0x19, 0x1a, 0x1b, 0x1c, 0x1d, 0x1e, 0x1f, 0x20,
|
||||
];
|
||||
let enc = Encryptor::from_raw_key(&raw_key);
|
||||
let plaintext = "raw-key-test";
|
||||
let encrypted = enc.encrypt(plaintext).unwrap();
|
||||
assert_eq!(enc.decrypt(&encrypted).unwrap(), plaintext);
|
||||
}
|
||||
}
|
||||
@@ -9,6 +9,7 @@
|
||||
//! - `sync` - 凭证与 YAML 配置文件的同步
|
||||
|
||||
mod balancer;
|
||||
pub mod encryption;
|
||||
mod quota;
|
||||
mod sync;
|
||||
|
||||
|
||||
@@ -0,0 +1,313 @@
|
||||
//! 对话历史管理器
|
||||
//!
|
||||
//! 提供对话历史修剪策略,防止上下文溢出
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// 修剪策略
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub enum TrimStrategy {
|
||||
/// 丢弃最旧的消息
|
||||
DropOldest,
|
||||
/// 滑动窗口(保留最近 N 条)
|
||||
SlidingWindow,
|
||||
}
|
||||
|
||||
impl Default for TrimStrategy {
|
||||
fn default() -> Self {
|
||||
Self::SlidingWindow
|
||||
}
|
||||
}
|
||||
|
||||
/// 修剪配置
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct TrimConfig {
|
||||
/// 是否启用
|
||||
#[serde(default)]
|
||||
pub enabled: bool,
|
||||
/// 最大消息数
|
||||
#[serde(default = "default_max_messages")]
|
||||
pub max_messages: usize,
|
||||
/// 是否保留 system 提示
|
||||
#[serde(default = "default_preserve_system")]
|
||||
pub preserve_system_prompt: bool,
|
||||
/// 修剪策略
|
||||
#[serde(default)]
|
||||
pub strategy: TrimStrategy,
|
||||
}
|
||||
|
||||
fn default_max_messages() -> usize {
|
||||
100
|
||||
}
|
||||
fn default_preserve_system() -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
impl Default for TrimConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
enabled: false,
|
||||
max_messages: default_max_messages(),
|
||||
preserve_system_prompt: default_preserve_system(),
|
||||
strategy: TrimStrategy::default(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 修剪结果
|
||||
#[derive(Debug)]
|
||||
pub struct TrimResult {
|
||||
/// 修剪后的消息
|
||||
pub messages: Vec<serde_json::Value>,
|
||||
/// 是否进行了修剪
|
||||
pub trimmed: bool,
|
||||
/// 被移除的消息数
|
||||
pub removed_count: usize,
|
||||
}
|
||||
|
||||
/// 对话修剪器
|
||||
pub struct ConversationTrimmer {
|
||||
config: TrimConfig,
|
||||
}
|
||||
|
||||
impl ConversationTrimmer {
|
||||
pub fn new(config: TrimConfig) -> Self {
|
||||
Self { config }
|
||||
}
|
||||
|
||||
/// 修剪消息列表
|
||||
///
|
||||
/// 兼容 Anthropic 和 OpenAI 消息格式(都使用 "role" 字段)
|
||||
pub fn trim_messages(&self, messages: Vec<serde_json::Value>) -> TrimResult {
|
||||
if !self.config.enabled || messages.len() <= self.config.max_messages {
|
||||
return TrimResult {
|
||||
messages,
|
||||
trimmed: false,
|
||||
removed_count: 0,
|
||||
};
|
||||
}
|
||||
|
||||
let original_count = messages.len();
|
||||
|
||||
// 分离 system 消息和非 system 消息
|
||||
let (system_msgs, non_system_msgs): (Vec<_>, Vec<_>) =
|
||||
if self.config.preserve_system_prompt {
|
||||
messages.into_iter().partition(|msg| {
|
||||
msg.get("role")
|
||||
.and_then(|r| r.as_str())
|
||||
.map(|r| r == "system")
|
||||
.unwrap_or(false)
|
||||
})
|
||||
} else {
|
||||
(Vec::new(), messages)
|
||||
};
|
||||
|
||||
// 计算非 system 消息的最大数量
|
||||
let max_non_system = self.config.max_messages.saturating_sub(system_msgs.len());
|
||||
|
||||
// 按策略修剪
|
||||
let trimmed_non_system = match self.config.strategy {
|
||||
TrimStrategy::DropOldest | TrimStrategy::SlidingWindow => {
|
||||
let len = non_system_msgs.len();
|
||||
if len > max_non_system {
|
||||
non_system_msgs
|
||||
.into_iter()
|
||||
.skip(len - max_non_system)
|
||||
.collect()
|
||||
} else {
|
||||
non_system_msgs
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
// 合并:system 消息在前,非 system 消息在后
|
||||
let mut result = system_msgs;
|
||||
result.extend(trimmed_non_system);
|
||||
|
||||
let removed_count = original_count - result.len();
|
||||
|
||||
TrimResult {
|
||||
messages: result,
|
||||
trimmed: removed_count > 0,
|
||||
removed_count,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use serde_json::json;
|
||||
|
||||
fn make_msg(role: &str, content: &str) -> serde_json::Value {
|
||||
json!({ "role": role, "content": content })
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_disabled_no_trim() {
|
||||
let trimmer = ConversationTrimmer::new(TrimConfig {
|
||||
enabled: false,
|
||||
max_messages: 2,
|
||||
..Default::default()
|
||||
});
|
||||
let msgs = vec![
|
||||
make_msg("user", "1"),
|
||||
make_msg("assistant", "2"),
|
||||
make_msg("user", "3"),
|
||||
];
|
||||
let result = trimmer.trim_messages(msgs);
|
||||
assert!(!result.trimmed);
|
||||
assert_eq!(result.removed_count, 0);
|
||||
assert_eq!(result.messages.len(), 3);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_within_limit_no_trim() {
|
||||
let trimmer = ConversationTrimmer::new(TrimConfig {
|
||||
enabled: true,
|
||||
max_messages: 5,
|
||||
..Default::default()
|
||||
});
|
||||
let msgs = vec![make_msg("user", "1"), make_msg("assistant", "2")];
|
||||
let result = trimmer.trim_messages(msgs);
|
||||
assert!(!result.trimmed);
|
||||
assert_eq!(result.messages.len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_trim_preserves_system() {
|
||||
let trimmer = ConversationTrimmer::new(TrimConfig {
|
||||
enabled: true,
|
||||
max_messages: 3,
|
||||
preserve_system_prompt: true,
|
||||
strategy: TrimStrategy::SlidingWindow,
|
||||
});
|
||||
let msgs = vec![
|
||||
make_msg("system", "You are helpful"),
|
||||
make_msg("user", "1"),
|
||||
make_msg("assistant", "2"),
|
||||
make_msg("user", "3"),
|
||||
make_msg("assistant", "4"),
|
||||
];
|
||||
let result = trimmer.trim_messages(msgs);
|
||||
assert!(result.trimmed);
|
||||
assert_eq!(result.removed_count, 2);
|
||||
assert_eq!(result.messages.len(), 3);
|
||||
assert_eq!(result.messages[0]["role"], "system");
|
||||
assert_eq!(result.messages[1]["content"], "3");
|
||||
assert_eq!(result.messages[2]["content"], "4");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_trim_drops_oldest() {
|
||||
let trimmer = ConversationTrimmer::new(TrimConfig {
|
||||
enabled: true,
|
||||
max_messages: 2,
|
||||
preserve_system_prompt: false,
|
||||
strategy: TrimStrategy::DropOldest,
|
||||
});
|
||||
let msgs = vec![
|
||||
make_msg("user", "old"),
|
||||
make_msg("assistant", "old-reply"),
|
||||
make_msg("user", "new"),
|
||||
make_msg("assistant", "new-reply"),
|
||||
];
|
||||
let result = trimmer.trim_messages(msgs);
|
||||
assert!(result.trimmed);
|
||||
assert_eq!(result.removed_count, 2);
|
||||
assert_eq!(result.messages.len(), 2);
|
||||
assert_eq!(result.messages[0]["content"], "new");
|
||||
assert_eq!(result.messages[1]["content"], "new-reply");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_trim_with_no_system() {
|
||||
let trimmer = ConversationTrimmer::new(TrimConfig {
|
||||
enabled: true,
|
||||
max_messages: 2,
|
||||
preserve_system_prompt: true,
|
||||
strategy: TrimStrategy::SlidingWindow,
|
||||
});
|
||||
let msgs = vec![
|
||||
make_msg("user", "1"),
|
||||
make_msg("assistant", "2"),
|
||||
make_msg("user", "3"),
|
||||
];
|
||||
let result = trimmer.trim_messages(msgs);
|
||||
assert!(result.trimmed);
|
||||
assert_eq!(result.removed_count, 1);
|
||||
assert_eq!(result.messages.len(), 2);
|
||||
assert_eq!(result.messages[0]["content"], "2");
|
||||
assert_eq!(result.messages[1]["content"], "3");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_trim_all_system_messages_preserved() {
|
||||
let trimmer = ConversationTrimmer::new(TrimConfig {
|
||||
enabled: true,
|
||||
max_messages: 3,
|
||||
preserve_system_prompt: true,
|
||||
strategy: TrimStrategy::SlidingWindow,
|
||||
});
|
||||
let msgs = vec![
|
||||
make_msg("system", "sys1"),
|
||||
make_msg("system", "sys2"),
|
||||
make_msg("user", "1"),
|
||||
make_msg("assistant", "2"),
|
||||
make_msg("user", "3"),
|
||||
];
|
||||
let result = trimmer.trim_messages(msgs);
|
||||
assert!(result.trimmed);
|
||||
assert_eq!(result.messages.len(), 3);
|
||||
assert_eq!(result.messages[0]["role"], "system");
|
||||
assert_eq!(result.messages[1]["role"], "system");
|
||||
assert_eq!(result.messages[2]["content"], "3");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_openai_format_compatibility() {
|
||||
let trimmer = ConversationTrimmer::new(TrimConfig {
|
||||
enabled: true,
|
||||
max_messages: 2,
|
||||
preserve_system_prompt: true,
|
||||
strategy: TrimStrategy::SlidingWindow,
|
||||
});
|
||||
// OpenAI 格式:system/user/assistant + content 字符串
|
||||
let msgs = vec![
|
||||
json!({ "role": "system", "content": "You are a helpful assistant." }),
|
||||
json!({ "role": "user", "content": "Hello" }),
|
||||
json!({ "role": "assistant", "content": "Hi there!" }),
|
||||
json!({ "role": "user", "content": "How are you?" }),
|
||||
];
|
||||
let result = trimmer.trim_messages(msgs);
|
||||
assert!(result.trimmed);
|
||||
assert_eq!(result.messages.len(), 2);
|
||||
assert_eq!(result.messages[0]["role"], "system");
|
||||
assert_eq!(result.messages[1]["content"], "How are you?");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_anthropic_format_compatibility() {
|
||||
let trimmer = ConversationTrimmer::new(TrimConfig {
|
||||
enabled: true,
|
||||
max_messages: 3,
|
||||
preserve_system_prompt: true,
|
||||
strategy: TrimStrategy::SlidingWindow,
|
||||
});
|
||||
// Anthropic 格式:content 可以是数组
|
||||
let msgs = vec![
|
||||
json!({ "role": "system", "content": "You are Claude." }),
|
||||
json!({ "role": "user", "content": [{"type": "text", "text": "msg1"}] }),
|
||||
json!({ "role": "assistant", "content": [{"type": "text", "text": "reply1"}] }),
|
||||
json!({ "role": "user", "content": [{"type": "text", "text": "msg2"}] }),
|
||||
json!({ "role": "assistant", "content": [{"type": "text", "text": "reply2"}] }),
|
||||
];
|
||||
let result = trimmer.trim_messages(msgs);
|
||||
assert!(result.trimmed);
|
||||
assert_eq!(result.messages.len(), 3);
|
||||
assert_eq!(result.messages[0]["role"], "system");
|
||||
assert_eq!(result.messages[1]["role"], "user");
|
||||
assert_eq!(result.messages[2]["role"], "assistant");
|
||||
assert_eq!(result.messages[2]["content"][0]["text"], "reply2");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,381 @@
|
||||
//! 对话摘要器
|
||||
//!
|
||||
//! 当对话历史过长时,使用 LLM 生成简洁摘要替代旧消息,
|
||||
//! 保留关键上下文同时减少 token 消耗。
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// 摘要配置
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct SummaryConfig {
|
||||
/// 是否启用
|
||||
#[serde(default)]
|
||||
pub enabled: bool,
|
||||
/// 触发摘要的消息数阈值
|
||||
#[serde(default = "default_threshold")]
|
||||
pub threshold_messages: usize,
|
||||
/// 摘要后保留的最近消息数
|
||||
#[serde(default = "default_keep_recent")]
|
||||
pub keep_recent_messages: usize,
|
||||
/// 摘要最大要点数
|
||||
#[serde(default = "default_max_points")]
|
||||
pub max_summary_points: usize,
|
||||
}
|
||||
|
||||
fn default_threshold() -> usize {
|
||||
50
|
||||
}
|
||||
fn default_keep_recent() -> usize {
|
||||
20
|
||||
}
|
||||
fn default_max_points() -> usize {
|
||||
12
|
||||
}
|
||||
|
||||
impl Default for SummaryConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
enabled: false,
|
||||
threshold_messages: default_threshold(),
|
||||
keep_recent_messages: default_keep_recent(),
|
||||
max_summary_points: default_max_points(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 摘要请求
|
||||
///
|
||||
/// 包含需要发送给 LLM 的摘要请求信息
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SummaryRequest {
|
||||
/// 摘要 prompt(system 消息)
|
||||
pub system_prompt: String,
|
||||
/// 需要摘要的消息(作为 user 消息发送)
|
||||
pub messages_to_summarize: String,
|
||||
}
|
||||
|
||||
/// 摘要结果
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SummaryResult {
|
||||
/// 摘要后的消息列表(摘要 system 消息 + 保留的最近消息)
|
||||
pub messages: Vec<serde_json::Value>,
|
||||
/// 是否进行了摘要
|
||||
pub summarized: bool,
|
||||
/// 被摘要的消息数
|
||||
pub summarized_count: usize,
|
||||
}
|
||||
|
||||
/// 对话摘要器
|
||||
pub struct ConversationSummarizer {
|
||||
config: SummaryConfig,
|
||||
}
|
||||
|
||||
impl ConversationSummarizer {
|
||||
pub fn new(config: SummaryConfig) -> Self {
|
||||
Self { config }
|
||||
}
|
||||
|
||||
/// 判断是否需要摘要
|
||||
pub fn should_summarize(&self, message_count: usize) -> bool {
|
||||
self.config.enabled && message_count > self.config.threshold_messages
|
||||
}
|
||||
|
||||
/// 构建摘要请求
|
||||
///
|
||||
/// 将需要摘要的旧消息格式化为 LLM 请求
|
||||
pub fn build_summary_request(
|
||||
&self,
|
||||
messages: &[serde_json::Value],
|
||||
) -> Option<SummaryRequest> {
|
||||
if !self.should_summarize(messages.len()) {
|
||||
return None;
|
||||
}
|
||||
|
||||
let non_system_msgs: Vec<_> = messages
|
||||
.iter()
|
||||
.filter(|msg| {
|
||||
msg.get("role")
|
||||
.and_then(|r| r.as_str())
|
||||
.map(|r| r != "system")
|
||||
.unwrap_or(true)
|
||||
})
|
||||
.collect();
|
||||
|
||||
let keep = self.config.keep_recent_messages.min(non_system_msgs.len());
|
||||
let to_summarize = non_system_msgs.len().saturating_sub(keep);
|
||||
|
||||
if to_summarize == 0 {
|
||||
return None;
|
||||
}
|
||||
|
||||
let msgs_text: Vec<String> = non_system_msgs[..to_summarize]
|
||||
.iter()
|
||||
.map(|msg| {
|
||||
let role = msg
|
||||
.get("role")
|
||||
.and_then(|r| r.as_str())
|
||||
.unwrap_or("unknown");
|
||||
let content = extract_content_text(msg);
|
||||
format!("[{role}]: {content}")
|
||||
})
|
||||
.collect();
|
||||
|
||||
let messages_text = msgs_text.join("\n\n");
|
||||
|
||||
let system_prompt = format!(
|
||||
"你是一个对话摘要助手。请将以下对话历史总结为最多 {} 个关键要点。\n\
|
||||
要求:\n\
|
||||
- 保留重要的决策、结论和上下文\n\
|
||||
- 保留关键的技术细节和代码引用\n\
|
||||
- 使用简洁的要点格式\n\
|
||||
- 按时间顺序组织\n\
|
||||
- 不要遗漏用户的关键需求",
|
||||
self.config.max_summary_points
|
||||
);
|
||||
|
||||
Some(SummaryRequest {
|
||||
system_prompt,
|
||||
messages_to_summarize: messages_text,
|
||||
})
|
||||
}
|
||||
|
||||
/// 将摘要文本组装为最终消息列表
|
||||
///
|
||||
/// 结构:原始 system 消息 + 摘要 system 消息 + 保留的最近消息
|
||||
pub fn assemble_with_summary(
|
||||
&self,
|
||||
original_messages: &[serde_json::Value],
|
||||
summary_text: &str,
|
||||
) -> SummaryResult {
|
||||
let (system_msgs, non_system_msgs): (Vec<_>, Vec<_>) =
|
||||
original_messages.iter().partition(|msg| {
|
||||
msg.get("role")
|
||||
.and_then(|r| r.as_str())
|
||||
.map(|r| r == "system")
|
||||
.unwrap_or(false)
|
||||
});
|
||||
|
||||
let keep = self.config.keep_recent_messages.min(non_system_msgs.len());
|
||||
let summarized_count = non_system_msgs.len().saturating_sub(keep);
|
||||
|
||||
let mut result = Vec::new();
|
||||
|
||||
// 1. 原始 system 消息
|
||||
for msg in &system_msgs {
|
||||
result.push((*msg).clone());
|
||||
}
|
||||
|
||||
// 2. 摘要 system 消息
|
||||
if summarized_count > 0 && !summary_text.is_empty() {
|
||||
result.push(serde_json::json!({
|
||||
"role": "system",
|
||||
"content": format!(
|
||||
"[对话摘要 - 以下是之前 {} 条消息的摘要]\n\n{}",
|
||||
summarized_count, summary_text
|
||||
)
|
||||
}));
|
||||
}
|
||||
|
||||
// 3. 保留的最近消息
|
||||
let start = non_system_msgs.len().saturating_sub(keep);
|
||||
for msg in &non_system_msgs[start..] {
|
||||
result.push((*msg).clone());
|
||||
}
|
||||
|
||||
SummaryResult {
|
||||
messages: result,
|
||||
summarized: summarized_count > 0,
|
||||
summarized_count,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 从消息中提取文本内容
|
||||
///
|
||||
/// 兼容 OpenAI 格式(content 为字符串)和 Anthropic 格式(content 为数组)
|
||||
fn extract_content_text(msg: &serde_json::Value) -> String {
|
||||
match msg.get("content") {
|
||||
Some(serde_json::Value::String(s)) => s.clone(),
|
||||
Some(serde_json::Value::Array(arr)) => arr
|
||||
.iter()
|
||||
.filter_map(|item| {
|
||||
if item.get("type").and_then(|t| t.as_str()) == Some("text") {
|
||||
item.get("text").and_then(|t| t.as_str()).map(String::from)
|
||||
} else {
|
||||
None
|
||||
}
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n"),
|
||||
_ => String::new(),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn test_default_config() {
|
||||
let config = SummaryConfig::default();
|
||||
assert!(!config.enabled);
|
||||
assert_eq!(config.threshold_messages, 50);
|
||||
assert_eq!(config.keep_recent_messages, 20);
|
||||
assert_eq!(config.max_summary_points, 12);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_should_summarize_disabled() {
|
||||
let s = ConversationSummarizer::new(SummaryConfig {
|
||||
enabled: false,
|
||||
threshold_messages: 5,
|
||||
..Default::default()
|
||||
});
|
||||
assert!(!s.should_summarize(100));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_should_summarize_below_threshold() {
|
||||
let s = ConversationSummarizer::new(SummaryConfig {
|
||||
enabled: true,
|
||||
threshold_messages: 50,
|
||||
..Default::default()
|
||||
});
|
||||
assert!(!s.should_summarize(30));
|
||||
assert!(!s.should_summarize(50)); // 等于阈值不触发
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_should_summarize_above_threshold() {
|
||||
let s = ConversationSummarizer::new(SummaryConfig {
|
||||
enabled: true,
|
||||
threshold_messages: 5,
|
||||
..Default::default()
|
||||
});
|
||||
assert!(s.should_summarize(6));
|
||||
assert!(s.should_summarize(100));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_summary_request_none_when_disabled() {
|
||||
let s = ConversationSummarizer::new(SummaryConfig {
|
||||
enabled: false,
|
||||
threshold_messages: 2,
|
||||
keep_recent_messages: 1,
|
||||
..Default::default()
|
||||
});
|
||||
let msgs = vec![
|
||||
json!({"role": "user", "content": "a"}),
|
||||
json!({"role": "assistant", "content": "b"}),
|
||||
json!({"role": "user", "content": "c"}),
|
||||
];
|
||||
assert!(s.build_summary_request(&msgs).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_build_summary_request_with_messages() {
|
||||
let s = ConversationSummarizer::new(SummaryConfig {
|
||||
enabled: true,
|
||||
threshold_messages: 2,
|
||||
keep_recent_messages: 1,
|
||||
max_summary_points: 5,
|
||||
});
|
||||
let msgs = vec![
|
||||
json!({"role": "system", "content": "You are helpful."}),
|
||||
json!({"role": "user", "content": "Hello"}),
|
||||
json!({"role": "assistant", "content": "Hi!"}),
|
||||
json!({"role": "user", "content": "Latest"}),
|
||||
];
|
||||
// 总消息数 4 > threshold 2,非 system 消息 3 条,保留 1 条,摘要 2 条
|
||||
let req = s.build_summary_request(&msgs).unwrap();
|
||||
assert!(req.system_prompt.contains("5"));
|
||||
assert!(req.messages_to_summarize.contains("[user]: Hello"));
|
||||
assert!(req.messages_to_summarize.contains("[assistant]: Hi!"));
|
||||
assert!(!req.messages_to_summarize.contains("Latest"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_assemble_with_summary() {
|
||||
let s = ConversationSummarizer::new(SummaryConfig {
|
||||
enabled: true,
|
||||
threshold_messages: 2,
|
||||
keep_recent_messages: 1,
|
||||
..Default::default()
|
||||
});
|
||||
let msgs = vec![
|
||||
json!({"role": "user", "content": "old1"}),
|
||||
json!({"role": "assistant", "content": "old2"}),
|
||||
json!({"role": "user", "content": "recent"}),
|
||||
];
|
||||
let result = s.assemble_with_summary(&msgs, "摘要内容");
|
||||
assert!(result.summarized);
|
||||
assert_eq!(result.summarized_count, 2);
|
||||
// 摘要 system + 保留的 1 条 = 2
|
||||
assert_eq!(result.messages.len(), 2);
|
||||
assert_eq!(result.messages[0]["role"], "system");
|
||||
assert!(result.messages[0]["content"]
|
||||
.as_str()
|
||||
.unwrap()
|
||||
.contains("摘要内容"));
|
||||
assert_eq!(result.messages[1]["content"], "recent");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_assemble_preserves_system_messages() {
|
||||
let s = ConversationSummarizer::new(SummaryConfig {
|
||||
enabled: true,
|
||||
threshold_messages: 2,
|
||||
keep_recent_messages: 1,
|
||||
..Default::default()
|
||||
});
|
||||
let msgs = vec![
|
||||
json!({"role": "system", "content": "You are helpful."}),
|
||||
json!({"role": "user", "content": "old"}),
|
||||
json!({"role": "assistant", "content": "old reply"}),
|
||||
json!({"role": "user", "content": "recent"}),
|
||||
];
|
||||
let result = s.assemble_with_summary(&msgs, "summary");
|
||||
// system 原始 + 摘要 system + 保留 1 条 = 3
|
||||
assert_eq!(result.messages.len(), 3);
|
||||
assert_eq!(result.messages[0]["content"], "You are helpful.");
|
||||
assert_eq!(result.messages[1]["role"], "system");
|
||||
assert!(result.messages[1]["content"]
|
||||
.as_str()
|
||||
.unwrap()
|
||||
.contains("summary"));
|
||||
assert_eq!(result.messages[2]["content"], "recent");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_content_text_string() {
|
||||
let msg = json!({"role": "user", "content": "hello world"});
|
||||
assert_eq!(extract_content_text(&msg), "hello world");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_content_text_array() {
|
||||
// Anthropic 格式
|
||||
let msg = json!({
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": "part1"},
|
||||
{"type": "image", "source": {}},
|
||||
{"type": "text", "text": "part2"}
|
||||
]
|
||||
});
|
||||
assert_eq!(extract_content_text(&msg), "part1\npart2");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_extract_content_text_empty() {
|
||||
let msg = json!({"role": "user"});
|
||||
assert_eq!(extract_content_text(&msg), "");
|
||||
|
||||
let msg2 = json!({"role": "user", "content": null});
|
||||
assert_eq!(extract_content_text(&msg2), "");
|
||||
|
||||
let msg3 = json!({"role": "user", "content": 42});
|
||||
assert_eq!(extract_content_text(&msg3), "");
|
||||
}
|
||||
}
|
||||
@@ -6,6 +6,8 @@
|
||||
//!
|
||||
//! - `steps` - 管道步骤(认证、注入、路由、插件、Provider、遥测)
|
||||
|
||||
pub mod conversation_manager;
|
||||
pub mod conversation_summarizer;
|
||||
pub mod processor;
|
||||
pub mod steps;
|
||||
|
||||
|
||||
@@ -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