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:
coso
2026-02-18 14:58:51 +08:00
parent 904ae548b3
commit 0c357c8bf1
21 changed files with 2506 additions and 1 deletions
+31
View File
@@ -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` |
+70
View File
@@ -8,6 +8,16 @@ version = "2.0.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "320119579fcad9c21884f5c4861d16174d0e06250625266f50fe6898340abefa"
[[package]]
name = "aead"
version = "0.5.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d122413f284cf2d62fb1b7db97e02edb8cda96d769b16e443a4f6195e35662b0"
dependencies = [
"crypto-common",
"generic-array",
]
[[package]]
name = "aes"
version = "0.8.4"
@@ -1654,6 +1664,30 @@ version = "0.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724"
[[package]]
name = "chacha20"
version = "0.9.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c3613f74bd2eac03dad61bd53dbe620703d4371614fe0bc3b9f04dd36fe4e818"
dependencies = [
"cfg-if",
"cipher",
"cpufeatures",
]
[[package]]
name = "chacha20poly1305"
version = "0.10.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "10cd79432192d1c0f4e1a0fef9527696cc039165d729fb41b3f4f4f354c2dc35"
dependencies = [
"aead",
"chacha20",
"cipher",
"poly1305",
"zeroize",
]
[[package]]
name = "chrono"
version = "0.4.43"
@@ -1686,6 +1720,7 @@ checksum = "773f3b9af64447d2ce9850330c473515014aa235e6a783b02db81ff39e4a3dad"
dependencies = [
"crypto-common",
"inout",
"zeroize",
]
[[package]]
@@ -2161,6 +2196,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a"
dependencies = [
"generic-array",
"rand_core 0.6.4",
"typenum",
]
@@ -5769,6 +5805,12 @@ version = "1.70.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe"
[[package]]
name = "opaque-debug"
version = "0.3.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381"
[[package]]
name = "open"
version = "5.3.3"
@@ -6423,6 +6465,17 @@ dependencies = [
"windows-sys 0.61.2",
]
[[package]]
name = "poly1305"
version = "0.8.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8159bd90725d2df49889a078b54f4f79e87f1f8a8444194cdca81d38f5393abf"
dependencies = [
"cpufeatures",
"opaque-debug",
"universal-hash",
]
[[package]]
name = "portable-atomic"
version = "1.13.1"
@@ -6811,14 +6864,18 @@ name = "proxycast-credential"
version = "0.68.0"
dependencies = [
"axum 0.7.9",
"base64 0.22.1",
"chacha20poly1305",
"chrono",
"dashmap 5.5.3",
"proptest",
"proxycast-core",
"proxycast-infra",
"rand 0.8.5",
"reqwest 0.12.28",
"serde",
"serde_json",
"sha2",
"tempfile",
"tokio",
"tracing",
@@ -6970,6 +7027,7 @@ dependencies = [
"chrono",
"dirs 5.0.1",
"futures",
"hex",
"parking_lot",
"proptest",
"proxycast-agent",
@@ -6983,11 +7041,13 @@ dependencies = [
"proxycast-server-utils",
"proxycast-services",
"proxycast-websocket",
"rand 0.8.5",
"regex",
"reqwest 0.12.28",
"rusqlite",
"serde",
"serde_json",
"sha2",
"subtle",
"tokio",
"tokio-util",
@@ -10305,6 +10365,16 @@ version = "0.1.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "39ec24b3121d976906ece63c9daad25b85969647682eee313cb5779fdd69e14e"
[[package]]
name = "universal-hash"
version = "0.5.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "fc1de2c688dc15305988b563c3854064043356019f97a4b46276fe734c4f07ea"
dependencies = [
"crypto-common",
"subtle",
]
[[package]]
name = "unsafe-libyaml"
version = "0.2.11"
+1 -1
View File
@@ -55,7 +55,7 @@ tracing-subscriber = "0.3"
axum = { version = "0.7", features = ["ws"] }
axum-server = { version = "0.7", features = ["tls-rustls"] }
tower = "0.5"
tower-http = { version = "0.6", features = ["limit", "cors"] }
tower-http = { version = "0.6", features = ["limit", "cors", "timeout"] }
# HTTP 客户端
reqwest = { version = "0.12", features = ["json", "stream", "gzip", "brotli", "deflate"] }
+3
View File
@@ -59,6 +59,9 @@ pub mod event_emit;
// 网络工具
pub mod network;
// 凭证清理(敏感信息过滤)
pub mod sanitizer;
// 数据层
pub mod content;
pub mod database;
@@ -0,0 +1,285 @@
//! 提示路由器
//!
//! 支持通过消息前缀提示(hint)将请求路由到不同的 Provider 和模型。
//!
//! 提示格式:`[hint] 消息内容`
//! 例如:`[reasoning] 请分析这段代码的复杂度`
//! `[fast] 翻译这句话`
//! `[code] 实现一个排序算法`
use crate::ProviderType;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
/// 提示路由配置
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct HintRouterConfig {
/// 是否启用提示路由
#[serde(default)]
pub enabled: bool,
/// 提示路由规则
#[serde(default)]
pub routes: Vec<HintRouteEntry>,
}
impl Default for HintRouterConfig {
fn default() -> Self {
Self {
enabled: false,
routes: Vec::new(),
}
}
}
/// 单条提示路由配置
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct HintRouteEntry {
/// 提示关键词(如 "reasoning", "fast", "code")
pub hint: String,
/// 目标 Provider
pub provider: ProviderType,
/// 目标模型
pub model: String,
}
/// 已解析的提示路由
#[derive(Debug, Clone, PartialEq)]
pub struct HintRoute {
/// 提示关键词
pub hint: String,
/// 目标 Provider
pub provider: ProviderType,
/// 目标模型
pub model: String,
}
/// 提示匹配结果
#[derive(Debug, Clone, PartialEq)]
pub struct HintMatch {
/// 匹配到的路由
pub route: HintRoute,
/// 去除提示前缀后的消息内容
pub stripped_message: String,
}
/// 提示路由器
#[derive(Debug, Clone)]
pub struct HintRouter {
enabled: bool,
/// hint 关键词 -> 路由(小写匹配)
routes: HashMap<String, HintRoute>,
}
impl HintRouter {
/// 从配置创建提示路由器
pub fn from_config(config: &HintRouterConfig) -> Self {
let mut routes = HashMap::new();
if config.enabled {
for entry in &config.routes {
let key = entry.hint.to_lowercase();
routes.insert(
key,
HintRoute {
hint: entry.hint.clone(),
provider: entry.provider,
model: entry.model.clone(),
},
);
}
}
Self {
enabled: config.enabled,
routes,
}
}
/// 是否启用
pub fn is_enabled(&self) -> bool {
self.enabled
}
/// 获取已注册的路由数量
pub fn route_count(&self) -> usize {
self.routes.len()
}
/// 根据提示关键词查找路由
pub fn route_by_hint(&self, hint: &str) -> Option<&HintRoute> {
if !self.enabled {
return None;
}
self.routes.get(&hint.to_lowercase())
}
/// 从消息中提取提示并匹配路由
///
/// 支持格式:`[hint] 消息内容` 或 `[hint]消息内容`
/// 提示匹配不区分大小写
pub fn match_message(&self, message: &str) -> Option<HintMatch> {
if !self.enabled {
return None;
}
let trimmed = message.trim_start();
if !trimmed.starts_with('[') {
return None;
}
let close_bracket = trimmed.find(']')?;
let hint = trimmed[1..close_bracket].trim();
if hint.is_empty() {
return None;
}
let route = self.routes.get(&hint.to_lowercase())?;
// 提取去除前缀后的消息
let rest = trimmed[close_bracket + 1..].trim_start();
Some(HintMatch {
route: route.clone(),
stripped_message: rest.to_string(),
})
}
}
impl Default for HintRouter {
fn default() -> Self {
Self {
enabled: false,
routes: HashMap::new(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn test_config() -> HintRouterConfig {
HintRouterConfig {
enabled: true,
routes: vec![
HintRouteEntry {
hint: "reasoning".to_string(),
provider: ProviderType::Kiro,
model: "claude-sonnet-4-5-20250514".to_string(),
},
HintRouteEntry {
hint: "fast".to_string(),
provider: ProviderType::Gemini,
model: "gemini-2.0-flash".to_string(),
},
HintRouteEntry {
hint: "code".to_string(),
provider: ProviderType::Kiro,
model: "claude-sonnet-4-5-20250514".to_string(),
},
],
}
}
#[test]
fn test_disabled_router() {
let config = HintRouterConfig::default();
let router = HintRouter::from_config(&config);
assert!(!router.is_enabled());
assert_eq!(router.route_count(), 0);
assert!(router.route_by_hint("reasoning").is_none());
assert!(router.match_message("[reasoning] test").is_none());
}
#[test]
fn test_route_by_hint() {
let router = HintRouter::from_config(&test_config());
assert!(router.is_enabled());
assert_eq!(router.route_count(), 3);
let route = router.route_by_hint("reasoning").unwrap();
assert_eq!(route.provider, ProviderType::Kiro);
let route = router.route_by_hint("fast").unwrap();
assert_eq!(route.provider, ProviderType::Gemini);
assert!(router.route_by_hint("unknown").is_none());
}
#[test]
fn test_case_insensitive_hint() {
let router = HintRouter::from_config(&test_config());
assert!(router.route_by_hint("Reasoning").is_some());
assert!(router.route_by_hint("FAST").is_some());
assert!(router.route_by_hint("Code").is_some());
}
#[test]
fn test_match_message_basic() {
let router = HintRouter::from_config(&test_config());
let m = router.match_message("[reasoning] 请分析这段代码").unwrap();
assert_eq!(m.route.hint, "reasoning");
assert_eq!(m.route.provider, ProviderType::Kiro);
assert_eq!(m.stripped_message, "请分析这段代码");
}
#[test]
fn test_match_message_no_space() {
let router = HintRouter::from_config(&test_config());
let m = router.match_message("[fast]翻译这句话").unwrap();
assert_eq!(m.route.hint, "fast");
assert_eq!(m.stripped_message, "翻译这句话");
}
#[test]
fn test_match_message_case_insensitive() {
let router = HintRouter::from_config(&test_config());
let m = router.match_message("[REASONING] test").unwrap();
assert_eq!(m.route.hint, "reasoning");
}
#[test]
fn test_match_message_leading_whitespace() {
let router = HintRouter::from_config(&test_config());
let m = router.match_message(" [fast] hello").unwrap();
assert_eq!(m.route.hint, "fast");
assert_eq!(m.stripped_message, "hello");
}
#[test]
fn test_match_message_no_hint() {
let router = HintRouter::from_config(&test_config());
assert!(router.match_message("普通消息").is_none());
assert!(router.match_message("").is_none());
}
#[test]
fn test_match_message_unknown_hint() {
let router = HintRouter::from_config(&test_config());
assert!(router.match_message("[unknown] test").is_none());
}
#[test]
fn test_match_message_empty_hint() {
let router = HintRouter::from_config(&test_config());
assert!(router.match_message("[] test").is_none());
}
#[test]
fn test_match_message_no_closing_bracket() {
let router = HintRouter::from_config(&test_config());
assert!(router.match_message("[reasoning test").is_none());
}
#[test]
fn test_default_config() {
let config = HintRouterConfig::default();
assert!(!config.enabled);
assert!(config.routes.is_empty());
}
}
+5
View File
@@ -10,13 +10,18 @@
//!
//! 模型映射:
//! - 支持模型别名映射(如 `gpt-4` -> `claude-sonnet-4-5-20250514`)
//!
//! 提示路由:
//! - 支持消息前缀提示路由(如 `[reasoning] 请分析...`)
mod amp_router;
mod hint_router;
mod mapper;
mod provider_router;
mod route_registry;
mod rules;
pub use amp_router::AmpRouter;
pub use hint_router::{HintMatch, HintRoute, HintRouter, HintRouterConfig, HintRouteEntry};
pub use mapper::ModelMapper;
pub use rules::Router;
+238
View File
@@ -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}");
}
}
}
+6
View File
@@ -30,6 +30,12 @@ chrono.workspace = true
# 并发
dashmap.workspace = true
# 加密
chacha20poly1305 = "0.10"
base64.workspace = true
rand.workspace = true
sha2.workspace = true
[dev-dependencies]
proptest.workspace = true
tempfile.workspace = true
@@ -0,0 +1,273 @@
//! ChaCha20-Poly1305 AEAD 加密模块
//!
//! 提供凭证加密/解密功能:
//! - ChaCha20-Poly1305 认证加密(防篡改)
//! - 随机 nonce(每次加密生成新的 12 字节 nonce)
//! - 密钥派生(SHA-256)
//! - 格式:enc2:base64(nonce || ciphertext || tag)
use base64::{engine::general_purpose::STANDARD as BASE64, Engine};
use chacha20poly1305::{
aead::{Aead, KeyInit, OsRng},
ChaCha20Poly1305, Nonce,
};
use sha2::{Digest, Sha256};
/// 加密前缀标识
const ENCRYPTED_PREFIX: &str = "enc2:";
/// Nonce 长度(12 字节)
const NONCE_SIZE: usize = 12;
/// 加密器
pub struct Encryptor {
cipher: ChaCha20Poly1305,
}
impl Encryptor {
/// 从密码/密钥创建加密器
///
/// 使用 SHA-256 将任意长度的密钥派生为 256-bit 密钥
pub fn new(key: &str) -> Self {
let derived_key = Self::derive_key(key);
let cipher = ChaCha20Poly1305::new(&derived_key.into());
Self { cipher }
}
/// 从原始 32 字节密钥创建加密器
pub fn from_raw_key(key: &[u8; 32]) -> Self {
let cipher = ChaCha20Poly1305::new(key.into());
Self { cipher }
}
/// 使用 SHA-256 派生 256-bit 密钥
fn derive_key(password: &str) -> [u8; 32] {
let mut hasher = Sha256::new();
hasher.update(password.as_bytes());
let result = hasher.finalize();
let mut key = [0u8; 32];
key.copy_from_slice(&result);
key
}
/// 加密明文
///
/// 返回格式:enc2:base64(nonce || ciphertext)
pub fn encrypt(&self, plaintext: &str) -> Result<String, EncryptionError> {
use chacha20poly1305::aead::AeadCore;
// 生成随机 nonce
let nonce = ChaCha20Poly1305::generate_nonce(&mut OsRng);
// 加密
let ciphertext = self
.cipher
.encrypt(&nonce, plaintext.as_bytes())
.map_err(|_| EncryptionError::EncryptionFailed)?;
// 组合 nonce + ciphertext
let mut combined = Vec::with_capacity(NONCE_SIZE + ciphertext.len());
combined.extend_from_slice(&nonce);
combined.extend_from_slice(&ciphertext);
// Base64 编码并添加前缀
Ok(format!("{}{}", ENCRYPTED_PREFIX, BASE64.encode(&combined)))
}
/// 解密密文
///
/// 输入格式:enc2:base64(nonce || ciphertext)
pub fn decrypt(&self, encrypted: &str) -> Result<String, EncryptionError> {
// 检查前缀
let encoded = encrypted
.strip_prefix(ENCRYPTED_PREFIX)
.ok_or(EncryptionError::InvalidFormat)?;
// Base64 解码
let combined = BASE64
.decode(encoded)
.map_err(|_| EncryptionError::InvalidBase64)?;
// 分离 nonce 和 ciphertext
if combined.len() < NONCE_SIZE {
return Err(EncryptionError::InvalidFormat);
}
let (nonce_bytes, ciphertext) = combined.split_at(NONCE_SIZE);
let nonce = Nonce::from_slice(nonce_bytes);
// 解密
let plaintext = self
.cipher
.decrypt(nonce, ciphertext)
.map_err(|_| EncryptionError::DecryptionFailed)?;
String::from_utf8(plaintext).map_err(|_| EncryptionError::InvalidUtf8)
}
/// 检查文本是否已加密
pub fn is_encrypted(text: &str) -> bool {
text.starts_with(ENCRYPTED_PREFIX)
}
/// 加密(如果尚未加密)
pub fn encrypt_if_needed(&self, text: &str) -> Result<String, EncryptionError> {
if Self::is_encrypted(text) {
Ok(text.to_string())
} else {
self.encrypt(text)
}
}
/// 解密(如果已加密)
pub fn decrypt_if_needed(&self, text: &str) -> Result<String, EncryptionError> {
if Self::is_encrypted(text) {
self.decrypt(text)
} else {
Ok(text.to_string())
}
}
}
/// 加密错误
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum EncryptionError {
/// 加密失败
EncryptionFailed,
/// 解密失败(密钥错误或数据被篡改)
DecryptionFailed,
/// 无效的格式(缺少 enc2: 前缀)
InvalidFormat,
/// 无效的 Base64 编码
InvalidBase64,
/// 无效的 UTF-8
InvalidUtf8,
}
impl std::fmt::Display for EncryptionError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::EncryptionFailed => write!(f, "加密失败"),
Self::DecryptionFailed => write!(f, "解密失败:密钥错误或数据被篡改"),
Self::InvalidFormat => write!(f, "无效的加密格式"),
Self::InvalidBase64 => write!(f, "无效的 Base64 编码"),
Self::InvalidUtf8 => write!(f, "无效的 UTF-8 编码"),
}
}
}
impl std::error::Error for EncryptionError {}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_encrypt_decrypt_roundtrip() {
let enc = Encryptor::new("test-password");
let plaintext = "sk-abc123-secret-api-key";
let encrypted = enc.encrypt(plaintext).unwrap();
assert!(encrypted.starts_with(ENCRYPTED_PREFIX));
let decrypted = enc.decrypt(&encrypted).unwrap();
assert_eq!(decrypted, plaintext);
}
#[test]
fn test_different_nonces() {
let enc = Encryptor::new("test-password");
let plaintext = "same-plaintext";
let encrypted1 = enc.encrypt(plaintext).unwrap();
let encrypted2 = enc.encrypt(plaintext).unwrap();
assert_ne!(encrypted1, encrypted2);
// 两者都能正确解密
assert_eq!(enc.decrypt(&encrypted1).unwrap(), plaintext);
assert_eq!(enc.decrypt(&encrypted2).unwrap(), plaintext);
}
#[test]
fn test_wrong_key_fails() {
let enc1 = Encryptor::new("correct-password");
let enc2 = Encryptor::new("wrong-password");
let encrypted = enc1.encrypt("secret").unwrap();
assert_eq!(enc2.decrypt(&encrypted), Err(EncryptionError::DecryptionFailed));
}
#[test]
fn test_is_encrypted() {
assert!(Encryptor::is_encrypted("enc2:abc123"));
assert!(!Encryptor::is_encrypted("plain-text"));
assert!(!Encryptor::is_encrypted("enc1:old-format"));
assert!(!Encryptor::is_encrypted(""));
}
#[test]
fn test_encrypt_if_needed_already_encrypted() {
let enc = Encryptor::new("key");
let already = "enc2:already-encrypted-data";
let result = enc.encrypt_if_needed(already).unwrap();
assert_eq!(result, already);
}
#[test]
fn test_decrypt_if_needed_not_encrypted() {
let enc = Encryptor::new("key");
let plain = "not-encrypted";
let result = enc.decrypt_if_needed(plain).unwrap();
assert_eq!(result, plain);
}
#[test]
fn test_invalid_format() {
let enc = Encryptor::new("key");
assert_eq!(enc.decrypt("no-prefix"), Err(EncryptionError::InvalidFormat));
}
#[test]
fn test_invalid_base64() {
let enc = Encryptor::new("key");
assert_eq!(enc.decrypt("enc2:!!!invalid-base64!!!"), Err(EncryptionError::InvalidBase64));
}
#[test]
fn test_tampered_data() {
let enc = Encryptor::new("key");
let encrypted = enc.encrypt("secret").unwrap();
// 篡改密文中的一个字符
let encoded = encrypted.strip_prefix(ENCRYPTED_PREFIX).unwrap();
let mut bytes = BASE64.decode(encoded).unwrap();
if let Some(last) = bytes.last_mut() {
*last ^= 0xFF;
}
let tampered = format!("{}{}", ENCRYPTED_PREFIX, BASE64.encode(&bytes));
assert_eq!(enc.decrypt(&tampered), Err(EncryptionError::DecryptionFailed));
}
#[test]
fn test_empty_string() {
let enc = Encryptor::new("key");
let encrypted = enc.encrypt("").unwrap();
assert_eq!(enc.decrypt(&encrypted).unwrap(), "");
}
#[test]
fn test_unicode_content() {
let enc = Encryptor::new("密钥");
let plaintext = "你好世界 🌍 こんにちは";
let encrypted = enc.encrypt(plaintext).unwrap();
assert_eq!(enc.decrypt(&encrypted).unwrap(), plaintext);
}
#[test]
fn test_from_raw_key() {
let raw_key: [u8; 32] = [
0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08,
0x09, 0x0a, 0x0b, 0x0c, 0x0d, 0x0e, 0x0f, 0x10,
0x11, 0x12, 0x13, 0x14, 0x15, 0x16, 0x17, 0x18,
0x19, 0x1a, 0x1b, 0x1c, 0x1d, 0x1e, 0x1f, 0x20,
];
let enc = Encryptor::from_raw_key(&raw_key);
let plaintext = "raw-key-test";
let encrypted = enc.encrypt(plaintext).unwrap();
assert_eq!(enc.decrypt(&encrypted).unwrap(), plaintext);
}
}
+1
View File
@@ -9,6 +9,7 @@
//! - `sync` - 凭证与 YAML 配置文件的同步
mod balancer;
pub mod encryption;
mod quota;
mod sync;
@@ -0,0 +1,313 @@
//! 对话历史管理器
//!
//! 提供对话历史修剪策略,防止上下文溢出
use serde::{Deserialize, Serialize};
/// 修剪策略
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum TrimStrategy {
/// 丢弃最旧的消息
DropOldest,
/// 滑动窗口(保留最近 N 条)
SlidingWindow,
}
impl Default for TrimStrategy {
fn default() -> Self {
Self::SlidingWindow
}
}
/// 修剪配置
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TrimConfig {
/// 是否启用
#[serde(default)]
pub enabled: bool,
/// 最大消息数
#[serde(default = "default_max_messages")]
pub max_messages: usize,
/// 是否保留 system 提示
#[serde(default = "default_preserve_system")]
pub preserve_system_prompt: bool,
/// 修剪策略
#[serde(default)]
pub strategy: TrimStrategy,
}
fn default_max_messages() -> usize {
100
}
fn default_preserve_system() -> bool {
true
}
impl Default for TrimConfig {
fn default() -> Self {
Self {
enabled: false,
max_messages: default_max_messages(),
preserve_system_prompt: default_preserve_system(),
strategy: TrimStrategy::default(),
}
}
}
/// 修剪结果
#[derive(Debug)]
pub struct TrimResult {
/// 修剪后的消息
pub messages: Vec<serde_json::Value>,
/// 是否进行了修剪
pub trimmed: bool,
/// 被移除的消息数
pub removed_count: usize,
}
/// 对话修剪器
pub struct ConversationTrimmer {
config: TrimConfig,
}
impl ConversationTrimmer {
pub fn new(config: TrimConfig) -> Self {
Self { config }
}
/// 修剪消息列表
///
/// 兼容 Anthropic 和 OpenAI 消息格式(都使用 "role" 字段)
pub fn trim_messages(&self, messages: Vec<serde_json::Value>) -> TrimResult {
if !self.config.enabled || messages.len() <= self.config.max_messages {
return TrimResult {
messages,
trimmed: false,
removed_count: 0,
};
}
let original_count = messages.len();
// 分离 system 消息和非 system 消息
let (system_msgs, non_system_msgs): (Vec<_>, Vec<_>) =
if self.config.preserve_system_prompt {
messages.into_iter().partition(|msg| {
msg.get("role")
.and_then(|r| r.as_str())
.map(|r| r == "system")
.unwrap_or(false)
})
} else {
(Vec::new(), messages)
};
// 计算非 system 消息的最大数量
let max_non_system = self.config.max_messages.saturating_sub(system_msgs.len());
// 按策略修剪
let trimmed_non_system = match self.config.strategy {
TrimStrategy::DropOldest | TrimStrategy::SlidingWindow => {
let len = non_system_msgs.len();
if len > max_non_system {
non_system_msgs
.into_iter()
.skip(len - max_non_system)
.collect()
} else {
non_system_msgs
}
}
};
// 合并:system 消息在前,非 system 消息在后
let mut result = system_msgs;
result.extend(trimmed_non_system);
let removed_count = original_count - result.len();
TrimResult {
messages: result,
trimmed: removed_count > 0,
removed_count,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn make_msg(role: &str, content: &str) -> serde_json::Value {
json!({ "role": role, "content": content })
}
#[test]
fn test_disabled_no_trim() {
let trimmer = ConversationTrimmer::new(TrimConfig {
enabled: false,
max_messages: 2,
..Default::default()
});
let msgs = vec![
make_msg("user", "1"),
make_msg("assistant", "2"),
make_msg("user", "3"),
];
let result = trimmer.trim_messages(msgs);
assert!(!result.trimmed);
assert_eq!(result.removed_count, 0);
assert_eq!(result.messages.len(), 3);
}
#[test]
fn test_within_limit_no_trim() {
let trimmer = ConversationTrimmer::new(TrimConfig {
enabled: true,
max_messages: 5,
..Default::default()
});
let msgs = vec![make_msg("user", "1"), make_msg("assistant", "2")];
let result = trimmer.trim_messages(msgs);
assert!(!result.trimmed);
assert_eq!(result.messages.len(), 2);
}
#[test]
fn test_trim_preserves_system() {
let trimmer = ConversationTrimmer::new(TrimConfig {
enabled: true,
max_messages: 3,
preserve_system_prompt: true,
strategy: TrimStrategy::SlidingWindow,
});
let msgs = vec![
make_msg("system", "You are helpful"),
make_msg("user", "1"),
make_msg("assistant", "2"),
make_msg("user", "3"),
make_msg("assistant", "4"),
];
let result = trimmer.trim_messages(msgs);
assert!(result.trimmed);
assert_eq!(result.removed_count, 2);
assert_eq!(result.messages.len(), 3);
assert_eq!(result.messages[0]["role"], "system");
assert_eq!(result.messages[1]["content"], "3");
assert_eq!(result.messages[2]["content"], "4");
}
#[test]
fn test_trim_drops_oldest() {
let trimmer = ConversationTrimmer::new(TrimConfig {
enabled: true,
max_messages: 2,
preserve_system_prompt: false,
strategy: TrimStrategy::DropOldest,
});
let msgs = vec![
make_msg("user", "old"),
make_msg("assistant", "old-reply"),
make_msg("user", "new"),
make_msg("assistant", "new-reply"),
];
let result = trimmer.trim_messages(msgs);
assert!(result.trimmed);
assert_eq!(result.removed_count, 2);
assert_eq!(result.messages.len(), 2);
assert_eq!(result.messages[0]["content"], "new");
assert_eq!(result.messages[1]["content"], "new-reply");
}
#[test]
fn test_trim_with_no_system() {
let trimmer = ConversationTrimmer::new(TrimConfig {
enabled: true,
max_messages: 2,
preserve_system_prompt: true,
strategy: TrimStrategy::SlidingWindow,
});
let msgs = vec![
make_msg("user", "1"),
make_msg("assistant", "2"),
make_msg("user", "3"),
];
let result = trimmer.trim_messages(msgs);
assert!(result.trimmed);
assert_eq!(result.removed_count, 1);
assert_eq!(result.messages.len(), 2);
assert_eq!(result.messages[0]["content"], "2");
assert_eq!(result.messages[1]["content"], "3");
}
#[test]
fn test_trim_all_system_messages_preserved() {
let trimmer = ConversationTrimmer::new(TrimConfig {
enabled: true,
max_messages: 3,
preserve_system_prompt: true,
strategy: TrimStrategy::SlidingWindow,
});
let msgs = vec![
make_msg("system", "sys1"),
make_msg("system", "sys2"),
make_msg("user", "1"),
make_msg("assistant", "2"),
make_msg("user", "3"),
];
let result = trimmer.trim_messages(msgs);
assert!(result.trimmed);
assert_eq!(result.messages.len(), 3);
assert_eq!(result.messages[0]["role"], "system");
assert_eq!(result.messages[1]["role"], "system");
assert_eq!(result.messages[2]["content"], "3");
}
#[test]
fn test_openai_format_compatibility() {
let trimmer = ConversationTrimmer::new(TrimConfig {
enabled: true,
max_messages: 2,
preserve_system_prompt: true,
strategy: TrimStrategy::SlidingWindow,
});
// OpenAI 格式:system/user/assistant + content 字符串
let msgs = vec![
json!({ "role": "system", "content": "You are a helpful assistant." }),
json!({ "role": "user", "content": "Hello" }),
json!({ "role": "assistant", "content": "Hi there!" }),
json!({ "role": "user", "content": "How are you?" }),
];
let result = trimmer.trim_messages(msgs);
assert!(result.trimmed);
assert_eq!(result.messages.len(), 2);
assert_eq!(result.messages[0]["role"], "system");
assert_eq!(result.messages[1]["content"], "How are you?");
}
#[test]
fn test_anthropic_format_compatibility() {
let trimmer = ConversationTrimmer::new(TrimConfig {
enabled: true,
max_messages: 3,
preserve_system_prompt: true,
strategy: TrimStrategy::SlidingWindow,
});
// Anthropic 格式:content 可以是数组
let msgs = vec![
json!({ "role": "system", "content": "You are Claude." }),
json!({ "role": "user", "content": [{"type": "text", "text": "msg1"}] }),
json!({ "role": "assistant", "content": [{"type": "text", "text": "reply1"}] }),
json!({ "role": "user", "content": [{"type": "text", "text": "msg2"}] }),
json!({ "role": "assistant", "content": [{"type": "text", "text": "reply2"}] }),
];
let result = trimmer.trim_messages(msgs);
assert!(result.trimmed);
assert_eq!(result.messages.len(), 3);
assert_eq!(result.messages[0]["role"], "system");
assert_eq!(result.messages[1]["role"], "user");
assert_eq!(result.messages[2]["role"], "assistant");
assert_eq!(result.messages[2]["content"][0]["text"], "reply2");
}
}
@@ -0,0 +1,381 @@
//! 对话摘要器
//!
//! 当对话历史过长时,使用 LLM 生成简洁摘要替代旧消息,
//! 保留关键上下文同时减少 token 消耗。
use serde::{Deserialize, Serialize};
/// 摘要配置
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SummaryConfig {
/// 是否启用
#[serde(default)]
pub enabled: bool,
/// 触发摘要的消息数阈值
#[serde(default = "default_threshold")]
pub threshold_messages: usize,
/// 摘要后保留的最近消息数
#[serde(default = "default_keep_recent")]
pub keep_recent_messages: usize,
/// 摘要最大要点数
#[serde(default = "default_max_points")]
pub max_summary_points: usize,
}
fn default_threshold() -> usize {
50
}
fn default_keep_recent() -> usize {
20
}
fn default_max_points() -> usize {
12
}
impl Default for SummaryConfig {
fn default() -> Self {
Self {
enabled: false,
threshold_messages: default_threshold(),
keep_recent_messages: default_keep_recent(),
max_summary_points: default_max_points(),
}
}
}
/// 摘要请求
///
/// 包含需要发送给 LLM 的摘要请求信息
#[derive(Debug, Clone)]
pub struct SummaryRequest {
/// 摘要 prompt(system 消息)
pub system_prompt: String,
/// 需要摘要的消息(作为 user 消息发送)
pub messages_to_summarize: String,
}
/// 摘要结果
#[derive(Debug, Clone)]
pub struct SummaryResult {
/// 摘要后的消息列表(摘要 system 消息 + 保留的最近消息)
pub messages: Vec<serde_json::Value>,
/// 是否进行了摘要
pub summarized: bool,
/// 被摘要的消息数
pub summarized_count: usize,
}
/// 对话摘要器
pub struct ConversationSummarizer {
config: SummaryConfig,
}
impl ConversationSummarizer {
pub fn new(config: SummaryConfig) -> Self {
Self { config }
}
/// 判断是否需要摘要
pub fn should_summarize(&self, message_count: usize) -> bool {
self.config.enabled && message_count > self.config.threshold_messages
}
/// 构建摘要请求
///
/// 将需要摘要的旧消息格式化为 LLM 请求
pub fn build_summary_request(
&self,
messages: &[serde_json::Value],
) -> Option<SummaryRequest> {
if !self.should_summarize(messages.len()) {
return None;
}
let non_system_msgs: Vec<_> = messages
.iter()
.filter(|msg| {
msg.get("role")
.and_then(|r| r.as_str())
.map(|r| r != "system")
.unwrap_or(true)
})
.collect();
let keep = self.config.keep_recent_messages.min(non_system_msgs.len());
let to_summarize = non_system_msgs.len().saturating_sub(keep);
if to_summarize == 0 {
return None;
}
let msgs_text: Vec<String> = non_system_msgs[..to_summarize]
.iter()
.map(|msg| {
let role = msg
.get("role")
.and_then(|r| r.as_str())
.unwrap_or("unknown");
let content = extract_content_text(msg);
format!("[{role}]: {content}")
})
.collect();
let messages_text = msgs_text.join("\n\n");
let system_prompt = format!(
"你是一个对话摘要助手。请将以下对话历史总结为最多 {} 个关键要点。\n\
要求:\n\
- 保留重要的决策、结论和上下文\n\
- 保留关键的技术细节和代码引用\n\
- 使用简洁的要点格式\n\
- 按时间顺序组织\n\
- 不要遗漏用户的关键需求",
self.config.max_summary_points
);
Some(SummaryRequest {
system_prompt,
messages_to_summarize: messages_text,
})
}
/// 将摘要文本组装为最终消息列表
///
/// 结构:原始 system 消息 + 摘要 system 消息 + 保留的最近消息
pub fn assemble_with_summary(
&self,
original_messages: &[serde_json::Value],
summary_text: &str,
) -> SummaryResult {
let (system_msgs, non_system_msgs): (Vec<_>, Vec<_>) =
original_messages.iter().partition(|msg| {
msg.get("role")
.and_then(|r| r.as_str())
.map(|r| r == "system")
.unwrap_or(false)
});
let keep = self.config.keep_recent_messages.min(non_system_msgs.len());
let summarized_count = non_system_msgs.len().saturating_sub(keep);
let mut result = Vec::new();
// 1. 原始 system 消息
for msg in &system_msgs {
result.push((*msg).clone());
}
// 2. 摘要 system 消息
if summarized_count > 0 && !summary_text.is_empty() {
result.push(serde_json::json!({
"role": "system",
"content": format!(
"[对话摘要 - 以下是之前 {} 条消息的摘要]\n\n{}",
summarized_count, summary_text
)
}));
}
// 3. 保留的最近消息
let start = non_system_msgs.len().saturating_sub(keep);
for msg in &non_system_msgs[start..] {
result.push((*msg).clone());
}
SummaryResult {
messages: result,
summarized: summarized_count > 0,
summarized_count,
}
}
}
/// 从消息中提取文本内容
///
/// 兼容 OpenAI 格式(content 为字符串)和 Anthropic 格式(content 为数组)
fn extract_content_text(msg: &serde_json::Value) -> String {
match msg.get("content") {
Some(serde_json::Value::String(s)) => s.clone(),
Some(serde_json::Value::Array(arr)) => arr
.iter()
.filter_map(|item| {
if item.get("type").and_then(|t| t.as_str()) == Some("text") {
item.get("text").and_then(|t| t.as_str()).map(String::from)
} else {
None
}
})
.collect::<Vec<_>>()
.join("\n"),
_ => String::new(),
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn test_default_config() {
let config = SummaryConfig::default();
assert!(!config.enabled);
assert_eq!(config.threshold_messages, 50);
assert_eq!(config.keep_recent_messages, 20);
assert_eq!(config.max_summary_points, 12);
}
#[test]
fn test_should_summarize_disabled() {
let s = ConversationSummarizer::new(SummaryConfig {
enabled: false,
threshold_messages: 5,
..Default::default()
});
assert!(!s.should_summarize(100));
}
#[test]
fn test_should_summarize_below_threshold() {
let s = ConversationSummarizer::new(SummaryConfig {
enabled: true,
threshold_messages: 50,
..Default::default()
});
assert!(!s.should_summarize(30));
assert!(!s.should_summarize(50)); // 等于阈值不触发
}
#[test]
fn test_should_summarize_above_threshold() {
let s = ConversationSummarizer::new(SummaryConfig {
enabled: true,
threshold_messages: 5,
..Default::default()
});
assert!(s.should_summarize(6));
assert!(s.should_summarize(100));
}
#[test]
fn test_build_summary_request_none_when_disabled() {
let s = ConversationSummarizer::new(SummaryConfig {
enabled: false,
threshold_messages: 2,
keep_recent_messages: 1,
..Default::default()
});
let msgs = vec![
json!({"role": "user", "content": "a"}),
json!({"role": "assistant", "content": "b"}),
json!({"role": "user", "content": "c"}),
];
assert!(s.build_summary_request(&msgs).is_none());
}
#[test]
fn test_build_summary_request_with_messages() {
let s = ConversationSummarizer::new(SummaryConfig {
enabled: true,
threshold_messages: 2,
keep_recent_messages: 1,
max_summary_points: 5,
});
let msgs = vec![
json!({"role": "system", "content": "You are helpful."}),
json!({"role": "user", "content": "Hello"}),
json!({"role": "assistant", "content": "Hi!"}),
json!({"role": "user", "content": "Latest"}),
];
// 总消息数 4 > threshold 2,非 system 消息 3 条,保留 1 条,摘要 2 条
let req = s.build_summary_request(&msgs).unwrap();
assert!(req.system_prompt.contains("5"));
assert!(req.messages_to_summarize.contains("[user]: Hello"));
assert!(req.messages_to_summarize.contains("[assistant]: Hi!"));
assert!(!req.messages_to_summarize.contains("Latest"));
}
#[test]
fn test_assemble_with_summary() {
let s = ConversationSummarizer::new(SummaryConfig {
enabled: true,
threshold_messages: 2,
keep_recent_messages: 1,
..Default::default()
});
let msgs = vec![
json!({"role": "user", "content": "old1"}),
json!({"role": "assistant", "content": "old2"}),
json!({"role": "user", "content": "recent"}),
];
let result = s.assemble_with_summary(&msgs, "摘要内容");
assert!(result.summarized);
assert_eq!(result.summarized_count, 2);
// 摘要 system + 保留的 1 条 = 2
assert_eq!(result.messages.len(), 2);
assert_eq!(result.messages[0]["role"], "system");
assert!(result.messages[0]["content"]
.as_str()
.unwrap()
.contains("摘要内容"));
assert_eq!(result.messages[1]["content"], "recent");
}
#[test]
fn test_assemble_preserves_system_messages() {
let s = ConversationSummarizer::new(SummaryConfig {
enabled: true,
threshold_messages: 2,
keep_recent_messages: 1,
..Default::default()
});
let msgs = vec![
json!({"role": "system", "content": "You are helpful."}),
json!({"role": "user", "content": "old"}),
json!({"role": "assistant", "content": "old reply"}),
json!({"role": "user", "content": "recent"}),
];
let result = s.assemble_with_summary(&msgs, "summary");
// system 原始 + 摘要 system + 保留 1 条 = 3
assert_eq!(result.messages.len(), 3);
assert_eq!(result.messages[0]["content"], "You are helpful.");
assert_eq!(result.messages[1]["role"], "system");
assert!(result.messages[1]["content"]
.as_str()
.unwrap()
.contains("summary"));
assert_eq!(result.messages[2]["content"], "recent");
}
#[test]
fn test_extract_content_text_string() {
let msg = json!({"role": "user", "content": "hello world"});
assert_eq!(extract_content_text(&msg), "hello world");
}
#[test]
fn test_extract_content_text_array() {
// Anthropic 格式
let msg = json!({
"role": "user",
"content": [
{"type": "text", "text": "part1"},
{"type": "image", "source": {}},
{"type": "text", "text": "part2"}
]
});
assert_eq!(extract_content_text(&msg), "part1\npart2");
}
#[test]
fn test_extract_content_text_empty() {
let msg = json!({"role": "user"});
assert_eq!(extract_content_text(&msg), "");
let msg2 = json!({"role": "user", "content": null});
assert_eq!(extract_content_text(&msg2), "");
let msg3 = json!({"role": "user", "content": 42});
assert_eq!(extract_content_text(&msg3), "");
}
}
+2
View File
@@ -6,6 +6,8 @@
//!
//! - `steps` - 管道步骤(认证、注入、路由、插件、Provider、遥测)
pub mod conversation_manager;
pub mod conversation_summarizer;
pub mod processor;
pub mod steps;
+3
View File
@@ -20,6 +20,7 @@ serde.workspace = true
serde_json.workspace = true
tokio.workspace = true
futures.workspace = true
hex.workspace = true
axum.workspace = true
tower.workspace = true
tower-http.workspace = true
@@ -35,6 +36,8 @@ subtle.workspace = true
async-stream.workspace = true
urlencoding.workspace = true
parking_lot.workspace = true
rand.workspace = true
sha2.workspace = true
tokio-util.workspace = true
dirs.workspace = true
+3
View File
@@ -0,0 +1,3 @@
//! 认证模块
pub mod pairing;
+309
View File
@@ -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);
}
}
+7
View File
@@ -1,6 +1,8 @@
//! HTTP API 服务器
pub mod auth;
pub mod client_detector;
pub mod middleware;
use axum::{
extract::{DefaultBodyLimit, Path, State},
@@ -43,6 +45,7 @@ use std::path::PathBuf;
use std::sync::Arc;
use tokio::sync::{oneshot, RwLock};
use tower_http::cors::CorsLayer;
use tower_http::timeout::TimeoutLayer;
/// 记录请求统计到遥测系统
pub fn record_request_telemetry(
@@ -1053,6 +1056,10 @@ async fn run_server(
.merge(batch_api_routes)
.layer(cors_layer)
.layer(DefaultBodyLimit::max(body_limit))
.layer(TimeoutLayer::with_status_code(
StatusCode::REQUEST_TIMEOUT,
std::time::Duration::from_secs(300),
))
.with_state(state);
let addr: std::net::SocketAddr = format!("{host}:{port}")
@@ -0,0 +1,270 @@
//! 幂等性中间件
//!
//! 通过 Idempotency-Key header 防止重复请求
use parking_lot::Mutex;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::time::{Duration, Instant};
/// 幂等性配置
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct IdempotencyConfig {
/// 是否启用
#[serde(default)]
pub enabled: bool,
/// 缓存 TTL(秒)
#[serde(default = "default_ttl_secs")]
pub ttl_secs: u64,
/// Header 名称
#[serde(default = "default_header_name")]
pub header_name: String,
}
fn default_ttl_secs() -> u64 {
86400 // 24 小时
}
fn default_header_name() -> String {
"Idempotency-Key".to_string()
}
impl Default for IdempotencyConfig {
fn default() -> Self {
Self {
enabled: false,
ttl_secs: default_ttl_secs(),
header_name: default_header_name(),
}
}
}
/// 幂等性检查结果
#[derive(Debug, PartialEq)]
pub enum IdempotencyCheck {
/// 新请求,可以处理
New,
/// 正在处理中(返回 409 Conflict)
InProgress,
/// 已完成,有缓存响应
Completed { status: u16, body: String },
}
/// 请求状态
#[derive(Debug, Clone)]
enum RequestState {
/// 正在处理
InProgress { started_at: Instant },
/// 已完成
Completed {
status: u16,
body: String,
completed_at: Instant,
},
}
/// 幂等性存储
pub struct IdempotencyStore {
config: IdempotencyConfig,
entries: Mutex<HashMap<String, RequestState>>,
}
impl IdempotencyStore {
pub fn new(config: IdempotencyConfig) -> Self {
Self {
config,
entries: Mutex::new(HashMap::new()),
}
}
/// 检查幂等性键
pub fn check(&self, key: &str) -> IdempotencyCheck {
if !self.config.enabled {
return IdempotencyCheck::New;
}
let mut entries = self.entries.lock();
let ttl = Duration::from_secs(self.config.ttl_secs);
let now = Instant::now();
match entries.get(key) {
Some(RequestState::InProgress { started_at }) => {
// 如果处理超过 TTL,视为过期
if now.duration_since(*started_at) > ttl {
entries.insert(
key.to_string(),
RequestState::InProgress { started_at: now },
);
IdempotencyCheck::New
} else {
IdempotencyCheck::InProgress
}
}
Some(RequestState::Completed {
status,
body,
completed_at,
}) => {
if now.duration_since(*completed_at) > ttl {
entries.insert(
key.to_string(),
RequestState::InProgress { started_at: now },
);
IdempotencyCheck::New
} else {
IdempotencyCheck::Completed {
status: *status,
body: body.clone(),
}
}
}
None => {
entries.insert(
key.to_string(),
RequestState::InProgress { started_at: now },
);
IdempotencyCheck::New
}
}
}
/// 标记请求完成
pub fn complete(&self, key: &str, status: u16, body: String) {
if !self.config.enabled {
return;
}
let mut entries = self.entries.lock();
entries.insert(
key.to_string(),
RequestState::Completed {
status,
body,
completed_at: Instant::now(),
},
);
}
/// 移除键(请求失败时调用,允许重试)
pub fn remove(&self, key: &str) {
let mut entries = self.entries.lock();
entries.remove(key);
}
/// 清理过期条目
pub fn cleanup(&self) {
let ttl = Duration::from_secs(self.config.ttl_secs);
let now = Instant::now();
let mut entries = self.entries.lock();
entries.retain(|_, state| match state {
RequestState::InProgress { started_at } => now.duration_since(*started_at) < ttl,
RequestState::Completed { completed_at, .. } => now.duration_since(*completed_at) < ttl,
});
}
/// 获取当前条目数
pub fn len(&self) -> usize {
self.entries.lock().len()
}
pub fn is_empty(&self) -> bool {
self.entries.lock().is_empty()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::thread;
fn enabled_config(ttl_secs: u64) -> IdempotencyConfig {
IdempotencyConfig {
enabled: true,
ttl_secs,
header_name: "Idempotency-Key".to_string(),
}
}
#[test]
fn test_disabled_always_new() {
let store = IdempotencyStore::new(IdempotencyConfig::default());
assert_eq!(store.check("key1"), IdempotencyCheck::New);
assert_eq!(store.check("key1"), IdempotencyCheck::New);
assert!(store.is_empty());
}
#[test]
fn test_new_request() {
let store = IdempotencyStore::new(enabled_config(60));
assert_eq!(store.check("key1"), IdempotencyCheck::New);
assert_eq!(store.len(), 1);
}
#[test]
fn test_in_progress_request() {
let store = IdempotencyStore::new(enabled_config(60));
assert_eq!(store.check("key1"), IdempotencyCheck::New);
// 同一 key 再次检查应返回 InProgress
assert_eq!(store.check("key1"), IdempotencyCheck::InProgress);
}
#[test]
fn test_completed_request() {
let store = IdempotencyStore::new(enabled_config(60));
assert_eq!(store.check("key1"), IdempotencyCheck::New);
store.complete("key1", 200, "ok".to_string());
assert_eq!(
store.check("key1"),
IdempotencyCheck::Completed {
status: 200,
body: "ok".to_string(),
}
);
}
#[test]
fn test_expired_entry() {
let store = IdempotencyStore::new(enabled_config(1)); // 1 秒 TTL
assert_eq!(store.check("key1"), IdempotencyCheck::New);
store.complete("key1", 200, "ok".to_string());
// 等待过期
thread::sleep(Duration::from_millis(1100));
// 过期后应视为新请求
assert_eq!(store.check("key1"), IdempotencyCheck::New);
}
#[test]
fn test_cleanup() {
let store = IdempotencyStore::new(enabled_config(1));
assert_eq!(store.check("key1"), IdempotencyCheck::New);
assert_eq!(store.check("key2"), IdempotencyCheck::New);
store.complete("key1", 200, "ok".to_string());
thread::sleep(Duration::from_millis(1100));
store.cleanup();
assert!(store.is_empty(), "清理后应无过期条目");
}
#[test]
fn test_remove_allows_retry() {
let store = IdempotencyStore::new(enabled_config(60));
assert_eq!(store.check("key1"), IdempotencyCheck::New);
assert_eq!(store.check("key1"), IdempotencyCheck::InProgress);
// 移除后应可重试
store.remove("key1");
assert_eq!(store.check("key1"), IdempotencyCheck::New);
}
#[test]
fn test_default_config() {
let config = IdempotencyConfig::default();
assert!(!config.enabled);
assert_eq!(config.ttl_secs, 86400);
assert_eq!(config.header_name, "Idempotency-Key");
}
}
@@ -0,0 +1,5 @@
//! 服务器中间件模块
pub mod idempotency;
pub mod rate_limit;
pub mod security;
@@ -0,0 +1,238 @@
//! 滑动窗口速率限制中间件
//!
//! 基于客户端 IP 的请求速率限制,防止 API 滥用
use parking_lot::Mutex;
use serde::{Deserialize, Serialize};
use std::{
collections::HashMap,
time::{Duration, Instant},
};
/// 速率限制配置
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RateLimitConfig {
/// 是否启用
#[serde(default = "default_enabled")]
pub enabled: bool,
/// 窗口内最大请求数
#[serde(default = "default_requests_per_minute")]
pub requests_per_minute: u32,
/// 窗口大小(秒)
#[serde(default = "default_window_secs")]
pub window_secs: u64,
}
fn default_enabled() -> bool {
false
}
fn default_requests_per_minute() -> u32 {
60
}
fn default_window_secs() -> u64 {
60
}
impl Default for RateLimitConfig {
fn default() -> Self {
Self {
enabled: false,
requests_per_minute: 60,
window_secs: 60,
}
}
}
/// 滑动窗口速率限制器
pub struct SlidingWindowRateLimiter {
config: RateLimitConfig,
/// 客户端 IP -> 请求时间戳列表
requests: Mutex<HashMap<String, Vec<Instant>>>,
}
impl SlidingWindowRateLimiter {
pub fn new(config: RateLimitConfig) -> Self {
Self {
config,
requests: Mutex::new(HashMap::new()),
}
}
/// 检查是否允许请求
pub fn check_rate_limit(&self, client_id: &str) -> RateLimitResult {
if !self.config.enabled {
return RateLimitResult::Allowed;
}
let now = Instant::now();
let window = Duration::from_secs(self.config.window_secs);
let mut requests = self.requests.lock();
let timestamps = requests.entry(client_id.to_string()).or_default();
// 清理窗口外的请求
timestamps.retain(|t| now.duration_since(*t) < window);
if timestamps.len() >= self.config.requests_per_minute as usize {
// 计算最早请求到窗口结束的剩余时间
let oldest = timestamps.first().copied();
let retry_after = oldest
.map(|t| window.saturating_sub(now.duration_since(t)))
.unwrap_or(window);
RateLimitResult::Limited { retry_after }
} else {
timestamps.push(now);
RateLimitResult::Allowed
}
}
/// 清理过期条目(应定期调用)
pub fn cleanup(&self) {
let now = Instant::now();
let window = Duration::from_secs(self.config.window_secs);
let mut requests = self.requests.lock();
requests.retain(|_, timestamps| {
timestamps.retain(|t| now.duration_since(*t) < window);
!timestamps.is_empty()
});
}
}
/// 速率限制检查结果
#[derive(Debug)]
pub enum RateLimitResult {
/// 允许
Allowed,
/// 被限制
Limited {
/// 建议重试等待时间
retry_after: Duration,
},
}
#[cfg(test)]
mod tests {
use super::*;
use std::thread;
#[test]
fn test_disabled_allows_all() {
let limiter = SlidingWindowRateLimiter::new(RateLimitConfig {
enabled: false,
requests_per_minute: 1,
window_secs: 60,
});
// 即使超过限制,禁用时也应全部允许
for _ in 0..100 {
assert!(matches!(
limiter.check_rate_limit("client1"),
RateLimitResult::Allowed
));
}
}
#[test]
fn test_within_limit() {
let limiter = SlidingWindowRateLimiter::new(RateLimitConfig {
enabled: true,
requests_per_minute: 5,
window_secs: 60,
});
for _ in 0..5 {
assert!(matches!(
limiter.check_rate_limit("client1"),
RateLimitResult::Allowed
));
}
}
#[test]
fn test_exceeds_limit() {
let limiter = SlidingWindowRateLimiter::new(RateLimitConfig {
enabled: true,
requests_per_minute: 3,
window_secs: 60,
});
// 前 3 个请求应允许
for _ in 0..3 {
assert!(matches!(
limiter.check_rate_limit("client1"),
RateLimitResult::Allowed
));
}
// 第 4 个应被限制
match limiter.check_rate_limit("client1") {
RateLimitResult::Limited { retry_after } => {
assert!(retry_after.as_secs() <= 60);
}
RateLimitResult::Allowed => panic!("应该被限制"),
}
}
#[test]
fn test_window_expiry() {
let limiter = SlidingWindowRateLimiter::new(RateLimitConfig {
enabled: true,
requests_per_minute: 2,
window_secs: 1, // 1 秒窗口,方便测试过期
});
// 用完配额
assert!(matches!(
limiter.check_rate_limit("client1"),
RateLimitResult::Allowed
));
assert!(matches!(
limiter.check_rate_limit("client1"),
RateLimitResult::Allowed
));
assert!(matches!(
limiter.check_rate_limit("client1"),
RateLimitResult::Limited { .. }
));
// 等待窗口过期
thread::sleep(Duration::from_millis(1100));
// 窗口过期后应重新允许
assert!(matches!(
limiter.check_rate_limit("client1"),
RateLimitResult::Allowed
));
}
#[test]
fn test_cleanup() {
let limiter = SlidingWindowRateLimiter::new(RateLimitConfig {
enabled: true,
requests_per_minute: 10,
window_secs: 1,
});
// 添加一些请求
limiter.check_rate_limit("client1");
limiter.check_rate_limit("client2");
// 等待窗口过期
thread::sleep(Duration::from_millis(1100));
// 清理应移除过期条目
limiter.cleanup();
let requests = limiter.requests.lock();
assert!(requests.is_empty(), "清理后应无过期条目");
}
#[test]
fn test_default_config() {
let config = RateLimitConfig::default();
assert!(!config.enabled);
assert_eq!(config.requests_per_minute, 60);
assert_eq!(config.window_secs, 60);
}
}
@@ -0,0 +1,62 @@
//! 安全中间件
//!
//! 提供请求体大小限制和请求超时控制
use serde::{Deserialize, Serialize};
use std::time::Duration;
/// 安全中间件配置
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SecurityMiddlewareConfig {
/// 最大请求体大小(字节),默认 10MB
#[serde(default = "default_max_body_size")]
pub max_body_size: usize,
/// 请求超时(秒),默认 300 秒(LLM 请求可能很长)
#[serde(default = "default_request_timeout_secs")]
pub request_timeout_secs: u64,
}
fn default_max_body_size() -> usize {
10 * 1024 * 1024 // 10MB
}
fn default_request_timeout_secs() -> u64 {
300 // 5 分钟
}
impl Default for SecurityMiddlewareConfig {
fn default() -> Self {
Self {
max_body_size: default_max_body_size(),
request_timeout_secs: default_request_timeout_secs(),
}
}
}
impl SecurityMiddlewareConfig {
/// 获取请求超时 Duration
pub fn request_timeout(&self) -> Duration {
Duration::from_secs(self.request_timeout_secs)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_default_config() {
let config = SecurityMiddlewareConfig::default();
assert_eq!(config.max_body_size, 10 * 1024 * 1024);
assert_eq!(config.request_timeout_secs, 300);
}
#[test]
fn test_request_timeout() {
let config = SecurityMiddlewareConfig {
max_body_size: 1024,
request_timeout_secs: 60,
};
assert_eq!(config.request_timeout(), Duration::from_secs(60));
}
}