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
+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));
}
}