diff --git a/IMPLEMENTATION_PLAN.md b/IMPLEMENTATION_PLAN.md new file mode 100644 index 000000000..782c34b04 --- /dev/null +++ b/IMPLEMENTATION_PLAN.md @@ -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` | diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 7ada76db7..8f2907f29 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -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" diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 8aa3d7251..c8ce4c505 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -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"] } diff --git a/src-tauri/crates/core/src/lib.rs b/src-tauri/crates/core/src/lib.rs index be89a26f5..382fb490e 100644 --- a/src-tauri/crates/core/src/lib.rs +++ b/src-tauri/crates/core/src/lib.rs @@ -59,6 +59,9 @@ pub mod event_emit; // 网络工具 pub mod network; +// 凭证清理(敏感信息过滤) +pub mod sanitizer; + // 数据层 pub mod content; pub mod database; diff --git a/src-tauri/crates/core/src/router/hint_router.rs b/src-tauri/crates/core/src/router/hint_router.rs new file mode 100644 index 000000000..021547d5f --- /dev/null +++ b/src-tauri/crates/core/src/router/hint_router.rs @@ -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, +} + +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, +} + +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 { + 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()); + } +} diff --git a/src-tauri/crates/core/src/router/mod.rs b/src-tauri/crates/core/src/router/mod.rs index 8953074cc..c67a24ca2 100644 --- a/src-tauri/crates/core/src/router/mod.rs +++ b/src-tauri/crates/core/src/router/mod.rs @@ -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; diff --git a/src-tauri/crates/core/src/sanitizer.rs b/src-tauri/crates/core/src/sanitizer.rs new file mode 100644 index 000000000..defb31632 --- /dev/null +++ b/src-tauri/crates/core/src/sanitizer.rs @@ -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, +} + +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, +} + +/// 内置的敏感信息正则模式 +fn builtin_patterns() -> &'static [Regex] { + static PATTERNS: OnceLock> = 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}"); + } + } +} diff --git a/src-tauri/crates/credential/Cargo.toml b/src-tauri/crates/credential/Cargo.toml index 8ea47bde8..767d6c853 100644 --- a/src-tauri/crates/credential/Cargo.toml +++ b/src-tauri/crates/credential/Cargo.toml @@ -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 diff --git a/src-tauri/crates/credential/src/encryption.rs b/src-tauri/crates/credential/src/encryption.rs new file mode 100644 index 000000000..99ba46d03 --- /dev/null +++ b/src-tauri/crates/credential/src/encryption.rs @@ -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 { + 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 { + // 检查前缀 + 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 { + if Self::is_encrypted(text) { + Ok(text.to_string()) + } else { + self.encrypt(text) + } + } + + /// 解密(如果已加密) + pub fn decrypt_if_needed(&self, text: &str) -> Result { + 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); + } +} diff --git a/src-tauri/crates/credential/src/lib.rs b/src-tauri/crates/credential/src/lib.rs index cd79ff1cf..aa59782c0 100644 --- a/src-tauri/crates/credential/src/lib.rs +++ b/src-tauri/crates/credential/src/lib.rs @@ -9,6 +9,7 @@ //! - `sync` - 凭证与 YAML 配置文件的同步 mod balancer; +pub mod encryption; mod quota; mod sync; diff --git a/src-tauri/crates/processor/src/conversation_manager.rs b/src-tauri/crates/processor/src/conversation_manager.rs new file mode 100644 index 000000000..d7c45707f --- /dev/null +++ b/src-tauri/crates/processor/src/conversation_manager.rs @@ -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, + /// 是否进行了修剪 + 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) -> 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"); + } +} \ No newline at end of file diff --git a/src-tauri/crates/processor/src/conversation_summarizer.rs b/src-tauri/crates/processor/src/conversation_summarizer.rs new file mode 100644 index 000000000..46c04edea --- /dev/null +++ b/src-tauri/crates/processor/src/conversation_summarizer.rs @@ -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, + /// 是否进行了摘要 + 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 { + 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 = 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::>() + .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), ""); + } +} diff --git a/src-tauri/crates/processor/src/lib.rs b/src-tauri/crates/processor/src/lib.rs index abcbbf9c1..d668afd96 100644 --- a/src-tauri/crates/processor/src/lib.rs +++ b/src-tauri/crates/processor/src/lib.rs @@ -6,6 +6,8 @@ //! //! - `steps` - 管道步骤(认证、注入、路由、插件、Provider、遥测) +pub mod conversation_manager; +pub mod conversation_summarizer; pub mod processor; pub mod steps; diff --git a/src-tauri/crates/server/Cargo.toml b/src-tauri/crates/server/Cargo.toml index 1364f226d..6c1df6c6f 100644 --- a/src-tauri/crates/server/Cargo.toml +++ b/src-tauri/crates/server/Cargo.toml @@ -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 diff --git a/src-tauri/crates/server/src/auth/mod.rs b/src-tauri/crates/server/src/auth/mod.rs new file mode 100644 index 000000000..e0475fead --- /dev/null +++ b/src-tauri/crates/server/src/auth/mod.rs @@ -0,0 +1,3 @@ +//! 认证模块 + +pub mod pairing; diff --git a/src-tauri/crates/server/src/auth/pairing.rs b/src-tauri/crates/server/src/auth/pairing.rs new file mode 100644 index 000000000..a85bf7747 --- /dev/null +++ b/src-tauri/crates/server/src/auth/pairing.rs @@ -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, + blocked_until: Option, +} + +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>, + /// 已配对的 token(存储 SHA-256 哈希) + paired_tokens: Mutex>, + /// 失败尝试追踪 + failed_attempts: Mutex, +} + +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 = (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); + } +} diff --git a/src-tauri/crates/server/src/lib.rs b/src-tauri/crates/server/src/lib.rs index 1f9c0ce27..d8a9560ef 100644 --- a/src-tauri/crates/server/src/lib.rs +++ b/src-tauri/crates/server/src/lib.rs @@ -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}") diff --git a/src-tauri/crates/server/src/middleware/idempotency.rs b/src-tauri/crates/server/src/middleware/idempotency.rs new file mode 100644 index 000000000..31a10facb --- /dev/null +++ b/src-tauri/crates/server/src/middleware/idempotency.rs @@ -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>, +} + +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"); + } +} diff --git a/src-tauri/crates/server/src/middleware/mod.rs b/src-tauri/crates/server/src/middleware/mod.rs new file mode 100644 index 000000000..db4384fc6 --- /dev/null +++ b/src-tauri/crates/server/src/middleware/mod.rs @@ -0,0 +1,5 @@ +//! 服务器中间件模块 + +pub mod idempotency; +pub mod rate_limit; +pub mod security; diff --git a/src-tauri/crates/server/src/middleware/rate_limit.rs b/src-tauri/crates/server/src/middleware/rate_limit.rs new file mode 100644 index 000000000..c50f1bca2 --- /dev/null +++ b/src-tauri/crates/server/src/middleware/rate_limit.rs @@ -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>>, +} + +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); + } +} diff --git a/src-tauri/crates/server/src/middleware/security.rs b/src-tauri/crates/server/src/middleware/security.rs new file mode 100644 index 000000000..98b0ff412 --- /dev/null +++ b/src-tauri/crates/server/src/middleware/security.rs @@ -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)); + } +}