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"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "320119579fcad9c21884f5c4861d16174d0e06250625266f50fe6898340abefa"
|
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]]
|
[[package]]
|
||||||
name = "aes"
|
name = "aes"
|
||||||
version = "0.8.4"
|
version = "0.8.4"
|
||||||
@@ -1654,6 +1664,30 @@ version = "0.2.1"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724"
|
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]]
|
[[package]]
|
||||||
name = "chrono"
|
name = "chrono"
|
||||||
version = "0.4.43"
|
version = "0.4.43"
|
||||||
@@ -1686,6 +1720,7 @@ checksum = "773f3b9af64447d2ce9850330c473515014aa235e6a783b02db81ff39e4a3dad"
|
|||||||
dependencies = [
|
dependencies = [
|
||||||
"crypto-common",
|
"crypto-common",
|
||||||
"inout",
|
"inout",
|
||||||
|
"zeroize",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -2161,6 +2196,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||||||
checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a"
|
checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"generic-array",
|
"generic-array",
|
||||||
|
"rand_core 0.6.4",
|
||||||
"typenum",
|
"typenum",
|
||||||
]
|
]
|
||||||
|
|
||||||
@@ -5769,6 +5805,12 @@ version = "1.70.2"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe"
|
checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe"
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "opaque-debug"
|
||||||
|
version = "0.3.1"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381"
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "open"
|
name = "open"
|
||||||
version = "5.3.3"
|
version = "5.3.3"
|
||||||
@@ -6423,6 +6465,17 @@ dependencies = [
|
|||||||
"windows-sys 0.61.2",
|
"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]]
|
[[package]]
|
||||||
name = "portable-atomic"
|
name = "portable-atomic"
|
||||||
version = "1.13.1"
|
version = "1.13.1"
|
||||||
@@ -6811,14 +6864,18 @@ name = "proxycast-credential"
|
|||||||
version = "0.68.0"
|
version = "0.68.0"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"axum 0.7.9",
|
"axum 0.7.9",
|
||||||
|
"base64 0.22.1",
|
||||||
|
"chacha20poly1305",
|
||||||
"chrono",
|
"chrono",
|
||||||
"dashmap 5.5.3",
|
"dashmap 5.5.3",
|
||||||
"proptest",
|
"proptest",
|
||||||
"proxycast-core",
|
"proxycast-core",
|
||||||
"proxycast-infra",
|
"proxycast-infra",
|
||||||
|
"rand 0.8.5",
|
||||||
"reqwest 0.12.28",
|
"reqwest 0.12.28",
|
||||||
"serde",
|
"serde",
|
||||||
"serde_json",
|
"serde_json",
|
||||||
|
"sha2",
|
||||||
"tempfile",
|
"tempfile",
|
||||||
"tokio",
|
"tokio",
|
||||||
"tracing",
|
"tracing",
|
||||||
@@ -6970,6 +7027,7 @@ dependencies = [
|
|||||||
"chrono",
|
"chrono",
|
||||||
"dirs 5.0.1",
|
"dirs 5.0.1",
|
||||||
"futures",
|
"futures",
|
||||||
|
"hex",
|
||||||
"parking_lot",
|
"parking_lot",
|
||||||
"proptest",
|
"proptest",
|
||||||
"proxycast-agent",
|
"proxycast-agent",
|
||||||
@@ -6983,11 +7041,13 @@ dependencies = [
|
|||||||
"proxycast-server-utils",
|
"proxycast-server-utils",
|
||||||
"proxycast-services",
|
"proxycast-services",
|
||||||
"proxycast-websocket",
|
"proxycast-websocket",
|
||||||
|
"rand 0.8.5",
|
||||||
"regex",
|
"regex",
|
||||||
"reqwest 0.12.28",
|
"reqwest 0.12.28",
|
||||||
"rusqlite",
|
"rusqlite",
|
||||||
"serde",
|
"serde",
|
||||||
"serde_json",
|
"serde_json",
|
||||||
|
"sha2",
|
||||||
"subtle",
|
"subtle",
|
||||||
"tokio",
|
"tokio",
|
||||||
"tokio-util",
|
"tokio-util",
|
||||||
@@ -10305,6 +10365,16 @@ version = "0.1.1"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "39ec24b3121d976906ece63c9daad25b85969647682eee313cb5779fdd69e14e"
|
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]]
|
[[package]]
|
||||||
name = "unsafe-libyaml"
|
name = "unsafe-libyaml"
|
||||||
version = "0.2.11"
|
version = "0.2.11"
|
||||||
|
|||||||
@@ -55,7 +55,7 @@ tracing-subscriber = "0.3"
|
|||||||
axum = { version = "0.7", features = ["ws"] }
|
axum = { version = "0.7", features = ["ws"] }
|
||||||
axum-server = { version = "0.7", features = ["tls-rustls"] }
|
axum-server = { version = "0.7", features = ["tls-rustls"] }
|
||||||
tower = "0.5"
|
tower = "0.5"
|
||||||
tower-http = { version = "0.6", features = ["limit", "cors"] }
|
tower-http = { version = "0.6", features = ["limit", "cors", "timeout"] }
|
||||||
|
|
||||||
# HTTP 客户端
|
# HTTP 客户端
|
||||||
reqwest = { version = "0.12", features = ["json", "stream", "gzip", "brotli", "deflate"] }
|
reqwest = { version = "0.12", features = ["json", "stream", "gzip", "brotli", "deflate"] }
|
||||||
|
|||||||
@@ -59,6 +59,9 @@ pub mod event_emit;
|
|||||||
// 网络工具
|
// 网络工具
|
||||||
pub mod network;
|
pub mod network;
|
||||||
|
|
||||||
|
// 凭证清理(敏感信息过滤)
|
||||||
|
pub mod sanitizer;
|
||||||
|
|
||||||
// 数据层
|
// 数据层
|
||||||
pub mod content;
|
pub mod content;
|
||||||
pub mod database;
|
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`)
|
//! - 支持模型别名映射(如 `gpt-4` -> `claude-sonnet-4-5-20250514`)
|
||||||
|
//!
|
||||||
|
//! 提示路由:
|
||||||
|
//! - 支持消息前缀提示路由(如 `[reasoning] 请分析...`)
|
||||||
|
|
||||||
mod amp_router;
|
mod amp_router;
|
||||||
|
mod hint_router;
|
||||||
mod mapper;
|
mod mapper;
|
||||||
mod provider_router;
|
mod provider_router;
|
||||||
mod route_registry;
|
mod route_registry;
|
||||||
mod rules;
|
mod rules;
|
||||||
|
|
||||||
pub use amp_router::AmpRouter;
|
pub use amp_router::AmpRouter;
|
||||||
|
pub use hint_router::{HintMatch, HintRoute, HintRouter, HintRouterConfig, HintRouteEntry};
|
||||||
pub use mapper::ModelMapper;
|
pub use mapper::ModelMapper;
|
||||||
pub use rules::Router;
|
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
|
dashmap.workspace = true
|
||||||
|
|
||||||
|
# 加密
|
||||||
|
chacha20poly1305 = "0.10"
|
||||||
|
base64.workspace = true
|
||||||
|
rand.workspace = true
|
||||||
|
sha2.workspace = true
|
||||||
|
|
||||||
[dev-dependencies]
|
[dev-dependencies]
|
||||||
proptest.workspace = true
|
proptest.workspace = true
|
||||||
tempfile.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 配置文件的同步
|
//! - `sync` - 凭证与 YAML 配置文件的同步
|
||||||
|
|
||||||
mod balancer;
|
mod balancer;
|
||||||
|
pub mod encryption;
|
||||||
mod quota;
|
mod quota;
|
||||||
mod sync;
|
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、遥测)
|
//! - `steps` - 管道步骤(认证、注入、路由、插件、Provider、遥测)
|
||||||
|
|
||||||
|
pub mod conversation_manager;
|
||||||
|
pub mod conversation_summarizer;
|
||||||
pub mod processor;
|
pub mod processor;
|
||||||
pub mod steps;
|
pub mod steps;
|
||||||
|
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ serde.workspace = true
|
|||||||
serde_json.workspace = true
|
serde_json.workspace = true
|
||||||
tokio.workspace = true
|
tokio.workspace = true
|
||||||
futures.workspace = true
|
futures.workspace = true
|
||||||
|
hex.workspace = true
|
||||||
axum.workspace = true
|
axum.workspace = true
|
||||||
tower.workspace = true
|
tower.workspace = true
|
||||||
tower-http.workspace = true
|
tower-http.workspace = true
|
||||||
@@ -35,6 +36,8 @@ subtle.workspace = true
|
|||||||
async-stream.workspace = true
|
async-stream.workspace = true
|
||||||
urlencoding.workspace = true
|
urlencoding.workspace = true
|
||||||
parking_lot.workspace = true
|
parking_lot.workspace = true
|
||||||
|
rand.workspace = true
|
||||||
|
sha2.workspace = true
|
||||||
tokio-util.workspace = true
|
tokio-util.workspace = true
|
||||||
dirs.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 服务器
|
//! HTTP API 服务器
|
||||||
|
|
||||||
|
pub mod auth;
|
||||||
pub mod client_detector;
|
pub mod client_detector;
|
||||||
|
pub mod middleware;
|
||||||
|
|
||||||
use axum::{
|
use axum::{
|
||||||
extract::{DefaultBodyLimit, Path, State},
|
extract::{DefaultBodyLimit, Path, State},
|
||||||
@@ -43,6 +45,7 @@ use std::path::PathBuf;
|
|||||||
use std::sync::Arc;
|
use std::sync::Arc;
|
||||||
use tokio::sync::{oneshot, RwLock};
|
use tokio::sync::{oneshot, RwLock};
|
||||||
use tower_http::cors::CorsLayer;
|
use tower_http::cors::CorsLayer;
|
||||||
|
use tower_http::timeout::TimeoutLayer;
|
||||||
|
|
||||||
/// 记录请求统计到遥测系统
|
/// 记录请求统计到遥测系统
|
||||||
pub fn record_request_telemetry(
|
pub fn record_request_telemetry(
|
||||||
@@ -1053,6 +1056,10 @@ async fn run_server(
|
|||||||
.merge(batch_api_routes)
|
.merge(batch_api_routes)
|
||||||
.layer(cors_layer)
|
.layer(cors_layer)
|
||||||
.layer(DefaultBodyLimit::max(body_limit))
|
.layer(DefaultBodyLimit::max(body_limit))
|
||||||
|
.layer(TimeoutLayer::with_status_code(
|
||||||
|
StatusCode::REQUEST_TIMEOUT,
|
||||||
|
std::time::Duration::from_secs(300),
|
||||||
|
))
|
||||||
.with_state(state);
|
.with_state(state);
|
||||||
|
|
||||||
let addr: std::net::SocketAddr = format!("{host}:{port}")
|
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