diff --git a/IMPLEMENTATION_PLAN.md b/IMPLEMENTATION_PLAN.md deleted file mode 100644 index d7f5bbc67..000000000 --- a/IMPLEMENTATION_PLAN.md +++ /dev/null @@ -1,105 +0,0 @@ -# 图文海报功能实现计划 - -## 当前进度 - -### Phase 1: 品牌人设系统扩展 ✅ 完成 -**Goal**: 扩展现有人设系统,支持海报设计专用字段(配色、字体、品牌调性) -**Status**: Complete - -#### 已完成任务 -- [x] 1.1 扩展数据模型 (`project_model.rs`) - - 新增 BrandPersonality, DesignStyle 枚举 - - 新增 ColorScheme, Typography, LogoPlacement, ImageStyle, IconStyle 结构体 - - 新增 BrandTone, DesignConfig, VisualConfig 结构体 - - 新增 BrandPersonaExtension, BrandPersona, BrandPersonaTemplate 结构体 - - 新增相关请求类型 - -- [x] 1.2 扩展数据库 Schema (`schema.rs`) - - 新增 `brand_persona_extensions` 表 - - 包含 persona_id, brand_tone_json, design_json, visual_json 字段 - -- [x] 1.3 创建 BrandPersona DAO (`brand_persona_dao.rs`) - - 实现 create, get, update, delete 方法 - - 实现 get_brand_persona 获取完整品牌人设 - - 实现 list_templates 获取预设模板 - -- [x] 1.4 扩展 PersonaService (`persona_service.rs`) - - 新增 get_brand_persona, get_brand_extension 方法 - - 新增 save_brand_extension, update_brand_extension 方法 - - 新增 delete_brand_extension, list_brand_persona_templates 方法 - -- [x] 1.5 扩展 Tauri 命令 (`persona_cmd.rs`) - - 新增 get_brand_persona, get_brand_extension 命令 - - 新增 save_brand_extension, update_brand_extension 命令 - - 新增 delete_brand_extension, list_brand_persona_templates 命令 - - 在 runner.rs 中注册新命令 - -- [x] 1.6 新增前端类型 (`brand-persona.ts`) - - 定义所有品牌人设相关的 TypeScript 类型 - - 包含预设配色方案、字体列表、默认值等常量 - -- [x] 1.7 新增 useBrandPersona Hook (`useBrandPersona.ts`) - - 实现品牌人设的 CRUD 操作 - - 支持模板应用功能 - -- [x] 1.8 新增 BrandPersonaDialog 组件 (`BrandPersonaDialog.tsx`) - - 分步骤创建品牌人设(品牌调性 → 配色方案 → 字体设置 → 预览确认) - - 支持模板快速应用 - - 支持预设配色方案选择 - - 实时预览效果 - -#### 验证标准 -- [x] 能够创建包含配色方案的品牌人设 -- [x] 品牌人设能够正确保存和加载 -- [x] 在项目详情页能够管理品牌人设 - ---- - -### Phase 2: 素材库扩展 -**Goal**: 扩展素材库支持 icon, color, layout 类型 -**Status**: Not Started - ---- - -### Phase 3: 海报 Agent 系统 -**Goal**: 实现 6 个专用 Agent,支持对话式海报设计 -**Status**: Not Started - ---- - -### Phase 4: 工作流系统 -**Goal**: 实现 6 步引导工作流 -**Status**: Not Started - ---- - -### Phase 5: 多平台导出 -**Goal**: 实现多平台尺寸适配和导出 -**Status**: Not Started - ---- - -## 新增文件清单 - -### 后端 (Rust) -- `src-tauri/src/database/dao/brand_persona_dao.rs` - 品牌人设 DAO - -### 前端 (TypeScript/React) -- `src/types/brand-persona.ts` - 品牌人设类型定义 -- `src/hooks/useBrandPersona.ts` - 品牌人设 Hook -- `src/components/projects/dialogs/BrandPersonaDialog.tsx` - 品牌人设对话框 - -## 修改文件清单 - -### 后端 (Rust) -- `src-tauri/src/models/project_model.rs` - 新增品牌人设数据模型 -- `src-tauri/src/database/schema.rs` - 新增品牌人设扩展表 -- `src-tauri/src/database/dao/mod.rs` - 导出新 DAO -- `src-tauri/src/services/persona_service.rs` - 新增品牌人设服务方法 -- `src-tauri/src/commands/persona_cmd.rs` - 新增品牌人设命令 -- `src-tauri/src/app/runner.rs` - 注册新命令 - -### 前端 (TypeScript/React) -- `src/types/index.ts` - 导出新类型 -- `src/hooks/index.ts` - 导出新 Hook -- `src/components/projects/dialogs/index.ts` - 导出新组件 diff --git a/docs/aiprompts/flow-monitor.md b/docs/aiprompts/flow-monitor.md deleted file mode 100644 index 08a136cc9..000000000 --- a/docs/aiprompts/flow-monitor.md +++ /dev/null @@ -1,261 +0,0 @@ -# 流量监控 - -## 概述 - -流量监控模块拦截和记录所有 LLM API 请求,提供 Token 统计、历史查询和分析功能。 - -## 目录结构 - -``` -src-tauri/src/flow_monitor/ -├── mod.rs # 模块入口 -├── interceptor.rs # 请求拦截器 -├── storage.rs # 存储层 -├── query.rs # 查询接口 -└── stats.rs # 统计分析 -``` - -## 架构 - -``` -┌─────────────────────────────────────────────────────────────────┐ -│ HTTP 请求 │ -└─────────────────────────────────────────────────────────────────┘ - │ - ▼ -┌─────────────────────────────────────────────────────────────────┐ -│ Flow Interceptor │ -│ ┌─────────────────────────────────────────────────────────────┐│ -│ │ 请求拦截 ││ -│ │ - 请求 ID 生成 ││ -│ │ - 请求体捕获 ││ -│ │ - 时间戳记录 ││ -│ └─────────────────────────────────────────────────────────────┘│ -└─────────────────────────────────────────────────────────────────┘ - │ - ▼ -┌─────────────────────────────────────────────────────────────────┐ -│ Provider 处理 │ -└─────────────────────────────────────────────────────────────────┘ - │ - ▼ -┌─────────────────────────────────────────────────────────────────┐ -│ Flow Interceptor │ -│ ┌─────────────────────────────────────────────────────────────┐│ -│ │ 响应拦截 ││ -│ │ - 响应体捕获 ││ -│ │ - Token 计数 ││ -│ │ - 延迟计算 ││ -│ └─────────────────────────────────────────────────────────────┘│ -└─────────────────────────────────────────────────────────────────┘ - │ - ▼ -┌─────────────────────────────────────────────────────────────────┐ -│ Flow Storage │ -│ ┌─────────────┐ ┌─────────────┐ ┌─────────────────────────┐ │ -│ │ SQLite │ │ 内存缓存 │ │ 事件发送 │ │ -│ │ 持久化 │ │ (最近 N 条) │ │ (前端通知) │ │ -│ └─────────────┘ └─────────────┘ └─────────────────────────┘ │ -└─────────────────────────────────────────────────────────────────┘ -``` - -## 数据模型 - -### FlowRecord - -```rust -pub struct FlowRecord { - pub id: String, // 请求 ID - pub timestamp: i64, // 时间戳 - pub provider: String, // Provider 类型 - pub model: String, // 模型名称 - pub request: FlowRequest, // 请求数据 - pub response: Option, // 响应数据 - pub status: FlowStatus, // 状态 - pub latency_ms: Option, // 延迟 (毫秒) -} - -pub struct FlowRequest { - pub messages: Vec, // 消息列表 - pub tools: Option>, // 工具定义 - pub stream: bool, // 是否流式 -} - -pub struct FlowResponse { - pub content: String, // 响应内容 - pub tool_calls: Option>, // 工具调用 - pub usage: TokenUsage, // Token 使用 -} - -pub struct TokenUsage { - pub prompt_tokens: u32, // 输入 Token - pub completion_tokens: u32, // 输出 Token - pub total_tokens: u32, // 总 Token -} -``` - -### 数据库表 - -```sql -CREATE TABLE flow_records ( - id TEXT PRIMARY KEY, - timestamp INTEGER NOT NULL, - provider TEXT NOT NULL, - model TEXT NOT NULL, - request_json TEXT NOT NULL, - response_json TEXT, - status TEXT NOT NULL, - latency_ms INTEGER, - prompt_tokens INTEGER, - completion_tokens INTEGER, - total_tokens INTEGER, - created_at INTEGER NOT NULL -); - -CREATE INDEX idx_flow_timestamp ON flow_records(timestamp); -CREATE INDEX idx_flow_provider ON flow_records(provider); -CREATE INDEX idx_flow_model ON flow_records(model); -``` - -## 拦截器实现 - -```rust -pub struct FlowInterceptor { - storage: Arc, - event_sender: mpsc::Sender, -} - -impl FlowInterceptor { - pub async fn intercept_request(&self, req: &Request) -> String { - let request_id = generate_uuid(); - - let record = FlowRecord { - id: request_id.clone(), - timestamp: current_timestamp(), - provider: extract_provider(req), - model: extract_model(req), - request: parse_request(req), - response: None, - status: FlowStatus::Pending, - latency_ms: None, - }; - - self.storage.insert(&record).await; - self.event_sender.send(FlowEvent::RequestStarted(record)).await; - - request_id - } - - pub async fn intercept_response( - &self, - request_id: &str, - response: &Response, - latency: Duration, - ) { - let flow_response = parse_response(response); - - self.storage.update_response( - request_id, - &flow_response, - latency.as_millis() as u64, - ).await; - - self.event_sender.send(FlowEvent::ResponseReceived { - request_id: request_id.to_string(), - response: flow_response, - }).await; - } -} -``` - -## 查询接口 - -### 分页查询 - -```rust -pub struct FlowQuery { - pub provider: Option, - pub model: Option, - pub start_time: Option, - pub end_time: Option, - pub status: Option, - pub page: u32, - pub page_size: u32, -} - -pub async fn query_flows(query: FlowQuery) -> Result> { - // 构建 SQL 查询 - // 执行分页查询 - // 返回结果 -} -``` - -### 统计查询 - -```rust -pub struct FlowStats { - pub total_requests: u64, - pub total_tokens: u64, - pub avg_latency_ms: f64, - pub by_provider: HashMap, - pub by_model: HashMap, -} - -pub async fn get_stats(time_range: TimeRange) -> Result { - // 聚合统计 -} -``` - -## 前端事件 - -### 事件类型 - -```typescript -interface FlowEvent { - type: 'request_started' | 'response_received' | 'error'; - data: FlowRecord; -} -``` - -### 事件监听 - -```typescript -// 前端监听 -import { listen } from '@tauri-apps/api/event'; - -listen('flow-event', (event) => { - switch (event.payload.type) { - case 'request_started': - addPendingRequest(event.payload.data); - break; - case 'response_received': - updateRequest(event.payload.data); - break; - } -}); -``` - -## Tauri Commands - -```rust -#[tauri::command] -async fn get_flow_records(query: FlowQuery) -> Result>; - -#[tauri::command] -async fn get_flow_stats(time_range: TimeRange) -> Result; - -#[tauri::command] -async fn get_flow_detail(id: String) -> Result; - -#[tauri::command] -async fn clear_flow_records(before: Option) -> Result; - -#[tauri::command] -async fn export_flow_records(format: ExportFormat) -> Result; -``` - -## 相关文档 - -- [server.md](server.md) - HTTP 服务器 -- [database.md](database.md) - 数据库层 -- [components.md](components.md) - 前端组件 diff --git a/package.json b/package.json index a1125373b..3a2b4c1bb 100644 --- a/package.json +++ b/package.json @@ -1,7 +1,7 @@ { "name": "proxycast", "private": true, - "version": "0.59.0", + "version": "0.60.0", "type": "module", "repository": { "type": "git", diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 6814aa6e0..818150cde 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -203,7 +203,6 @@ checksum = "7c02d123df017efcdfbd739ef81735b36c5ba83ec3c59c80a9d7ecc718f92e50" [[package]] name = "aster" version = "0.11.0" -source = "git+https://github.com/astercloud/aster-rust?tag=v0.11.0#c4a9c40b48bcb0c77375b10ccb1313366fb4e7ba" dependencies = [ "ahash", "anyhow", @@ -6636,7 +6635,7 @@ dependencies = [ [[package]] name = "proxycast" -version = "0.59.0" +version = "0.60.0" dependencies = [ "anyhow", "arboard", @@ -6672,6 +6671,7 @@ dependencies = [ "proptest", "proxycast-core", "proxycast-infra", + "proxycast-providers", "rand 0.8.5", "regex", "reqwest 0.12.28", @@ -6719,23 +6719,44 @@ dependencies = [ [[package]] name = "proxycast-core" -version = "0.59.0" +version = "0.60.0" dependencies = [ + "async-trait", + "axum 0.7.9", + "bytes", "chrono", + "dashmap 5.5.3", "dirs 5.0.1", + "flate2", + "futures", "indexmap 2.13.0", + "notify 6.1.1", "parking_lot", "proptest", + "rand 0.8.5", + "reqwest 0.12.28", + "rusqlite", "serde", "serde_json", + "serde_urlencoded", + "serde_yaml", "sha2", + "subtle", + "tar", + "tempfile", + "thiserror 1.0.69", + "tokio", + "tower 0.5.3", "tracing", + "url", + "urlencoding", "uuid", + "zip", ] [[package]] name = "proxycast-infra" -version = "0.59.0" +version = "0.60.0" dependencies = [ "chrono", "dashmap 5.5.3", @@ -6753,6 +6774,40 @@ dependencies = [ "uuid", ] +[[package]] +name = "proxycast-providers" +version = "0.60.0" +dependencies = [ + "anyhow", + "async-stream", + "async-trait", + "axum 0.7.9", + "base64 0.22.1", + "bytes", + "chrono", + "dirs 5.0.1", + "flate2", + "futures", + "once_cell", + "open", + "proptest", + "proxycast-core", + "rand 0.8.5", + "regex", + "reqwest 0.12.28", + "serde", + "serde_json", + "serde_urlencoded", + "sha2", + "tempfile", + "thiserror 1.0.69", + "tokio", + "tracing", + "url", + "urlencoding", + "uuid", +] + [[package]] name = "psl-types" version = "2.0.11" diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 496663157..28e8b2502 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -3,7 +3,7 @@ members = ["crates/*"] resolver = "2" [workspace.package] -version = "0.59.0" +version = "0.60.0" edition = "2021" authors = ["you"] repository = "https://github.com/aiclientproxy/proxycast" @@ -13,6 +13,7 @@ homepage = "https://github.com/aiclientproxy/proxycast" # 项目内 crate 依赖 proxycast-core = { path = "crates/core" } proxycast-infra = { path = "crates/infra" } +proxycast-providers = { path = "crates/providers" } voice-core = { path = "crates/voice-core" } # 序列化 @@ -167,7 +168,7 @@ version = "2.4" [package] name = "proxycast" -version = "0.59.0" +version = "0.60.0" description = "AI API Proxy Desktop App" authors = ["you"] edition = "2021" @@ -185,6 +186,7 @@ tauri-build.workspace = true # 项目内 crate proxycast-core.workspace = true proxycast-infra.workspace = true +proxycast-providers.workspace = true voice-core.workspace = true # Tauri diff --git a/src-tauri/capabilities/default.json b/src-tauri/capabilities/default.json deleted file mode 100644 index fbf6c3c3a..000000000 --- a/src-tauri/capabilities/default.json +++ /dev/null @@ -1,57 +0,0 @@ -{ - "$schema": "https://schemas.tauri.app/config/2/capability", - "identifier": "default", - "description": "Default capabilities for ProxyCast", - "windows": ["main", "smart-input", "update-notification"], - "permissions": [ - "core:default", - "core:webview:default", - "core:webview:allow-webview-close", - "core:webview:allow-webview-position", - "core:webview:allow-webview-size", - "core:webview:allow-set-webview-position", - "core:webview:allow-set-webview-size", - "core:window:default", - "core:window:allow-close", - "core:window:allow-show", - "core:window:allow-hide", - "core:window:allow-set-focus", - "core:window:allow-center", - "core:window:allow-start-dragging", - "shell:allow-open", - "shell:allow-spawn", - "shell:allow-execute", - "shell:allow-kill", - "shell:allow-stdin-write", - "dialog:default", - "global-shortcut:default", - "global-shortcut:allow-is-registered", - "global-shortcut:allow-register", - "global-shortcut:allow-unregister", - { - "identifier": "shell:allow-execute", - "allow": [ - { - "name": "binaries/aster-server", - "sidecar": true, - "args": true - }, - { - "name": "open", - "cmd": "open", - "args": true - } - ] - }, - { - "identifier": "shell:allow-spawn", - "allow": [ - { - "name": "binaries/aster-server", - "sidecar": true, - "args": true - } - ] - } - ] -} diff --git a/src-tauri/crates/core/Cargo.toml b/src-tauri/crates/core/Cargo.toml index dbf317f36..0526eba8d 100644 --- a/src-tauri/crates/core/Cargo.toml +++ b/src-tauri/crates/core/Cargo.toml @@ -9,10 +9,21 @@ repository.workspace = true # 序列化 serde.workspace = true serde_json.workspace = true +serde_urlencoded.workspace = true + +# 异步运行时 +tokio.workspace = true +async-trait.workspace = true + +# 错误处理 +thiserror.workspace = true # 日志 tracing.workspace = true +# HTTP 客户端 +reqwest.workspace = true + # 时间和 UUID chrono.workspace = true uuid.workspace = true @@ -22,6 +33,30 @@ indexmap.workspace = true parking_lot.workspace = true dirs.workspace = true sha2.workspace = true +url.workspace = true +urlencoding.workspace = true +bytes.workspace = true +futures.workspace = true +dashmap.workspace = true +notify.workspace = true +rand.workspace = true + +# HTTP 服务器(middleware 模块需要) +axum.workspace = true +tower.workspace = true +subtle.workspace = true + +# 压缩/归档(plugin installer 需要) +flate2.workspace = true +tar.workspace = true +zip.workspace = true + +# YAML 配置 +serde_yaml.workspace = true + +# 数据库(errors 模块需要 rusqlite::Error) +rusqlite.workspace = true [dev-dependencies] -proptest.workspace = true \ No newline at end of file +proptest.workspace = true +tempfile.workspace = true \ No newline at end of file diff --git a/src-tauri/src/backends/mod.rs b/src-tauri/crates/core/src/backends/mod.rs similarity index 100% rename from src-tauri/src/backends/mod.rs rename to src-tauri/crates/core/src/backends/mod.rs diff --git a/src-tauri/src/backends/traits.rs b/src-tauri/crates/core/src/backends/traits.rs similarity index 100% rename from src-tauri/src/backends/traits.rs rename to src-tauri/crates/core/src/backends/traits.rs diff --git a/src-tauri/src/config/export.rs b/src-tauri/crates/core/src/config/export.rs similarity index 100% rename from src-tauri/src/config/export.rs rename to src-tauri/crates/core/src/config/export.rs diff --git a/src-tauri/src/config/hot_reload.rs b/src-tauri/crates/core/src/config/hot_reload.rs similarity index 100% rename from src-tauri/src/config/hot_reload.rs rename to src-tauri/crates/core/src/config/hot_reload.rs diff --git a/src-tauri/src/config/import.rs b/src-tauri/crates/core/src/config/import.rs similarity index 100% rename from src-tauri/src/config/import.rs rename to src-tauri/crates/core/src/config/import.rs diff --git a/src-tauri/crates/core/src/config/mod.rs b/src-tauri/crates/core/src/config/mod.rs new file mode 100644 index 000000000..6847cc1f2 --- /dev/null +++ b/src-tauri/crates/core/src/config/mod.rs @@ -0,0 +1,34 @@ +//! 配置管理模块 +//! +//! 提供 YAML 配置文件支持、热重载和配置导入导出功能 +//! 同时保持与旧版 JSON 配置的向后兼容性 + +#![allow(unused_imports)] + +mod export; +mod hot_reload; +mod import; +mod path_utils; +mod types; +mod yaml; + +pub use export::{ExportBundle, ExportOptions, ExportService, REDACTED_PLACEHOLDER}; +pub use hot_reload::{ + ConfigChangeEvent as FileChangeEvent, ConfigChangeKind, FileWatcher, HotReloadManager, + ReloadResult, +}; +pub use import::{ImportOptions, ImportService, ValidationResult}; +pub use path_utils::{collapse_tilde, contains_tilde, expand_tilde}; +pub use types::{ + generate_secure_api_key, AmpConfig, AmpModelMapping, ApiKeyEntry, AsrCredentialEntry, + AsrProviderType, BaiduConfig, Config, ContentCreatorConfig, CredentialEntry, + CredentialPoolConfig, CustomProviderConfig, EndpointProvidersConfig, ExperimentalFeatures, + GeminiApiKeyEntry, InjectionRuleConfig, InjectionSettings, LoggingConfig, ModelInfo, + ModelsConfig, NativeAgentConfig, NavigationConfig, OpenAIAsrConfig, ProviderConfig, + ProviderModelsConfig, ProvidersConfig, QuotaExceededConfig, RemoteManagementConfig, + RetrySettings, RoutingConfig, ScreenshotChatConfig, ServerConfig, TlsConfig, UpdateCheckConfig, + VertexApiKeyEntry, VertexModelAlias, VoiceInputConfig, VoiceInstruction, VoiceOutputConfig, + VoiceOutputMode, VoiceProcessorConfig, WhisperLocalConfig, WhisperModelSize, XunfeiConfig, + DEFAULT_API_KEY, +}; +pub use yaml::{load_config, save_config, ConfigError, ConfigManager, YamlService}; diff --git a/src-tauri/src/config/path_utils.rs b/src-tauri/crates/core/src/config/path_utils.rs similarity index 100% rename from src-tauri/src/config/path_utils.rs rename to src-tauri/crates/core/src/config/path_utils.rs diff --git a/src-tauri/crates/core/src/config/tests.rs b/src-tauri/crates/core/src/config/tests.rs new file mode 100644 index 000000000..b72b69a68 --- /dev/null +++ b/src-tauri/crates/core/src/config/tests.rs @@ -0,0 +1,2685 @@ +//! 配置模块属性测试 +//! +//! 使用 proptest 进行属性测试 + +use crate::config::types::{ContentCreatorConfig, NavigationConfig}; +use crate::config::{ + collapse_tilde, contains_tilde, expand_tilde, Config, ConfigManager, CustomProviderConfig, + HotReloadManager, InjectionSettings, LoggingConfig, ProviderConfig, ProvidersConfig, + ReloadResult, RetrySettings, RoutingConfig, ServerConfig, YamlService, +}; +use proptest::prelude::*; +use std::io::Write; +use tempfile::NamedTempFile; + +/// 生成随机的主机地址 +fn arb_host() -> impl Strategy { + prop_oneof![ + Just("127.0.0.1".to_string()), + Just("localhost".to_string()), + Just("::1".to_string()), + "[0-9]{1,3}\\.[0-9]{1,3}\\.[0-9]{1,3}\\.[0-9]{1,3}".prop_map(|s| s), + ] +} + +/// 生成随机的端口号 +fn arb_port() -> impl Strategy { + 1024u16..65535u16 +} + +/// 生成随机的 API 密钥 +fn arb_api_key() -> impl Strategy { + "[a-zA-Z0-9_-]{8,32}".prop_map(|s| s) +} + +/// 生成随机的服务器配置 +fn arb_server_config() -> impl Strategy { + (arb_host(), arb_port(), arb_api_key()).prop_map(|(host, port, api_key)| ServerConfig { + host, + port, + api_key, + tls: crate::config::TlsConfig::default(), + }) +} + +/// 生成随机的 Provider 配置 +fn arb_provider_config() -> impl Strategy { + ( + any::(), + proptest::option::of("[a-zA-Z0-9/_.-]{5,50}".prop_map(|s| s)), + proptest::option::of(prop_oneof![ + Just("us-east-1".to_string()), + Just("us-west-2".to_string()), + Just("eu-west-1".to_string()), + ]), + proptest::option::of("[a-zA-Z0-9-]{5,20}".prop_map(|s| s)), + ) + .prop_map( + |(enabled, credentials_path, region, project_id)| ProviderConfig { + enabled, + credentials_path, + region, + project_id, + }, + ) +} + +/// 生成随机的自定义 Provider 配置 +fn arb_custom_provider_config() -> impl Strategy { + ( + any::(), + proptest::option::of(arb_api_key()), + proptest::option::of(prop_oneof![ + Just("https://api.openai.com/v1".to_string()), + Just("https://api.anthropic.com".to_string()), + Just("https://custom.api.com".to_string()), + ]), + ) + .prop_map(|(enabled, api_key, base_url)| CustomProviderConfig { + enabled, + api_key, + base_url, + }) +} + +/// 生成随机的 Providers 配置 +fn arb_providers_config() -> impl Strategy { + ( + arb_provider_config(), + arb_provider_config(), + arb_provider_config(), + arb_custom_provider_config(), + arb_custom_provider_config(), + ) + .prop_map(|(kiro, gemini, qwen, openai, claude)| ProvidersConfig { + kiro, + gemini, + qwen, + openai, + claude, + }) +} + +/// 生成随机的路由配置 +fn arb_routing_config() -> impl Strategy { + ( + prop_oneof![ + Just("kiro".to_string()), + Just("gemini".to_string()), + Just("qwen".to_string()), + ], + proptest::collection::hash_map( + "[a-z]+-[a-z0-9]+".prop_map(|s| s), + "[a-z]+-[a-z0-9-]+".prop_map(|s| s), + 0..5, + ), + ) + .prop_map(|(default_provider, model_aliases)| RoutingConfig { + default_provider, + model_aliases, + }) +} + +/// 生成随机的重试配置 +fn arb_retry_settings() -> impl Strategy { + ( + 1u32..10u32, + 100u64..5000u64, + 5000u64..60000u64, + any::(), + ) + .prop_map( + |(max_retries, base_delay_ms, max_delay_ms, auto_switch_provider)| RetrySettings { + max_retries, + base_delay_ms, + max_delay_ms, + auto_switch_provider, + }, + ) +} + +/// 生成随机的日志配置 +fn arb_logging_config() -> impl Strategy { + ( + any::(), + prop_oneof![ + Just("debug".to_string()), + Just("info".to_string()), + Just("warn".to_string()), + Just("error".to_string()), + ], + 1u32..30u32, + any::(), + ) + .prop_map( + |(enabled, level, retention_days, include_request_body)| LoggingConfig { + enabled, + level, + retention_days, + include_request_body, + }, + ) +} + +/// 生成随机的完整配置 +fn arb_config() -> impl Strategy { + ( + arb_server_config(), + arb_providers_config(), + arb_routing_config(), + arb_retry_settings(), + arb_logging_config(), + ) + .prop_map(|(server, providers, routing, retry, logging)| Config { + server, + providers, + default_provider: routing.default_provider.clone(), + routing, + retry, + logging, + injection: InjectionSettings::default(), + auth_dir: "~/.proxycast/auth".to_string(), + credential_pool: crate::config::CredentialPoolConfig::default(), + remote_management: crate::config::RemoteManagementConfig::default(), + quota_exceeded: crate::config::QuotaExceededConfig::default(), + proxy_url: None, + ampcode: crate::config::AmpConfig::default(), + endpoint_providers: crate::config::EndpointProvidersConfig::default(), + minimize_to_tray: true, + models: crate::config::ModelsConfig::default(), + agent: crate::config::NativeAgentConfig::default(), + language: "zh".to_string(), + experimental: crate::config::ExperimentalFeatures::default(), + content_creator: ContentCreatorConfig::default(), + navigation: NavigationConfig::default(), + }) +} + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: enhancement-roadmap, Property 11: 配置往返一致性** + /// *对于任意* 有效配置,序列化后再反序列化应得到等价的配置 + /// **Validates: Requirements 4.1** + #[test] + fn prop_config_roundtrip(config in arb_config()) { + // 序列化为 YAML + let yaml = ConfigManager::to_yaml(&config) + .expect("序列化应成功"); + + // 反序列化回 Config + let parsed = ConfigManager::parse_yaml(&yaml) + .expect("反序列化应成功"); + + // 验证往返一致性 + prop_assert_eq!( + config.server, + parsed.server, + "服务器配置往返不一致" + ); + prop_assert_eq!( + config.providers, + parsed.providers, + "Provider 配置往返不一致" + ); + prop_assert_eq!( + config.routing.default_provider, + parsed.routing.default_provider, + "默认 Provider 往返不一致" + ); + prop_assert_eq!( + config.routing.model_aliases, + parsed.routing.model_aliases, + "模型别名往返不一致" + ); + prop_assert_eq!( + config.retry, + parsed.retry, + "重试配置往返不一致" + ); + prop_assert_eq!( + config.logging, + parsed.logging, + "日志配置往返不一致" + ); + } + + /// **Feature: enhancement-roadmap, Property 11: 配置往返一致性(服务器配置)** + /// *对于任意* 服务器配置,序列化后再反序列化应得到等价的配置 + /// **Validates: Requirements 4.1** + #[test] + fn prop_server_config_roundtrip(server in arb_server_config()) { + let config = Config { + server: server.clone(), + ..Config::default() + }; + + let yaml = ConfigManager::to_yaml(&config).expect("序列化应成功"); + let parsed = ConfigManager::parse_yaml(&yaml).expect("反序列化应成功"); + + prop_assert_eq!( + server, + parsed.server, + "服务器配置往返不一致" + ); + } + + /// **Feature: enhancement-roadmap, Property 11: 配置往返一致性(Provider 配置)** + /// *对于任意* Provider 配置,序列化后再反序列化应得到等价的配置 + /// **Validates: Requirements 4.1** + #[test] + fn prop_providers_config_roundtrip(providers in arb_providers_config()) { + let config = Config { + providers: providers.clone(), + ..Config::default() + }; + + let yaml = ConfigManager::to_yaml(&config).expect("序列化应成功"); + let parsed = ConfigManager::parse_yaml(&yaml).expect("反序列化应成功"); + + prop_assert_eq!( + providers, + parsed.providers, + "Provider 配置往返不一致" + ); + } + + /// **Feature: enhancement-roadmap, Property 11: 配置往返一致性(重试配置)** + /// *对于任意* 重试配置,序列化后再反序列化应得到等价的配置 + /// **Validates: Requirements 4.1** + #[test] + fn prop_retry_settings_roundtrip(retry in arb_retry_settings()) { + let config = Config { + retry: retry.clone(), + ..Config::default() + }; + + let yaml = ConfigManager::to_yaml(&config).expect("序列化应成功"); + let parsed = ConfigManager::parse_yaml(&yaml).expect("反序列化应成功"); + + prop_assert_eq!( + retry, + parsed.retry, + "重试配置往返不一致" + ); + } + + /// **Feature: enhancement-roadmap, Property 11: 配置往返一致性(日志配置)** + /// *对于任意* 日志配置,序列化后再反序列化应得到等价的配置 + /// **Validates: Requirements 4.1** + #[test] + fn prop_logging_config_roundtrip(logging in arb_logging_config()) { + let config = Config { + logging: logging.clone(), + ..Config::default() + }; + + let yaml = ConfigManager::to_yaml(&config).expect("序列化应成功"); + let parsed = ConfigManager::parse_yaml(&yaml).expect("反序列化应成功"); + + prop_assert_eq!( + logging, + parsed.logging, + "日志配置往返不一致" + ); + } + + /// **Feature: enhancement-roadmap, Property 11: 配置往返一致性(路由配置)** + /// *对于任意* 路由配置,序列化后再反序列化应得到等价的配置 + /// **Validates: Requirements 4.1** + #[test] + fn prop_routing_config_roundtrip(routing in arb_routing_config()) { + let config = Config { + routing: routing.clone(), + ..Config::default() + }; + + let yaml = ConfigManager::to_yaml(&config).expect("序列化应成功"); + let parsed = ConfigManager::parse_yaml(&yaml).expect("反序列化应成功"); + + prop_assert_eq!( + routing.default_provider, + parsed.routing.default_provider, + "默认 Provider 往返不一致" + ); + prop_assert_eq!( + routing.model_aliases, + parsed.routing.model_aliases, + "模型别名往返不一致" + ); + } +} + +/// 生成有效的服务器配置(端口非零) +fn arb_valid_server_config() -> impl Strategy { + (arb_host(), 1u16..65535u16, arb_api_key()).prop_map(|(host, port, api_key)| ServerConfig { + host, + port, + api_key, + tls: crate::config::TlsConfig::default(), + }) +} + +/// 生成有效的重试配置(通过验证) +fn arb_valid_retry_settings() -> impl Strategy { + ( + 1u32..100u32, // max_retries <= 100 + 1u64..5000u64, // base_delay_ms > 0 + 5000u64..60000u64, // max_delay_ms + any::(), + ) + .prop_map( + |(max_retries, base_delay_ms, max_delay_ms, auto_switch_provider)| RetrySettings { + max_retries, + base_delay_ms, + max_delay_ms, + auto_switch_provider, + }, + ) +} + +/// 生成有效的日志配置(保留天数非零) +fn arb_valid_logging_config() -> impl Strategy { + ( + any::(), + prop_oneof![ + Just("debug".to_string()), + Just("info".to_string()), + Just("warn".to_string()), + Just("error".to_string()), + ], + 1u32..30u32, // retention_days > 0 + any::(), + ) + .prop_map( + |(enabled, level, retention_days, include_request_body)| LoggingConfig { + enabled, + level, + retention_days, + include_request_body, + }, + ) +} + +/// 生成有效的配置(通过验证的配置) +fn arb_valid_config() -> impl Strategy { + ( + arb_valid_server_config(), + arb_providers_config(), + arb_routing_config(), + arb_valid_retry_settings(), + arb_valid_logging_config(), + ) + .prop_map(|(server, providers, routing, retry, logging)| Config { + server, + providers, + default_provider: routing.default_provider.clone(), + routing, + retry, + logging, + injection: InjectionSettings::default(), + auth_dir: "~/.proxycast/auth".to_string(), + credential_pool: crate::config::CredentialPoolConfig::default(), + remote_management: crate::config::RemoteManagementConfig::default(), + quota_exceeded: crate::config::QuotaExceededConfig::default(), + proxy_url: None, + ampcode: crate::config::AmpConfig::default(), + endpoint_providers: crate::config::EndpointProvidersConfig::default(), + minimize_to_tray: true, + models: crate::config::ModelsConfig::default(), + agent: crate::config::NativeAgentConfig::default(), + language: "zh".to_string(), + experimental: crate::config::ExperimentalFeatures::default(), + content_creator: ContentCreatorConfig::default(), + navigation: NavigationConfig::default(), + }) +} + +/// 生成无效配置的类型 +#[derive(Debug, Clone, Copy)] +enum InvalidConfigType { + ZeroPort, + TooManyRetries, + ZeroBaseDelay, + ZeroRetentionDays, +} + +/// 生成无效的配置(不通过验证的配置) +fn arb_invalid_config() -> impl Strategy { + ( + arb_valid_server_config(), + arb_providers_config(), + arb_routing_config(), + arb_valid_retry_settings(), + arb_valid_logging_config(), + prop_oneof![ + Just(InvalidConfigType::ZeroPort), + Just(InvalidConfigType::TooManyRetries), + Just(InvalidConfigType::ZeroBaseDelay), + Just(InvalidConfigType::ZeroRetentionDays), + ], + ) + .prop_map( + |(server, providers, routing, retry, logging, invalid_type)| { + let mut config = Config { + server, + providers, + default_provider: routing.default_provider.clone(), + routing, + retry, + logging, + injection: InjectionSettings::default(), + auth_dir: "~/.proxycast/auth".to_string(), + credential_pool: crate::config::CredentialPoolConfig::default(), + remote_management: crate::config::RemoteManagementConfig::default(), + quota_exceeded: crate::config::QuotaExceededConfig::default(), + proxy_url: None, + ampcode: crate::config::AmpConfig::default(), + endpoint_providers: crate::config::EndpointProvidersConfig::default(), + minimize_to_tray: true, + models: crate::config::ModelsConfig::default(), + agent: crate::config::NativeAgentConfig::default(), + language: "zh".to_string(), + experimental: crate::config::ExperimentalFeatures::default(), + content_creator: ContentCreatorConfig::default(), + navigation: NavigationConfig::default(), + }; + // 根据类型使配置无效 + match invalid_type { + InvalidConfigType::ZeroPort => config.server.port = 0, + InvalidConfigType::TooManyRetries => config.retry.max_retries = 101, + InvalidConfigType::ZeroBaseDelay => config.retry.base_delay_ms = 0, + InvalidConfigType::ZeroRetentionDays => config.logging.retention_days = 0, + } + config + }, + ) +} + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: enhancement-roadmap, Property 12: 热重载原子性** + /// *对于任意* 配置变更,要么完全应用成功,要么回滚到之前状态 + /// **Validates: Requirements 4.2 (验收标准 3)** + #[test] + fn prop_hot_reload_atomicity_success( + initial_config in arb_valid_config(), + new_config in arb_valid_config() + ) { + // 创建临时配置文件 + let mut temp_file = NamedTempFile::new().expect("创建临时文件失败"); + let yaml = ConfigManager::to_yaml(&new_config).expect("序列化失败"); + temp_file.write_all(yaml.as_bytes()).expect("写入文件失败"); + + // 创建热重载管理器 + let manager = HotReloadManager::new(initial_config.clone(), temp_file.path().to_path_buf()); + + // 执行热重载 + let result = manager.reload(); + + // 验证原子性:成功时配置应完全更新 + match result { + ReloadResult::Success { .. } => { + let current = manager.config(); + // 验证配置已完全更新 + prop_assert_eq!( + current.server, + new_config.server, + "服务器配置未正确更新" + ); + prop_assert_eq!( + current.providers, + new_config.providers, + "Provider 配置未正确更新" + ); + prop_assert_eq!( + current.retry, + new_config.retry, + "重试配置未正确更新" + ); + prop_assert_eq!( + current.logging, + new_config.logging, + "日志配置未正确更新" + ); + } + _ => { + // 如果失败,应该保持原始配置 + let current = manager.config(); + prop_assert_eq!( + current, + initial_config, + "失败时配置应保持不变" + ); + } + } + } + + /// **Feature: enhancement-roadmap, Property 12: 热重载原子性(失败回滚)** + /// *对于任意* 无效配置变更,配置应回滚到之前状态 + /// **Validates: Requirements 4.2 (验收标准 3)** + #[test] + fn prop_hot_reload_atomicity_rollback( + initial_config in arb_valid_config(), + invalid_config in arb_invalid_config() + ) { + // 创建临时配置文件(包含无效配置) + let mut temp_file = NamedTempFile::new().expect("创建临时文件失败"); + let yaml = ConfigManager::to_yaml(&invalid_config).expect("序列化失败"); + temp_file.write_all(yaml.as_bytes()).expect("写入文件失败"); + + // 创建热重载管理器 + let manager = HotReloadManager::new(initial_config.clone(), temp_file.path().to_path_buf()); + + // 执行热重载 + let result = manager.reload(); + + // 验证原子性:失败时配置应回滚到之前状态 + match result { + ReloadResult::RolledBack { .. } => { + let current = manager.config(); + // 验证配置已回滚到初始状态 + prop_assert_eq!( + current, + initial_config, + "配置应回滚到初始状态" + ); + } + ReloadResult::Success { .. } => { + // 如果意外成功(不应该发生),验证配置一致性 + let current = manager.config(); + prop_assert_eq!( + current, + invalid_config, + "成功时配置应完全更新" + ); + } + ReloadResult::Failed { .. } => { + // 完全失败的情况,配置应保持不变 + let current = manager.config(); + prop_assert_eq!( + current, + initial_config, + "失败时配置应保持不变" + ); + } + } + } + + /// **Feature: enhancement-roadmap, Property 12: 热重载原子性(文件不存在)** + /// *对于任意* 初始配置,当配置文件不存在时,配置应保持不变 + /// **Validates: Requirements 4.2 (验收标准 3)** + #[test] + fn prop_hot_reload_atomicity_file_not_exists(initial_config in arb_valid_config()) { + // 使用不存在的文件路径 + let nonexistent_path = std::path::PathBuf::from("/tmp/nonexistent_config_test_12345.yaml"); + + // 创建热重载管理器 + let manager = HotReloadManager::new(initial_config.clone(), nonexistent_path); + + // 执行热重载 + let result = manager.reload(); + + // 验证原子性:文件不存在时配置应保持不变 + match result { + ReloadResult::RolledBack { .. } => { + let current = manager.config(); + prop_assert_eq!( + current, + initial_config, + "文件不存在时配置应保持不变" + ); + } + _ => { + // 其他情况也应保持配置不变 + let current = manager.config(); + prop_assert_eq!( + current, + initial_config, + "配置应保持不变" + ); + } + } + } + + /// **Feature: enhancement-roadmap, Property 12: 热重载原子性(无效 YAML)** + /// *对于任意* 初始配置,当配置文件包含无效 YAML 时,配置应保持不变 + /// **Validates: Requirements 4.2 (验收标准 3)** + #[test] + fn prop_hot_reload_atomicity_invalid_yaml(initial_config in arb_valid_config()) { + // 创建包含无效 YAML 的临时文件 + let mut temp_file = NamedTempFile::new().expect("创建临时文件失败"); + temp_file.write_all(b"invalid: yaml: content: [").expect("写入文件失败"); + + // 创建热重载管理器 + let manager = HotReloadManager::new(initial_config.clone(), temp_file.path().to_path_buf()); + + // 执行热重载 + let result = manager.reload(); + + // 验证原子性:无效 YAML 时配置应保持不变 + match result { + ReloadResult::RolledBack { .. } => { + let current = manager.config(); + prop_assert_eq!( + current, + initial_config, + "无效 YAML 时配置应保持不变" + ); + } + _ => { + // 其他情况也应保持配置不变 + let current = manager.config(); + prop_assert_eq!( + current, + initial_config, + "配置应保持不变" + ); + } + } + } +} + +// ============================================================================ +// Property 3: Tilde Path Expansion +// ============================================================================ + +/// 生成有效的 tilde 路径(~/path 格式) +/// 排除 "." 和 ".." 路径段,因为这些会导致路径规范化问题 +fn arb_tilde_path() -> impl Strategy { + // 生成路径段:字母数字、下划线、连字符 + // 排除单独的 "." 和 ".." 以避免路径规范化问题 + let path_segment = "[a-zA-Z0-9_-]{1,20}"; + + // 生成 0-5 个路径段 + proptest::collection::vec(path_segment, 0..6).prop_map(|segments| { + if segments.is_empty() { + "~".to_string() + } else { + format!("~/{}", segments.join("/")) + } + }) +} + +/// 生成不包含 tilde 的绝对路径 +fn arb_absolute_path() -> impl Strategy { + let path_segment = "[a-zA-Z0-9_.-]{1,20}"; + + proptest::collection::vec(path_segment, 1..6) + .prop_map(|segments| format!("/{}", segments.join("/"))) +} + +/// 生成不包含 tilde 的相对路径 +/// 排除单独的 "." 和 ".." 以避免路径规范化问题 +fn arb_relative_path() -> impl Strategy { + // 使用至少2个字符的路径段,或者不以单独的点开头 + // 这样可以避免生成 "." 或 ".." 这样的特殊路径 + let path_segment = "[a-zA-Z0-9_-][a-zA-Z0-9_.-]{0,19}"; + + proptest::collection::vec(path_segment, 1..6).prop_map(|segments| segments.join("/")) +} + +/// 生成 ~user/path 格式的路径(不支持的格式) +fn arb_tilde_user_path() -> impl Strategy { + let username = "[a-z]{3,10}"; + let path_segment = "[a-zA-Z0-9_.-]{1,20}"; + + (username, proptest::collection::vec(path_segment, 0..4)).prop_map(|(user, segments)| { + if segments.is_empty() { + format!("~{user}") + } else { + format!("~{}/{}", user, segments.join("/")) + } + }) +} + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: config-credential-export, Property 3: Tilde Path Expansion** + /// *For any* valid tilde path (~/path format), expanding and then collapsing + /// should produce the original path. + /// **Validates: Requirements 2.3** + #[test] + fn prop_tilde_path_roundtrip(path in arb_tilde_path()) { + // 展开 tilde 路径 + let expanded = expand_tilde(&path); + + // 收缩回 tilde 格式 + let collapsed = collapse_tilde(&expanded); + + // 验证往返一致性 + prop_assert_eq!( + &collapsed, + &path, + "Tilde 路径往返不一致: 原始={}, 展开={:?}, 收缩={}", + path, + expanded, + collapsed + ); + } + + /// **Feature: config-credential-export, Property 3: Tilde Path Expansion** + /// *For any* tilde path, the expanded path should start with the user's home directory. + /// **Validates: Requirements 2.3** + #[test] + fn prop_tilde_expansion_starts_with_home(path in arb_tilde_path()) { + let home_dir = dirs::home_dir().expect("应该能获取主目录"); + let expanded = expand_tilde(&path); + + prop_assert!( + expanded.starts_with(&home_dir), + "展开后的路径应以主目录开头: 路径={}, 展开={:?}, 主目录={:?}", + path, + expanded, + home_dir + ); + } + + /// **Feature: config-credential-export, Property 3: Tilde Path Expansion** + /// *For any* tilde path, contains_tilde should return true before expansion + /// and false after expansion. + /// **Validates: Requirements 2.3** + #[test] + fn prop_contains_tilde_before_expansion(path in arb_tilde_path()) { + // 展开前应包含 tilde + prop_assert!( + contains_tilde(&path), + "展开前路径应包含 tilde: {}", + path + ); + + // 展开后不应包含 tilde + let expanded = expand_tilde(&path); + prop_assert!( + !contains_tilde(&expanded), + "展开后路径不应包含 tilde: {:?}", + expanded + ); + } + + /// **Feature: config-credential-export, Property 3: Tilde Path Expansion** + /// *For any* absolute path (not starting with ~), expand_tilde should return + /// the path unchanged. + /// **Validates: Requirements 2.3** + #[test] + fn prop_absolute_path_unchanged(path in arb_absolute_path()) { + let expanded = expand_tilde(&path); + let expanded_str = expanded.to_string_lossy().to_string(); + + prop_assert_eq!( + &expanded_str, + &path, + "绝对路径应保持不变: 原始={}, 展开={:?}", + path, + expanded + ); + } + + /// **Feature: config-credential-export, Property 3: Tilde Path Expansion** + /// *For any* relative path (not starting with ~ or /), expand_tilde should + /// return the path unchanged. + /// **Validates: Requirements 2.3** + #[test] + fn prop_relative_path_unchanged(path in arb_relative_path()) { + let expanded = expand_tilde(&path); + let expanded_str = expanded.to_string_lossy().to_string(); + + prop_assert_eq!( + &expanded_str, + &path, + "相对路径应保持不变: 原始={}, 展开={:?}", + path, + expanded + ); + } + + /// **Feature: config-credential-export, Property 3: Tilde Path Expansion** + /// *For any* ~user/path format (unsupported), expand_tilde should return + /// the path unchanged. + /// **Validates: Requirements 2.3** + #[test] + fn prop_tilde_user_path_unchanged(path in arb_tilde_user_path()) { + let expanded = expand_tilde(&path); + let expanded_str = expanded.to_string_lossy().to_string(); + + prop_assert_eq!( + &expanded_str, + &path, + "~user/path 格式应保持不变: 原始={}, 展开={:?}", + path, + expanded + ); + } + + /// **Feature: config-credential-export, Property 3: Tilde Path Expansion** + /// *For any* path under the home directory, collapse_tilde should produce + /// a path starting with ~. + /// **Validates: Requirements 2.3** + #[test] + fn prop_collapse_home_path_starts_with_tilde(subpath in arb_relative_path()) { + let home_dir = dirs::home_dir().expect("应该能获取主目录"); + let full_path = home_dir.join(&subpath); + + let collapsed = collapse_tilde(&full_path); + + prop_assert!( + collapsed.starts_with("~/"), + "主目录下的路径收缩后应以 ~/ 开头: 路径={:?}, 收缩={}", + full_path, + collapsed + ); + } + + /// **Feature: config-credential-export, Property 3: Tilde Path Expansion** + /// *For any* path not under the home directory, collapse_tilde should return + /// the path unchanged. + /// **Validates: Requirements 2.3** + #[test] + fn prop_collapse_non_home_path_unchanged(path in arb_absolute_path()) { + // 确保路径不在主目录下(使用 /tmp 或类似路径) + let test_path = format!("/tmp{path}"); + let collapsed = collapse_tilde(&test_path); + + prop_assert_eq!( + &collapsed, + &test_path, + "非主目录路径应保持不变: 原始={}, 收缩={}", + test_path, + collapsed + ); + } +} + +// ============================================================================ +// Property 2: YAML Comment Preservation +// ============================================================================ + +/// 生成有效的 YAML 注释(以 # 开头) +fn arb_yaml_comment() -> impl Strategy { + // 生成注释内容:字母、数字、空格、中文字符 + "[a-zA-Z0-9 ]{1,50}".prop_map(|s| format!("# {s}")) +} + +/// 生成带注释的 YAML 配置字符串 +fn arb_yaml_with_comments() -> impl Strategy)> { + ( + arb_valid_config(), + proptest::collection::vec(arb_yaml_comment(), 1..5), + ) + .prop_map(|(config, comments)| { + // 序列化配置 + let yaml = ConfigManager::to_yaml(&config).expect("序列化应成功"); + let lines: Vec<&str> = yaml.lines().collect(); + + // 在 YAML 中插入注释 + let mut result_lines: Vec = Vec::new(); + let mut comment_iter = comments.iter(); + + // 在文件开头添加一个注释 + if let Some(comment) = comment_iter.next() { + result_lines.push(comment.clone()); + } + + for (i, line) in lines.iter().enumerate() { + result_lines.push(line.to_string()); + + // 在某些行后添加注释 + if i % 5 == 0 { + if let Some(comment) = comment_iter.next() { + result_lines.push(comment.clone()); + } + } + } + + // 收集实际插入的注释 + let inserted_comments: Vec = result_lines + .iter() + .filter(|line| line.trim().starts_with('#')) + .cloned() + .collect(); + + (result_lines.join("\n"), inserted_comments) + }) +} + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: config-credential-export, Property 2: YAML Comment Preservation** + /// *For any* YAML file with comments, saving configuration changes should preserve + /// all existing comments in their original positions. + /// **Validates: Requirements 1.3** + #[test] + fn prop_yaml_comment_preservation( + (yaml_with_comments, original_comments) in arb_yaml_with_comments(), + new_config in arb_valid_config() + ) { + // 创建临时文件 + let temp_dir = tempfile::tempdir().expect("创建临时目录失败"); + let config_path = temp_dir.path().join("config.yaml"); + + // 写入带注释的原始 YAML + std::fs::write(&config_path, &yaml_with_comments).expect("写入文件失败"); + + // 使用 YamlService 保存新配置(应保留注释) + YamlService::save_preserve_comments(&config_path, &new_config) + .expect("保存配置失败"); + + // 读取保存后的内容 + let saved_content = std::fs::read_to_string(&config_path).expect("读取文件失败"); + + // 提取保存后的注释 + let saved_comments: Vec = saved_content + .lines() + .filter(|line| line.trim().starts_with('#')) + .map(|s| s.to_string()) + .collect(); + + // 验证注释被保留 + // 注意:由于 YAML 结构可能变化,我们只验证注释内容被保留,不验证位置 + for original_comment in &original_comments { + let comment_content = original_comment.trim(); + let found = saved_comments.iter().any(|c| c.trim() == comment_content); + prop_assert!( + found, + "注释应被保留: 原始注释='{}', 保存后的注释={:?}", + comment_content, + saved_comments + ); + } + } + + /// **Feature: config-credential-export, Property 2: YAML Comment Preservation** + /// *For any* configuration saved with YamlService, the configuration should be + /// correctly parseable and equivalent to the original. + /// **Validates: Requirements 1.3** + #[test] + fn prop_yaml_save_preserves_config(config in arb_valid_config()) { + // 创建临时文件 + let temp_dir = tempfile::tempdir().expect("创建临时目录失败"); + let config_path = temp_dir.path().join("config.yaml"); + + // 使用 YamlService 保存配置 + YamlService::save_preserve_comments(&config_path, &config) + .expect("保存配置失败"); + + // 读取并解析保存后的配置 + let saved_content = std::fs::read_to_string(&config_path).expect("读取文件失败"); + let parsed_config = ConfigManager::parse_yaml(&saved_content).expect("解析配置失败"); + + // 验证配置一致性 + prop_assert_eq!( + config.server, + parsed_config.server, + "服务器配置应一致" + ); + prop_assert_eq!( + config.providers, + parsed_config.providers, + "Provider 配置应一致" + ); + prop_assert_eq!( + config.retry, + parsed_config.retry, + "重试配置应一致" + ); + prop_assert_eq!( + config.logging, + parsed_config.logging, + "日志配置应一致" + ); + } + + /// **Feature: config-credential-export, Property 2: YAML Comment Preservation** + /// *For any* YAML file with header comments, saving should preserve header comments. + /// **Validates: Requirements 1.3** + #[test] + fn prop_yaml_header_comment_preservation( + header_comment in arb_yaml_comment(), + config in arb_valid_config() + ) { + // 创建临时文件 + let temp_dir = tempfile::tempdir().expect("创建临时目录失败"); + let config_path = temp_dir.path().join("config.yaml"); + + // 创建带头部注释的 YAML + let yaml = ConfigManager::to_yaml(&config).expect("序列化应成功"); + let yaml_with_header = format!("{header_comment}\n{yaml}"); + + // 写入文件 + std::fs::write(&config_path, &yaml_with_header).expect("写入文件失败"); + + // 使用 YamlService 保存新配置 + YamlService::save_preserve_comments(&config_path, &config) + .expect("保存配置失败"); + + // 读取保存后的内容 + let saved_content = std::fs::read_to_string(&config_path).expect("读取文件失败"); + + // 验证头部注释被保留 + let header_content = header_comment.trim(); + let has_header = saved_content.lines().any(|line| line.trim() == header_content); + + prop_assert!( + has_header, + "头部注释应被保留: 原始='{}', 保存后内容前100字符='{}'", + header_content, + &saved_content[..saved_content.len().min(100)] + ); + } +} + +// ============================================================================ +// Unit Tests for YamlService::update_field +// ============================================================================ + +#[test] +fn test_update_field_simple() { + // 创建临时文件 + let temp_dir = tempfile::tempdir().expect("创建临时目录失败"); + let config_path = temp_dir.path().join("config.yaml"); + + // 写入初始 YAML + let initial_yaml = r#"server: + host: 127.0.0.1 + port: 8999 + api_key: test_key +"#; + std::fs::write(&config_path, initial_yaml).expect("写入文件失败"); + + // 更新 port 字段 + YamlService::update_field(&config_path, &["server", "port"], "9000").expect("更新字段失败"); + + // 读取并验证 + let content = std::fs::read_to_string(&config_path).expect("读取文件失败"); + assert!(content.contains("port: 9000"), "端口应被更新为 9000"); + assert!(content.contains("host: 127.0.0.1"), "其他字段应保持不变"); +} + +#[test] +fn test_update_field_preserves_comments() { + // 创建临时文件 + let temp_dir = tempfile::tempdir().expect("创建临时目录失败"); + let config_path = temp_dir.path().join("config.yaml"); + + // 写入带注释的 YAML + let initial_yaml = r#"# 服务器配置 +server: + # 监听地址 + host: 127.0.0.1 + # 监听端口 + port: 8999 + api_key: test_key +"#; + std::fs::write(&config_path, initial_yaml).expect("写入文件失败"); + + // 更新 port 字段 + YamlService::update_field(&config_path, &["server", "port"], "9000").expect("更新字段失败"); + + // 读取并验证 + let content = std::fs::read_to_string(&config_path).expect("读取文件失败"); + assert!(content.contains("port: 9000"), "端口应被更新为 9000"); + assert!(content.contains("# 服务器配置"), "头部注释应保留"); + assert!(content.contains("# 监听地址"), "字段注释应保留"); + assert!(content.contains("# 监听端口"), "字段注释应保留"); +} + +#[test] +fn test_update_field_not_found() { + // 创建临时文件 + let temp_dir = tempfile::tempdir().expect("创建临时目录失败"); + let config_path = temp_dir.path().join("config.yaml"); + + // 写入初始 YAML + let initial_yaml = r#"server: + host: 127.0.0.1 + port: 8999 +"#; + std::fs::write(&config_path, initial_yaml).expect("写入文件失败"); + + // 尝试更新不存在的字段 + let result = YamlService::update_field(&config_path, &["server", "nonexistent"], "value"); + assert!(result.is_err(), "更新不存在的字段应返回错误"); +} + +#[test] +fn test_update_field_nested() { + // 创建临时文件 + let temp_dir = tempfile::tempdir().expect("创建临时目录失败"); + let config_path = temp_dir.path().join("config.yaml"); + + // 写入初始 YAML + let initial_yaml = r#"server: + host: 127.0.0.1 + port: 8999 +providers: + kiro: + enabled: true + region: us-east-1 +"#; + std::fs::write(&config_path, initial_yaml).expect("写入文件失败"); + + // 更新嵌套字段 + YamlService::update_field(&config_path, &["providers", "kiro", "region"], "us-west-2") + .expect("更新字段失败"); + + // 读取并验证 + let content = std::fs::read_to_string(&config_path).expect("读取文件失败"); + assert!( + content.contains("region: us-west-2"), + "region 应被更新为 us-west-2" + ); + assert!(content.contains("enabled: true"), "其他字段应保持不变"); +} + +// ============================================================================ +// Property 4: Export Scope Filtering +// ============================================================================ + +use crate::config::{ + ApiKeyEntry, CredentialEntry, CredentialPoolConfig, ExportOptions, ExportService, +}; + +/// 生成随机的 OAuth 凭证条目 +fn arb_credential_entry() -> impl Strategy { + ( + "[a-z]{3,10}-[0-9]{1,5}".prop_map(|s| s), + "[a-z]+/token-[0-9]{1,5}\\.json".prop_map(|s| s), + any::(), + proptest::option::of("socks5://proxy\\.[a-z]+\\.com:[0-9]{4}".prop_map(|s| s)), + ) + .prop_map(|(id, token_file, disabled, proxy_url)| CredentialEntry { + id, + token_file, + disabled, + proxy_url, + }) +} + +/// 生成随机的 API Key 凭证条目 +fn arb_api_key_entry() -> impl Strategy { + ( + "[a-z]{3,10}-[0-9]{1,5}".prop_map(|s| s), + "sk-[a-zA-Z0-9]{20,40}".prop_map(|s| s), + proptest::option::of("https://api\\.[a-z]+\\.com/v[0-9]".prop_map(|s| s)), + any::(), + proptest::option::of("http://proxy\\.[a-z]+\\.com:[0-9]{4}".prop_map(|s| s)), + ) + .prop_map(|(id, api_key, base_url, disabled, proxy_url)| ApiKeyEntry { + id, + api_key, + base_url, + disabled, + proxy_url, + }) +} + +/// 生成随机的凭证池配置 +fn arb_credential_pool_config() -> impl Strategy { + ( + proptest::collection::vec(arb_credential_entry(), 0..3), + proptest::collection::vec(arb_credential_entry(), 0..3), + proptest::collection::vec(arb_credential_entry(), 0..3), + proptest::collection::vec(arb_api_key_entry(), 0..3), + proptest::collection::vec(arb_api_key_entry(), 0..3), + ) + .prop_map( + |(kiro, gemini, qwen, openai, claude)| CredentialPoolConfig { + kiro, + gemini, + qwen, + openai, + claude, + gemini_api_keys: vec![], + vertex_api_keys: vec![], + codex: vec![], + asr: vec![], + }, + ) +} + +/// 生成带凭证池的配置 +fn arb_config_with_credentials() -> impl Strategy { + (arb_valid_config(), arb_credential_pool_config()).prop_map(|(mut config, pool)| { + config.credential_pool = pool; + config + }) +} + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: config-credential-export, Property 4: Export Scope Filtering** + /// *For any* export operation with config-only scope, the resulting bundle + /// should contain only configuration data and no credential token files. + /// **Validates: Requirements 3.2** + #[test] + fn prop_export_scope_config_only(config in arb_config_with_credentials()) { + let options = ExportOptions::config_only(); + let bundle = ExportService::export(&config, &options, "1.0.0") + .expect("导出应成功"); + + // 验证只包含配置 + prop_assert!( + bundle.has_config(), + "config-only 导出应包含配置" + ); + prop_assert!( + !bundle.has_credentials(), + "config-only 导出不应包含凭证 token 文件" + ); + } + + /// **Feature: config-credential-export, Property 4: Export Scope Filtering** + /// *For any* export operation with credentials-only scope, the resulting bundle + /// should contain only credential data and no configuration YAML. + /// **Validates: Requirements 3.2** + #[test] + fn prop_export_scope_credentials_only(config in arb_config_with_credentials()) { + let options = ExportOptions::credentials_only(); + let bundle = ExportService::export(&config, &options, "1.0.0") + .expect("导出应成功"); + + // 验证不包含配置 + prop_assert!( + !bundle.has_config(), + "credentials-only 导出不应包含配置 YAML" + ); + // 注意:token_files 可能为空(如果没有实际的 token 文件存在) + // 但 config_yaml 必须为 None + prop_assert!( + bundle.config_yaml.is_none(), + "credentials-only 导出的 config_yaml 应为 None" + ); + } + + /// **Feature: config-credential-export, Property 4: Export Scope Filtering** + /// *For any* export operation with full scope, the resulting bundle + /// should contain both configuration and credential data. + /// **Validates: Requirements 3.2** + #[test] + fn prop_export_scope_full(config in arb_config_with_credentials()) { + let options = ExportOptions::full(); + let bundle = ExportService::export(&config, &options, "1.0.0") + .expect("导出应成功"); + + // 验证包含配置 + prop_assert!( + bundle.has_config(), + "full 导出应包含配置" + ); + // token_files 可能为空(如果没有实际的 token 文件存在) + // 但 config_yaml 必须存在 + prop_assert!( + bundle.config_yaml.is_some(), + "full 导出的 config_yaml 应存在" + ); + } + + /// **Feature: config-credential-export, Property 4: Export Scope Filtering** + /// *For any* export operation, the bundle should correctly reflect the + /// include_config and include_credentials options. + /// **Validates: Requirements 3.2** + #[test] + fn prop_export_scope_matches_options( + config in arb_config_with_credentials(), + include_config in any::(), + include_credentials in any::() + ) { + let options = ExportOptions { + include_config, + include_credentials, + redact_secrets: false, + }; + + let bundle = ExportService::export(&config, &options, "1.0.0") + .expect("导出应成功"); + + // 验证配置包含状态与选项一致 + prop_assert_eq!( + bundle.has_config(), + include_config, + "配置包含状态应与 include_config 选项一致" + ); + + // 验证 config_yaml 存在性与选项一致 + prop_assert_eq!( + bundle.config_yaml.is_some(), + include_config, + "config_yaml 存在性应与 include_config 选项一致" + ); + } +} + +// ============================================================================ +// Property 5: Redaction Completeness +// ============================================================================ + +use crate::config::REDACTED_PLACEHOLDER; + +/// 生成包含敏感信息的配置 +fn arb_config_with_secrets() -> impl Strategy { + ( + arb_valid_config(), + arb_credential_pool_config(), + // 生成看起来像真实 API 密钥的字符串 + "sk-[a-zA-Z0-9]{20,40}".prop_map(|s| s), + proptest::option::of("sk-[a-zA-Z0-9]{20,40}".prop_map(|s| s)), + proptest::option::of("sk-ant-[a-zA-Z0-9]{20,40}".prop_map(|s| s)), + ) + .prop_map(|(mut config, pool, server_key, openai_key, claude_key)| { + config.server.api_key = server_key; + config.providers.openai.api_key = openai_key; + config.providers.claude.api_key = claude_key; + config.credential_pool = pool; + config + }) +} + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: config-credential-export, Property 5: Redaction Completeness** + /// *For any* export with redaction enabled, all sensitive values (API keys, tokens, + /// secrets) should be replaced with placeholder markers, and no original sensitive + /// data should remain. + /// **Validates: Requirements 3.4** + #[test] + fn prop_redaction_removes_all_secrets(config in arb_config_with_secrets()) { + // 脱敏配置 + let redacted = ExportService::redact_config(&config); + + // 验证脱敏后不包含敏感信息 + prop_assert!( + !ExportService::contains_secrets(&redacted), + "脱敏后的配置不应包含敏感信息" + ); + + // 验证服务器 API 密钥已脱敏 + prop_assert_eq!( + &redacted.server.api_key, + REDACTED_PLACEHOLDER, + "服务器 API 密钥应被脱敏" + ); + + // 验证 OpenAI API 密钥已脱敏(如果存在) + if config.providers.openai.api_key.is_some() { + prop_assert_eq!( + redacted.providers.openai.api_key.as_deref(), + Some(REDACTED_PLACEHOLDER), + "OpenAI API 密钥应被脱敏" + ); + } + + // 验证 Claude API 密钥已脱敏(如果存在) + if config.providers.claude.api_key.is_some() { + prop_assert_eq!( + redacted.providers.claude.api_key.as_deref(), + Some(REDACTED_PLACEHOLDER), + "Claude API 密钥应被脱敏" + ); + } + } + + /// **Feature: config-credential-export, Property 5: Redaction Completeness** + /// *For any* configuration with API keys in credential pool, redaction should + /// replace all API keys with placeholder markers. + /// **Validates: Requirements 3.4** + #[test] + fn prop_redaction_credential_pool_api_keys(config in arb_config_with_secrets()) { + let redacted = ExportService::redact_config(&config); + + // 验证 OpenAI 凭证池中的 API 密钥已脱敏 + for (i, entry) in redacted.credential_pool.openai.iter().enumerate() { + prop_assert_eq!( + &entry.api_key, + REDACTED_PLACEHOLDER, + "OpenAI 凭证池条目 {} 的 API 密钥应被脱敏", + i + ); + } + + // 验证 Claude 凭证池中的 API 密钥已脱敏 + for (i, entry) in redacted.credential_pool.claude.iter().enumerate() { + prop_assert_eq!( + &entry.api_key, + REDACTED_PLACEHOLDER, + "Claude 凭证池条目 {} 的 API 密钥应被脱敏", + i + ); + } + } + + /// **Feature: config-credential-export, Property 5: Redaction Completeness** + /// *For any* export with redaction enabled, the exported YAML should not contain + /// any original sensitive values. + /// **Validates: Requirements 3.4** + #[test] + fn prop_redaction_yaml_no_secrets(config in arb_config_with_secrets()) { + // 导出带脱敏的 YAML + let yaml = ExportService::export_yaml(&config, true) + .expect("导出应成功"); + + // 验证 YAML 中不包含原始敏感值 + // 检查原始服务器 API 密钥 + if !config.server.api_key.is_empty() && config.server.api_key != REDACTED_PLACEHOLDER { + prop_assert!( + !yaml.contains(&config.server.api_key), + "YAML 不应包含原始服务器 API 密钥: {}", + config.server.api_key + ); + } + + // 检查原始 OpenAI API 密钥 + if let Some(ref key) = config.providers.openai.api_key { + if !key.is_empty() && key != REDACTED_PLACEHOLDER { + prop_assert!( + !yaml.contains(key), + "YAML 不应包含原始 OpenAI API 密钥" + ); + } + } + + // 检查原始 Claude API 密钥 + if let Some(ref key) = config.providers.claude.api_key { + if !key.is_empty() && key != REDACTED_PLACEHOLDER { + prop_assert!( + !yaml.contains(key), + "YAML 不应包含原始 Claude API 密钥" + ); + } + } + + // 检查凭证池中的 API 密钥 + for entry in &config.credential_pool.openai { + if !entry.api_key.is_empty() && entry.api_key != REDACTED_PLACEHOLDER { + prop_assert!( + !yaml.contains(&entry.api_key), + "YAML 不应包含原始 OpenAI 凭证池 API 密钥" + ); + } + } + + for entry in &config.credential_pool.claude { + if !entry.api_key.is_empty() && entry.api_key != REDACTED_PLACEHOLDER { + prop_assert!( + !yaml.contains(&entry.api_key), + "YAML 不应包含原始 Claude 凭证池 API 密钥" + ); + } + } + } + + /// **Feature: config-credential-export, Property 5: Redaction Completeness** + /// *For any* configuration, redaction should preserve non-sensitive data unchanged. + /// **Validates: Requirements 3.4** + #[test] + fn prop_redaction_preserves_non_sensitive_data(config in arb_config_with_secrets()) { + let redacted = ExportService::redact_config(&config); + + // 验证非敏感数据保持不变 + prop_assert_eq!( + config.server.host, + redacted.server.host, + "服务器主机应保持不变" + ); + prop_assert_eq!( + config.server.port, + redacted.server.port, + "服务器端口应保持不变" + ); + prop_assert_eq!( + config.providers.kiro.enabled, + redacted.providers.kiro.enabled, + "Kiro 启用状态应保持不变" + ); + prop_assert_eq!( + config.routing.default_provider, + redacted.routing.default_provider, + "默认 Provider 应保持不变" + ); + prop_assert_eq!( + config.retry, + redacted.retry, + "重试配置应保持不变" + ); + prop_assert_eq!( + config.logging, + redacted.logging, + "日志配置应保持不变" + ); + + // 验证 OAuth 凭证条目保持不变(它们不包含敏感信息) + prop_assert_eq!( + config.credential_pool.kiro, + redacted.credential_pool.kiro, + "Kiro 凭证条目应保持不变" + ); + prop_assert_eq!( + config.credential_pool.gemini, + redacted.credential_pool.gemini, + "Gemini 凭证条目应保持不变" + ); + prop_assert_eq!( + config.credential_pool.qwen, + redacted.credential_pool.qwen, + "Qwen 凭证条目应保持不变" + ); + } + + /// **Feature: config-credential-export, Property 5: Redaction Completeness** + /// *For any* export bundle with redaction, the redacted flag should be true. + /// **Validates: Requirements 3.4** + #[test] + fn prop_redaction_bundle_flag(config in arb_config_with_secrets()) { + let options = ExportOptions::redacted(); + let bundle = ExportService::export(&config, &options, "1.0.0") + .expect("导出应成功"); + + prop_assert!( + bundle.is_redacted(), + "脱敏导出的 bundle 应标记为已脱敏" + ); + } +} + +// ============================================================================ +// Property 6: Import Validation +// ============================================================================ + +use crate::config::{ExportBundle, ImportService}; + +/// 生成有效的导出包 +fn arb_valid_export_bundle() -> impl Strategy { + ( + arb_valid_config(), + any::(), // redacted + "[0-9]+\\.[0-9]+\\.[0-9]+".prop_map(|s| s), // app_version + ) + .prop_map(|(config, redacted, app_version)| { + let yaml = ConfigManager::to_yaml(&config).expect("序列化应成功"); + let mut bundle = ExportBundle::new(&app_version); + bundle.config_yaml = Some(yaml); + bundle.redacted = redacted; + bundle + }) +} + +/// 生成无效的导入内容(既不是有效的 JSON ExportBundle,也不是有效的 YAML Config) +/// 注意:YAML 解析器非常宽松,大多数内容都可以解析为某种 YAML 结构 +/// 因此我们只测试语法错误的内容 +fn arb_invalid_import_content() -> impl Strategy { + prop_oneof![ + // 无效的 JSON/YAML 语法 + Just("{invalid json".to_string()), + Just("invalid: yaml: content: [".to_string()), + Just(" - bad\n indentation".to_string()), + Just("key: value\n invalid: indent".to_string()), + // 有效的 JSON 但不是 ExportBundle 或 Config(数组类型) + Just("[1, 2, 3]".to_string()), + ] +} + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: config-credential-export, Property 6: Import Validation** + /// *For any* valid export bundle, validation should return valid=true and + /// correctly identify format, version, and redaction status. + /// **Validates: Requirements 4.1, 4.2** + #[test] + fn prop_import_validation_valid_bundle(bundle in arb_valid_export_bundle()) { + let json = bundle.to_json().expect("序列化应成功"); + let result = ImportService::validate(&json); + + // 验证结果应为有效 + prop_assert!( + result.valid, + "有效的导出包应通过验证: errors={:?}", + result.errors + ); + + // 验证版本被正确识别 + prop_assert_eq!( + result.version, + Some(bundle.version.clone()), + "版本应被正确识别" + ); + + // 验证脱敏状态被正确识别 + prop_assert_eq!( + result.redacted, + bundle.redacted, + "脱敏状态应被正确识别" + ); + + // 验证配置存在性被正确识别 + prop_assert_eq!( + result.has_config, + bundle.has_config(), + "配置存在性应被正确识别" + ); + } + + /// **Feature: config-credential-export, Property 6: Import Validation** + /// *For any* valid YAML configuration, validation should return valid=true + /// and identify it as config-only (no credentials). + /// **Validates: Requirements 4.1, 4.2** + #[test] + fn prop_import_validation_valid_yaml(config in arb_valid_config()) { + let yaml = ConfigManager::to_yaml(&config).expect("序列化应成功"); + let result = ImportService::validate(&yaml); + + // 验证结果应为有效 + prop_assert!( + result.valid, + "有效的 YAML 配置应通过验证: errors={:?}", + result.errors + ); + + // 验证识别为配置 + prop_assert!( + result.has_config, + "应识别为包含配置" + ); + + // YAML 配置不包含凭证 token 文件 + prop_assert!( + !result.has_credentials, + "YAML 配置不应包含凭证 token 文件" + ); + + // YAML 配置不是脱敏的 + prop_assert!( + !result.redacted, + "YAML 配置不应标记为脱敏" + ); + } + + /// **Feature: config-credential-export, Property 6: Import Validation** + /// *For any* invalid import content (neither valid ExportBundle JSON nor valid Config YAML), + /// validation should return valid=false with appropriate error messages. + /// **Validates: Requirements 4.1, 4.2** + #[test] + fn prop_import_validation_invalid_content(content in arb_invalid_import_content()) { + let result = ImportService::validate(&content); + + // 无效内容应验证失败 + prop_assert!( + !result.valid, + "无效的导入内容应验证失败: content={}", content + ); + + // 应有错误信息 + prop_assert!( + !result.errors.is_empty(), + "应有错误信息" + ); + } + + /// **Feature: config-credential-export, Property 6: Import Validation** + /// *For any* redacted export bundle, validation should warn about + /// credentials that cannot be restored. + /// **Validates: Requirements 4.1, 4.2** + #[test] + fn prop_import_validation_redacted_warning(config in arb_valid_config()) { + let yaml = ConfigManager::to_yaml(&config).expect("序列化应成功"); + let mut bundle = ExportBundle::new("1.0.0"); + bundle.config_yaml = Some(yaml); + bundle.redacted = true; + + let json = bundle.to_json().expect("序列化应成功"); + let result = ImportService::validate(&json); + + // 验证结果应为有效(脱敏不影响有效性) + prop_assert!( + result.valid, + "脱敏的导出包应通过验证" + ); + + // 应有脱敏警告 + prop_assert!( + !result.warnings.is_empty(), + "脱敏的导出包应有警告信息" + ); + + // 警告应提及脱敏 + let has_redaction_warning = result.warnings.iter().any(|w| + w.contains("脱敏") || w.contains("redact") + ); + prop_assert!( + has_redaction_warning, + "应有关于脱敏的警告: {:?}", + result.warnings + ); + } +} + +// ============================================================================ +// Property 7: Import Merge vs Replace +// ============================================================================ + +use crate::config::ImportOptions; + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: config-credential-export, Property 7: Import Merge vs Replace** + /// *For any* import operation in replace mode, the resulting configuration + /// should be exactly the imported configuration (not merged with current). + /// **Validates: Requirements 4.3** + #[test] + fn prop_import_replace_mode( + current_config in arb_config_with_credentials(), + imported_config in arb_config_with_credentials() + ) { + let yaml = ConfigManager::to_yaml(&imported_config).expect("序列化应成功"); + let options = ImportOptions::replace(); + + let result = ImportService::import_yaml(&yaml, ¤t_config, &options) + .expect("导入应成功"); + + // 替换模式下,结果应等于导入的配置 + prop_assert_eq!( + result.config.server, + imported_config.server, + "替换模式下服务器配置应等于导入的配置" + ); + prop_assert_eq!( + result.config.providers, + imported_config.providers, + "替换模式下 Provider 配置应等于导入的配置" + ); + prop_assert_eq!( + result.config.routing.default_provider, + imported_config.routing.default_provider, + "替换模式下默认 Provider 应等于导入的配置" + ); + prop_assert_eq!( + result.config.retry, + imported_config.retry, + "替换模式下重试配置应等于导入的配置" + ); + prop_assert_eq!( + result.config.logging, + imported_config.logging, + "替换模式下日志配置应等于导入的配置" + ); + } + + /// **Feature: config-credential-export, Property 7: Import Merge vs Replace** + /// *For any* import operation in merge mode, the resulting configuration + /// should combine new data with existing data. + /// **Validates: Requirements 4.3** + #[test] + fn prop_import_merge_mode_combines_credentials( + current_config in arb_config_with_credentials(), + imported_config in arb_config_with_credentials() + ) { + let yaml = ConfigManager::to_yaml(&imported_config).expect("序列化应成功"); + let options = ImportOptions::merge(); + + let result = ImportService::import_yaml(&yaml, ¤t_config, &options) + .expect("导入应成功"); + + // 合并模式下,凭证池应包含两边的凭证(按 ID 去重) + // 计算预期的凭证数量(去重后) + let expected_kiro_ids: std::collections::HashSet<_> = current_config + .credential_pool + .kiro + .iter() + .chain(imported_config.credential_pool.kiro.iter()) + .map(|e| e.id.clone()) + .collect(); + + prop_assert_eq!( + result.config.credential_pool.kiro.len(), + expected_kiro_ids.len(), + "合并模式下 Kiro 凭证数量应为去重后的总数" + ); + + let expected_openai_ids: std::collections::HashSet<_> = current_config + .credential_pool + .openai + .iter() + .chain(imported_config.credential_pool.openai.iter()) + .filter(|e| e.api_key != REDACTED_PLACEHOLDER) + .map(|e| e.id.clone()) + .collect(); + + // OpenAI 凭证数量应包含所有非脱敏的凭证 + prop_assert!( + result.config.credential_pool.openai.len() >= expected_openai_ids.len().saturating_sub( + imported_config.credential_pool.openai.iter() + .filter(|e| e.api_key == REDACTED_PLACEHOLDER) + .count() + ), + "合并模式下 OpenAI 凭证应包含所有非脱敏的凭证" + ); + } + + /// **Feature: config-credential-export, Property 7: Import Merge vs Replace** + /// *For any* import operation in merge mode, imported values should override + /// current values for the same keys. + /// **Validates: Requirements 4.3** + #[test] + fn prop_import_merge_mode_overrides_config( + current_config in arb_config_with_credentials(), + imported_config in arb_config_with_credentials() + ) { + let yaml = ConfigManager::to_yaml(&imported_config).expect("序列化应成功"); + let options = ImportOptions::merge(); + + let result = ImportService::import_yaml(&yaml, ¤t_config, &options) + .expect("导入应成功"); + + // 合并模式下,配置值应被导入的值覆盖 + prop_assert_eq!( + result.config.server, + imported_config.server, + "合并模式下服务器配置应被导入的值覆盖" + ); + prop_assert_eq!( + result.config.providers, + imported_config.providers, + "合并模式下 Provider 配置应被导入的值覆盖" + ); + prop_assert_eq!( + result.config.retry, + imported_config.retry, + "合并模式下重试配置应被导入的值覆盖" + ); + } + + /// **Feature: config-credential-export, Property 7: Import Merge vs Replace** + /// *For any* import operation in replace mode with empty imported credentials, + /// the result should have empty credentials (not preserve current). + /// **Validates: Requirements 4.3** + #[test] + fn prop_import_replace_mode_clears_credentials( + current_config in arb_config_with_credentials() + ) { + // 创建一个没有凭证的配置 + let mut imported_config = Config::default(); + imported_config.server.port = 9999; // 修改一个值以区分 + + let yaml = ConfigManager::to_yaml(&imported_config).expect("序列化应成功"); + let options = ImportOptions::replace(); + + let result = ImportService::import_yaml(&yaml, ¤t_config, &options) + .expect("导入应成功"); + + // 替换模式下,凭证池应为空(因为导入的配置没有凭证) + prop_assert!( + result.config.credential_pool.kiro.is_empty(), + "替换模式下 Kiro 凭证应为空" + ); + prop_assert!( + result.config.credential_pool.openai.is_empty(), + "替换模式下 OpenAI 凭证应为空" + ); + prop_assert_eq!( + result.config.server.port, + 9999, + "替换模式下端口应为导入的值" + ); + } +} + +// ============================================================================ +// Property 8: Export-Import Round Trip +// ============================================================================ + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: config-credential-export, Property 8: Export-Import Round Trip** + /// *For any* valid configuration, exporting (without redaction) and then importing + /// should produce an equivalent configuration. + /// **Validates: Requirements 5.5** + #[test] + fn prop_export_import_roundtrip_yaml(config in arb_config_with_credentials()) { + // 导出为 YAML(不脱敏) + let yaml = ExportService::export_yaml(&config, false) + .expect("导出应成功"); + + // 导入 YAML(替换模式) + let empty_config = Config::default(); + let options = ImportOptions::replace(); + let result = ImportService::import_yaml(&yaml, &empty_config, &options) + .expect("导入应成功"); + + // 验证往返一致性 + prop_assert_eq!( + config.server, + result.config.server, + "服务器配置往返不一致" + ); + prop_assert_eq!( + config.providers, + result.config.providers, + "Provider 配置往返不一致" + ); + prop_assert_eq!( + config.routing.default_provider, + result.config.routing.default_provider, + "默认 Provider 往返不一致" + ); + prop_assert_eq!( + config.retry, + result.config.retry, + "重试配置往返不一致" + ); + prop_assert_eq!( + config.logging, + result.config.logging, + "日志配置往返不一致" + ); + prop_assert_eq!( + config.auth_dir, + result.config.auth_dir, + "auth_dir 往返不一致" + ); + + // 验证凭证池往返一致性 + prop_assert_eq!( + config.credential_pool.kiro, + result.config.credential_pool.kiro, + "Kiro 凭证池往返不一致" + ); + prop_assert_eq!( + config.credential_pool.gemini, + result.config.credential_pool.gemini, + "Gemini 凭证池往返不一致" + ); + prop_assert_eq!( + config.credential_pool.qwen, + result.config.credential_pool.qwen, + "Qwen 凭证池往返不一致" + ); + prop_assert_eq!( + config.credential_pool.openai, + result.config.credential_pool.openai, + "OpenAI 凭证池往返不一致" + ); + prop_assert_eq!( + config.credential_pool.claude, + result.config.credential_pool.claude, + "Claude 凭证池往返不一致" + ); + } + + /// **Feature: config-credential-export, Property 8: Export-Import Round Trip** + /// *For any* valid configuration, exporting as a bundle (without redaction) and + /// then importing should produce an equivalent configuration. + /// **Validates: Requirements 5.5** + #[test] + fn prop_export_import_roundtrip_bundle(config in arb_config_with_credentials()) { + // 导出为 bundle(不脱敏,仅配置) + let options = ExportOptions { + include_config: true, + include_credentials: false, // 不包含 token 文件,因为测试环境没有实际文件 + redact_secrets: false, + }; + let bundle = ExportService::export(&config, &options, "1.0.0") + .expect("导出应成功"); + + // 序列化为 JSON + let json = bundle.to_json().expect("序列化应成功"); + + // 反序列化 + let parsed_bundle = ExportBundle::from_json(&json).expect("反序列化应成功"); + + // 导入 bundle + let empty_config = Config::default(); + let import_options = ImportOptions::replace(); + let result = ImportService::import( + &parsed_bundle, + &empty_config, + &import_options, + &config.auth_dir, + ) + .expect("导入应成功"); + + // 验证往返一致性 + prop_assert_eq!( + config.server, + result.config.server, + "服务器配置往返不一致" + ); + prop_assert_eq!( + config.providers, + result.config.providers, + "Provider 配置往返不一致" + ); + prop_assert_eq!( + config.retry, + result.config.retry, + "重试配置往返不一致" + ); + prop_assert_eq!( + config.logging, + result.config.logging, + "日志配置往返不一致" + ); + } + + /// **Feature: config-credential-export, Property 8: Export-Import Round Trip** + /// *For any* configuration with API keys, exporting with redaction and then + /// importing should NOT restore the original API keys. + /// **Validates: Requirements 5.5** + #[test] + fn prop_export_import_redacted_loses_secrets(config in arb_config_with_secrets()) { + // 导出为脱敏 bundle + let options = ExportOptions::redacted(); + let bundle = ExportService::export(&config, &options, "1.0.0") + .expect("导出应成功"); + + // 导入 bundle(脱敏数据会触发清理) + let empty_config = Config::default(); + let import_options = ImportOptions::replace(); + let result = ImportService::import( + &bundle, + &empty_config, + &import_options, + &config.auth_dir, + ) + .expect("导入应成功"); + + // 验证脱敏后的配置不包含原始敏感信息 + // 服务器 API 密钥应被清空 + prop_assert_eq!( + result.config.server.api_key, + "", + "脱敏后服务器 API 密钥应被清空" + ); + + // 如果原始配置有 OpenAI API 密钥,导入后应为脱敏占位符 + if config.providers.openai.api_key.is_some() { + prop_assert_eq!( + result.config.providers.openai.api_key, + None, + "脱敏后 OpenAI API 密钥应被清空" + ); + } + } + + /// **Feature: config-credential-export, Property 8: Export-Import Round Trip** + /// *For any* configuration, the export bundle should be valid JSON that can + /// be parsed back. + /// **Validates: Requirements 5.5** + #[test] + fn prop_export_bundle_json_roundtrip(config in arb_config_with_credentials()) { + let options = ExportOptions { + include_config: true, + include_credentials: false, + redact_secrets: false, + }; + let bundle = ExportService::export(&config, &options, "1.0.0") + .expect("导出应成功"); + + // 序列化为 JSON + let json = bundle.to_json().expect("序列化应成功"); + + // 反序列化 + let parsed = ExportBundle::from_json(&json).expect("反序列化应成功"); + + // 验证往返一致性 + prop_assert_eq!( + bundle.version, + parsed.version, + "版本往返不一致" + ); + prop_assert_eq!( + bundle.app_version, + parsed.app_version, + "应用版本往返不一致" + ); + prop_assert_eq!( + bundle.redacted, + parsed.redacted, + "脱敏状态往返不一致" + ); + prop_assert_eq!( + bundle.config_yaml, + parsed.config_yaml, + "配置 YAML 往返不一致" + ); + prop_assert_eq!( + bundle.token_files, + parsed.token_files, + "Token 文件往返不一致" + ); + } +} + +// ============================================================================ +// Property 1: OAuth Token Storage Round-Trip (CLIProxyAPI Parity) +// ============================================================================ + +/// 生成随机的 OAuth 凭证条目(用于 Codex/iFlow) +fn arb_oauth_credential_entry() -> impl Strategy { + ( + "[a-z]{3,10}-[0-9]{1,5}".prop_map(|s| s), + "[a-z]+/oauth-token-[0-9]{1,5}\\.json".prop_map(|s| s), + any::(), + proptest::option::of("socks5://proxy\\.[a-z]+\\.com:[0-9]{4}".prop_map(|s| s)), + ) + .prop_map(|(id, token_file, disabled, proxy_url)| CredentialEntry { + id, + token_file, + disabled, + proxy_url, + }) +} + +/// 生成随机的 Gemini API Key 条目 +fn arb_gemini_api_key_entry() -> impl Strategy { + ( + "[a-z]{3,10}-[0-9]{1,5}".prop_map(|s| s), + "AIzaSy[a-zA-Z0-9_-]{33}".prop_map(|s| s), + proptest::option::of("https://generativelanguage\\.googleapis\\.com".prop_map(|s| s)), + proptest::option::of("http://proxy\\.[a-z]+\\.com:[0-9]{4}".prop_map(|s| s)), + proptest::collection::vec("[a-z]+-[0-9]+\\.[0-9]+-pro".prop_map(|s| s), 0..3), + any::(), + ) + .prop_map( + |(id, api_key, base_url, proxy_url, excluded_models, disabled)| { + crate::config::GeminiApiKeyEntry { + id, + api_key, + base_url, + proxy_url, + excluded_models, + disabled, + } + }, + ) +} + +/// 生成随机的 Vertex AI 条目 +fn arb_vertex_api_key_entry() -> impl Strategy { + ( + "[a-z]{3,10}-[0-9]{1,5}".prop_map(|s| s), + "vk-[a-zA-Z0-9]{20,40}".prop_map(|s| s), + proptest::option::of("https://[a-z]+-aiplatform\\.googleapis\\.com".prop_map(|s| s)), + proptest::collection::vec( + ( + "[a-z]+-[0-9]+\\.[0-9]+".prop_map(|s| s), + "[a-z]+-alias".prop_map(|s| s), + ), + 0..3, + ), + proptest::option::of("http://proxy\\.[a-z]+\\.com:[0-9]{4}".prop_map(|s| s)), + any::(), + ) + .prop_map(|(id, api_key, base_url, models, proxy_url, disabled)| { + crate::config::VertexApiKeyEntry { + id, + api_key, + base_url, + models: models + .into_iter() + .map(|(name, alias)| crate::config::VertexModelAlias { name, alias }) + .collect(), + proxy_url, + disabled, + } + }) +} + +/// 生成包含新 Provider 凭证的凭证池配置 +fn arb_extended_credential_pool_config() -> impl Strategy { + ( + proptest::collection::vec(arb_credential_entry(), 0..3), + proptest::collection::vec(arb_credential_entry(), 0..3), + proptest::collection::vec(arb_credential_entry(), 0..3), + proptest::collection::vec(arb_api_key_entry(), 0..3), + proptest::collection::vec(arb_api_key_entry(), 0..3), + proptest::collection::vec(arb_gemini_api_key_entry(), 0..3), + proptest::collection::vec(arb_vertex_api_key_entry(), 0..3), + proptest::collection::vec(arb_credential_entry(), 0..3), + ) + .prop_map( + |(kiro, gemini, qwen, openai, claude, gemini_api_keys, vertex_api_keys, codex)| { + CredentialPoolConfig { + kiro, + gemini, + qwen, + openai, + claude, + gemini_api_keys, + vertex_api_keys, + codex, + asr: vec![], + } + }, + ) +} + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: cliproxyapi-parity, Property 1: OAuth Token Storage Round-Trip** + /// *For any* valid OAuth response containing access_token, refresh_token, and expires_at, + /// storing and then loading the credentials SHALL produce equivalent values. + /// **Validates: Requirements 1.1, 2.1** + #[test] + fn prop_oauth_token_storage_roundtrip(pool in arb_extended_credential_pool_config()) { + let config = Config { + credential_pool: pool.clone(), + ..Config::default() + }; + + // 序列化为 YAML + let yaml = ConfigManager::to_yaml(&config) + .expect("序列化应成功"); + + // 反序列化回 Config + let parsed = ConfigManager::parse_yaml(&yaml) + .expect("反序列化应成功"); + + // 验证 OAuth 凭证往返一致性 + prop_assert_eq!( + pool.kiro.len(), + parsed.credential_pool.kiro.len(), + "Kiro OAuth 凭证数量往返不一致" + ); + prop_assert_eq!( + pool.gemini.len(), + parsed.credential_pool.gemini.len(), + "Gemini OAuth 凭证数量往返不一致" + ); + + // 验证 Gemini API Key 多账号配置往返一致性 + prop_assert_eq!( + pool.gemini_api_keys.len(), + parsed.credential_pool.gemini_api_keys.len(), + "Gemini API Key 凭证数量往返不一致" + ); + + // 验证 Vertex AI 配置往返一致性 + prop_assert_eq!( + pool.vertex_api_keys.len(), + parsed.credential_pool.vertex_api_keys.len(), + "Vertex AI 凭证数量往返不一致" + ); + + // 验证每个 Gemini API Key 的详细内容 + for (original, parsed_entry) in pool.gemini_api_keys.iter().zip(parsed.credential_pool.gemini_api_keys.iter()) { + prop_assert_eq!( + &original.id, + &parsed_entry.id, + "Gemini API Key ID 往返不一致" + ); + prop_assert_eq!( + &original.api_key, + &parsed_entry.api_key, + "Gemini API Key 往返不一致" + ); + prop_assert_eq!( + &original.excluded_models, + &parsed_entry.excluded_models, + "Gemini 排除模型列表往返不一致" + ); + } + + // 验证每个 Vertex AI 凭证的详细内容 + for (original, parsed_entry) in pool.vertex_api_keys.iter().zip(parsed.credential_pool.vertex_api_keys.iter()) { + prop_assert_eq!( + &original.id, + &parsed_entry.id, + "Vertex AI 凭证 ID 往返不一致" + ); + prop_assert_eq!( + original.models.len(), + parsed_entry.models.len(), + "Vertex AI 模型别名数量往返不一致" + ); + } + } +} + +// ============================================================================ +// Property 3: EndpointProvidersConfig 序列化往返一致性 +// ============================================================================ + +use crate::config::EndpointProvidersConfig; + +/// 生成随机的 Provider 名称 +fn arb_provider_name() -> impl Strategy { + prop_oneof![ + Just("kiro".to_string()), + Just("gemini".to_string()), + Just("qwen".to_string()), + Just("openai".to_string()), + Just("claude".to_string()), + Just("codex".to_string()), + ] +} + +/// 生成随机的可选 Provider 名称 +fn arb_optional_provider() -> impl Strategy> { + proptest::option::of(arb_provider_name()) +} + +/// 生成随机的 EndpointProvidersConfig +fn arb_endpoint_providers_config() -> impl Strategy { + ( + arb_optional_provider(), + arb_optional_provider(), + arb_optional_provider(), + arb_optional_provider(), + arb_optional_provider(), + arb_optional_provider(), + ) + .prop_map(|(cursor, claude_code, codex, windsurf, kiro, other)| { + EndpointProvidersConfig { + cursor, + claude_code, + codex, + windsurf, + kiro, + other, + } + }) +} + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: endpoint-provider-config, Property 3: 配置序列化往返一致性** + /// *对于任意* 有效的 EndpointProvidersConfig 对象,序列化后再反序列化应产生等价的对象。 + /// **Validates: Requirements 1.1** + #[test] + fn prop_endpoint_providers_config_yaml_roundtrip(config in arb_endpoint_providers_config()) { + // 序列化为 YAML + let yaml = serde_yaml::to_string(&config) + .expect("YAML 序列化应成功"); + + // 反序列化回 EndpointProvidersConfig + let parsed: EndpointProvidersConfig = serde_yaml::from_str(&yaml) + .expect("YAML 反序列化应成功"); + + // 验证往返一致性 + prop_assert_eq!( + config.cursor, + parsed.cursor, + "cursor 字段往返不一致" + ); + prop_assert_eq!( + config.claude_code, + parsed.claude_code, + "claude_code 字段往返不一致" + ); + prop_assert_eq!( + config.codex, + parsed.codex, + "codex 字段往返不一致" + ); + prop_assert_eq!( + config.windsurf, + parsed.windsurf, + "windsurf 字段往返不一致" + ); + prop_assert_eq!( + config.kiro, + parsed.kiro, + "kiro 字段往返不一致" + ); + prop_assert_eq!( + config.other, + parsed.other, + "other 字段往返不一致" + ); + } + + /// **Feature: endpoint-provider-config, Property 3: 配置序列化往返一致性(JSON)** + /// *对于任意* 有效的 EndpointProvidersConfig 对象,JSON 序列化后再反序列化应产生等价的对象。 + /// **Validates: Requirements 1.1** + #[test] + fn prop_endpoint_providers_config_json_roundtrip(config in arb_endpoint_providers_config()) { + // 序列化为 JSON + let json = serde_json::to_string(&config) + .expect("JSON 序列化应成功"); + + // 反序列化回 EndpointProvidersConfig + let parsed: EndpointProvidersConfig = serde_json::from_str(&json) + .expect("JSON 反序列化应成功"); + + // 验证往返一致性 + prop_assert_eq!( + config, + parsed, + "EndpointProvidersConfig JSON 往返不一致" + ); + } + + /// **Feature: endpoint-provider-config, Property 3: 配置序列化往返一致性(完整配置)** + /// *对于任意* 包含 EndpointProvidersConfig 的完整配置,序列化后再反序列化应保持 endpoint_providers 一致。 + /// **Validates: Requirements 1.1** + #[test] + fn prop_config_with_endpoint_providers_roundtrip( + endpoint_providers in arb_endpoint_providers_config() + ) { + // 创建包含 endpoint_providers 的完整配置 + let config = Config { + endpoint_providers: endpoint_providers.clone(), + ..Config::default() + }; + + // 序列化为 YAML + let yaml = ConfigManager::to_yaml(&config) + .expect("序列化应成功"); + + // 反序列化回 Config + let parsed = ConfigManager::parse_yaml(&yaml) + .expect("反序列化应成功"); + + // 验证 endpoint_providers 往返一致性 + prop_assert_eq!( + endpoint_providers, + parsed.endpoint_providers, + "endpoint_providers 往返不一致" + ); + } +} + +// ============================================================================ +// Property 4: Provider 类型验证 +// ============================================================================ + +use crate::ProviderType; + +/// 生成有效的 Provider 类型字符串 +/// 注意:只包含往返一致的 Provider 类型(即 parse().to_string() == 原值) +/// qwen 等第三方 Provider 会被映射到 openai,不满足往返一致性 +fn arb_valid_provider_type() -> impl Strategy { + prop_oneof![ + Just("kiro".to_string()), + Just("gemini".to_string()), + Just("openai".to_string()), + Just("claude".to_string()), + Just("antigravity".to_string()), + Just("vertex".to_string()), + Just("gemini_api_key".to_string()), + Just("codex".to_string()), + Just("claude_oauth".to_string()), + Just("anthropic".to_string()), + Just("anthropic_compatible".to_string()), + Just("azure_openai".to_string()), + Just("aws_bedrock".to_string()), + Just("ollama".to_string()), + ] +} + +/// 生成无效的 Provider 类型字符串 +fn arb_invalid_provider_type() -> impl Strategy { + // 生成不在有效列表中的字符串 + // 注意:需要排除所有在 ProviderType::from_str 中有效的字符串 + "[a-z]{3,15}".prop_filter("排除有效的 Provider 类型", |s| { + // 排除所有在 ProviderType::from_str 中能成功解析的字符串 + use std::str::FromStr; + crate::ProviderType::from_str(s).is_err() + }) +} + +/// 生成有效的客户端类型字符串 +fn arb_valid_client_type() -> impl Strategy { + prop_oneof![ + Just("cursor".to_string()), + Just("claude_code".to_string()), + Just("codex".to_string()), + Just("windsurf".to_string()), + Just("kiro".to_string()), + Just("other".to_string()), + ] +} + +/// 生成无效的客户端类型字符串 +fn arb_invalid_client_type() -> impl Strategy { + // 生成不在有效列表中的字符串 + "[a-z]{3,15}".prop_filter("排除有效的客户端类型", |s| { + !matches!( + s.as_str(), + "cursor" | "claude_code" | "codex" | "windsurf" | "kiro" | "other" + ) + }) +} + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: endpoint-provider-config, Property 4: Provider 类型验证** + /// *对于任意* 有效的 Provider 类型字符串,解析应成功并返回正确的 ProviderType。 + /// **Validates: Requirements 5.1, 5.2** + #[test] + fn prop_valid_provider_type_parsing(provider in arb_valid_provider_type()) { + // 解析 Provider 类型 + let result: Result = provider.parse(); + + // 验证解析成功 + prop_assert!( + result.is_ok(), + "有效的 Provider 类型应解析成功: {}", + provider + ); + + // 验证往返一致性 + let parsed = result.unwrap(); + prop_assert_eq!( + parsed.to_string(), + provider, + "Provider 类型往返不一致" + ); + } + + /// **Feature: endpoint-provider-config, Property 4: Provider 类型验证** + /// *对于任意* 无效的 Provider 类型字符串,解析应失败并返回描述性错误消息。 + /// **Validates: Requirements 5.1, 5.2** + #[test] + fn prop_invalid_provider_type_parsing(provider in arb_invalid_provider_type()) { + // 解析 Provider 类型 + let result: Result = provider.parse(); + + // 验证解析失败 + prop_assert!( + result.is_err(), + "无效的 Provider 类型应解析失败: {}", + provider + ); + + // 验证错误消息包含描述性信息 + let error = result.unwrap_err(); + prop_assert!( + error.contains("Invalid provider") || error.contains(&provider), + "错误消息应包含描述性信息: {}", + error + ); + } + + /// **Feature: endpoint-provider-config, Property 4: Provider 类型验证** + /// *对于任意* 有效的客户端类型,set_provider 应成功设置 Provider。 + /// **Validates: Requirements 5.1, 5.2** + #[test] + fn prop_valid_client_type_set_provider( + client_type in arb_valid_client_type(), + provider in arb_valid_provider_type() + ) { + let mut config = EndpointProvidersConfig::default(); + + // 设置 Provider + let result = config.set_provider(&client_type, Some(provider.clone())); + + // 验证设置成功 + prop_assert!( + result, + "有效的客户端类型应设置成功: {}", + client_type + ); + + // 验证 Provider 已正确设置 + let stored = config.get_provider(&client_type); + prop_assert_eq!( + stored, + Some(&provider), + "Provider 应正确存储" + ); + } + + /// **Feature: endpoint-provider-config, Property 4: Provider 类型验证** + /// *对于任意* 无效的客户端类型,set_provider 应返回 false。 + /// **Validates: Requirements 5.1, 5.2** + #[test] + fn prop_invalid_client_type_set_provider( + client_type in arb_invalid_client_type(), + provider in arb_valid_provider_type() + ) { + let mut config = EndpointProvidersConfig::default(); + + // 设置 Provider + let result = config.set_provider(&client_type, Some(provider)); + + // 验证设置失败 + prop_assert!( + !result, + "无效的客户端类型应设置失败: {}", + client_type + ); + } + + /// **Feature: endpoint-provider-config, Property 4: Provider 类型验证** + /// *对于任意* 有效的客户端类型,使用 None 或空字符串应清除 Provider 配置。 + /// **Validates: Requirements 5.1, 5.2** + #[test] + fn prop_clear_provider_config( + client_type in arb_valid_client_type(), + provider in arb_valid_provider_type() + ) { + let mut config = EndpointProvidersConfig::default(); + + // 先设置 Provider + config.set_provider(&client_type, Some(provider)); + + // 使用 None 清除 + let result = config.set_provider(&client_type, None); + prop_assert!(result, "清除操作应成功"); + prop_assert_eq!( + config.get_provider(&client_type), + None, + "Provider 应被清除" + ); + + // 重新设置后使用空字符串清除 + config.set_provider(&client_type, Some("kiro".to_string())); + let result = config.set_provider(&client_type, Some("".to_string())); + prop_assert!(result, "空字符串清除操作应成功"); + prop_assert_eq!( + config.get_provider(&client_type), + None, + "Provider 应被清除(空字符串)" + ); + } +} diff --git a/src-tauri/src/config/types.rs b/src-tauri/crates/core/src/config/types.rs similarity index 98% rename from src-tauri/src/config/types.rs rename to src-tauri/crates/core/src/config/types.rs index 03e932add..330121a12 100644 --- a/src-tauri/src/config/types.rs +++ b/src-tauri/crates/core/src/config/types.rs @@ -3,7 +3,7 @@ //! 定义 ProxyCast 的配置结构,支持 YAML 和 JSON 序列化/反序列化 //! 保持与旧版 JSON 配置的向后兼容性 -use crate::injection::{InjectionMode, InjectionRule}; +use crate::models::injection_types::{InjectionMode, InjectionRule}; use serde::{Deserialize, Serialize}; use std::collections::HashMap; @@ -192,34 +192,8 @@ pub struct GeminiApiKeyEntry { } /// Vertex AI 模型别名映射 -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] -pub struct VertexModelAlias { - /// 上游模型名称 - pub name: String, - /// 客户端可见的别名 - pub alias: String, -} - -/// Vertex AI 凭证条目 -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] -pub struct VertexApiKeyEntry { - /// 凭证 ID - pub id: String, - /// API Key - pub api_key: String, - /// Base URL - #[serde(default, skip_serializing_if = "Option::is_none")] - pub base_url: Option, - /// 模型别名映射 - #[serde(default, skip_serializing_if = "Vec::is_empty")] - pub models: Vec, - /// 单独的代理 URL - #[serde(default, skip_serializing_if = "Option::is_none")] - pub proxy_url: Option, - /// 是否禁用 - #[serde(default)] - pub disabled: bool, -} +// Vertex AI 类型从 core crate 重新导出 +pub use crate::models::vertex_model::{VertexApiKeyEntry, VertexModelAlias}; #[allow(dead_code)] fn default_auth_type() -> String { @@ -871,7 +845,7 @@ impl NativeAgentConfig { pub fn get_effective_system_prompt(&self) -> Option { // 优先从文件加载 if let Some(file_path) = &self.system_prompt_file { - let expanded_path = crate::config::expand_tilde(file_path); + let expanded_path = super::path_utils::expand_tilde(file_path); if let Ok(content) = std::fs::read_to_string(&expanded_path) { let trimmed = content.trim(); if !trimmed.is_empty() { diff --git a/src-tauri/src/config/yaml.rs b/src-tauri/crates/core/src/config/yaml.rs similarity index 100% rename from src-tauri/src/config/yaml.rs rename to src-tauri/crates/core/src/config/yaml.rs diff --git a/src-tauri/src/connect/README.md b/src-tauri/crates/core/src/connect/README.md similarity index 100% rename from src-tauri/src/connect/README.md rename to src-tauri/crates/core/src/connect/README.md diff --git a/src-tauri/src/connect/deep_link.rs b/src-tauri/crates/core/src/connect/deep_link.rs similarity index 98% rename from src-tauri/src/connect/deep_link.rs rename to src-tauri/crates/core/src/connect/deep_link.rs index 351fe3450..5c0a4826b 100644 --- a/src-tauri/src/connect/deep_link.rs +++ b/src-tauri/crates/core/src/connect/deep_link.rs @@ -11,7 +11,7 @@ //! ## 使用示例 //! //! ```rust -//! use proxycast_lib::connect::deep_link::{parse_deep_link, ConnectPayload, DeepLinkError}; +//! use proxycast_core::connect::deep_link::{parse_deep_link, ConnectPayload, DeepLinkError}; //! //! let url = "proxycast://connect?relay=example&key=sk-xxx&name=MyKey"; //! match parse_deep_link(url) { @@ -80,7 +80,7 @@ impl std::error::Error for DeepLinkError {} /// # 示例 /// /// ```rust -/// use proxycast_lib::connect::deep_link::parse_deep_link; +/// use proxycast_core::connect::deep_link::parse_deep_link; /// /// // 完整 URL /// let result = parse_deep_link("proxycast://connect?relay=example&key=sk-xxx&name=MyKey&ref=abc"); diff --git a/src-tauri/src/connect/mod.rs b/src-tauri/crates/core/src/connect/mod.rs similarity index 100% rename from src-tauri/src/connect/mod.rs rename to src-tauri/crates/core/src/connect/mod.rs diff --git a/src-tauri/src/connect/registry.rs b/src-tauri/crates/core/src/connect/registry.rs similarity index 99% rename from src-tauri/src/connect/registry.rs rename to src-tauri/crates/core/src/connect/registry.rs index 0fe616920..0e5e87dae 100644 --- a/src-tauri/src/connect/registry.rs +++ b/src-tauri/crates/core/src/connect/registry.rs @@ -11,7 +11,7 @@ //! ## 使用示例 //! //! ```rust,ignore -//! use proxycast_lib::connect::registry::{RelayRegistry, RelayInfo}; +//! use proxycast_core::connect::registry::{RelayRegistry, RelayInfo}; //! //! let registry = RelayRegistry::new(cache_path); //! registry.load_from_remote().await?; diff --git a/src-tauri/src/connect/webhook.rs b/src-tauri/crates/core/src/connect/webhook.rs similarity index 100% rename from src-tauri/src/connect/webhook.rs rename to src-tauri/crates/core/src/connect/webhook.rs diff --git a/src-tauri/crates/core/src/data/mod.rs b/src-tauri/crates/core/src/data/mod.rs index 60a799247..802b30544 100644 --- a/src-tauri/crates/core/src/data/mod.rs +++ b/src-tauri/crates/core/src/data/mod.rs @@ -1,4 +1 @@ //! 静态数据模块 -//! -//! 模型数据现在从 aiclientproxy/models 仓库获取 -//! 本地硬编码数据已迁移到独立仓库: https://github.com/aiclientproxy/models diff --git a/src-tauri/src/errors/README.md b/src-tauri/crates/core/src/errors/README.md similarity index 100% rename from src-tauri/src/errors/README.md rename to src-tauri/crates/core/src/errors/README.md diff --git a/src-tauri/src/errors/mod.rs b/src-tauri/crates/core/src/errors/mod.rs similarity index 100% rename from src-tauri/src/errors/mod.rs rename to src-tauri/crates/core/src/errors/mod.rs diff --git a/src-tauri/src/errors/project_error.rs b/src-tauri/crates/core/src/errors/project_error.rs similarity index 100% rename from src-tauri/src/errors/project_error.rs rename to src-tauri/crates/core/src/errors/project_error.rs diff --git a/src-tauri/crates/core/src/lib.rs b/src-tauri/crates/core/src/lib.rs index 089201fd2..cae70e149 100644 --- a/src-tauri/crates/core/src/lib.rs +++ b/src-tauri/crates/core/src/lib.rs @@ -1,13 +1,36 @@ -//! 核心类型模块 +//! ProxyCast Core Crate //! -//! 包含纯数据类型(models)、静态数据(data)、日志配置(logger) +//! 包含纯数据类型、基础模块和无外部业务依赖的独立模块。 //! -//! 本 crate 不包含任何业务逻辑,只提供基础类型定义。 +//! ## 模块结构 +//! - `models`: 核心数据模型定义 +//! - `data`: 静态数据 +//! - `logger`: 日志配置 +//! - `errors`: 错误类型定义 +//! - `backends`: 后端调用层 Trait +//! - `config`: 配置管理(类型、YAML、热重载、导入导出) +//! - `connect`: Deep Link 协议和中转商注册表 +//! - `middleware`: HTTP 中间件(认证、限速) +//! - `orchestrator`: 模型选择编排器 +//! - `plugin`: 插件系统(加载、管理、UI、安装) +//! - `session`: 会话管理(限速、粘性路由) +//! - `session_files`: 会话文件存储 pub mod data; pub mod logger; pub mod models; +// 独立业务模块(无主 crate 依赖) +pub mod backends; +pub mod config; +pub mod connect; +pub mod errors; +pub mod middleware; +pub mod orchestrator; +pub mod plugin; +pub mod session; +pub mod session_files; + // 重新导出常用类型 pub use logger::{LogEntry, LogStore, LogStoreConfig, SharedLogStore}; pub use models::provider_type::ProviderType; diff --git a/src-tauri/src/middleware/management_auth.rs b/src-tauri/crates/core/src/middleware/management_auth.rs similarity index 99% rename from src-tauri/src/middleware/management_auth.rs rename to src-tauri/crates/core/src/middleware/management_auth.rs index a9ffdb651..1b06169d4 100644 --- a/src-tauri/src/middleware/management_auth.rs +++ b/src-tauri/crates/core/src/middleware/management_auth.rs @@ -48,7 +48,7 @@ fn failure_map() -> &'static Mutex impl Strategy { "[a-zA-Z0-9_-]{8,32}".prop_map(|s| s) } -/// 生成随机的无效 secret_key(与有效 key 不同) -fn arb_invalid_secret_key(valid_key: String) -> impl Strategy { - "[a-zA-Z0-9_-]{8,32}".prop_filter_map("must differ from valid key", move |s| { - if s != valid_key { - Some(s) - } else { - None - } - }) -} - /// 生成随机的 IP 地址 fn arb_ip_addr() -> impl Strategy { prop_oneof![ - // localhost IPv4 Just("127.0.0.1".to_string()), - // localhost IPv6 Just("::1".to_string()), - // remote IPv4 (1u8..255u8, 0u8..255u8, 0u8..255u8, 1u8..255u8).prop_filter_map( "not localhost", |(a, b, c, d)| { @@ -55,11 +41,6 @@ fn arb_ip_addr() -> impl Strategy { ] } -/// 生成随机端口 -fn arb_port() -> impl Strategy { - 1024u16..65535u16 -} - /// Mock service that always returns 200 OK #[derive(Clone)] struct MockService; @@ -88,40 +69,18 @@ impl Service> for MockService { /// Helper to create a request with optional Authorization header fn create_request_with_auth(auth_header: Option<&str>) -> Request { let mut builder = Request::builder().uri("/v0/management/status"); - if let Some(auth) = auth_header { builder = builder.header("authorization", auth); } - builder.body(Body::empty()).unwrap() } /// Helper to create a request with X-Management-Key header fn create_request_with_management_key(key: Option<&str>) -> Request { let mut builder = Request::builder().uri("/v0/management/status"); - if let Some(k) = key { builder = builder.header("x-management-key", k); } - - builder.body(Body::empty()).unwrap() -} - -/// Helper to create a request with X-Management-Key and X-Forwarded-For headers -fn create_request_with_management_key_and_forwarded( - key: Option<&str>, - forwarded_for: Option<&str>, -) -> Request { - let mut builder = Request::builder().uri("/v0/management/status"); - - if let Some(k) = key { - builder = builder.header("x-management-key", k); - } - - if let Some(addr) = forwarded_for { - builder = builder.header("x-forwarded-for", addr); - } - builder.body(Body::empty()).unwrap() } @@ -137,21 +96,17 @@ fn test_management_auth_rate_limit_after_failures() { let mut service = layer.layer(MockService); let rt = tokio::runtime::Runtime::new().unwrap(); - // 使用唯一的 IP 地址避免测试间干扰 - // 直接使用原子计数器确保唯一性,避免与其他测试冲突 use std::sync::atomic::{AtomicU32, Ordering}; static RATE_LIMIT_TEST_COUNTER: AtomicU32 = AtomicU32::new(1); let unique_id = RATE_LIMIT_TEST_COUNTER.fetch_add(1, Ordering::SeqCst); - // 使用 TEST-NET-2 (198.51.100.0/24) 范围,确保与其他测试不冲突 let octet3 = ((unique_id >> 8) & 0xFF) as u8; let octet4 = (unique_id & 0xFF) as u8; let client_ip = format!("198.51.{}.{}", 100 + (octet3 % 155), octet4.max(1)); let addr: SocketAddr = format!("{client_ip}:12345").parse().unwrap(); - // 发送 5 次失败请求,每次都应该返回 401 + // 发送 5 次失败请求 for i in 0..5 { let mut req = create_request_with_management_key(Some("invalid")); - // 安全修复后不再信任 X-Forwarded-For,需要注入 ConnectInfo req.extensions_mut().insert(ConnectInfo(addr)); let response = rt.block_on(async { service.call(req).await.unwrap() }); assert_eq!( @@ -162,7 +117,7 @@ fn test_management_auth_rate_limit_after_failures() { ); } - // 第 6 次请求应该被限速,返回 429 + // 第 6 次请求应该被限速 let mut req = create_request_with_management_key(Some("invalid")); req.extensions_mut().insert(ConnectInfo(addr)); let response = rt.block_on(async { service.call(req).await.unwrap() }); @@ -176,191 +131,96 @@ fn test_management_auth_rate_limit_after_failures() { proptest! { #![proptest_config(ProptestConfig::with_cases(100))] - /// **Feature: cliproxyapi-parity, Property 19: Management Auth Rejection** - /// *For any* management API request without valid secret_key, the response SHALL be 401 Unauthorized. - /// **Validates: Requirements 9.3** #[test] fn prop_management_auth_rejection_missing_key( secret_key in arb_secret_key() ) { - // 只清除 "unknown" 客户端的状态,避免影响并行测试 clear_auth_failure_state_for("unknown"); - // Create config with a valid secret_key let config = RemoteManagementConfig { allow_remote: true, secret_key: Some(secret_key), disable_control_panel: false, }; - - // Create the auth layer and service let layer = ManagementAuthLayer::new(config); let mut service = layer.layer(MockService); - - // Create request WITHOUT any auth header let req = create_request_with_auth(None); - - // Execute the service let rt = tokio::runtime::Runtime::new().unwrap(); - let response = rt.block_on(async { - service.call(req).await.unwrap() - }); - - // Verify: should return 401 Unauthorized - prop_assert_eq!( - response.status(), - StatusCode::UNAUTHORIZED, - "Request without secret_key should return 401 Unauthorized" - ); + let response = rt.block_on(async { service.call(req).await.unwrap() }); + prop_assert_eq!(response.status(), StatusCode::UNAUTHORIZED); } - /// **Feature: cliproxyapi-parity, Property 19: Management Auth Rejection** - /// *For any* management API request with invalid secret_key, the response SHALL be 401 Unauthorized. - /// **Validates: Requirements 9.3** #[test] fn prop_management_auth_rejection_invalid_key( secret_key in arb_secret_key() ) { - // 只清除 "unknown" 客户端的状态,避免影响并行测试 clear_auth_failure_state_for("unknown"); - // Create config with a valid secret_key let config = RemoteManagementConfig { allow_remote: true, secret_key: Some(secret_key.clone()), disable_control_panel: false, }; - - // Create the auth layer and service let layer = ManagementAuthLayer::new(config); let mut service = layer.layer(MockService); - - // Create request with WRONG auth header (append "wrong" to make it different) let wrong_key = format!("{secret_key}wrong"); let req = create_request_with_auth(Some(&format!("Bearer {wrong_key}"))); - - // Execute the service let rt = tokio::runtime::Runtime::new().unwrap(); - let response = rt.block_on(async { - service.call(req).await.unwrap() - }); - - // Verify: should return 401 Unauthorized - prop_assert_eq!( - response.status(), - StatusCode::UNAUTHORIZED, - "Request with invalid secret_key should return 401 Unauthorized" - ); + let response = rt.block_on(async { service.call(req).await.unwrap() }); + prop_assert_eq!(response.status(), StatusCode::UNAUTHORIZED); } - /// **Feature: cliproxyapi-parity, Property 19: Management Auth Rejection** - /// *For any* management API request with valid secret_key, the response SHALL NOT be 401 Unauthorized. - /// **Validates: Requirements 9.3** #[test] fn prop_management_auth_acceptance_valid_key( secret_key in arb_secret_key() ) { - // 只清除 "unknown" 客户端的状态,避免影响并行测试 clear_auth_failure_state_for("unknown"); - // Create config with a valid secret_key let config = RemoteManagementConfig { allow_remote: true, secret_key: Some(secret_key.clone()), disable_control_panel: false, }; - - // Create the auth layer and service let layer = ManagementAuthLayer::new(config); let mut service = layer.layer(MockService); - - // Create request with CORRECT auth header let req = create_request_with_auth(Some(&format!("Bearer {secret_key}"))); - - // Execute the service let rt = tokio::runtime::Runtime::new().unwrap(); - let response = rt.block_on(async { - service.call(req).await.unwrap() - }); - - // Verify: should return 200 OK (passed through to MockService) - prop_assert_eq!( - response.status(), - StatusCode::OK, - "Request with valid secret_key should pass through (200 OK)" - ); + let response = rt.block_on(async { service.call(req).await.unwrap() }); + prop_assert_eq!(response.status(), StatusCode::OK); } - /// **Feature: cliproxyapi-parity, Property 19: Management Auth Rejection** - /// *For any* management API request with valid X-Management-Key header, the response SHALL NOT be 401 Unauthorized. - /// **Validates: Requirements 9.3** #[test] fn prop_management_auth_acceptance_x_management_key( secret_key in arb_secret_key() ) { - // 只清除 "unknown" 客户端的状态,避免影响并行测试 clear_auth_failure_state_for("unknown"); - // Create config with a valid secret_key let config = RemoteManagementConfig { allow_remote: true, secret_key: Some(secret_key.clone()), disable_control_panel: false, }; - - // Create the auth layer and service let layer = ManagementAuthLayer::new(config); let mut service = layer.layer(MockService); - - // Create request with X-Management-Key header let req = create_request_with_management_key(Some(&secret_key)); - - // Execute the service let rt = tokio::runtime::Runtime::new().unwrap(); - let response = rt.block_on(async { - service.call(req).await.unwrap() - }); - - // Verify: should return 200 OK (passed through to MockService) - prop_assert_eq!( - response.status(), - StatusCode::OK, - "Request with valid X-Management-Key should pass through (200 OK)" - ); + let response = rt.block_on(async { service.call(req).await.unwrap() }); + prop_assert_eq!(response.status(), StatusCode::OK); } - /// **Feature: cliproxyapi-parity, Property 19: Management Auth Rejection** - /// *For any* management API request with invalid X-Management-Key header, the response SHALL be 401 Unauthorized. - /// **Validates: Requirements 9.3** #[test] fn prop_management_auth_rejection_invalid_x_management_key( secret_key in arb_secret_key() ) { - // 只清除 "unknown" 客户端的状态,避免影响并行测试 clear_auth_failure_state_for("unknown"); - // Create config with a valid secret_key let config = RemoteManagementConfig { allow_remote: true, secret_key: Some(secret_key.clone()), disable_control_panel: false, }; - - // Create the auth layer and service let layer = ManagementAuthLayer::new(config); let mut service = layer.layer(MockService); - - // Create request with WRONG X-Management-Key header let wrong_key = format!("{secret_key}wrong"); let req = create_request_with_management_key(Some(&wrong_key)); - - // Execute the service let rt = tokio::runtime::Runtime::new().unwrap(); - let response = rt.block_on(async { - service.call(req).await.unwrap() - }); - - // Verify: should return 401 Unauthorized - prop_assert_eq!( - response.status(), - StatusCode::UNAUTHORIZED, - "Request with invalid X-Management-Key should return 401 Unauthorized" - ); + let response = rt.block_on(async { service.call(req).await.unwrap() }); + prop_assert_eq!(response.status(), StatusCode::UNAUTHORIZED); } } @@ -375,13 +235,10 @@ mod unit_tests { secret_key: Some("test-secret-key".to_string()), disable_control_panel: false, }; - let layer = ManagementAuthLayer::new(config); let mut service = layer.layer(MockService); - let req = create_request_with_auth(None); let response = service.call(req).await.unwrap(); - assert_eq!(response.status(), StatusCode::UNAUTHORIZED); } @@ -392,13 +249,10 @@ mod unit_tests { secret_key: Some("correct-key".to_string()), disable_control_panel: false, }; - let layer = ManagementAuthLayer::new(config); let mut service = layer.layer(MockService); - let req = create_request_with_auth(Some("Bearer wrong-key")); let response = service.call(req).await.unwrap(); - assert_eq!(response.status(), StatusCode::UNAUTHORIZED); } @@ -409,13 +263,10 @@ mod unit_tests { secret_key: Some("correct-key".to_string()), disable_control_panel: false, }; - let layer = ManagementAuthLayer::new(config); let mut service = layer.layer(MockService); - let req = create_request_with_auth(Some("Bearer correct-key")); let response = service.call(req).await.unwrap(); - assert_eq!(response.status(), StatusCode::OK); } } diff --git a/src-tauri/crates/core/src/models/anthropic.rs b/src-tauri/crates/core/src/models/anthropic.rs index 34a956162..87926833f 100644 --- a/src-tauri/crates/core/src/models/anthropic.rs +++ b/src-tauri/crates/core/src/models/anthropic.rs @@ -1,5 +1,4 @@ //! Anthropic/Claude API 数据模型 - use serde::{Deserialize, Serialize}; #[derive(Debug, Clone, Serialize, Deserialize)] @@ -40,7 +39,7 @@ pub struct ImageSource { #[derive(Debug, Clone, Serialize, Deserialize)] pub struct AnthropicMessage { pub role: String, - pub content: serde_json::Value, + pub content: serde_json::Value, // Can be string or array } #[derive(Debug, Clone, Serialize, Deserialize)] @@ -89,6 +88,7 @@ pub struct AnthropicMessagesResponse { pub usage: AnthropicUsage, } +// Streaming events #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(tag = "type")] pub enum AnthropicStreamEvent { diff --git a/src-tauri/crates/core/src/models/app_type.rs b/src-tauri/crates/core/src/models/app_type.rs index dd48833fb..4ca757550 100644 --- a/src-tauri/crates/core/src/models/app_type.rs +++ b/src-tauri/crates/core/src/models/app_type.rs @@ -1,5 +1,3 @@ -//! 应用类型定义 - use serde::{Deserialize, Serialize}; #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] diff --git a/src-tauri/crates/core/src/models/codewhisperer.rs b/src-tauri/crates/core/src/models/codewhisperer.rs index 948635247..98197e0ab 100644 --- a/src-tauri/crates/core/src/models/codewhisperer.rs +++ b/src-tauri/crates/core/src/models/codewhisperer.rs @@ -1,7 +1,10 @@ //! CodeWhisperer/Kiro API 数据模型 //! //! 支持标准工具和特殊工具类型(如 web_search)。 - +//! +//! # 更新日志 +//! +//! - 2025-12-27: 添加 CWWebSearchTool 支持,修复 Issue #49 use serde::{Deserialize, Serialize}; #[derive(Debug, Clone, Serialize, Deserialize)] @@ -50,6 +53,10 @@ pub struct UserInputMessageContext { } /// CodeWhisperer 工具项 +/// +/// 支持两种类型: +/// - 标准工具(带 tool_specification) +/// - 联网搜索工具(仅 type 字段) #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(untagged)] pub enum CWToolItem { @@ -67,6 +74,9 @@ pub struct CWTool { } /// 联网搜索工具 +/// +/// Codex/Kiro API 支持的特殊工具类型,用于联网搜索。 +/// 格式:`{"type": "web_search"}` #[derive(Debug, Clone, Serialize, Deserialize)] pub struct CWWebSearchTool { #[serde(rename = "type")] @@ -145,6 +155,7 @@ pub struct CWToolUse { pub tool_use_id: String, } +// Response types #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(rename_all = "camelCase")] pub struct CWStreamEvent { diff --git a/src-tauri/crates/core/src/models/injection_types.rs b/src-tauri/crates/core/src/models/injection_types.rs index 446871a08..57bbe0fa7 100644 --- a/src-tauri/crates/core/src/models/injection_types.rs +++ b/src-tauri/crates/core/src/models/injection_types.rs @@ -72,6 +72,43 @@ impl InjectionRule { pub fn is_exact(&self) -> bool { !self.pattern.contains('*') } + + /// 检查模型是否匹配此规则 + /// + /// 支持的通配符模式: + /// - 精确匹配: `claude-sonnet-4-5` + /// - 前缀匹配: `claude-*` + /// - 后缀匹配: `*-preview` + /// - 包含匹配: `*flash*` + pub fn matches(&self, model: &str) -> bool { + if !self.enabled { + return false; + } + pattern_matches(&self.pattern, model) + } +} + +/// 检查模式是否匹配模型名 +/// +/// 支持的通配符模式: +/// - 精确匹配: `claude-sonnet-4-5` +/// - 前缀匹配: `claude-*` +/// - 后缀匹配: `*-preview` +/// - 包含匹配: `*flash*` +pub fn pattern_matches(pattern: &str, model: &str) -> bool { + if !pattern.contains('*') { + return pattern == model; + } + + let parts: Vec<&str> = pattern.split('*').collect(); + + match parts.as_slice() { + [prefix, ""] => model.starts_with(prefix), + ["", suffix] => model.ends_with(suffix), + ["", middle, ""] => model.contains(middle), + [prefix, suffix] => model.starts_with(prefix) && model.ends_with(suffix), + _ => false, + } } /// 规则排序:精确匹配优先,然后按优先级 diff --git a/src-tauri/crates/core/src/models/kiro_fingerprint.rs b/src-tauri/crates/core/src/models/kiro_fingerprint.rs index 00628f642..384eb743c 100644 --- a/src-tauri/crates/core/src/models/kiro_fingerprint.rs +++ b/src-tauri/crates/core/src/models/kiro_fingerprint.rs @@ -37,6 +37,7 @@ impl KiroFingerprintStore { .ok_or_else(|| "无法获取应用数据目录".to_string())? .join("proxycast"); + // 确保目录存在 if !app_data_dir.exists() { fs::create_dir_all(&app_data_dir).map_err(|e| format!("创建应用数据目录失败: {e}"))?; } @@ -73,6 +74,8 @@ impl KiroFingerprintStore { } /// 获取或创建凭证的指纹绑定 + /// + /// 如果凭证没有绑定指纹,会基于凭证信息生成一个新的 Machine ID pub fn get_or_create_binding( &mut self, credential_uuid: &str, @@ -80,6 +83,7 @@ impl KiroFingerprintStore { client_id: Option<&str>, ) -> Result<&KiroFingerprintBinding, String> { if !self.bindings.contains_key(credential_uuid) { + // 生成基于凭证的 Machine ID let machine_id = generate_stable_machine_id(credential_uuid, profile_arn, client_id); let binding = KiroFingerprintBinding { @@ -113,6 +117,9 @@ impl KiroFingerprintStore { } /// 生成稳定的 Machine ID +/// +/// 基于凭证信息生成一个稳定的 UUID 格式 Machine ID。 +/// 同一凭证每次生成的 Machine ID 相同,确保账号身份一致。 fn generate_stable_machine_id( credential_uuid: &str, profile_arn: Option<&str>, @@ -120,6 +127,7 @@ fn generate_stable_machine_id( ) -> String { use sha2::{Digest, Sha256}; + // 使用凭证相关信息作为种子 let seed = format!( "kiro_fingerprint:{}:{}:{}", credential_uuid, @@ -131,6 +139,7 @@ fn generate_stable_machine_id( hasher.update(seed.as_bytes()); let result = hasher.finalize(); + // 将哈希结果转换为 UUID 格式 let hex = format!("{result:x}"); format!( "{}-{}-{}-{}-{}", @@ -145,10 +154,15 @@ fn generate_stable_machine_id( /// 切换到本地的结果 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct SwitchToLocalResult { + /// 是否成功 pub success: bool, + /// 结果消息 pub message: String, + /// 是否需要用户操作(如需管理员权限) pub requires_action: bool, + /// 切换的 Machine ID pub machine_id: Option, + /// 是否需要重启 Kiro IDE pub requires_kiro_restart: bool, } diff --git a/src-tauri/crates/core/src/models/machine_id.rs b/src-tauri/crates/core/src/models/machine_id.rs index edc7b0928..0bc295d26 100644 --- a/src-tauri/crates/core/src/models/machine_id.rs +++ b/src-tauri/crates/core/src/models/machine_id.rs @@ -1,35 +1,49 @@ -//! 机器码相关数据模型 - use serde::{Deserialize, Serialize}; -/// 机器码信息结构 +/// 机器码信息结构 - v0.20.0 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct MachineIdInfo { + /// 当前机器码 pub current_id: String, + /// 原始机器码(如果有备份) pub original_id: Option, + /// 操作系统平台 pub platform: String, + /// 是否可以修改 pub can_modify: bool, + /// 是否需要管理员权限 pub requires_admin: bool, + /// 是否存在备份 pub backup_exists: bool, + /// 机器码格式类型 pub format_type: MachineIdFormat, } /// 机器码操作结果 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct MachineIdResult { + /// 操作是否成功 pub success: bool, + /// 结果消息 pub message: String, + /// 是否需要重启 pub requires_restart: bool, + /// 是否需要管理员权限 pub requires_admin: bool, + /// 新的机器码(如果操作成功) pub new_machine_id: Option, } /// 管理员权限状态 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct AdminStatus { + /// 是否具有管理员权限 pub is_admin: bool, + /// 操作系统平台 pub platform: String, + /// 权限提升方法说明 pub elevation_method: Option, + /// 权限检查是否成功 pub check_success: bool, } @@ -37,9 +51,12 @@ pub struct AdminStatus { #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(rename_all = "lowercase")] pub enum MachineIdFormat { + /// UUID 格式 (xxxxxxxx-xxxx-xxxx-xxxx-xxxxxxxxxxxx) Uuid, + /// 32位十六进制格式 (xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx) #[serde(rename = "hex32")] Hex32, + /// 其他格式 #[serde(rename = "unknown")] Unknown, } @@ -47,20 +64,30 @@ pub enum MachineIdFormat { /// 机器码备份信息 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct MachineIdBackup { + /// 备份的机器码 pub machine_id: String, + /// 备份时间戳 pub timestamp: i64, + /// 操作系统平台 pub platform: String, + /// 机器码格式 pub format: MachineIdFormat, + /// 备份描述 pub description: Option, } /// 机器码历史记录 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct MachineIdHistory { + /// 记录ID pub id: String, + /// 机器码 pub machine_id: String, + /// 操作时间戳 pub timestamp: String, + /// 操作系统平台 pub platform: String, + /// 备份路径(可选) pub backup_path: Option, } @@ -68,20 +95,30 @@ pub struct MachineIdHistory { #[derive(Debug, Clone, Serialize, Deserialize)] #[allow(dead_code)] pub enum MachineIdOperation { + /// 获取当前机器码 Get, + /// 设置新机器码 Set, + /// 生成随机机器码 Generate, + /// 备份机器码 Backup, + /// 恢复机器码 Restore, + /// 重置为原始机器码 Reset, } /// 机器码验证结果 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct MachineIdValidation { + /// 是否有效 pub is_valid: bool, + /// 检测到的格式 pub detected_format: MachineIdFormat, + /// 验证错误信息 pub error_message: Option, + /// 格式化后的机器码(如果有效) pub formatted_id: Option, } @@ -90,6 +127,7 @@ impl MachineIdFormat { pub fn detect(machine_id: &str) -> Self { let cleaned = machine_id.replace("-", "").replace(" ", "").to_lowercase(); + // 检查UUID格式:8-4-4-4-12个十六进制字符 if machine_id.contains("-") && machine_id.len() == 36 { let parts: Vec<&str> = machine_id.split('-').collect(); if parts.len() == 5 @@ -104,6 +142,7 @@ impl MachineIdFormat { } } + // 检查32位十六进制格式 if cleaned.len() == 32 && cleaned.chars().all(|c| c.is_ascii_hexdigit()) { return MachineIdFormat::Hex32; } diff --git a/src-tauri/crates/core/src/models/mcp_model.rs b/src-tauri/crates/core/src/models/mcp_model.rs index e38dd22b8..fe257cc13 100644 --- a/src-tauri/crates/core/src/models/mcp_model.rs +++ b/src-tauri/crates/core/src/models/mcp_model.rs @@ -1,7 +1,48 @@ -//! MCP Server 数据模型 - use serde::{Deserialize, Serialize}; use serde_json::Value; +use std::collections::HashMap; + +/// MCP 服务器配置(类型化) +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct McpServerConfigTyped { + /// 启动命令 + pub command: String, + /// 命令参数 + #[serde(default)] + pub args: Vec, + /// 环境变量 + #[serde(default)] + pub env: HashMap, + /// 工作目录 + #[serde(skip_serializing_if = "Option::is_none")] + pub cwd: Option, + /// 超时时间(秒) + #[serde(default = "default_timeout")] + pub timeout: u64, +} + +fn default_timeout() -> u64 { + 30 +} + +impl Default for McpServerConfigTyped { + fn default() -> Self { + Self { + command: String::new(), + args: Vec::new(), + env: HashMap::new(), + cwd: None, + timeout: 30, + } + } +} + +/// 配置验证错误 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ConfigValidationError { + pub field: String, + pub message: String, +} #[derive(Debug, Clone, Serialize, Deserialize)] pub struct McpServer { @@ -37,4 +78,104 @@ impl McpServer { created_at: Some(chrono::Utc::now().timestamp()), } } + + /// 解析 server_config 为类型化配置 + /// + /// 将 JSON Value 解析为 McpServerConfigTyped 结构。 + /// 如果解析失败,返回默认配置并尝试提取基本字段。 + pub fn parse_config(&self) -> McpServerConfigTyped { + serde_json::from_value(self.server_config.clone()).unwrap_or_else(|_| { + // 尝试手动提取字段 + McpServerConfigTyped { + command: self + .server_config + .get("command") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(), + args: self + .server_config + .get("args") + .and_then(|v| v.as_array()) + .map(|arr| { + arr.iter() + .filter_map(|v| v.as_str().map(|s| s.to_string())) + .collect() + }) + .unwrap_or_default(), + env: self + .server_config + .get("env") + .and_then(|v| v.as_object()) + .map(|obj| { + obj.iter() + .filter_map(|(k, v)| v.as_str().map(|s| (k.clone(), s.to_string()))) + .collect() + }) + .unwrap_or_default(), + cwd: self + .server_config + .get("cwd") + .and_then(|v| v.as_str()) + .map(|s| s.to_string()), + timeout: self + .server_config + .get("timeout") + .and_then(|v| v.as_u64()) + .unwrap_or(30), + } + }) + } + + /// 验证服务器配置 + /// + /// 检查配置是否有效,返回验证错误列表。 + /// 空列表表示配置有效。 + pub fn validate_config(&self) -> Vec { + let mut errors = Vec::new(); + let config = self.parse_config(); + + // 验证 command 不为空 + if config.command.trim().is_empty() { + errors.push(ConfigValidationError { + field: "command".to_string(), + message: "启动命令不能为空".to_string(), + }); + } + + // 验证 name 不为空 + if self.name.trim().is_empty() { + errors.push(ConfigValidationError { + field: "name".to_string(), + message: "服务器名称不能为空".to_string(), + }); + } + + // 验证 name 不包含特殊字符(用于工具名称前缀) + if !self + .name + .chars() + .all(|c| c.is_alphanumeric() || c == '-' || c == '_') + { + errors.push(ConfigValidationError { + field: "name".to_string(), + message: "服务器名称只能包含字母、数字、连字符和下划线".to_string(), + }); + } + + // 验证 timeout 在合理范围内 + if config.timeout == 0 || config.timeout > 300 { + errors.push(ConfigValidationError { + field: "timeout".to_string(), + message: "超时时间必须在 1-300 秒之间".to_string(), + }); + } + + errors + } + + /// 检查配置是否有效 + pub fn is_valid(&self) -> bool { + self.validate_config().is_empty() + } } diff --git a/src-tauri/crates/core/src/models/mod.rs b/src-tauri/crates/core/src/models/mod.rs index 7a833854c..2e36cf900 100644 --- a/src-tauri/crates/core/src/models/mod.rs +++ b/src-tauri/crates/core/src/models/mod.rs @@ -17,6 +17,7 @@ pub mod provider_pool_model; pub mod provider_type; pub mod route_model; pub mod skill_model; +pub mod vertex_model; #[allow(unused_imports)] pub use anthropic::*; @@ -33,3 +34,4 @@ pub use provider_model::Provider; pub use provider_pool_model::*; pub use provider_type::ProviderType; pub use skill_model::{Skill, SkillMetadata, SkillRepo, SkillState, SkillStates}; +pub use vertex_model::{VertexApiKeyEntry, VertexModelAlias}; diff --git a/src-tauri/crates/core/src/models/model_registry.rs b/src-tauri/crates/core/src/models/model_registry.rs index 9b3ca74f9..e2418904e 100644 --- a/src-tauri/crates/core/src/models/model_registry.rs +++ b/src-tauri/crates/core/src/models/model_registry.rs @@ -1,25 +1,38 @@ //! 模型注册表数据结构 +//! +//! 借鉴 opencode 的模型管理方式,定义增强的模型元数据结构 use serde::{Deserialize, Serialize}; /// 模型能力 #[derive(Debug, Clone, Serialize, Deserialize, Default)] pub struct ModelCapabilities { + /// 是否支持视觉输入 pub vision: bool, + /// 是否支持工具调用 pub tools: bool, + /// 是否支持流式输出 pub streaming: bool, + /// 是否支持 JSON 模式 pub json_mode: bool, + /// 是否支持函数调用 pub function_calling: bool, + /// 是否支持推理/思考 pub reasoning: bool, } /// 模型定价 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ModelPricing { + /// 输入价格(每百万 token) pub input_per_million: Option, + /// 输出价格(每百万 token) pub output_per_million: Option, + /// 缓存读取价格(每百万 token) pub cache_read_per_million: Option, + /// 缓存写入价格(每百万 token) pub cache_write_per_million: Option, + /// 货币单位 ("USD" | "CNY") pub currency: String, } @@ -38,9 +51,13 @@ impl Default for ModelPricing { /// 模型限制 #[derive(Debug, Clone, Serialize, Deserialize, Default)] pub struct ModelLimits { + /// 上下文长度 pub context_length: Option, + /// 最大输出 token 数 pub max_output_tokens: Option, + /// 每分钟请求数限制 pub requests_per_minute: Option, + /// 每分钟 token 数限制 pub tokens_per_minute: Option, } @@ -48,11 +65,17 @@ pub struct ModelLimits { #[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] #[serde(rename_all = "lowercase")] pub enum ModelStatus { + /// 活跃可用 Active, + /// 预览版 Preview, + /// Alpha 测试 Alpha, + /// Beta 测试 Beta, + /// 已弃用 Deprecated, + /// 旧版本 Legacy, } @@ -95,8 +118,11 @@ impl std::str::FromStr for ModelStatus { #[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] #[serde(rename_all = "lowercase")] pub enum ModelTier { + /// 快速响应,适合简单任务 Mini, + /// 均衡性能,适合大多数任务 Pro, + /// 最强能力,适合复杂任务 Max, } @@ -133,10 +159,16 @@ impl std::str::FromStr for ModelTier { #[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] #[serde(rename_all = "lowercase")] pub enum ModelSource { + /// 从内嵌资源加载(构建时打包) Embedded, + /// 从 models.dev API 获取(已弃用) ModelsDev, + /// 本地硬编码(国内模型等) Local, + /// 用户自定义 Custom, + /// 从 Provider API 获取 + Api, } impl Default for ModelSource { @@ -152,6 +184,7 @@ impl std::fmt::Display for ModelSource { Self::ModelsDev => write!(f, "models.dev"), Self::Local => write!(f, "local"), Self::Custom => write!(f, "custom"), + Self::Api => write!(f, "api"), } } } @@ -165,6 +198,7 @@ impl std::str::FromStr for ModelSource { "models.dev" | "modelsdev" => Ok(Self::ModelsDev), "local" => Ok(Self::Local), "custom" => Ok(Self::Custom), + "api" => Ok(Self::Api), _ => Err(format!("Unknown model source: {s}")), } } @@ -173,25 +207,42 @@ impl std::str::FromStr for ModelSource { /// 增强的模型元数据 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct EnhancedModelMetadata { + /// 模型 ID (如 "claude-sonnet-4-5-20250514") pub id: String, + /// 显示名称 (如 "Claude Sonnet 4.5") pub display_name: String, + /// Provider ID (如 "anthropic", "openai", "dashscope") pub provider_id: String, + /// Provider 显示名称 pub provider_name: String, + /// 模型家族 (如 "sonnet", "gpt-4", "qwen") pub family: Option, + /// 服务等级 pub tier: ModelTier, + /// 模型能力 pub capabilities: ModelCapabilities, + /// 定价信息 pub pricing: Option, + /// 限制信息 pub limits: ModelLimits, + /// 模型状态 pub status: ModelStatus, + /// 发布日期 pub release_date: Option, + /// 是否为最新版本 pub is_latest: bool, + /// 描述 pub description: Option, + /// 数据来源 pub source: ModelSource, + /// 创建时间 (Unix 时间戳) pub created_at: i64, + /// 最后更新时间 (Unix 时间戳) pub updated_at: i64, } impl EnhancedModelMetadata { + /// 创建新的模型元数据 pub fn new( id: String, display_name: String, @@ -219,51 +270,61 @@ impl EnhancedModelMetadata { } } + /// 设置模型家族 pub fn with_family(mut self, family: impl Into) -> Self { self.family = Some(family.into()); self } + /// 设置服务等级 pub fn with_tier(mut self, tier: ModelTier) -> Self { self.tier = tier; self } + /// 设置模型能力 pub fn with_capabilities(mut self, capabilities: ModelCapabilities) -> Self { self.capabilities = capabilities; self } + /// 设置定价信息 pub fn with_pricing(mut self, pricing: ModelPricing) -> Self { self.pricing = Some(pricing); self } + /// 设置限制信息 pub fn with_limits(mut self, limits: ModelLimits) -> Self { self.limits = limits; self } + /// 设置模型状态 pub fn with_status(mut self, status: ModelStatus) -> Self { self.status = status; self } + /// 设置发布日期 pub fn with_release_date(mut self, date: impl Into) -> Self { self.release_date = Some(date.into()); self } + /// 设置是否为最新版本 pub fn with_is_latest(mut self, is_latest: bool) -> Self { self.is_latest = is_latest; self } + /// 设置描述 pub fn with_description(mut self, description: impl Into) -> Self { self.description = Some(description.into()); self } + /// 设置数据来源 pub fn with_source(mut self, source: ModelSource) -> Self { self.source = source; self @@ -273,17 +334,26 @@ impl EnhancedModelMetadata { /// 用户模型偏好 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct UserModelPreference { + /// 模型 ID pub model_id: String, + /// 是否收藏 pub is_favorite: bool, + /// 是否隐藏 pub is_hidden: bool, + /// 自定义别名 pub custom_alias: Option, + /// 使用次数 pub usage_count: u32, + /// 最后使用时间 (Unix 时间戳) pub last_used_at: Option, + /// 创建时间 (Unix 时间戳) pub created_at: i64, + /// 更新时间 (Unix 时间戳) pub updated_at: i64, } impl UserModelPreference { + /// 创建新的用户偏好 pub fn new(model_id: String) -> Self { let now = chrono::Utc::now().timestamp(); Self { @@ -302,45 +372,63 @@ impl UserModelPreference { /// 模型同步状态 #[derive(Debug, Clone, Serialize, Deserialize, Default)] pub struct ModelSyncState { + /// 最后同步时间 (Unix 时间戳) pub last_sync_at: Option, + /// 同步的模型数量 pub model_count: u32, + /// 是否正在同步 pub is_syncing: bool, + /// 最后同步错误 pub last_error: Option, } -// Provider Alias 相关类型 +// ============================================================================ +// Provider Alias 相关类型(用于 Kiro、Antigravity 等中转服务) +// ============================================================================ /// 单个模型别名映射 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ModelAlias { + /// 实际模型 ID(如 "claude-sonnet-4-5-20250929") pub actual: String, + /// 内部 API 名称(如 "CLAUDE_SONNET_4_5_20250929_V1_0") pub internal_name: Option, + /// 原始 Provider(如 "anthropic") pub provider: Option, + /// 描述 pub description: Option, } /// Provider 的别名配置 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ProviderAliasConfig { + /// Provider ID(如 "kiro"、"antigravity") pub provider: String, + /// 描述 pub description: Option, + /// 支持的模型列表 #[serde(default)] pub models: Vec, + /// 别名映射(模型名 -> 别名配置) pub aliases: std::collections::HashMap, + /// 更新时间 pub updated_at: Option, } impl ProviderAliasConfig { + /// 检查是否支持指定模型 pub fn supports_model(&self, model: &str) -> bool { self.models.contains(&model.to_string()) || self.aliases.contains_key(model) } + /// 获取模型的内部名称 pub fn get_internal_name(&self, model: &str) -> Option<&str> { self.aliases .get(model) .and_then(|a| a.internal_name.as_deref()) } + /// 获取模型的实际 ID pub fn get_actual_model(&self, model: &str) -> Option<&str> { self.aliases.get(model).map(|a| a.actual.as_str()) } @@ -421,6 +509,8 @@ pub struct ModelsDevModalities { } impl ModelsDevModel { + /// 转换为 EnhancedModelMetadata + /// 预留:用于从 models.dev API 导入模型数据 #[allow(dead_code)] pub fn to_enhanced_metadata( &self, @@ -429,6 +519,7 @@ impl ModelsDevModel { ) -> EnhancedModelMetadata { let now = chrono::Utc::now().timestamp(); + // 判断是否支持视觉 let supports_vision = self .modalities .as_ref() @@ -436,14 +527,17 @@ impl ModelsDevModel { .unwrap_or(false) || self.attachment; + // 根据模型名称推断服务等级 let tier = infer_model_tier(&self.id, &self.name); + // 解析状态 let status = self .status .as_ref() .and_then(|s| s.parse().ok()) .unwrap_or(ModelStatus::Active); + // 判断是否为最新版本 let is_latest = self.id.contains("latest"); EnhancedModelMetadata { @@ -456,8 +550,8 @@ impl ModelsDevModel { capabilities: ModelCapabilities { vision: supports_vision, tools: self.tool_call, - streaming: true, - json_mode: true, + streaming: true, // 大多数模型都支持流式 + json_mode: true, // 大多数模型都支持 JSON 模式 function_calling: self.tool_call, reasoning: self.reasoning, }, @@ -486,28 +580,13 @@ impl ModelsDevModel { } /// 根据模型 ID 和名称推断服务等级 +/// 用于 to_enhanced_metadata 和测试 #[allow(dead_code)] fn infer_model_tier(model_id: &str, model_name: &str) -> ModelTier { let id_lower = model_id.to_lowercase(); let name_lower = model_name.to_lowercase(); - let max_patterns = [ - "opus", - "gpt-4o", - "gpt-4-turbo", - "gemini-2.5-pro", - "gemini-ultra", - "claude-3-opus", - "qwen-max", - "glm-4-plus", - "deepseek-v3", - ]; - for pattern in max_patterns { - if id_lower.contains(pattern) || name_lower.contains(pattern) { - return ModelTier::Max; - } - } - + // Mini 等级模型(优先检查,因为 gpt-4o-mini 包含 gpt-4o) let mini_patterns = [ "mini", "nano", @@ -525,6 +604,25 @@ fn infer_model_tier(model_id: &str, model_name: &str) -> ModelTier { } } + // Max 等级模型 + let max_patterns = [ + "opus", + "gpt-4o", + "gpt-4-turbo", + "gemini-2.5-pro", + "gemini-ultra", + "claude-3-opus", + "qwen-max", + "glm-4-plus", + "deepseek-v3", + ]; + for pattern in max_patterns { + if id_lower.contains(pattern) || name_lower.contains(pattern) { + return ModelTier::Max; + } + } + + // 默认为 Pro 等级 ModelTier::Pro } diff --git a/src-tauri/crates/core/src/models/openai.rs b/src-tauri/crates/core/src/models/openai.rs index 0f4927748..2119cb667 100644 --- a/src-tauri/crates/core/src/models/openai.rs +++ b/src-tauri/crates/core/src/models/openai.rs @@ -1,7 +1,15 @@ //! OpenAI API 数据模型 //! //! 支持标准 OpenAI 格式以及扩展的工具类型(如 web_search)。 - +//! +//! # 工具类型支持 +//! +//! - `function`: 标准函数调用工具 +//! - `web_search`: 联网搜索工具(Claude Code 使用 `web_search_20250305`) +//! +//! # 更新日志 +//! +//! - 2025-12-27: 添加 web_search 工具支持,修复 Issue #49 use serde::{Deserialize, Serialize}; #[derive(Debug, Clone, Serialize, Deserialize)] @@ -50,6 +58,8 @@ pub struct ChatMessage { pub tool_calls: Option>, #[serde(skip_serializing_if = "Option::is_none")] pub tool_call_id: Option, + /// 推理内容(DeepSeek R1 等模型的思维链内容) + /// DeepSeek Reasoner 在 Tool Calls 场景下要求此字段 #[serde(skip_serializing_if = "Option::is_none")] pub reasoning_content: Option, } @@ -74,21 +84,25 @@ impl ChatMessage { } /// 提取消息中的图片 URL 列表 + /// 返回 (format, base64_data) 元组列表 pub fn get_images(&self) -> Vec<(String, String)> { match &self.content { Some(MessageContent::Parts(parts)) => parts .iter() .filter_map(|p| { if let ContentPart::ImageUrl { image_url } = p { + // 解析 data URL: data:image/jpeg;base64,xxxxx if image_url.url.starts_with("data:") { let parts: Vec<&str> = image_url.url.splitn(2, ',').collect(); if parts.len() == 2 { + // 提取 media_type: data:image/jpeg;base64 -> image/jpeg let header = parts[0]; let data = parts[1]; let media_type = header .strip_prefix("data:") .and_then(|s| s.split(';').next()) .unwrap_or("image/jpeg"); + // 提取格式: image/jpeg -> jpeg let format = media_type.split('/').nth(1).unwrap_or("jpeg").to_string(); return Some((format, data.to_string())); @@ -115,13 +129,21 @@ pub struct FunctionDef { } /// 工具定义 +/// +/// 支持多种工具类型: +/// - `function`: 标准函数调用工具,包含 function 字段 +/// - `web_search`: 联网搜索工具,无需额外字段 +/// - `web_search_20250305`: Claude Code 的联网搜索工具类型 #[derive(Debug, Clone, Serialize, Deserialize)] #[serde(tag = "type")] pub enum Tool { + /// 标准函数调用工具 #[serde(rename = "function")] Function { function: FunctionDef }, + /// 联网搜索工具(Codex/Kiro 格式) #[serde(rename = "web_search")] WebSearch, + /// 联网搜索工具(Claude Code 格式) #[serde(rename = "web_search_20250305")] WebSearch20250305, } @@ -142,6 +164,7 @@ pub struct ChatCompletionRequest { pub tools: Option>, #[serde(skip_serializing_if = "Option::is_none")] pub tool_choice: Option, + /// 思维链强度:none, low, medium, high #[serde(skip_serializing_if = "Option::is_none")] pub reasoning_effort: Option, } @@ -206,24 +229,43 @@ pub struct ChatCompletionChunk { pub choices: Vec, } +// ============================================================================ // 图像生成 API 数据模型 +// ============================================================================ /// OpenAI 图像生成请求 +/// +/// 兼容 OpenAI Images API,支持通过 Antigravity 生成图像。 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ImageGenerationRequest { + /// 图像生成提示词 pub prompt: String, + + /// 模型名称 (默认: gemini-3-pro-image-preview) #[serde(default = "default_image_model")] pub model: String, + + /// 生成图像数量 (默认: 1) #[serde(default = "default_n")] pub n: u32, + + /// 图像尺寸 (可选,Antigravity 可能忽略) #[serde(skip_serializing_if = "Option::is_none")] pub size: Option, + + /// 响应格式: "url" 或 "b64_json" (默认: "url") #[serde(default = "default_response_format")] pub response_format: String, + + /// 图像质量 (可选,Antigravity 可能忽略) #[serde(skip_serializing_if = "Option::is_none")] pub quality: Option, + + /// 图像风格 (可选,Antigravity 可能忽略) #[serde(skip_serializing_if = "Option::is_none")] pub style: Option, + + /// 用户标识 (可选) #[serde(skip_serializing_if = "Option::is_none")] pub user: Option, } @@ -243,17 +285,25 @@ fn default_response_format() -> String { /// OpenAI 图像生成响应 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ImageGenerationResponse { + /// 创建时间戳 (Unix epoch seconds) pub created: i64, + + /// 生成的图像数组 pub data: Vec, } /// 单个图像数据 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ImageData { + /// Base64 编码的图像数据 (当 response_format="b64_json") #[serde(skip_serializing_if = "Option::is_none")] pub b64_json: Option, + + /// 图像 URL (当 response_format="url",返回 data URL) #[serde(skip_serializing_if = "Option::is_none")] pub url: Option, + + /// 修订后的提示词 (如果 Antigravity 返回了文本) #[serde(skip_serializing_if = "Option::is_none")] pub revised_prompt: Option, } diff --git a/src-tauri/crates/core/src/models/prompt_model.rs b/src-tauri/crates/core/src/models/prompt_model.rs index c2313533d..b77d529b9 100644 --- a/src-tauri/crates/core/src/models/prompt_model.rs +++ b/src-tauri/crates/core/src/models/prompt_model.rs @@ -1,5 +1,3 @@ -//! Prompt 数据模型 - use serde::{Deserialize, Serialize}; #[derive(Debug, Clone, Serialize, Deserialize)] @@ -10,6 +8,7 @@ pub struct Prompt { pub content: String, #[serde(skip_serializing_if = "Option::is_none")] pub description: Option, + /// Whether this prompt is currently enabled (synced to live file) #[serde(default)] pub enabled: bool, #[serde(rename = "createdAt", skip_serializing_if = "Option::is_none")] diff --git a/src-tauri/crates/core/src/models/provider_model.rs b/src-tauri/crates/core/src/models/provider_model.rs index a0176d8bf..354cee392 100644 --- a/src-tauri/crates/core/src/models/provider_model.rs +++ b/src-tauri/crates/core/src/models/provider_model.rs @@ -1,5 +1,3 @@ -//! Provider 数据模型 - use serde::{Deserialize, Serialize}; use serde_json::Value; diff --git a/src-tauri/crates/core/src/models/provider_pool_model.rs b/src-tauri/crates/core/src/models/provider_pool_model.rs index cdcb73896..2326a4827 100644 --- a/src-tauri/crates/core/src/models/provider_pool_model.rs +++ b/src-tauri/crates/core/src/models/provider_pool_model.rs @@ -27,7 +27,7 @@ pub enum CredentialSource { /// /// 为了向后兼容,PoolProviderType 是 crate::ProviderType 的类型别名。 /// 所有 Provider 类型定义已统一到 lib.rs 中的 ProviderType。 -pub type PoolProviderType = crate::ProviderType; +pub type PoolProviderType = super::provider_type::ProviderType; /// 凭证数据,根据 Provider 类型不同而不同 #[derive(Debug, Clone, Serialize, Deserialize)] diff --git a/src-tauri/crates/core/src/models/route_model.rs b/src-tauri/crates/core/src/models/route_model.rs index aad72620a..5a8bfe24c 100644 --- a/src-tauri/crates/core/src/models/route_model.rs +++ b/src-tauri/crates/core/src/models/route_model.rs @@ -7,38 +7,53 @@ use serde::{Deserialize, Serialize}; /// 单个路由信息 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct RouteInfo { + /// 路由选择器 (provider 类型或凭证名称) pub selector: String, + /// Provider 类型 pub provider_type: String, + /// 关联的凭证数量 pub credential_count: usize, + /// 可用的端点列表 pub endpoints: Vec, + /// 标签 (如 "突破限制", "官方API/三方") pub tags: Vec, + /// 是否启用 pub enabled: bool, } /// 路由端点 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct RouteEndpoint { + /// 端点路径 pub path: String, - pub protocol: String, + /// 协议类型 + pub protocol: String, // "openai" 或 "claude" + /// 完整 URL pub url: String, } /// 路由列表响应 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct RouteListResponse { + /// 服务器基础 URL pub base_url: String, + /// 默认 Provider pub default_provider: String, + /// 所有可用路由 pub routes: Vec, } /// curl 示例 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct CurlExample { + /// 描述 pub description: String, + /// curl 命令 pub command: String, } impl RouteInfo { + /// 创建新的路由信息 pub fn new(selector: String, provider_type: String) -> Self { Self { selector, @@ -50,6 +65,7 @@ impl RouteInfo { } } + /// 添加端点 pub fn add_endpoint(&mut self, base_url: &str, protocol: &str) { let path = match protocol { "claude" => format!("/{}/v1/messages", self.selector), @@ -64,6 +80,7 @@ impl RouteInfo { }); } + /// 生成 curl 示例 pub fn generate_curl_examples(&self, api_key: &str) -> Vec { let mut examples = Vec::new(); diff --git a/src-tauri/crates/core/src/models/skill_model.rs b/src-tauri/crates/core/src/models/skill_model.rs index 8b9e862b4..09c69d9eb 100644 --- a/src-tauri/crates/core/src/models/skill_model.rs +++ b/src-tauri/crates/core/src/models/skill_model.rs @@ -1,5 +1,3 @@ -//! Skill 数据模型 - use chrono::{DateTime, Utc}; use serde::{Deserialize, Serialize}; use std::collections::HashMap; @@ -77,6 +75,7 @@ impl SkillRepo { pub fn get_default_skill_repos() -> Vec { vec![ + // ProxyCast 官方仓库(排第一位) SkillRepo { owner: "proxycast".to_string(), name: "skills".to_string(), @@ -109,13 +108,18 @@ pub type SkillStates = HashMap; #[cfg(test)] mod tests { use super::*; + use proptest::prelude::*; + /// Feature: skills-platform-mvp, Property 1: Default Repositories Include ProxyCast Official + /// Validates: Requirements 1.1, 1.2, 1.3 #[test] fn test_default_repos_include_proxycast_official() { let repos = get_default_skill_repos(); + // 验证列表非空 assert!(!repos.is_empty(), "默认仓库列表不应为空"); + // 验证第一个仓库是 ProxyCast 官方仓库 let first_repo = &repos[0]; assert_eq!( first_repo.owner, "proxycast", @@ -126,10 +130,34 @@ mod tests { assert!(first_repo.enabled, "ProxyCast 官方仓库应默认启用"); } + // Property 1: Default Repositories Include ProxyCast Official (Property-Based Test) + // For any call to get_default_skill_repos(), the returned list SHALL contain + // a SkillRepo with owner="proxycast", name="skills", branch="main", and enabled=true, + // and this repo SHALL be the first item in the list. + // Validates: Requirements 1.1, 1.2, 1.3 + proptest! { + #[test] + fn prop_default_repos_proxycast_first(_seed in 0u64..1000) { + // 无论调用多少次,结果应该一致 + let repos = get_default_skill_repos(); + + // Property: 列表非空 + prop_assert!(!repos.is_empty()); + + // Property: 第一个仓库是 ProxyCast 官方仓库 + let first = &repos[0]; + prop_assert_eq!(&first.owner, "proxycast"); + prop_assert_eq!(&first.name, "skills"); + prop_assert_eq!(&first.branch, "main"); + prop_assert!(first.enabled); + } + } + #[test] fn test_proxycast_repo_exists_in_list() { let repos = get_default_skill_repos(); + // 验证 ProxyCast 仓库存在于列表中 let proxycast_repo = repos .iter() .find(|r| r.owner == "proxycast" && r.name == "skills"); diff --git a/src-tauri/crates/core/src/models/vertex_model.rs b/src-tauri/crates/core/src/models/vertex_model.rs new file mode 100644 index 000000000..79ee45001 --- /dev/null +++ b/src-tauri/crates/core/src/models/vertex_model.rs @@ -0,0 +1,35 @@ +//! Vertex AI 配置模型 +//! +//! 定义 Vertex AI 相关的配置类型,供 providers 和 config 模块共享。 + +use serde::{Deserialize, Serialize}; + +/// Vertex AI 模型别名映射 +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +pub struct VertexModelAlias { + /// 上游模型名称 + pub name: String, + /// 客户端可见的别名 + pub alias: String, +} + +/// Vertex AI 凭证条目 +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +pub struct VertexApiKeyEntry { + /// 凭证 ID + pub id: String, + /// API Key + pub api_key: String, + /// Base URL + #[serde(default, skip_serializing_if = "Option::is_none")] + pub base_url: Option, + /// 模型别名映射 + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub models: Vec, + /// 单独的代理 URL + #[serde(default, skip_serializing_if = "Option::is_none")] + pub proxy_url: Option, + /// 是否禁用 + #[serde(default)] + pub disabled: bool, +} diff --git a/src-tauri/src/orchestrator/fallback.rs b/src-tauri/crates/core/src/orchestrator/fallback.rs similarity index 100% rename from src-tauri/src/orchestrator/fallback.rs rename to src-tauri/crates/core/src/orchestrator/fallback.rs diff --git a/src-tauri/src/orchestrator/mod.rs b/src-tauri/crates/core/src/orchestrator/mod.rs similarity index 100% rename from src-tauri/src/orchestrator/mod.rs rename to src-tauri/crates/core/src/orchestrator/mod.rs diff --git a/src-tauri/src/orchestrator/orchestrator.rs b/src-tauri/crates/core/src/orchestrator/orchestrator.rs similarity index 98% rename from src-tauri/src/orchestrator/orchestrator.rs rename to src-tauri/crates/core/src/orchestrator/orchestrator.rs index b93be3ca8..2927ad777 100644 --- a/src-tauri/src/orchestrator/orchestrator.rs +++ b/src-tauri/crates/core/src/orchestrator/orchestrator.rs @@ -297,8 +297,8 @@ pub struct PoolStats { } /// 全局编排器实例 -static GLOBAL_ORCHESTRATOR: once_cell::sync::OnceCell> = - once_cell::sync::OnceCell::new(); +static GLOBAL_ORCHESTRATOR: std::sync::OnceLock> = + std::sync::OnceLock::new(); /// 初始化全局编排器 pub fn init_global_orchestrator() -> Arc { diff --git a/src-tauri/src/orchestrator/pool_builder.rs b/src-tauri/crates/core/src/orchestrator/pool_builder.rs similarity index 100% rename from src-tauri/src/orchestrator/pool_builder.rs rename to src-tauri/crates/core/src/orchestrator/pool_builder.rs diff --git a/src-tauri/src/orchestrator/selector.rs b/src-tauri/crates/core/src/orchestrator/selector.rs similarity index 100% rename from src-tauri/src/orchestrator/selector.rs rename to src-tauri/crates/core/src/orchestrator/selector.rs diff --git a/src-tauri/src/orchestrator/strategies/cost_optimized.rs b/src-tauri/crates/core/src/orchestrator/strategies/cost_optimized.rs similarity index 100% rename from src-tauri/src/orchestrator/strategies/cost_optimized.rs rename to src-tauri/crates/core/src/orchestrator/strategies/cost_optimized.rs diff --git a/src-tauri/src/orchestrator/strategies/load_balanced.rs b/src-tauri/crates/core/src/orchestrator/strategies/load_balanced.rs similarity index 100% rename from src-tauri/src/orchestrator/strategies/load_balanced.rs rename to src-tauri/crates/core/src/orchestrator/strategies/load_balanced.rs diff --git a/src-tauri/src/orchestrator/strategies/mod.rs b/src-tauri/crates/core/src/orchestrator/strategies/mod.rs similarity index 100% rename from src-tauri/src/orchestrator/strategies/mod.rs rename to src-tauri/crates/core/src/orchestrator/strategies/mod.rs diff --git a/src-tauri/src/orchestrator/strategies/round_robin.rs b/src-tauri/crates/core/src/orchestrator/strategies/round_robin.rs similarity index 100% rename from src-tauri/src/orchestrator/strategies/round_robin.rs rename to src-tauri/crates/core/src/orchestrator/strategies/round_robin.rs diff --git a/src-tauri/src/orchestrator/strategies/speed_optimized.rs b/src-tauri/crates/core/src/orchestrator/strategies/speed_optimized.rs similarity index 100% rename from src-tauri/src/orchestrator/strategies/speed_optimized.rs rename to src-tauri/crates/core/src/orchestrator/strategies/speed_optimized.rs diff --git a/src-tauri/src/orchestrator/strategies/task_based.rs b/src-tauri/crates/core/src/orchestrator/strategies/task_based.rs similarity index 100% rename from src-tauri/src/orchestrator/strategies/task_based.rs rename to src-tauri/crates/core/src/orchestrator/strategies/task_based.rs diff --git a/src-tauri/src/orchestrator/strategy.rs b/src-tauri/crates/core/src/orchestrator/strategy.rs similarity index 100% rename from src-tauri/src/orchestrator/strategy.rs rename to src-tauri/crates/core/src/orchestrator/strategy.rs diff --git a/src-tauri/src/orchestrator/tier.rs b/src-tauri/crates/core/src/orchestrator/tier.rs similarity index 100% rename from src-tauri/src/orchestrator/tier.rs rename to src-tauri/crates/core/src/orchestrator/tier.rs diff --git a/src-tauri/src/plugin/binary_downloader.rs b/src-tauri/crates/core/src/plugin/binary_downloader.rs similarity index 100% rename from src-tauri/src/plugin/binary_downloader.rs rename to src-tauri/crates/core/src/plugin/binary_downloader.rs diff --git a/src-tauri/src/plugin/examples/credential_monitor.rs b/src-tauri/crates/core/src/plugin/examples/credential_monitor.rs similarity index 100% rename from src-tauri/src/plugin/examples/credential_monitor.rs rename to src-tauri/crates/core/src/plugin/examples/credential_monitor.rs diff --git a/src-tauri/src/plugin/examples/mod.rs b/src-tauri/crates/core/src/plugin/examples/mod.rs similarity index 100% rename from src-tauri/src/plugin/examples/mod.rs rename to src-tauri/crates/core/src/plugin/examples/mod.rs diff --git a/src-tauri/src/plugin/installer/README.md b/src-tauri/crates/core/src/plugin/installer/README.md similarity index 100% rename from src-tauri/src/plugin/installer/README.md rename to src-tauri/crates/core/src/plugin/installer/README.md diff --git a/src-tauri/src/plugin/installer/downloader.rs b/src-tauri/crates/core/src/plugin/installer/downloader.rs similarity index 100% rename from src-tauri/src/plugin/installer/downloader.rs rename to src-tauri/crates/core/src/plugin/installer/downloader.rs diff --git a/src-tauri/src/plugin/installer/installer.rs b/src-tauri/crates/core/src/plugin/installer/installer.rs similarity index 100% rename from src-tauri/src/plugin/installer/installer.rs rename to src-tauri/crates/core/src/plugin/installer/installer.rs diff --git a/src-tauri/src/plugin/installer/mod.rs b/src-tauri/crates/core/src/plugin/installer/mod.rs similarity index 100% rename from src-tauri/src/plugin/installer/mod.rs rename to src-tauri/crates/core/src/plugin/installer/mod.rs diff --git a/src-tauri/src/plugin/installer/registry.rs b/src-tauri/crates/core/src/plugin/installer/registry.rs similarity index 100% rename from src-tauri/src/plugin/installer/registry.rs rename to src-tauri/crates/core/src/plugin/installer/registry.rs diff --git a/src-tauri/src/plugin/installer/tests.rs b/src-tauri/crates/core/src/plugin/installer/tests.rs similarity index 100% rename from src-tauri/src/plugin/installer/tests.rs rename to src-tauri/crates/core/src/plugin/installer/tests.rs diff --git a/src-tauri/src/plugin/installer/types.rs b/src-tauri/crates/core/src/plugin/installer/types.rs similarity index 100% rename from src-tauri/src/plugin/installer/types.rs rename to src-tauri/crates/core/src/plugin/installer/types.rs diff --git a/src-tauri/src/plugin/installer/validator.rs b/src-tauri/crates/core/src/plugin/installer/validator.rs similarity index 100% rename from src-tauri/src/plugin/installer/validator.rs rename to src-tauri/crates/core/src/plugin/installer/validator.rs diff --git a/src-tauri/src/plugin/loader.rs b/src-tauri/crates/core/src/plugin/loader.rs similarity index 100% rename from src-tauri/src/plugin/loader.rs rename to src-tauri/crates/core/src/plugin/loader.rs diff --git a/src-tauri/src/plugin/manager.rs b/src-tauri/crates/core/src/plugin/manager.rs similarity index 100% rename from src-tauri/src/plugin/manager.rs rename to src-tauri/crates/core/src/plugin/manager.rs diff --git a/src-tauri/crates/core/src/plugin/mod.rs b/src-tauri/crates/core/src/plugin/mod.rs new file mode 100644 index 000000000..1c877fc81 --- /dev/null +++ b/src-tauri/crates/core/src/plugin/mod.rs @@ -0,0 +1,36 @@ +//! 插件系统模块 +//! +//! 提供插件扩展功能,支持: +//! - 插件加载和初始化 +//! - 请求前/响应后钩子 +//! - 插件隔离和错误处理 +//! - 插件配置管理 +//! - 二进制组件下载和管理 +//! - 声明式插件 UI 系统 +//! - 插件安装和卸载 + +pub mod binary_downloader; +pub mod examples; +pub mod installer; +mod loader; +mod manager; +mod types; +pub mod ui_builder; +pub mod ui_trait; +pub mod ui_types; + +pub use binary_downloader::BinaryDownloader; +pub use loader::PluginLoader; +pub use manager::PluginManager; +pub use types::{ + BinaryComponentStatus, BinaryManifest, HookResult, PlatformBinaries, Plugin, PluginConfig, + PluginContext, PluginError, PluginInfo, PluginManifest, PluginState, PluginStatus, PluginType, +}; +pub use ui_trait::{NoUI, PluginUI}; +pub use ui_types::{ + Action, BoundValue, ChildrenDef, ComponentDef, ComponentType, DataEntry, DataModelUpdate, + SurfaceDefinition, SurfaceUpdate, UIMessage, UserAction, +}; + +#[cfg(test)] +mod tests; diff --git a/src-tauri/src/plugin/tests.rs b/src-tauri/crates/core/src/plugin/tests.rs similarity index 100% rename from src-tauri/src/plugin/tests.rs rename to src-tauri/crates/core/src/plugin/tests.rs diff --git a/src-tauri/src/plugin/types.rs b/src-tauri/crates/core/src/plugin/types.rs similarity index 100% rename from src-tauri/src/plugin/types.rs rename to src-tauri/crates/core/src/plugin/types.rs diff --git a/src-tauri/src/plugin/ui_builder.rs b/src-tauri/crates/core/src/plugin/ui_builder.rs similarity index 100% rename from src-tauri/src/plugin/ui_builder.rs rename to src-tauri/crates/core/src/plugin/ui_builder.rs diff --git a/src-tauri/src/plugin/ui_trait.rs b/src-tauri/crates/core/src/plugin/ui_trait.rs similarity index 100% rename from src-tauri/src/plugin/ui_trait.rs rename to src-tauri/crates/core/src/plugin/ui_trait.rs diff --git a/src-tauri/src/plugin/ui_types.rs b/src-tauri/crates/core/src/plugin/ui_types.rs similarity index 100% rename from src-tauri/src/plugin/ui_types.rs rename to src-tauri/crates/core/src/plugin/ui_types.rs diff --git a/src-tauri/crates/core/src/session/mod.rs b/src-tauri/crates/core/src/session/mod.rs new file mode 100644 index 000000000..b143a3b24 --- /dev/null +++ b/src-tauri/crates/core/src/session/mod.rs @@ -0,0 +1,16 @@ +//! 会话管理核心模块 +//! +//! 提供以下功能: +//! - 增强的限流处理(Duration 解析、指数退避) +//! - 会话粘性管理(会话与账号映射) +//! - 调度模式配置 + +pub mod rate_limit; +pub mod sticky_config; +pub mod sticky_manager; + +pub use rate_limit::{ + extract_retry_delay, parse_duration_string, RateLimitReason, RateLimitRecord, RateLimitTracker, +}; +pub use sticky_config::{SchedulingMode, StickySessionConfig}; +pub use sticky_manager::{AccountInfo, StickySessionManager}; diff --git a/src-tauri/src/session/rate_limit.rs b/src-tauri/crates/core/src/session/rate_limit.rs similarity index 100% rename from src-tauri/src/session/rate_limit.rs rename to src-tauri/crates/core/src/session/rate_limit.rs diff --git a/src-tauri/src/session/sticky_config.rs b/src-tauri/crates/core/src/session/sticky_config.rs similarity index 100% rename from src-tauri/src/session/sticky_config.rs rename to src-tauri/crates/core/src/session/sticky_config.rs diff --git a/src-tauri/src/session/sticky_manager.rs b/src-tauri/crates/core/src/session/sticky_manager.rs similarity index 100% rename from src-tauri/src/session/sticky_manager.rs rename to src-tauri/crates/core/src/session/sticky_manager.rs diff --git a/src-tauri/src/session_files/mod.rs b/src-tauri/crates/core/src/session_files/mod.rs similarity index 97% rename from src-tauri/src/session_files/mod.rs rename to src-tauri/crates/core/src/session_files/mod.rs index 97c758312..7e9ebd41e 100644 --- a/src-tauri/src/session_files/mod.rs +++ b/src-tauri/crates/core/src/session_files/mod.rs @@ -3,7 +3,7 @@ //! 为每个 Agent 会话提供独立的临时工作目录, //! //! ## 目录结构 -//! ``` +//! ```text //! ~/.proxycast/sessions/ //! ├── {session-id}/ //! │ ├── .meta.json # 会话元数据 diff --git a/src-tauri/src/session_files/storage.rs b/src-tauri/crates/core/src/session_files/storage.rs similarity index 100% rename from src-tauri/src/session_files/storage.rs rename to src-tauri/crates/core/src/session_files/storage.rs diff --git a/src-tauri/src/session_files/types.rs b/src-tauri/crates/core/src/session_files/types.rs similarity index 100% rename from src-tauri/src/session_files/types.rs rename to src-tauri/crates/core/src/session_files/types.rs diff --git a/src-tauri/crates/infra/src/injection/types.rs b/src-tauri/crates/infra/src/injection/types.rs index 178253c23..27ecb9f0b 100644 --- a/src-tauri/crates/infra/src/injection/types.rs +++ b/src-tauri/crates/infra/src/injection/types.rs @@ -1,9 +1,13 @@ //! 参数注入类型定义 //! -//! 定义注入规则、注入模式和注入器 +//! 基础类型(InjectionMode, InjectionRule)从 proxycast-core 重新导出。 +//! 本模块定义注入器(Injector)和注入结果等 infra 层特有类型。 use serde::{Deserialize, Serialize}; +// 从 core 重新导出基础类型 +pub use proxycast_core::models::injection_types::{InjectionMode, InjectionRule}; + /// 允许注入的参数白名单 /// 这些参数是安全的,不会影响请求的核心行为 const ALLOWED_INJECTION_PARAMS: &[&str] = &[ @@ -28,110 +32,6 @@ const BLOCKED_OVERRIDE_PARAMS: &[&str] = &[ "response_format", ]; -/// 注入模式 -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] -#[serde(rename_all = "lowercase")] -pub enum InjectionMode { - /// 合并模式:不覆盖已有参数 - #[default] - Merge, - /// 覆盖模式:覆盖已有参数 - Override, -} - -/// 注入规则 -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] -pub struct InjectionRule { - /// 规则 ID - pub id: String, - /// 模型匹配模式(支持通配符) - pub pattern: String, - /// 要注入的参数 - pub parameters: serde_json::Value, - /// 注入模式 - #[serde(default)] - pub mode: InjectionMode, - /// 优先级(数字越小优先级越高) - #[serde(default = "default_priority")] - pub priority: i32, - /// 是否启用 - #[serde(default = "default_enabled")] - pub enabled: bool, -} - -fn default_priority() -> i32 { - 100 -} - -fn default_enabled() -> bool { - true -} - -impl InjectionRule { - /// 创建新的注入规则 - pub fn new(id: &str, pattern: &str, parameters: serde_json::Value) -> Self { - Self { - id: id.to_string(), - pattern: pattern.to_string(), - parameters, - mode: InjectionMode::Merge, - priority: default_priority(), - enabled: true, - } - } - - /// 设置注入模式 - pub fn with_mode(mut self, mode: InjectionMode) -> Self { - self.mode = mode; - self - } - - /// 设置优先级 - pub fn with_priority(mut self, priority: i32) -> Self { - self.priority = priority; - self - } - - /// 检查模型是否匹配此规则 - /// - /// 支持的通配符模式: - /// - 精确匹配: `claude-sonnet-4-5` - /// - 前缀匹配: `claude-*` - /// - 后缀匹配: `*-preview` - /// - 包含匹配: `*flash*` - pub fn matches(&self, model: &str) -> bool { - if !self.enabled { - return false; - } - pattern_matches(&self.pattern, model) - } - - /// 检查是否为精确匹配规则 - pub fn is_exact(&self) -> bool { - !self.pattern.contains('*') - } -} - -/// 规则排序:精确匹配优先,然后按优先级 -impl Ord for InjectionRule { - fn cmp(&self, other: &Self) -> std::cmp::Ordering { - match (self.is_exact(), other.is_exact()) { - (true, false) => return std::cmp::Ordering::Less, - (false, true) => return std::cmp::Ordering::Greater, - _ => {} - } - self.priority.cmp(&other.priority) - } -} - -impl PartialOrd for InjectionRule { - fn partial_cmp(&self, other: &Self) -> Option { - Some(self.cmp(other)) - } -} - -impl Eq for InjectionRule {} - /// 注入结果 #[derive(Debug, Clone, Default, Serialize, Deserialize)] pub struct InjectionResult { @@ -164,6 +64,10 @@ pub struct InjectionConfig { pub rules: Vec, } +fn default_enabled() -> bool { + true +} + /// 参数注入器 #[derive(Debug, Clone, Default)] pub struct Injector { @@ -273,26 +177,3 @@ impl Injector { result } } - -/// 检查模式是否匹配模型名 -/// -/// 支持的通配符模式: -/// - 精确匹配: `claude-sonnet-4-5` -/// - 前缀匹配: `claude-*` -/// - 后缀匹配: `*-preview` -/// - 包含匹配: `*flash*` -fn pattern_matches(pattern: &str, model: &str) -> bool { - if !pattern.contains('*') { - return pattern == model; - } - - let parts: Vec<&str> = pattern.split('*').collect(); - - match parts.as_slice() { - [prefix, ""] => model.starts_with(prefix), - ["", suffix] => model.ends_with(suffix), - ["", middle, ""] => model.contains(middle), - [prefix, suffix] => model.starts_with(prefix) && model.ends_with(suffix), - _ => false, - } -} diff --git a/src-tauri/crates/providers/Cargo.toml b/src-tauri/crates/providers/Cargo.toml new file mode 100644 index 000000000..55568fae1 --- /dev/null +++ b/src-tauri/crates/providers/Cargo.toml @@ -0,0 +1,55 @@ +[package] +name = "proxycast-providers" +version.workspace = true +edition.workspace = true +authors.workspace = true +repository.workspace = true + +[dependencies] +# 项目内 crate +proxycast-core.workspace = true + +# 序列化 +serde.workspace = true +serde_json.workspace = true +serde_urlencoded.workspace = true + +# 异步运行时 +tokio.workspace = true +futures.workspace = true +async-stream.workspace = true +async-trait.workspace = true + +# 错误处理 +anyhow.workspace = true +thiserror.workspace = true + +# 日志 +tracing.workspace = true + +# HTTP 客户端 +reqwest.workspace = true + +# 时间和 UUID +chrono.workspace = true +uuid.workspace = true + +# 工具库 +base64.workspace = true +bytes.workspace = true +dirs.workspace = true +flate2.workspace = true +sha2.workspace = true +rand.workspace = true +open.workspace = true +urlencoding.workspace = true +regex.workspace = true +once_cell.workspace = true +url.workspace = true + +# HTTP 服务器(streaming 模块需要 axum 类型) +axum.workspace = true + +[dev-dependencies] +proptest.workspace = true +tempfile.workspace = true diff --git a/src-tauri/src/converter/README.md b/src-tauri/crates/providers/src/converter/README.md similarity index 100% rename from src-tauri/src/converter/README.md rename to src-tauri/crates/providers/src/converter/README.md diff --git a/src-tauri/src/converter/anthropic_to_openai.rs b/src-tauri/crates/providers/src/converter/anthropic_to_openai.rs similarity index 98% rename from src-tauri/src/converter/anthropic_to_openai.rs rename to src-tauri/crates/providers/src/converter/anthropic_to_openai.rs index 230c33e9b..f217c6711 100644 --- a/src-tauri/src/converter/anthropic_to_openai.rs +++ b/src-tauri/crates/providers/src/converter/anthropic_to_openai.rs @@ -1,6 +1,6 @@ //! Anthropic 格式转换为 OpenAI 格式 (支持 Claude Code) -use crate::models::anthropic::*; -use crate::models::openai::*; +use proxycast_core::models::anthropic::*; +use proxycast_core::models::openai::*; use uuid::Uuid; /// 将 Anthropic MessagesRequest 转换为 OpenAI ChatCompletionRequest diff --git a/src-tauri/src/converter/cw_to_openai.rs b/src-tauri/crates/providers/src/converter/cw_to_openai.rs similarity index 98% rename from src-tauri/src/converter/cw_to_openai.rs rename to src-tauri/crates/providers/src/converter/cw_to_openai.rs index 77c9ec0f8..a15753c0f 100644 --- a/src-tauri/src/converter/cw_to_openai.rs +++ b/src-tauri/crates/providers/src/converter/cw_to_openai.rs @@ -1,8 +1,8 @@ //! CodeWhisperer 响应转换为 OpenAI 格式 #![allow(dead_code)] -use crate::models::codewhisperer::*; -use crate::models::openai::*; +use proxycast_core::models::codewhisperer::*; +use proxycast_core::models::openai::*; use std::time::{SystemTime, UNIX_EPOCH}; use uuid::Uuid; diff --git a/src-tauri/src/converter/mod.rs b/src-tauri/crates/providers/src/converter/mod.rs similarity index 100% rename from src-tauri/src/converter/mod.rs rename to src-tauri/crates/providers/src/converter/mod.rs diff --git a/src-tauri/src/converter/openai_to_antigravity.rs b/src-tauri/crates/providers/src/converter/openai_to_antigravity.rs similarity index 99% rename from src-tauri/src/converter/openai_to_antigravity.rs rename to src-tauri/crates/providers/src/converter/openai_to_antigravity.rs index a128b44bd..89db19758 100644 --- a/src-tauri/src/converter/openai_to_antigravity.rs +++ b/src-tauri/crates/providers/src/converter/openai_to_antigravity.rs @@ -14,8 +14,8 @@ //! ## 更新日志 //! - 2025-12-28: 修复请求格式,对齐 CLIProxyAPI 实现 -use crate::models::openai::*; use crate::session::{get_thought_signature, SessionManager}; +use proxycast_core::models::openai::*; use serde::{Deserialize, Serialize}; use uuid::Uuid; @@ -1061,7 +1061,7 @@ pub fn convert_antigravity_to_openai_response( // 图像生成 API 转换函数 // ============================================================================ -use crate::models::openai::{ImageData, ImageGenerationRequest, ImageGenerationResponse}; +use proxycast_core::models::openai::{ImageData, ImageGenerationRequest, ImageGenerationResponse}; /// 图像生成模型名称映射 /// diff --git a/src-tauri/src/converter/openai_to_cw.rs b/src-tauri/crates/providers/src/converter/openai_to_cw.rs similarity index 99% rename from src-tauri/src/converter/openai_to_cw.rs rename to src-tauri/crates/providers/src/converter/openai_to_cw.rs index b6183fca7..68ebb0d27 100644 --- a/src-tauri/src/converter/openai_to_cw.rs +++ b/src-tauri/crates/providers/src/converter/openai_to_cw.rs @@ -8,8 +8,8 @@ #![allow(dead_code)] -use crate::models::codewhisperer::*; -use crate::models::openai::*; +use proxycast_core::models::codewhisperer::*; +use proxycast_core::models::openai::*; use std::collections::HashMap; use uuid::Uuid; diff --git a/src-tauri/src/converter/protocol_selector.rs b/src-tauri/crates/providers/src/converter/protocol_selector.rs similarity index 99% rename from src-tauri/src/converter/protocol_selector.rs rename to src-tauri/crates/providers/src/converter/protocol_selector.rs index f97acb656..c77b7f724 100644 --- a/src-tauri/src/converter/protocol_selector.rs +++ b/src-tauri/crates/providers/src/converter/protocol_selector.rs @@ -4,7 +4,7 @@ #![allow(dead_code)] -use crate::models::provider_pool_model::PoolProviderType; +use proxycast_core::models::provider_pool_model::PoolProviderType; /// 协议类型 #[derive(Debug, Clone, Copy, PartialEq, Eq)] diff --git a/src-tauri/src/converter/reasoning_handler.rs b/src-tauri/crates/providers/src/converter/reasoning_handler.rs similarity index 98% rename from src-tauri/src/converter/reasoning_handler.rs rename to src-tauri/crates/providers/src/converter/reasoning_handler.rs index 4eebdf735..3aa53898e 100644 --- a/src-tauri/src/converter/reasoning_handler.rs +++ b/src-tauri/crates/providers/src/converter/reasoning_handler.rs @@ -24,7 +24,7 @@ // 预留功能模块,暂未在主流程中调用 #![allow(dead_code)] -use crate::models::openai::ChatMessage; +use proxycast_core::models::openai::ChatMessage; /// 模型类型,用于确定推理内容处理策略 #[derive(Debug, Clone, PartialEq)] @@ -147,7 +147,7 @@ impl ReasoningHandler { #[cfg(test)] mod tests { use super::*; - use crate::models::openai::MessageContent; + use proxycast_core::models::openai::MessageContent; #[test] fn test_model_type_detection() { diff --git a/src-tauri/crates/providers/src/lib.rs b/src-tauri/crates/providers/src/lib.rs new file mode 100644 index 000000000..b5f2de577 --- /dev/null +++ b/src-tauri/crates/providers/src/lib.rs @@ -0,0 +1,18 @@ +//! ProxyCast Providers Crate +//! +//! 包含所有 Provider 实现、协议转换、流式传输等核心业务模块。 +//! +//! ## 模块结构 +//! - `providers`: Provider 实现(Kiro、Gemini、Claude、OpenAI、Vertex 等) +//! - `converter`: 协议转换(OpenAI ↔ CW、OpenAI ↔ Antigravity 等) +//! - `streaming`: 流式传输管理 +//! - `translator`: 请求/响应翻译层 +//! - `stream`: 流事件解析和生成 +//! - `session`: 会话管理(签名存储、会话 ID 生成) + +pub mod converter; +pub mod providers; +pub mod session; +pub mod stream; +pub mod streaming; +pub mod translator; diff --git a/src-tauri/src/providers/README.md b/src-tauri/crates/providers/src/providers/README.md similarity index 100% rename from src-tauri/src/providers/README.md rename to src-tauri/crates/providers/src/providers/README.md diff --git a/src-tauri/src/providers/antigravity.rs b/src-tauri/crates/providers/src/providers/antigravity.rs similarity index 99% rename from src-tauri/src/providers/antigravity.rs rename to src-tauri/crates/providers/src/providers/antigravity.rs index 33d572b2a..38ecade0c 100644 --- a/src-tauri/src/providers/antigravity.rs +++ b/src-tauri/crates/providers/src/providers/antigravity.rs @@ -2109,11 +2109,11 @@ impl CredentialProvider for AntigravityProvider { // ============================================================================ use crate::converter::openai_to_antigravity::convert_openai_to_antigravity_with_context; -use crate::models::openai::ChatCompletionRequest; use crate::providers::ProviderError; use crate::streaming::traits::{ reqwest_stream_to_stream_response, StreamFormat, StreamResponse, StreamingProvider, }; +use proxycast_core::models::openai::ChatCompletionRequest; #[async_trait] impl StreamingProvider for AntigravityProvider { diff --git a/src-tauri/src/providers/claude_custom.rs b/src-tauri/crates/providers/src/providers/claude_custom.rs similarity index 98% rename from src-tauri/src/providers/claude_custom.rs rename to src-tauri/crates/providers/src/providers/claude_custom.rs index 64af8bbbe..f1e882008 100644 --- a/src-tauri/src/providers/claude_custom.rs +++ b/src-tauri/crates/providers/src/providers/claude_custom.rs @@ -1,6 +1,6 @@ //! Claude Custom Provider (自定义 Claude API) -use crate::models::anthropic::AnthropicMessagesRequest; -use crate::models::openai::{ChatCompletionRequest, ContentPart, MessageContent}; +use proxycast_core::models::anthropic::AnthropicMessagesRequest; +use proxycast_core::models::openai::{ChatCompletionRequest, ContentPart, MessageContent}; use reqwest::Client; use serde::{Deserialize, Serialize}; use std::error::Error; @@ -560,7 +560,7 @@ impl StreamingProvider for ClaudeCustomProvider { .iter() .filter_map(|tool| { match tool { - crate::models::openai::Tool::Function { function } => { + proxycast_core::models::openai::Tool::Function { function } => { Some(serde_json::json!({ "name": function.name, "description": function.description.clone().unwrap_or_default(), diff --git a/src-tauri/src/providers/claude_oauth.rs b/src-tauri/crates/providers/src/providers/claude_oauth.rs similarity index 100% rename from src-tauri/src/providers/claude_oauth.rs rename to src-tauri/crates/providers/src/providers/claude_oauth.rs diff --git a/src-tauri/src/providers/codex.rs b/src-tauri/crates/providers/src/providers/codex.rs similarity index 99% rename from src-tauri/src/providers/codex.rs rename to src-tauri/crates/providers/src/providers/codex.rs index 958094488..e788bfde6 100644 --- a/src-tauri/src/providers/codex.rs +++ b/src-tauri/crates/providers/src/providers/codex.rs @@ -457,7 +457,7 @@ impl CodexProvider { .filter(|s| !s.is_empty()) } - pub(crate) fn build_responses_url(base_url: &str) -> String { + pub fn build_responses_url(base_url: &str) -> String { let base = base_url.trim_end_matches('/'); // 规则说明: diff --git a/src-tauri/src/providers/error.rs b/src-tauri/crates/providers/src/providers/error.rs similarity index 100% rename from src-tauri/src/providers/error.rs rename to src-tauri/crates/providers/src/providers/error.rs diff --git a/src-tauri/src/providers/gemini.rs b/src-tauri/crates/providers/src/providers/gemini.rs similarity index 100% rename from src-tauri/src/providers/gemini.rs rename to src-tauri/crates/providers/src/providers/gemini.rs diff --git a/src-tauri/src/providers/kiro.rs b/src-tauri/crates/providers/src/providers/kiro.rs similarity index 99% rename from src-tauri/src/providers/kiro.rs rename to src-tauri/crates/providers/src/providers/kiro.rs index 8f000e155..a78b1aba6 100644 --- a/src-tauri/src/providers/kiro.rs +++ b/src-tauri/crates/providers/src/providers/kiro.rs @@ -3,18 +3,18 @@ #![allow(dead_code)] // 使用新的 translator 模块替代旧的 converter -use crate::models::anthropic::AnthropicMessagesRequest; -use crate::models::openai::*; use crate::providers::traits::{CredentialProvider, ProviderResult}; use crate::translator::kiro::anthropic::request::convert_anthropic_to_codewhisperer; use crate::translator::kiro::openai::request::convert_openai_to_codewhisperer; use async_trait::async_trait; +use proxycast_core::models::anthropic::AnthropicMessagesRequest; +use proxycast_core::models::openai::*; use reqwest::Client; use serde::{Deserialize, Serialize}; use std::error::Error; use std::path::PathBuf; -/// 根据凭证信息生成唯一的 Machine ID(与 AIClient-2-API 保持一致) +/// 根据凭证信息生成唯一的 Machine ID /// /// 采用静态 UUID 方案:每个凭证生成固定的 Machine ID,不随时间变化 /// 优先级:uuid > profileArn > clientId > 系统硬件 ID @@ -269,6 +269,16 @@ impl Default for KiroProvider { } } +impl Clone for KiroProvider { + fn clone(&self) -> Self { + Self { + credentials: self.credentials.clone(), + client: reqwest::Client::new(), + creds_path: self.creds_path.clone(), + } + } +} + impl KiroProvider { pub fn new() -> Self { Self::default() diff --git a/src-tauri/src/providers/mod.rs b/src-tauri/crates/providers/src/providers/mod.rs similarity index 100% rename from src-tauri/src/providers/mod.rs rename to src-tauri/crates/providers/src/providers/mod.rs diff --git a/src-tauri/src/providers/openai_custom.rs b/src-tauri/crates/providers/src/providers/openai_custom.rs similarity index 99% rename from src-tauri/src/providers/openai_custom.rs rename to src-tauri/crates/providers/src/providers/openai_custom.rs index 5cb924ace..949ebb96b 100644 --- a/src-tauri/src/providers/openai_custom.rs +++ b/src-tauri/crates/providers/src/providers/openai_custom.rs @@ -1,5 +1,5 @@ //! OpenAI Custom Provider (自定义 OpenAI 兼容 API) -use crate::models::openai::ChatCompletionRequest; +use proxycast_core::models::openai::ChatCompletionRequest; use reqwest::Client; use reqwest::StatusCode; use serde::{Deserialize, Serialize}; diff --git a/src-tauri/src/providers/tests.rs b/src-tauri/crates/providers/src/providers/tests.rs similarity index 100% rename from src-tauri/src/providers/tests.rs rename to src-tauri/crates/providers/src/providers/tests.rs diff --git a/src-tauri/src/providers/traits.rs b/src-tauri/crates/providers/src/providers/traits.rs similarity index 100% rename from src-tauri/src/providers/traits.rs rename to src-tauri/crates/providers/src/providers/traits.rs diff --git a/src-tauri/src/providers/vertex.rs b/src-tauri/crates/providers/src/providers/vertex.rs similarity index 98% rename from src-tauri/src/providers/vertex.rs rename to src-tauri/crates/providers/src/providers/vertex.rs index d99d2cfba..d5ed248e5 100644 --- a/src-tauri/src/providers/vertex.rs +++ b/src-tauri/crates/providers/src/providers/vertex.rs @@ -5,7 +5,7 @@ #![allow(dead_code)] -use crate::config::VertexApiKeyEntry; +use proxycast_core::models::vertex_model::VertexApiKeyEntry; use reqwest::Client; use serde::{Deserialize, Serialize}; use std::collections::HashMap; @@ -290,7 +290,7 @@ mod tests { #[test] fn test_vertex_provider_from_entry() { - use crate::config::VertexModelAlias; + use proxycast_core::models::vertex_model::VertexModelAlias; let entry = VertexApiKeyEntry { id: "test-vertex".to_string(), diff --git a/src-tauri/crates/providers/src/session/mod.rs b/src-tauri/crates/providers/src/session/mod.rs new file mode 100644 index 000000000..2295fdd13 --- /dev/null +++ b/src-tauri/crates/providers/src/session/mod.rs @@ -0,0 +1,13 @@ +//! 会话管理模块(providers crate 部分) +//! +//! 包含 signature_store 和 session_manager, +//! 这两个模块被 converter 和 streaming 直接使用。 + +pub mod session_manager; +pub mod signature_store; + +pub use session_manager::SessionManager; +pub use signature_store::{ + clear_thought_signature, get_thought_signature, has_valid_signature, store_thought_signature, + take_thought_signature, +}; diff --git a/src-tauri/src/session/session_manager.rs b/src-tauri/crates/providers/src/session/session_manager.rs similarity index 95% rename from src-tauri/src/session/session_manager.rs rename to src-tauri/crates/providers/src/session/session_manager.rs index c3baf65db..7b90b07bc 100644 --- a/src-tauri/src/session/session_manager.rs +++ b/src-tauri/crates/providers/src/session/session_manager.rs @@ -3,7 +3,7 @@ //! 根据请求内容生成稳定的会话指纹(Session Fingerprint), //! 用于实现会话粘性和 Prompt Caching 优化。 -use crate::models::openai::ChatCompletionRequest; +use proxycast_core::models::openai::ChatCompletionRequest; use sha2::{Digest, Sha256}; /// 会话管理器 @@ -155,7 +155,7 @@ impl SessionManager { #[cfg(test)] mod tests { use super::*; - use crate::models::openai::{ChatCompletionRequest, ChatMessage}; + use proxycast_core::models::openai::{ChatCompletionRequest, ChatMessage}; #[test] fn test_session_id_stability() { @@ -163,7 +163,7 @@ mod tests { model: "gpt-4".to_string(), messages: vec![ChatMessage { role: "user".to_string(), - content: Some(crate::models::openai::MessageContent::Text( + content: Some(proxycast_core::models::openai::MessageContent::Text( "Hello, how are you?".to_string(), )), tool_calls: None, @@ -195,7 +195,7 @@ mod tests { model: "gpt-4".to_string(), messages: vec![ChatMessage { role: "user".to_string(), - content: Some(crate::models::openai::MessageContent::Text( + content: Some(proxycast_core::models::openai::MessageContent::Text( "Hello, how are you?".to_string(), )), tool_calls: None, @@ -215,7 +215,7 @@ mod tests { model: "gpt-4".to_string(), messages: vec![ChatMessage { role: "user".to_string(), - content: Some(crate::models::openai::MessageContent::Text( + content: Some(proxycast_core::models::openai::MessageContent::Text( "What is the weather today?".to_string(), )), tool_calls: None, diff --git a/src-tauri/src/session/signature_store.rs b/src-tauri/crates/providers/src/session/signature_store.rs similarity index 100% rename from src-tauri/src/session/signature_store.rs rename to src-tauri/crates/providers/src/session/signature_store.rs diff --git a/src-tauri/src/stream/events.rs b/src-tauri/crates/providers/src/stream/events.rs similarity index 100% rename from src-tauri/src/stream/events.rs rename to src-tauri/crates/providers/src/stream/events.rs diff --git a/src-tauri/src/stream/generators/anthropic_sse.rs b/src-tauri/crates/providers/src/stream/generators/anthropic_sse.rs similarity index 100% rename from src-tauri/src/stream/generators/anthropic_sse.rs rename to src-tauri/crates/providers/src/stream/generators/anthropic_sse.rs diff --git a/src-tauri/src/stream/generators/mod.rs b/src-tauri/crates/providers/src/stream/generators/mod.rs similarity index 100% rename from src-tauri/src/stream/generators/mod.rs rename to src-tauri/crates/providers/src/stream/generators/mod.rs diff --git a/src-tauri/src/stream/generators/openai_sse.rs b/src-tauri/crates/providers/src/stream/generators/openai_sse.rs similarity index 100% rename from src-tauri/src/stream/generators/openai_sse.rs rename to src-tauri/crates/providers/src/stream/generators/openai_sse.rs diff --git a/src-tauri/src/stream/mod.rs b/src-tauri/crates/providers/src/stream/mod.rs similarity index 100% rename from src-tauri/src/stream/mod.rs rename to src-tauri/crates/providers/src/stream/mod.rs diff --git a/src-tauri/src/stream/parsers/aws_event_stream.rs b/src-tauri/crates/providers/src/stream/parsers/aws_event_stream.rs similarity index 100% rename from src-tauri/src/stream/parsers/aws_event_stream.rs rename to src-tauri/crates/providers/src/stream/parsers/aws_event_stream.rs diff --git a/src-tauri/src/stream/parsers/mod.rs b/src-tauri/crates/providers/src/stream/parsers/mod.rs similarity index 100% rename from src-tauri/src/stream/parsers/mod.rs rename to src-tauri/crates/providers/src/stream/parsers/mod.rs diff --git a/src-tauri/src/stream/pipeline.rs b/src-tauri/crates/providers/src/stream/pipeline.rs similarity index 100% rename from src-tauri/src/stream/pipeline.rs rename to src-tauri/crates/providers/src/stream/pipeline.rs diff --git a/src-tauri/src/streaming/anthropic_sse.rs b/src-tauri/crates/providers/src/streaming/anthropic_sse.rs similarity index 100% rename from src-tauri/src/streaming/anthropic_sse.rs rename to src-tauri/crates/providers/src/streaming/anthropic_sse.rs diff --git a/src-tauri/src/streaming/aws_parser.rs b/src-tauri/crates/providers/src/streaming/aws_parser.rs similarity index 100% rename from src-tauri/src/streaming/aws_parser.rs rename to src-tauri/crates/providers/src/streaming/aws_parser.rs diff --git a/src-tauri/src/streaming/converter.rs b/src-tauri/crates/providers/src/streaming/converter.rs similarity index 100% rename from src-tauri/src/streaming/converter.rs rename to src-tauri/crates/providers/src/streaming/converter.rs diff --git a/src-tauri/src/streaming/error.rs b/src-tauri/crates/providers/src/streaming/error.rs similarity index 100% rename from src-tauri/src/streaming/error.rs rename to src-tauri/crates/providers/src/streaming/error.rs diff --git a/src-tauri/src/streaming/manager.rs b/src-tauri/crates/providers/src/streaming/manager.rs similarity index 100% rename from src-tauri/src/streaming/manager.rs rename to src-tauri/crates/providers/src/streaming/manager.rs diff --git a/src-tauri/src/streaming/metrics.rs b/src-tauri/crates/providers/src/streaming/metrics.rs similarity index 100% rename from src-tauri/src/streaming/metrics.rs rename to src-tauri/crates/providers/src/streaming/metrics.rs diff --git a/src-tauri/src/streaming/mod.rs b/src-tauri/crates/providers/src/streaming/mod.rs similarity index 100% rename from src-tauri/src/streaming/mod.rs rename to src-tauri/crates/providers/src/streaming/mod.rs diff --git a/src-tauri/src/streaming/traits.rs b/src-tauri/crates/providers/src/streaming/traits.rs similarity index 98% rename from src-tauri/src/streaming/traits.rs rename to src-tauri/crates/providers/src/streaming/traits.rs index f1db186fa..b319bfdb5 100644 --- a/src-tauri/src/streaming/traits.rs +++ b/src-tauri/crates/providers/src/streaming/traits.rs @@ -11,12 +11,12 @@ #![allow(dead_code)] -use crate::models::openai::ChatCompletionRequest; use crate::providers::ProviderError; use crate::streaming::StreamError; use async_trait::async_trait; use bytes::Bytes; use futures::Stream; +use proxycast_core::models::openai::ChatCompletionRequest; use std::pin::Pin; /// 流式响应类型别名 diff --git a/src-tauri/src/translator/kiro/anthropic/mod.rs b/src-tauri/crates/providers/src/translator/kiro/anthropic/mod.rs similarity index 100% rename from src-tauri/src/translator/kiro/anthropic/mod.rs rename to src-tauri/crates/providers/src/translator/kiro/anthropic/mod.rs diff --git a/src-tauri/src/translator/kiro/anthropic/request.rs b/src-tauri/crates/providers/src/translator/kiro/anthropic/request.rs similarity index 99% rename from src-tauri/src/translator/kiro/anthropic/request.rs rename to src-tauri/crates/providers/src/translator/kiro/anthropic/request.rs index d98b78b95..c80956a74 100644 --- a/src-tauri/src/translator/kiro/anthropic/request.rs +++ b/src-tauri/crates/providers/src/translator/kiro/anthropic/request.rs @@ -3,10 +3,10 @@ //! 直接将 Anthropic MessagesRequest 转换为 CodeWhisperer API 格式, //! 无需经过 OpenAI 中间格式,减少转换开销。 -use crate::models::anthropic::*; -use crate::models::codewhisperer::*; use crate::translator::kiro::openai::request::{get_model_map, DEFAULT_MODEL}; use crate::translator::traits::{RequestTranslator, TranslateError}; +use proxycast_core::models::anthropic::*; +use proxycast_core::models::codewhisperer::*; use std::collections::HashSet; use uuid::Uuid; diff --git a/src-tauri/src/translator/kiro/anthropic/response.rs b/src-tauri/crates/providers/src/translator/kiro/anthropic/response.rs similarity index 100% rename from src-tauri/src/translator/kiro/anthropic/response.rs rename to src-tauri/crates/providers/src/translator/kiro/anthropic/response.rs diff --git a/src-tauri/src/translator/kiro/mod.rs b/src-tauri/crates/providers/src/translator/kiro/mod.rs similarity index 100% rename from src-tauri/src/translator/kiro/mod.rs rename to src-tauri/crates/providers/src/translator/kiro/mod.rs diff --git a/src-tauri/src/translator/kiro/openai/mod.rs b/src-tauri/crates/providers/src/translator/kiro/openai/mod.rs similarity index 100% rename from src-tauri/src/translator/kiro/openai/mod.rs rename to src-tauri/crates/providers/src/translator/kiro/openai/mod.rs diff --git a/src-tauri/src/translator/kiro/openai/request.rs b/src-tauri/crates/providers/src/translator/kiro/openai/request.rs similarity index 99% rename from src-tauri/src/translator/kiro/openai/request.rs rename to src-tauri/crates/providers/src/translator/kiro/openai/request.rs index 66d78c6ff..6c12f91e7 100644 --- a/src-tauri/src/translator/kiro/openai/request.rs +++ b/src-tauri/crates/providers/src/translator/kiro/openai/request.rs @@ -9,9 +9,9 @@ //! - claude-sonnet-4-20250514 → CLAUDE_SONNET_4_20250514_V1_0 //! - claude-haiku-4-5 → claude-haiku-4.5 -use crate::models::codewhisperer::*; -use crate::models::openai::*; use crate::translator::traits::{RequestTranslator, TranslateError}; +use proxycast_core::models::codewhisperer::*; +use proxycast_core::models::openai::*; use std::collections::{HashMap, HashSet}; use uuid::Uuid; diff --git a/src-tauri/src/translator/kiro/openai/response.rs b/src-tauri/crates/providers/src/translator/kiro/openai/response.rs similarity index 100% rename from src-tauri/src/translator/kiro/openai/response.rs rename to src-tauri/crates/providers/src/translator/kiro/openai/response.rs diff --git a/src-tauri/src/translator/mod.rs b/src-tauri/crates/providers/src/translator/mod.rs similarity index 100% rename from src-tauri/src/translator/mod.rs rename to src-tauri/crates/providers/src/translator/mod.rs diff --git a/src-tauri/src/translator/traits.rs b/src-tauri/crates/providers/src/translator/traits.rs similarity index 100% rename from src-tauri/src/translator/traits.rs rename to src-tauri/crates/providers/src/translator/traits.rs diff --git a/src-tauri/proptest-regressions/config/tests.txt b/src-tauri/proptest-regressions/config/tests.txt deleted file mode 100644 index 848ff117f..000000000 --- a/src-tauri/proptest-regressions/config/tests.txt +++ /dev/null @@ -1,7 +0,0 @@ -# Seeds for failure cases proptest has generated in the past. It is -# automatically read and these particular cases re-run before any -# novel cases are generated. -# -# It is recommended to check this file in to source control so that -# everyone who runs the test benefits from these saved cases. -cc 690e03d4ddd600c7175d123787d49a9afacf6f13642e9f7ae0e9e766a2ba93c2 # shrinks to provider = "qwen" diff --git a/src-tauri/src/app/bootstrap.rs b/src-tauri/src/app/bootstrap.rs index 331b6ebe8..f99062c9c 100644 --- a/src-tauri/src/app/bootstrap.rs +++ b/src-tauri/src/app/bootstrap.rs @@ -9,11 +9,6 @@ use crate::agent::AsterAgentState; use crate::commands::api_key_provider_cmd::ApiKeyProviderServiceState; use crate::commands::connect_cmd::ConnectStateWrapper; use crate::commands::context_memory::ContextMemoryServiceState; -use crate::commands::flow_monitor_cmd::{ - BatchOperationsState, BookmarkManagerState, EnhancedStatsServiceState, FlowInterceptorState, - FlowMonitorState, FlowQueryServiceState, FlowReplayerState, QuickFilterManagerState, - SessionManagerState, -}; use crate::commands::machine_id_cmd::MachineIdState; use crate::commands::model_registry_cmd::ModelRegistryState; use crate::commands::orchestrator_cmd::OrchestratorState; @@ -28,11 +23,6 @@ use crate::commands::tool_hooks::ToolHooksServiceState; use crate::commands::webview_cmd::{WebviewManagerState, WebviewManagerWrapper}; use crate::config::{self, Config, ConfigManager, GlobalConfigManager, GlobalConfigManagerState}; use crate::database::{self, DbConnection}; -use crate::flow_monitor::{ - BatchOperations, BookmarkManager, EnhancedStatsService, FlowFileStore, FlowInterceptor, - FlowMonitor, FlowMonitorConfig, FlowQueryService, FlowReplayer, InterceptConfig, - QuickFilterManager, RotationConfig, SessionManager, -}; use crate::logger; use crate::mcp::McpManagerState; use crate::plugin; @@ -131,15 +121,6 @@ pub struct AppStates { pub plugin_installer: PluginInstallerState, pub plugin_rpc_manager: crate::commands::plugin_rpc_cmd::PluginRpcManagerState, pub telemetry: crate::commands::telemetry_cmd::TelemetryState, - pub flow_monitor: FlowMonitorState, - pub flow_query_service: FlowQueryServiceState, - pub flow_interceptor: FlowInterceptorState, - pub flow_replayer: FlowReplayerState, - pub session_manager: SessionManagerState, - pub quick_filter_manager: QuickFilterManagerState, - pub bookmark_manager: BookmarkManagerState, - pub enhanced_stats_service: EnhancedStatsServiceState, - pub batch_operations: BatchOperationsState, pub aster_agent: AsterAgentState, pub orchestrator: OrchestratorState, pub connect_state: ConnectStateWrapper, @@ -157,8 +138,6 @@ pub struct AppStates { pub shared_stats: Arc>, pub shared_tokens: Arc>, pub shared_logger: Arc, - pub flow_monitor_arc: Arc, - pub flow_interceptor_arc: Arc, } /// 初始化所有应用状态 @@ -205,21 +184,6 @@ pub fn init_states(config: &Config) -> Result { // 遥测系统 let (telemetry_state, shared_stats, shared_tokens, shared_logger) = init_telemetry(config)?; - // Flow Monitor 系统(根据插件安装状态启用/禁用) - let ( - flow_monitor_state, - flow_query_service_state, - flow_interceptor_state, - flow_replayer_state, - session_manager_state, - quick_filter_manager_state, - bookmark_manager_state, - enhanced_stats_service_state, - batch_operations_state, - flow_monitor_arc, - flow_interceptor_arc, - ) = init_flow_monitor(&provider_pool_service_state, &db, &plugin_installer_state)?; - // 其他状态 // 设置 Aster 全局 session store(使用 ProxyCast 数据库) let session_store = Arc::new(ProxyCastSessionStore::new(db.clone())); @@ -309,15 +273,6 @@ pub fn init_states(config: &Config) -> Result { plugin_installer: plugin_installer_state, plugin_rpc_manager: plugin_rpc_manager_state, telemetry: telemetry_state, - flow_monitor: flow_monitor_state, - flow_query_service: flow_query_service_state, - flow_interceptor: flow_interceptor_state, - flow_replayer: flow_replayer_state, - session_manager: session_manager_state, - quick_filter_manager: quick_filter_manager_state, - bookmark_manager: bookmark_manager_state, - enhanced_stats_service: enhanced_stats_service_state, - batch_operations: batch_operations_state, aster_agent: aster_agent_state, orchestrator: orchestrator_state, connect_state, @@ -334,8 +289,6 @@ pub fn init_states(config: &Config) -> Result { shared_stats, shared_tokens, shared_logger, - flow_monitor_arc, - flow_interceptor_arc, }) } @@ -416,132 +369,3 @@ fn init_telemetry( Ok((telemetry_state, shared_stats, shared_tokens, shared_logger)) } - -/// 初始化 Flow Monitor 系统 -/// -/// 如果 flow-monitor 插件已安装,则启用监控功能;否则禁用。 -#[allow(clippy::type_complexity)] -fn init_flow_monitor( - provider_pool_service_state: &ProviderPoolServiceState, - db: &DbConnection, - plugin_installer_state: &PluginInstallerState, -) -> Result< - ( - FlowMonitorState, - FlowQueryServiceState, - FlowInterceptorState, - FlowReplayerState, - SessionManagerState, - QuickFilterManagerState, - BookmarkManagerState, - EnhancedStatsServiceState, - BatchOperationsState, - Arc, - Arc, - ), - String, -> { - // 检查 flow-monitor 插件是否已安装 - let is_plugin_installed = { - let installer = plugin_installer_state.0.blocking_read(); - installer.is_installed("flow-monitor").unwrap_or(false) - }; - - // 根据插件安装状态设置 enabled - let mut flow_monitor_config = FlowMonitorConfig::default(); - flow_monitor_config.enabled = is_plugin_installed; - - if is_plugin_installed { - tracing::info!("[启动] flow-monitor 插件已安装,启用 Flow 监控"); - } else { - tracing::info!("[启动] flow-monitor 插件未安装,禁用 Flow 监控"); - } - - // 初始化文件存储 - let data_dir = dirs::data_dir() - .unwrap_or_else(|| std::path::PathBuf::from(".")) - .join("proxycast") - .join("flows"); - let _ = std::fs::create_dir_all(&data_dir); - - let rotation_config = RotationConfig::default(); - let flow_file_store = match FlowFileStore::new(data_dir, rotation_config.clone()) { - Ok(store) => Some(Arc::new(store)), - Err(e) => { - tracing::warn!("无法初始化 Flow 文件存储: {}", e); - None - } - }; - - let flow_monitor = Arc::new(FlowMonitor::new( - flow_monitor_config, - flow_file_store.clone(), - )); - let flow_monitor_state = FlowMonitorState(flow_monitor.clone()); - - let flow_interceptor = Arc::new(FlowInterceptor::new(InterceptConfig::default())); - let flow_interceptor_state = FlowInterceptorState(flow_interceptor.clone()); - - let flow_replayer = Arc::new(FlowReplayer::new( - flow_monitor.clone(), - provider_pool_service_state.0.clone(), - db.clone(), - )); - let flow_replayer_state = FlowReplayerState(flow_replayer); - - let db_path = database::get_db_path().map_err(|e| format!("获取数据库路径失败: {e}"))?; - - let session_manager = Arc::new( - SessionManager::new(db_path.clone()) - .map_err(|e| format!("SessionManager 初始化失败: {e}"))?, - ); - let session_manager_state = SessionManagerState(session_manager.clone()); - - let quick_filter_manager = Arc::new( - QuickFilterManager::new(db_path.clone()) - .map_err(|e| format!("QuickFilterManager 初始化失败: {e}"))?, - ); - let quick_filter_manager_state = QuickFilterManagerState(quick_filter_manager); - - let bookmark_manager = Arc::new( - BookmarkManager::new(db_path).map_err(|e| format!("BookmarkManager 初始化失败: {e}"))?, - ); - let bookmark_manager_state = BookmarkManagerState(bookmark_manager); - - let enhanced_stats_service = Arc::new(EnhancedStatsService::new(flow_monitor.memory_store())); - let enhanced_stats_service_state = EnhancedStatsServiceState(enhanced_stats_service); - - let batch_operations = Arc::new(BatchOperations::new( - flow_monitor.clone(), - Some(session_manager_state.0.clone()), - )); - let batch_operations_state = BatchOperationsState(batch_operations); - - // FlowQueryService - let flow_query_service_state = if let Some(file_store) = flow_file_store { - let query_service = FlowQueryService::new(flow_monitor.memory_store(), file_store); - FlowQueryServiceState(Arc::new(query_service)) - } else { - let temp_dir = std::env::temp_dir().join("proxycast_flows"); - let _ = std::fs::create_dir_all(&temp_dir); - let temp_store = FlowFileStore::new(temp_dir, rotation_config) - .map_err(|e| format!("临时 FlowFileStore 初始化失败: {e}"))?; - let query_service = - FlowQueryService::new(flow_monitor.memory_store(), Arc::new(temp_store)); - FlowQueryServiceState(Arc::new(query_service)) - }; - - Ok(( - flow_monitor_state, - flow_query_service_state, - flow_interceptor_state, - flow_replayer_state, - session_manager_state, - quick_filter_manager_state, - bookmark_manager_state, - enhanced_stats_service_state, - batch_operations_state, - flow_monitor, - flow_interceptor, - )) -} diff --git a/src-tauri/src/app/runner.rs b/src-tauri/src/app/runner.rs index e17bc1eeb..f18551a5b 100644 --- a/src-tauri/src/app/runner.rs +++ b/src-tauri/src/app/runner.rs @@ -61,15 +61,6 @@ pub fn run() { plugin_installer: plugin_installer_state, plugin_rpc_manager: plugin_rpc_manager_state, telemetry: telemetry_state, - flow_monitor: flow_monitor_state, - flow_query_service: flow_query_service_state, - flow_interceptor: flow_interceptor_state, - flow_replayer: flow_replayer_state, - session_manager: session_manager_state, - quick_filter_manager: quick_filter_manager_state, - bookmark_manager: bookmark_manager_state, - enhanced_stats_service: enhanced_stats_service_state, - batch_operations: batch_operations_state, aster_agent: aster_agent_state, orchestrator: orchestrator_state, connect_state, @@ -86,8 +77,6 @@ pub fn run() { shared_stats, shared_tokens, shared_logger, - flow_monitor_arc: flow_monitor, - flow_interceptor_arc: flow_interceptor, } = states; // Clone for setup hook @@ -99,8 +88,6 @@ pub fn run() { let shared_stats_clone = shared_stats.clone(); let shared_tokens_clone = shared_tokens.clone(); let shared_logger_clone = shared_logger.clone(); - let flow_monitor_clone = flow_monitor.clone(); - let flow_interceptor_clone = flow_interceptor.clone(); let update_check_service_clone = update_check_service_state.0.clone(); let mut builder = tauri::Builder::default() @@ -146,15 +133,6 @@ pub fn run() { .manage(plugin_manager_state) .manage(plugin_installer_state) .manage(plugin_rpc_manager_state) - .manage(flow_monitor_state) - .manage(flow_query_service_state) - .manage(flow_interceptor_state) - .manage(flow_replayer_state) - .manage(session_manager_state) - .manage(quick_filter_manager_state) - .manage(bookmark_manager_state) - .manage(enhanced_stats_service_state) - .manage(batch_operations_state) .manage(aster_agent_state) .manage(orchestrator_state) .manage(connect_state) @@ -469,7 +447,6 @@ pub fn run() { let shared_stats = shared_stats_clone.clone(); let shared_tokens = shared_tokens_clone.clone(); let shared_logger = shared_logger_clone.clone(); - let shared_flow_monitor = flow_monitor_clone.clone(); let app_handle = app.handle().clone(); tauri::async_runtime::spawn(async move { // 先加载凭证池中的凭证 @@ -527,7 +504,7 @@ pub fn run() { .add("debug", &format!("[启动] 旧版 Kiro 凭证加载失败: {e}")); } } - // 启动服务器(使用共享的遥测实例和 Flow Monitor) + // 启动服务器(使用共享的遥测实例) let server_started; let server_address; { @@ -536,7 +513,7 @@ pub fn run() { .await .add("info", "[启动] 正在自动启动服务器..."); match s - .start_with_telemetry_and_flow_monitor( + .start_with_telemetry( logs.clone(), pool_service, token_cache, @@ -544,8 +521,6 @@ pub fn run() { Some(shared_stats), Some(shared_tokens), Some(shared_logger), - Some(shared_flow_monitor), - Some(flow_interceptor_clone), ) .await { @@ -923,141 +898,12 @@ pub fn run() { commands::plugin_rpc_cmd::plugin_rpc_connect, commands::plugin_rpc_cmd::plugin_rpc_disconnect, commands::plugin_rpc_cmd::plugin_rpc_call, - // Flow Monitor commands - commands::flow_monitor_cmd::query_flows, - commands::flow_monitor_cmd::get_flow_detail, - commands::flow_monitor_cmd::search_flows, - commands::flow_monitor_cmd::get_flow_stats, - commands::flow_monitor_cmd::export_flows, - commands::flow_monitor_cmd::update_flow_annotations, - commands::flow_monitor_cmd::toggle_flow_starred, - commands::flow_monitor_cmd::add_flow_comment, - commands::flow_monitor_cmd::add_flow_tag, - commands::flow_monitor_cmd::remove_flow_tag, - commands::flow_monitor_cmd::set_flow_marker, - commands::flow_monitor_cmd::cleanup_flows, - commands::flow_monitor_cmd::get_recent_flows, - commands::flow_monitor_cmd::get_flow_monitor_status, - commands::flow_monitor_cmd::get_flow_monitor_debug_info, - commands::flow_monitor_cmd::create_test_flows, - commands::flow_monitor_cmd::enable_flow_monitor, - commands::flow_monitor_cmd::disable_flow_monitor, - commands::flow_monitor_cmd::subscribe_flow_events, - commands::flow_monitor_cmd::get_all_flow_tags, - // Flow Monitor filter expression commands - commands::flow_monitor_cmd::parse_filter, - commands::flow_monitor_cmd::validate_filter, - commands::flow_monitor_cmd::get_filter_help_items, - commands::flow_monitor_cmd::get_filter_help_text, - commands::flow_monitor_cmd::query_flows_with_expression, - // Flow Interceptor commands - commands::flow_monitor_cmd::intercept_config_get, - commands::flow_monitor_cmd::intercept_config_set, - commands::flow_monitor_cmd::intercept_continue, - commands::flow_monitor_cmd::intercept_cancel, - commands::flow_monitor_cmd::intercept_get_flow, - commands::flow_monitor_cmd::intercept_list_flows, - commands::flow_monitor_cmd::intercept_count, - commands::flow_monitor_cmd::intercept_is_enabled, - commands::flow_monitor_cmd::intercept_enable, - commands::flow_monitor_cmd::intercept_disable, - commands::flow_monitor_cmd::intercept_set_editing, - commands::flow_monitor_cmd::subscribe_intercept_events, - // Flow Monitor realtime enhancement commands - commands::flow_monitor_cmd::get_threshold_config, - commands::flow_monitor_cmd::update_threshold_config, - commands::flow_monitor_cmd::get_request_rate, - commands::flow_monitor_cmd::set_rate_window, - // Flow Replayer commands - commands::flow_monitor_cmd::replay_flow, - commands::flow_monitor_cmd::replay_flows_batch, - // Flow Diff commands - commands::flow_monitor_cmd::diff_flows, - // Session Management commands - commands::flow_monitor_cmd::create_session, - commands::flow_monitor_cmd::get_session, - commands::flow_monitor_cmd::list_sessions, - commands::flow_monitor_cmd::add_flow_to_session, - commands::flow_monitor_cmd::remove_flow_from_session, - commands::flow_monitor_cmd::update_session, - commands::flow_monitor_cmd::archive_session, - commands::flow_monitor_cmd::unarchive_session, - commands::flow_monitor_cmd::delete_session, - commands::flow_monitor_cmd::export_session, - commands::flow_monitor_cmd::get_session_flow_count, - commands::flow_monitor_cmd::is_flow_in_session, - commands::flow_monitor_cmd::get_sessions_for_flow, - commands::flow_monitor_cmd::get_auto_session_config, - commands::flow_monitor_cmd::set_auto_session_config, - commands::flow_monitor_cmd::register_active_session, - // Quick Filter commands - commands::flow_monitor_cmd::save_quick_filter, - commands::flow_monitor_cmd::get_quick_filter, - commands::flow_monitor_cmd::update_quick_filter, - commands::flow_monitor_cmd::delete_quick_filter, - commands::flow_monitor_cmd::list_quick_filters, - commands::flow_monitor_cmd::list_quick_filters_by_group, - commands::flow_monitor_cmd::list_quick_filter_groups, - commands::flow_monitor_cmd::export_quick_filters, - commands::flow_monitor_cmd::import_quick_filters, - commands::flow_monitor_cmd::find_quick_filter_by_name, - // Code Export commands - commands::flow_monitor_cmd::export_flow_as_code, - commands::flow_monitor_cmd::export_flows_as_code, - commands::flow_monitor_cmd::get_code_export_formats, - // Bookmark Management commands - commands::flow_monitor_cmd::add_bookmark, - commands::flow_monitor_cmd::get_bookmark, - commands::flow_monitor_cmd::get_bookmark_by_flow_id, - commands::flow_monitor_cmd::remove_bookmark, - commands::flow_monitor_cmd::remove_bookmark_by_flow_id, - commands::flow_monitor_cmd::update_bookmark, - commands::flow_monitor_cmd::list_bookmarks, - commands::flow_monitor_cmd::list_bookmark_groups, - commands::flow_monitor_cmd::is_flow_bookmarked, - commands::flow_monitor_cmd::get_bookmark_count, - commands::flow_monitor_cmd::export_bookmarks, - commands::flow_monitor_cmd::import_bookmarks, - commands::flow_monitor_cmd::toggle_bookmark, - // Enhanced Stats commands - commands::flow_monitor_cmd::get_enhanced_stats, - commands::flow_monitor_cmd::get_request_trend, - commands::flow_monitor_cmd::get_token_distribution, - commands::flow_monitor_cmd::get_latency_histogram, - commands::flow_monitor_cmd::export_stats_report, - // Batch Operations commands - commands::flow_monitor_cmd::batch_star_flows, - commands::flow_monitor_cmd::batch_unstar_flows, - commands::flow_monitor_cmd::batch_add_tags, - commands::flow_monitor_cmd::batch_remove_tags, - commands::flow_monitor_cmd::batch_export_flows, - commands::flow_monitor_cmd::batch_delete_flows, - commands::flow_monitor_cmd::batch_add_to_session, // Window control commands commands::window_cmd::get_window_size, commands::window_cmd::set_window_size, commands::window_cmd::center_window, commands::window_cmd::toggle_fullscreen, commands::window_cmd::is_fullscreen, - // Browser Interceptor commands - commands::browser_interceptor_cmd::get_browser_interceptor_state, - commands::browser_interceptor_cmd::start_browser_interceptor, - commands::browser_interceptor_cmd::stop_browser_interceptor, - commands::browser_interceptor_cmd::restore_normal_browser_behavior, - commands::browser_interceptor_cmd::temporary_disable_interceptor, - commands::browser_interceptor_cmd::get_intercepted_urls, - commands::browser_interceptor_cmd::get_interceptor_history, - commands::browser_interceptor_cmd::copy_intercepted_url_to_clipboard, - commands::browser_interceptor_cmd::open_url_in_fingerprint_browser, - commands::browser_interceptor_cmd::dismiss_intercepted_url, - commands::browser_interceptor_cmd::update_browser_interceptor_config, - commands::browser_interceptor_cmd::get_default_browser_interceptor_config, - commands::browser_interceptor_cmd::validate_browser_interceptor_config, - commands::browser_interceptor_cmd::is_browser_interceptor_running, - commands::browser_interceptor_cmd::get_browser_interceptor_statistics, - commands::browser_interceptor_cmd::show_notification, - commands::browser_interceptor_cmd::show_url_intercept_notification, - commands::browser_interceptor_cmd::show_status_notification, // Auto fix commands commands::auto_fix_cmd::auto_fix_configuration, // Machine ID commands diff --git a/src-tauri/src/app/setup.rs b/src-tauri/src/app/setup.rs index db44933ff..cc32b7ed9 100644 --- a/src-tauri/src/app/setup.rs +++ b/src-tauri/src/app/setup.rs @@ -8,7 +8,6 @@ use tauri::{App, Manager}; // use crate::agent::tools::{set_term_scrollback_tool_app_handle, set_terminal_tool_app_handle}; use crate::agent::AsterAgentState; use crate::database; -use crate::flow_monitor::FlowInterceptor; use crate::services::aster_session_store::ProxyCastSessionStore; use crate::services::provider_pool_service::ProviderPoolService; use crate::services::token_cache_service::TokenCacheService; @@ -30,8 +29,6 @@ pub fn setup_app( shared_stats: Arc>, shared_tokens: Arc>, shared_logger: Arc, - flow_monitor: Arc, - flow_interceptor: Arc, ) -> Result<(), Box> { // 注册全局 SessionStore(作为后备方案) // 注意:主要的 SessionStore 注入在 AsterAgentState::init_agent_with_db() 中完成 @@ -92,8 +89,6 @@ pub fn setup_app( shared_stats, shared_tokens, shared_logger, - flow_monitor, - flow_interceptor, app_handle, ) .await; @@ -112,8 +107,6 @@ async fn start_server_async( shared_stats: Arc>, shared_tokens: Arc>, shared_logger: Arc, - shared_flow_monitor: Arc, - flow_interceptor: Arc, app_handle: tauri::AppHandle, ) { // 先加载凭证池中的凭证 @@ -181,7 +174,7 @@ async fn start_server_async( .await .add("info", "[启动] 正在自动启动服务器..."); match s - .start_with_telemetry_and_flow_monitor( + .start_with_telemetry( logs.clone(), pool_service, token_cache, @@ -189,8 +182,6 @@ async fn start_server_async( Some(shared_stats), Some(shared_tokens), Some(shared_logger), - Some(shared_flow_monitor), - Some(flow_interceptor), ) .await { diff --git a/src-tauri/src/app/state.rs b/src-tauri/src/app/state.rs index ea32df9e6..edeb93ad6 100644 --- a/src-tauri/src/app/state.rs +++ b/src-tauri/src/app/state.rs @@ -7,11 +7,6 @@ use tokio::sync::RwLock; use crate::commands::api_key_provider_cmd::ApiKeyProviderServiceState; use crate::commands::context_memory::ContextMemoryServiceState; -use crate::commands::flow_monitor_cmd::{ - BatchOperationsState, BookmarkManagerState, EnhancedStatsServiceState, FlowInterceptorState, - FlowMonitorState, FlowQueryServiceState, FlowReplayerState, QuickFilterManagerState, - SessionManagerState, -}; use crate::commands::machine_id_cmd::MachineIdState; use crate::commands::orchestrator_cmd::OrchestratorState; use crate::commands::plugin_cmd::PluginManagerState; @@ -22,11 +17,6 @@ use crate::commands::skill_cmd::SkillServiceState; use crate::commands::tool_hooks::ToolHooksServiceState; use crate::config::{Config, ConfigManager, GlobalConfigManager, GlobalConfigManagerState}; use crate::database; -use crate::flow_monitor::{ - BatchOperations, BookmarkManager, EnhancedStatsService, FlowFileStore, FlowInterceptor, - FlowMonitor, FlowMonitorConfig, FlowQueryService, FlowReplayer, InterceptConfig, - QuickFilterManager, RotationConfig, SessionManager, -}; use crate::plugin; use crate::services::api_key_provider_service::ApiKeyProviderService; use crate::services::context_memory_service::{ContextMemoryConfig, ContextMemoryService}; @@ -216,141 +206,3 @@ pub fn init_telemetry_states(config: &Config) -> TelemetryStates { telemetry_state, } } - -/// Flow Monitor 状态 -pub struct FlowMonitorStates { - pub flow_monitor: Arc, - pub flow_monitor_state: FlowMonitorState, - pub flow_interceptor: Arc, - pub flow_interceptor_state: FlowInterceptorState, - pub flow_replayer_state: FlowReplayerState, - pub flow_query_service_state: FlowQueryServiceState, - pub session_manager_state: SessionManagerState, - pub quick_filter_manager_state: QuickFilterManagerState, - pub bookmark_manager_state: BookmarkManagerState, - pub enhanced_stats_service_state: EnhancedStatsServiceState, - pub batch_operations_state: BatchOperationsState, -} - -/// 初始化 Flow Monitor 状态 -/// -/// 如果 flow-monitor 插件已安装,则启用监控功能;否则禁用。 -pub fn init_flow_monitor_states( - provider_pool_service: Arc, - db: database::DbConnection, - plugin_installer_state: &PluginInstallerState, -) -> FlowMonitorStates { - // 检查 flow-monitor 插件是否已安装 - let is_plugin_installed = { - let installer = plugin_installer_state.0.blocking_read(); - installer.is_installed("flow-monitor").unwrap_or(false) - }; - - // 根据插件安装状态设置 enabled - let mut flow_monitor_config = FlowMonitorConfig::default(); - flow_monitor_config.enabled = is_plugin_installed; - - if is_plugin_installed { - tracing::info!("[启动] flow-monitor 插件已安装,启用 Flow 监控"); - } else { - tracing::info!("[启动] flow-monitor 插件未安装,禁用 Flow 监控"); - } - - let flow_file_store = init_flow_file_store(); - - let flow_monitor = Arc::new(FlowMonitor::new( - flow_monitor_config, - flow_file_store.clone(), - )); - let flow_monitor_state = FlowMonitorState(flow_monitor.clone()); - - // 初始化 Flow 拦截器 - let flow_interceptor = Arc::new(FlowInterceptor::new(InterceptConfig::default())); - let flow_interceptor_state = FlowInterceptorState(flow_interceptor.clone()); - - // 初始化 Flow 重放器 - let flow_replayer = Arc::new(FlowReplayer::new( - flow_monitor.clone(), - provider_pool_service, - db, - )); - let flow_replayer_state = FlowReplayerState(flow_replayer); - - // 初始化会话管理器 - let db_path = database::get_db_path().expect("Failed to get database path"); - let session_manager = - Arc::new(SessionManager::new(db_path.clone()).expect("Failed to create SessionManager")); - let session_manager_state = SessionManagerState(session_manager.clone()); - - // 初始化快速过滤器管理器 - let quick_filter_manager = Arc::new( - QuickFilterManager::new(db_path.clone()).expect("Failed to create QuickFilterManager"), - ); - let quick_filter_manager_state = QuickFilterManagerState(quick_filter_manager); - - // 初始化书签管理器 - let bookmark_manager = - Arc::new(BookmarkManager::new(db_path).expect("Failed to create BookmarkManager")); - let bookmark_manager_state = BookmarkManagerState(bookmark_manager); - - // 初始化增强统计服务 - let enhanced_stats_service = Arc::new(EnhancedStatsService::new(flow_monitor.memory_store())); - let enhanced_stats_service_state = EnhancedStatsServiceState(enhanced_stats_service); - - // 初始化批量操作服务 - let batch_operations = Arc::new(BatchOperations::new( - flow_monitor.clone(), - Some(session_manager_state.0.clone()), - )); - let batch_operations_state = BatchOperationsState(batch_operations); - - // FlowQueryService - let flow_query_service_state = if let Some(file_store) = flow_file_store { - let query_service = FlowQueryService::new(flow_monitor.memory_store(), file_store); - FlowQueryServiceState(Arc::new(query_service)) - } else { - let temp_dir = std::env::temp_dir().join("proxycast_flows"); - let _ = std::fs::create_dir_all(&temp_dir); - let rotation_config = RotationConfig::default(); - let temp_store = FlowFileStore::new(temp_dir, rotation_config) - .expect("Failed to create temp FlowFileStore"); - let query_service = - FlowQueryService::new(flow_monitor.memory_store(), Arc::new(temp_store)); - FlowQueryServiceState(Arc::new(query_service)) - }; - - FlowMonitorStates { - flow_monitor, - flow_monitor_state, - flow_interceptor, - flow_interceptor_state, - flow_replayer_state, - flow_query_service_state, - session_manager_state, - quick_filter_manager_state, - bookmark_manager_state, - enhanced_stats_service_state, - batch_operations_state, - } -} - -/// 初始化 Flow 文件存储 -fn init_flow_file_store() -> Option> { - let data_dir = dirs::data_dir() - .unwrap_or_else(|| std::path::PathBuf::from(".")) - .join("proxycast") - .join("flows"); - - if let Err(e) = std::fs::create_dir_all(&data_dir) { - tracing::warn!("无法创建 Flow 存储目录: {}", e); - } - - let rotation_config = RotationConfig::default(); - match FlowFileStore::new(data_dir, rotation_config) { - Ok(store) => Some(Arc::new(store)), - Err(e) => { - tracing::warn!("无法初始化 Flow 文件存储: {}", e); - None - } - } -} diff --git a/src-tauri/src/browser_interceptor/config.rs b/src-tauri/src/browser_interceptor/config.rs deleted file mode 100644 index cd4a9b213..000000000 --- a/src-tauri/src/browser_interceptor/config.rs +++ /dev/null @@ -1,241 +0,0 @@ -use chrono::{DateTime, Utc}; -use serde::{Deserialize, Serialize}; - -/// 拦截器状态 -#[derive(Debug, Clone, Serialize, Deserialize, Default)] -pub struct InterceptorState { - pub enabled: bool, - pub active_hooks: Vec, - pub intercepted_count: u32, - pub last_activity: Option>, - pub can_restore: bool, // 是否可以恢复正常状态 -} - -/// 被拦截的 URL 信息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct InterceptedUrl { - pub id: String, - pub url: String, - pub source_process: String, - pub timestamp: DateTime, - pub copied: bool, - pub opened_in_browser: bool, - pub dismissed: bool, -} - -impl InterceptedUrl { - pub fn new(url: String, source_process: String) -> Self { - Self { - id: uuid::Uuid::new_v4().to_string(), - url, - source_process, - timestamp: Utc::now(), - copied: false, - opened_in_browser: false, - dismissed: false, - } - } -} - -/// 指纹浏览器配置 -#[derive(Debug, Clone, Serialize, Deserialize, Default)] -pub struct FingerprintBrowserConfig { - pub enabled: bool, - pub executable_path: String, - pub profile_path: String, - pub additional_args: Vec, -} - -/// 恢复机制配置 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct RecoveryConfig { - pub backup_system_state: bool, - pub emergency_recovery_hotkey: String, - pub auto_recovery_on_crash: bool, - pub recovery_timeout: u64, // 恢复操作超时时间(秒) -} - -impl Default for RecoveryConfig { - fn default() -> Self { - Self { - backup_system_state: true, - emergency_recovery_hotkey: "Ctrl+Alt+Shift+R".to_string(), - auto_recovery_on_crash: true, - recovery_timeout: 30, - } - } -} - -/// 浏览器拦截器配置 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct BrowserInterceptorConfig { - pub enabled: bool, - pub target_processes: Vec, - pub url_patterns: Vec, - pub excluded_processes: Vec, - pub notification_enabled: bool, - pub auto_copy_to_clipboard: bool, - pub auto_launch_browser: bool, - pub restore_on_exit: bool, - pub temporary_disable_timeout: Option, // 临时禁用超时(秒) - pub fingerprint_browser: FingerprintBrowserConfig, - pub recovery: RecoveryConfig, -} - -impl Default for BrowserInterceptorConfig { - fn default() -> Self { - Self { - enabled: false, - target_processes: vec![ - "kiro".to_string(), - "kiro.exe".to_string(), - "cursor".to_string(), - "cursor.exe".to_string(), - "code".to_string(), - "code.exe".to_string(), - "Kiro".to_string(), // macOS 应用名称通常首字母大写 - "Cursor".to_string(), - "Visual Studio Code".to_string(), - ], - url_patterns: vec![ - "https://auth.*".to_string(), - "https://*/oauth/*".to_string(), - "https://accounts.google.com/*".to_string(), - "https://github.com/login/*".to_string(), - "https://login.microsoftonline.com/*".to_string(), - ], - excluded_processes: vec![ - "explorer.exe".to_string(), - "winlogon.exe".to_string(), - "system".to_string(), - "chrome.exe".to_string(), - "firefox.exe".to_string(), - "safari".to_string(), - "Safari".to_string(), // macOS Safari - "Google Chrome".to_string(), // macOS Chrome - "Firefox".to_string(), // macOS Firefox - ], - notification_enabled: true, - auto_copy_to_clipboard: true, - auto_launch_browser: false, - restore_on_exit: true, - temporary_disable_timeout: Some(300), // 5分钟 - fingerprint_browser: FingerprintBrowserConfig::default(), - recovery: RecoveryConfig::default(), - } - } -} - -impl BrowserInterceptorConfig { - /// 检查进程是否在目标列表中 - pub fn is_target_process(&self, process_name: &str) -> bool { - self.target_processes.iter().any(|pattern| { - // 支持简单的通配符匹配 - if pattern.contains('*') { - process_name.contains(&pattern.replace('*', "")) - } else { - process_name.eq_ignore_ascii_case(pattern) - } - }) - } - - /// 检查进程是否在排除列表中 - pub fn is_excluded_process(&self, process_name: &str) -> bool { - self.excluded_processes.iter().any(|pattern| { - if pattern.contains('*') { - process_name.contains(&pattern.replace('*', "")) - } else { - process_name.eq_ignore_ascii_case(pattern) - } - }) - } - - /// 检查 URL 是否匹配拦截模式 - pub fn matches_url_pattern(&self, url: &str) -> bool { - self.url_patterns.iter().any(|pattern| { - // 简单的模式匹配实现 - if pattern.contains('*') { - let parts: Vec<&str> = pattern.split('*').collect(); - if parts.len() == 2 { - url.starts_with(parts[0]) && url.ends_with(parts[1]) - } else { - // 更复杂的模式匹配 - url.contains(&pattern.replace('*', "")) - } - } else { - url.starts_with(pattern) - } - }) - } - - /// 验证配置的有效性 - pub fn validate(&self) -> std::result::Result<(), String> { - if self.target_processes.is_empty() { - return Err("目标进程列表不能为空".to_string()); - } - - if self.url_patterns.is_empty() { - return Err("URL 模式列表不能为空".to_string()); - } - - if self.fingerprint_browser.enabled && self.fingerprint_browser.executable_path.is_empty() { - return Err("启用指纹浏览器时必须指定可执行文件路径".to_string()); - } - - if let Some(timeout) = self.temporary_disable_timeout { - if timeout == 0 { - return Err("临时禁用超时时间必须大于 0".to_string()); - } - } - - Ok(()) - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_is_target_process() { - let config = BrowserInterceptorConfig::default(); - - assert!(config.is_target_process("kiro")); - assert!(config.is_target_process("KIRO.EXE")); - assert!(config.is_target_process("cursor")); - assert!(!config.is_target_process("notepad")); - } - - #[test] - fn test_is_excluded_process() { - let config = BrowserInterceptorConfig::default(); - - assert!(config.is_excluded_process("chrome.exe")); - assert!(config.is_excluded_process("FIREFOX.EXE")); - assert!(!config.is_excluded_process("kiro")); - } - - #[test] - fn test_matches_url_pattern() { - let config = BrowserInterceptorConfig::default(); - - assert!(config.matches_url_pattern("https://accounts.google.com/oauth/authorize")); - assert!(config.matches_url_pattern("https://github.com/login/oauth")); - assert!(config.matches_url_pattern("https://auth.example.com/login")); - assert!(!config.matches_url_pattern("https://example.com/normal-page")); - } - - #[test] - fn test_config_validation() { - let mut config = BrowserInterceptorConfig::default(); - assert!(config.validate().is_ok()); - - config.target_processes.clear(); - assert!(config.validate().is_err()); - - config = BrowserInterceptorConfig::default(); - config.fingerprint_browser.enabled = true; - config.fingerprint_browser.executable_path = String::new(); - assert!(config.validate().is_err()); - } -} diff --git a/src-tauri/src/browser_interceptor/interceptor.rs b/src-tauri/src/browser_interceptor/interceptor.rs deleted file mode 100644 index c0bca65f2..000000000 --- a/src-tauri/src/browser_interceptor/interceptor.rs +++ /dev/null @@ -1,457 +0,0 @@ -use crate::browser_interceptor::{ - BrowserInterceptorConfig, BrowserInterceptorError, InterceptedUrl, NotificationService, Result, - StateManager, UrlManager, -}; -use std::sync::Arc; -use tokio::sync::RwLock; - -/// 浏览器拦截器主结构 -pub struct BrowserInterceptor { - config: Arc>, - state_manager: Arc, - url_manager: Arc, - notification_service: Arc>, - #[cfg(target_os = "windows")] - windows_interceptor: Option, - #[cfg(target_os = "macos")] - macos_interceptor: Option, - #[cfg(target_os = "linux")] - linux_interceptor: Option, -} - -impl BrowserInterceptor { - /// 创建新的浏览器拦截器实例 - pub fn new(config: BrowserInterceptorConfig) -> Self { - Self { - config: Arc::new(RwLock::new(config)), - state_manager: Arc::new(StateManager::new()), - url_manager: Arc::new(UrlManager::new()), - notification_service: Arc::new(RwLock::new(NotificationService::new())), - #[cfg(target_os = "windows")] - windows_interceptor: None, - #[cfg(target_os = "macos")] - macos_interceptor: None, - #[cfg(target_os = "linux")] - linux_interceptor: None, - } - } - - /// 启动拦截器 - pub async fn start(&mut self) -> Result<()> { - tracing::info!("开始启动浏览器拦截器..."); - - let config = self.config.read().await; - - if !config.enabled { - tracing::error!("拦截器配置中 enabled 为 false"); - return Err(BrowserInterceptorError::InterceptorError( - "拦截器未启用".to_string(), - )); - } - - // 验证配置 - tracing::info!("验证配置..."); - config.validate().map_err(|e| { - tracing::error!("配置验证失败: {}", e); - BrowserInterceptorError::ConfigError(format!("配置验证失败: {e}")) - })?; - - drop(config); // 释放读锁 - - // 启用状态管理器 - tracing::info!("启用状态管理器..."); - self.state_manager.enable_interceptor().await.map_err(|e| { - tracing::error!("启用状态管理器失败: {}", e); - e - })?; - - // 启动平台特定的拦截器 - tracing::info!("启动平台拦截器..."); - self.start_platform_interceptor().await.map_err(|e| { - tracing::error!("启动平台拦截器失败: {}", e); - e - })?; - - // 发送启用通知 - tracing::info!("发送启用通知..."); - let notification_service = self.notification_service.read().await; - notification_service.notify_interceptor_enabled().await?; - - tracing::info!("浏览器拦截器已成功启动"); - Ok(()) - } - - /// 停止拦截器 - pub async fn stop(&mut self) -> Result<()> { - // 停止平台特定的拦截器 - self.stop_platform_interceptor().await?; - - // 禁用状态管理器 - self.state_manager.disable_interceptor().await?; - - // 发送禁用通知 - let notification_service = self.notification_service.read().await; - notification_service.notify_interceptor_disabled().await?; - - tracing::info!("浏览器拦截器已停止"); - Ok(()) - } - - /// 恢复正常浏览器行为 - pub async fn restore_normal_behavior(&mut self) -> Result<()> { - // 停止拦截 - self.stop_platform_interceptor().await?; - - // 恢复系统状态 - self.state_manager.restore_normal_behavior().await?; - - // 发送恢复通知 - let notification_service = self.notification_service.read().await; - notification_service.notify_system_restored().await?; - - tracing::info!("已恢复正常浏览器行为"); - Ok(()) - } - - /// 临时禁用拦截器 - pub async fn temporary_disable(&mut self, duration_seconds: u64) -> Result<()> { - self.state_manager - .temporary_disable(duration_seconds) - .await?; - - // 临时停止平台拦截器 - self.stop_platform_interceptor().await?; - - tracing::info!("拦截器已临时禁用 {} 秒", duration_seconds); - Ok(()) - } - - /// 处理拦截到的浏览器启动请求 - pub async fn handle_browser_launch(&self, url: String, source_process: String) -> Result { - let config = self.config.read().await; - - // 检查是否应该拦截这个进程 - if !self.should_intercept(&config, &source_process, &url) { - return Ok(false); // 不拦截,让浏览器正常启动 - } - - drop(config); // 释放读锁 - - // 添加到拦截列表 - let url_id = self - .url_manager - .add_intercepted_url(url.clone(), source_process.clone())?; - - // 更新拦截计数 - self.state_manager.increment_intercept_count()?; - - // 创建拦截的 URL 对象用于通知 - let intercepted_url = InterceptedUrl::new(url, source_process); - - // 发送通知 - let notification_service = self.notification_service.read().await; - notification_service - .notify_url_intercepted(&intercepted_url) - .await?; - - // 如果启用了自动复制到剪贴板 - let config = self.config.read().await; - if config.auto_copy_to_clipboard { - self.copy_to_clipboard(&intercepted_url.url).await?; - self.url_manager.mark_as_copied(&url_id)?; - } - - tracing::info!( - "已拦截浏览器启动: {} (来源: {})", - intercepted_url.url, - intercepted_url.source_process - ); - Ok(true) // 已拦截 - } - - /// 获取当前状态 - pub async fn get_state(&self) -> Result { - self.state_manager.get_state() - } - - /// 获取拦截的 URL 列表 - pub async fn get_intercepted_urls(&self) -> Result> { - self.url_manager.get_intercepted_urls() - } - - /// 获取历史记录 - pub async fn get_history(&self, limit: Option) -> Result> { - self.url_manager.get_history(limit) - } - - /// 复制 URL 到剪贴板 - pub async fn copy_url_to_clipboard(&self, url_id: &str) -> Result<()> { - if let Some(intercepted_url) = self.url_manager.get_intercepted_url(url_id)? { - self.copy_to_clipboard(&intercepted_url.url).await?; - self.url_manager.mark_as_copied(url_id)?; - tracing::info!("URL {} 已复制到剪贴板", url_id); - } - Ok(()) - } - - /// 在指纹浏览器中打开 URL - pub async fn open_in_fingerprint_browser(&self, url_id: &str) -> Result<()> { - let config = self.config.read().await; - - if !config.fingerprint_browser.enabled { - return Err(BrowserInterceptorError::InterceptorError( - "指纹浏览器未启用".to_string(), - )); - } - - if let Some(intercepted_url) = self.url_manager.get_intercepted_url(url_id)? { - self.launch_fingerprint_browser(&config.fingerprint_browser, &intercepted_url.url) - .await?; - self.url_manager.mark_as_opened(url_id)?; - tracing::info!("URL {} 已在指纹浏览器中打开", url_id); - } - - Ok(()) - } - - /// 忽略指定的 URL - pub async fn dismiss_url(&self, url_id: &str) -> Result<()> { - self.url_manager.dismiss_url(url_id)?; - tracing::info!("URL {} 已被忽略", url_id); - Ok(()) - } - - /// 更新配置 - pub async fn update_config(&self, new_config: BrowserInterceptorConfig) -> Result<()> { - // 验证新配置 - new_config - .validate() - .map_err(|e| BrowserInterceptorError::ConfigError(format!("配置验证失败: {e}")))?; - - let mut config = self.config.write().await; - *config = new_config; - - tracing::info!("浏览器拦截器配置已更新"); - Ok(()) - } - - /// 检查是否应该拦截指定的进程和 URL - fn should_intercept( - &self, - config: &BrowserInterceptorConfig, - process_name: &str, - url: &str, - ) -> bool { - // 检查是否在排除列表中 - if config.is_excluded_process(process_name) { - return false; - } - - // 检查是否是目标进程 - if !config.is_target_process(process_name) { - return false; - } - - // 检查 URL 是否匹配模式 - config.matches_url_pattern(url) - } - - /// 启动平台特定的拦截器 - async fn start_platform_interceptor(&mut self) -> Result<()> { - #[cfg(target_os = "windows")] - { - // 创建 URL 处理器闭包 - let url_manager = Arc::clone(&self.url_manager); - let state_manager = Arc::clone(&self.state_manager); - - let url_handler = move |intercepted_url: InterceptedUrl| { - let url_manager = Arc::clone(&url_manager); - let state_manager = Arc::clone(&state_manager); - - tokio::spawn(async move { - if let Err(e) = url_manager - .add_intercepted_url(intercepted_url.url, intercepted_url.source_process) - { - tracing::error!("添加拦截 URL 失败: {}", e); - } - if let Err(e) = state_manager.increment_intercept_count() { - tracing::error!("更新拦截计数失败: {}", e); - } - }); - }; - - let mut interceptor = - crate::browser_interceptor::platform::windows::WindowsInterceptor::new(url_handler); - interceptor.start().await?; - self.windows_interceptor = Some(interceptor); - } - - #[cfg(target_os = "macos")] - { - let url_manager = Arc::clone(&self.url_manager); - let state_manager = Arc::clone(&self.state_manager); - - let url_handler = move |intercepted_url: InterceptedUrl| { - let url_manager = Arc::clone(&url_manager); - let state_manager = Arc::clone(&state_manager); - - tokio::spawn(async move { - if let Err(e) = url_manager - .add_intercepted_url(intercepted_url.url, intercepted_url.source_process) - { - tracing::error!("添加拦截 URL 失败: {}", e); - } - if let Err(e) = state_manager.increment_intercept_count() { - tracing::error!("更新拦截计数失败: {}", e); - } - }); - }; - - let mut interceptor = - crate::browser_interceptor::platform::macos::MacOSInterceptor::new(url_handler); - interceptor.start().await?; - self.macos_interceptor = Some(interceptor); - } - - #[cfg(target_os = "linux")] - { - let url_manager = Arc::clone(&self.url_manager); - let state_manager = Arc::clone(&self.state_manager); - - let url_handler = move |intercepted_url: InterceptedUrl| { - let url_manager = Arc::clone(&url_manager); - let state_manager = Arc::clone(&state_manager); - - tokio::spawn(async move { - if let Err(e) = url_manager - .add_intercepted_url(intercepted_url.url, intercepted_url.source_process) - { - tracing::error!("添加拦截 URL 失败: {}", e); - } - if let Err(e) = state_manager.increment_intercept_count() { - tracing::error!("更新拦截计数失败: {}", e); - } - }); - }; - - let mut interceptor = - crate::browser_interceptor::platform::linux::LinuxInterceptor::new(url_handler); - interceptor.start().await?; - self.linux_interceptor = Some(interceptor); - } - - Ok(()) - } - - /// 停止平台特定的拦截器 - async fn stop_platform_interceptor(&mut self) -> Result<()> { - #[cfg(target_os = "windows")] - { - if let Some(ref mut interceptor) = self.windows_interceptor { - interceptor.stop().await?; - } - } - - #[cfg(target_os = "macos")] - { - if let Some(ref mut interceptor) = self.macos_interceptor { - interceptor.stop().await?; - } - } - - #[cfg(target_os = "linux")] - { - if let Some(ref mut interceptor) = self.linux_interceptor { - interceptor.stop().await?; - } - } - - Ok(()) - } - - /// 复制文本到剪贴板 - async fn copy_to_clipboard(&self, text: &str) -> Result<()> { - match arboard::Clipboard::new() { - Ok(mut clipboard) => { - if let Err(e) = clipboard.set_text(text) { - return Err(BrowserInterceptorError::InterceptorError(format!( - "复制到剪贴板失败: {e}" - ))); - } - tracing::info!("已复制到剪贴板: {}", text); - Ok(()) - } - Err(e) => Err(BrowserInterceptorError::InterceptorError(format!( - "创建剪贴板实例失败: {e}" - ))), - } - } - - /// 启动指纹浏览器 - async fn launch_fingerprint_browser( - &self, - browser_config: &crate::browser_interceptor::config::FingerprintBrowserConfig, - url: &str, - ) -> Result<()> { - if browser_config.executable_path.is_empty() { - return Err(BrowserInterceptorError::ConfigError( - "指纹浏览器可执行文件路径未配置".to_string(), - )); - } - - // 构建启动命令 - let mut command = std::process::Command::new(&browser_config.executable_path); - - // 添加 URL 参数 - command.arg(url); - - // 添加额外的参数 - for arg in &browser_config.additional_args { - command.arg(arg); - } - - // 如果配置了配置文件路径 - if !browser_config.profile_path.is_empty() { - command - .arg("--user-data-dir") - .arg(&browser_config.profile_path); - } - - // 异步启动进程 - match command.spawn() { - Ok(mut child) => { - // 在后台等待进程完成 - tokio::spawn(async move { - match child.wait() { - Ok(status) => { - if status.success() { - tracing::info!("指纹浏览器启动成功"); - } else { - tracing::error!("指纹浏览器退出异常: {}", status); - } - } - Err(e) => { - tracing::error!("等待指纹浏览器进程失败: {}", e); - } - } - }); - - tracing::info!( - "已启动指纹浏览器: {} -> {}", - browser_config.executable_path, - url - ); - Ok(()) - } - Err(e) => Err(BrowserInterceptorError::InterceptorError(format!( - "启动指纹浏览器失败: {e}" - ))), - } - } -} - -impl Default for BrowserInterceptor { - fn default() -> Self { - Self::new(BrowserInterceptorConfig::default()) - } -} diff --git a/src-tauri/src/browser_interceptor/mod.rs b/src-tauri/src/browser_interceptor/mod.rs deleted file mode 100644 index 914e650b7..000000000 --- a/src-tauri/src/browser_interceptor/mod.rs +++ /dev/null @@ -1,69 +0,0 @@ -pub mod config; -pub mod interceptor; -pub mod notification_service; -pub mod state_manager; -pub mod url_manager; - -#[cfg(target_os = "windows")] -pub mod platform { - pub mod windows; -} - -#[cfg(target_os = "macos")] -pub mod platform { - pub mod macos; -} - -#[cfg(target_os = "linux")] -pub mod platform { - pub mod linux; -} - -// 重新导出主要类型和函数 -pub use config::{BrowserInterceptorConfig, InterceptedUrl, InterceptorState}; -pub use interceptor::BrowserInterceptor; -pub use notification_service::NotificationService; -pub use state_manager::StateManager; -pub use url_manager::{UrlManager, UrlStatistics}; - -use serde::{Deserialize, Serialize}; -use std::error::Error; -use std::fmt; - -/// 浏览器拦截器错误类型 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub enum BrowserInterceptorError { - ConfigError(String), - InterceptorError(String), - StateError(String), - PlatformError(String), - NotificationError(String), - AlreadyRunning, - UnsupportedPlatform(String), - IoError(String), -} - -impl fmt::Display for BrowserInterceptorError { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - match self { - BrowserInterceptorError::ConfigError(msg) => write!(f, "配置错误: {msg}"), - BrowserInterceptorError::InterceptorError(msg) => write!(f, "拦截器错误: {msg}"), - BrowserInterceptorError::StateError(msg) => write!(f, "状态管理错误: {msg}"), - BrowserInterceptorError::PlatformError(msg) => write!(f, "平台错误: {msg}"), - BrowserInterceptorError::NotificationError(msg) => write!(f, "通知错误: {msg}"), - BrowserInterceptorError::AlreadyRunning => write!(f, "拦截器已在运行"), - BrowserInterceptorError::UnsupportedPlatform(msg) => write!(f, "不支持的平台: {msg}"), - BrowserInterceptorError::IoError(msg) => write!(f, "IO错误: {msg}"), - } - } -} - -impl Error for BrowserInterceptorError {} - -impl From for BrowserInterceptorError { - fn from(err: std::io::Error) -> Self { - BrowserInterceptorError::IoError(err.to_string()) - } -} - -pub type Result = std::result::Result; diff --git a/src-tauri/src/browser_interceptor/notification_service.rs b/src-tauri/src/browser_interceptor/notification_service.rs deleted file mode 100644 index 3dadb15e4..000000000 --- a/src-tauri/src/browser_interceptor/notification_service.rs +++ /dev/null @@ -1,234 +0,0 @@ -use crate::browser_interceptor::{InterceptedUrl, Result}; -use serde::{Deserialize, Serialize}; -use std::time::Duration; - -/// 通知类型 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub enum NotificationType { - UrlIntercepted, - InterceptorEnabled, - InterceptorDisabled, - SystemRestored, - Error, -} - -/// 通知消息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct NotificationMessage { - pub id: String, - pub notification_type: NotificationType, - pub title: String, - pub message: String, - pub url: Option, - pub source_process: Option, - pub timestamp: chrono::DateTime, - pub auto_dismiss_after: Option, -} - -impl NotificationMessage { - pub fn new(notification_type: NotificationType, title: String, message: String) -> Self { - Self { - id: uuid::Uuid::new_v4().to_string(), - notification_type, - title, - message, - url: None, - source_process: None, - timestamp: chrono::Utc::now(), - auto_dismiss_after: Some(Duration::from_secs(30)), - } - } - - pub fn with_url(mut self, url: String) -> Self { - self.url = Some(url); - self - } - - pub fn with_source_process(mut self, source_process: String) -> Self { - self.source_process = Some(source_process); - self - } - - pub fn with_auto_dismiss(mut self, duration: Option) -> Self { - self.auto_dismiss_after = duration; - self - } -} - -/// 通知服务 -pub struct NotificationService { - enabled: bool, - show_url_preview: bool, -} - -impl NotificationService { - pub fn new() -> Self { - Self { - enabled: true, - show_url_preview: true, - } - } - - /// 设置通知是否启用 - pub fn set_enabled(&mut self, enabled: bool) { - self.enabled = enabled; - } - - /// 设置是否显示 URL 预览 - pub fn set_show_url_preview(&mut self, show_preview: bool) { - self.show_url_preview = show_preview; - } - - /// 发送 URL 拦截通知 - pub async fn notify_url_intercepted(&self, intercepted_url: &InterceptedUrl) -> Result<()> { - if !self.enabled { - return Ok(()); - } - - let title = format!("已拦截来自 {} 的 URL", intercepted_url.source_process); - let message = if self.show_url_preview { - format!("URL: {}", self.truncate_url(&intercepted_url.url, 100)) - } else { - "点击查看详情".to_string() - }; - - let notification = - NotificationMessage::new(NotificationType::UrlIntercepted, title, message) - .with_url(intercepted_url.url.clone()) - .with_source_process(intercepted_url.source_process.clone()); - - self.send_notification(notification).await - } - - /// 发送拦截器启用通知 - pub async fn notify_interceptor_enabled(&self) -> Result<()> { - if !self.enabled { - return Ok(()); - } - - let notification = NotificationMessage::new( - NotificationType::InterceptorEnabled, - "浏览器拦截器已启用".to_string(), - "现在会拦截目标应用的浏览器启动请求".to_string(), - ); - - self.send_notification(notification).await - } - - /// 发送拦截器禁用通知 - pub async fn notify_interceptor_disabled(&self) -> Result<()> { - if !self.enabled { - return Ok(()); - } - - let notification = NotificationMessage::new( - NotificationType::InterceptorDisabled, - "浏览器拦截器已禁用".to_string(), - "应用将正常打开默认浏览器".to_string(), - ); - - self.send_notification(notification).await - } - - /// 发送系统恢复通知 - pub async fn notify_system_restored(&self) -> Result<()> { - if !self.enabled { - return Ok(()); - } - - let notification = NotificationMessage::new( - NotificationType::SystemRestored, - "系统已恢复正常".to_string(), - "浏览器行为已恢复到原始状态".to_string(), - ); - - self.send_notification(notification).await - } - - /// 发送错误通知 - pub async fn notify_error(&self, error_message: &str) -> Result<()> { - if !self.enabled { - return Ok(()); - } - - let notification = NotificationMessage::new( - NotificationType::Error, - "浏览器拦截器错误".to_string(), - error_message.to_string(), - ) - .with_auto_dismiss(Some(Duration::from_secs(60))); // 错误通知显示更长时间 - - self.send_notification(notification).await - } - - /// 发送通知的具体实现 - async fn send_notification(&self, notification: NotificationMessage) -> Result<()> { - // 记录日志 - tracing::info!( - "发送通知: {} - {}", - notification.title, - notification.message - ); - - // 发送系统通知 - self.send_system_notification(¬ification).await?; - - Ok(()) - } - - /// 发送系统通知 - async fn send_system_notification(&self, notification: &NotificationMessage) -> Result<()> { - #[cfg(target_os = "windows")] - { - self.send_windows_notification(notification).await?; - } - - #[cfg(target_os = "macos")] - { - self.send_macos_notification(notification).await?; - } - - #[cfg(target_os = "linux")] - { - self.send_linux_notification(notification).await?; - } - - Ok(()) - } - - /// Windows 系统通知 - #[cfg(target_os = "windows")] - async fn send_windows_notification(&self, notification: &NotificationMessage) -> Result<()> { - tracing::debug!("Windows 通知: {}", notification.title); - Ok(()) - } - - /// macOS 系统通知 - #[cfg(target_os = "macos")] - async fn send_macos_notification(&self, notification: &NotificationMessage) -> Result<()> { - tracing::debug!("macOS 通知: {}", notification.title); - Ok(()) - } - - /// Linux 系统通知 - #[cfg(target_os = "linux")] - async fn send_linux_notification(&self, notification: &NotificationMessage) -> Result<()> { - tracing::debug!("Linux 通知: {}", notification.title); - Ok(()) - } - - /// 截断 URL 以适应通知显示 - fn truncate_url(&self, url: &str, max_length: usize) -> String { - if url.len() <= max_length { - url.to_string() - } else { - format!("{}...", &url[..max_length.saturating_sub(3)]) - } - } -} - -impl Default for NotificationService { - fn default() -> Self { - Self::new() - } -} diff --git a/src-tauri/src/browser_interceptor/platform/linux.rs b/src-tauri/src/browser_interceptor/platform/linux.rs deleted file mode 100644 index 9c861bbaf..000000000 --- a/src-tauri/src/browser_interceptor/platform/linux.rs +++ /dev/null @@ -1,71 +0,0 @@ -use crate::browser_interceptor::{InterceptedUrl, Result}; - -/// Linux 平台的浏览器拦截器 -pub struct LinuxInterceptor { - running: bool, -} - -impl LinuxInterceptor { - pub fn new(_url_handler: F) -> Self - where - F: Fn(InterceptedUrl) + Send + Sync + 'static, - { - Self { running: false } - } - - /// 启动拦截 - pub async fn start(&mut self) -> Result<()> { - // TODO: 实现 Linux 平台的浏览器拦截 - // 可以使用 xdg-open 拦截或 D-Bus 监听 - self.running = true; - tracing::info!("Linux 浏览器拦截器已启动 (占位符)"); - Ok(()) - } - - /// 停止拦截 - pub async fn stop(&mut self) -> Result<()> { - if !self.running { - return Ok(()); - } - - self.running = false; - tracing::info!("Linux 浏览器拦截器已停止"); - Ok(()) - } - - /// 检查是否正在拦截 - pub fn is_running(&self) -> bool { - self.running - } - - /// 恢复系统默认设置 - pub async fn restore_system_defaults(&self) -> Result<()> { - tracing::info!("Linux 系统默认设置已恢复"); - Ok(()) - } - - /// 临时禁用拦截 - pub async fn temporarily_disable(&mut self) -> Result<()> { - tracing::info!("Linux 拦截器已临时禁用"); - Ok(()) - } - - /// 重新启用拦截 - pub async fn re_enable(&mut self) -> Result<()> { - tracing::info!("Linux 拦截器已重新启用"); - Ok(()) - } -} - -impl Drop for LinuxInterceptor { - fn drop(&mut self) { - if self.running { - let _ = tokio::runtime::Handle::try_current().map(|handle| { - handle.block_on(async { - let _ = self.stop().await; - let _ = self.restore_system_defaults().await; - }) - }); - } - } -} diff --git a/src-tauri/src/browser_interceptor/platform/macos.rs b/src-tauri/src/browser_interceptor/platform/macos.rs deleted file mode 100644 index 4d4a206f4..000000000 --- a/src-tauri/src/browser_interceptor/platform/macos.rs +++ /dev/null @@ -1,372 +0,0 @@ -#![allow(dead_code)] - -use crate::browser_interceptor::{BrowserInterceptorError, InterceptedUrl, Result}; -use once_cell::sync::Lazy; -use std::process::Command; -use std::sync::{Arc, Mutex}; -use tokio::sync::mpsc; - -/// 全局 URL sender,用于从 deep-link 事件接收 URL(使用 Mutex 保证线程安全) -static GLOBAL_URL_SENDER: Lazy>>> = - Lazy::new(|| Mutex::new(None)); - -/// macOS 平台的浏览器拦截器(基于设置默认浏览器 + Deep Link) -pub struct MacOSInterceptor { - running: bool, - url_sender: Option>, - original_default_browser: Option, - url_handler: Option>, -} - -impl MacOSInterceptor { - pub fn new(url_handler: F) -> Self - where - F: Fn(InterceptedUrl) + Send + Sync + 'static, - { - let (tx, mut rx) = mpsc::unbounded_channel(); - let handler = Arc::new(url_handler); - let handler_clone = handler.clone(); - - // 启动后台任务处理拦截的 URL - tokio::spawn(async move { - while let Some(intercepted_url) = rx.recv().await { - handler_clone(intercepted_url); - } - }); - - Self { - running: false, - url_sender: Some(tx), - original_default_browser: None, - url_handler: Some(handler), - } - } - - /// 启动拦截 - pub async fn start(&mut self) -> Result<()> { - if self.running { - tracing::info!("macOS 拦截器已在运行中,跳过启动"); - return Ok(()); - } - - tracing::info!("正在启动 macOS 浏览器拦截器..."); - - // 1. 保存当前默认浏览器 - self.original_default_browser = self.get_default_browser().await; - tracing::info!("当前默认浏览器: {:?}", self.original_default_browser); - - // 2. 设置全局 URL sender - if let Some(ref sender) = self.url_sender { - if let Ok(mut global_sender) = GLOBAL_URL_SENDER.lock() { - *global_sender = Some(sender.clone()); - } - } - - // 3. 将应用设置为默认浏览器 - self.set_as_default_browser().await?; - - self.running = true; - tracing::info!("macOS 浏览器拦截器已启动"); - tracing::info!("提示:现在所有 http/https URL 打开请求都会被拦截"); - - Ok(()) - } - - /// 获取当前默认浏览器的 Bundle ID - async fn get_default_browser(&self) -> Option { - // 使用 Swift 调用 LSCopyDefaultHandlerForURLScheme 获取真实的默认浏览器 - let swift_code = r#" -import Foundation -import CoreServices - -if let handler = LSCopyDefaultHandlerForURLScheme("https" as CFString) { - print(handler.takeRetainedValue() as String) -} else if let handler = LSCopyDefaultHandlerForURLScheme("http" as CFString) { - print(handler.takeRetainedValue() as String) -} else { - print("") -} -"#; - - let output = Command::new("swift") - .args(["-e", swift_code]) - .output() - .ok()?; - - if output.status.success() { - let bundle_id = String::from_utf8_lossy(&output.stdout).trim().to_string(); - if !bundle_id.is_empty() && !bundle_id.contains("browser-interception") { - tracing::info!("检测到当前默认浏览器: {}", bundle_id); - return Some(bundle_id); - } - // 如果当前是本应用,说明之前已设置过,需要检测用户常用的浏览器 - if bundle_id.contains("browser-interception") { - tracing::info!("当前默认浏览器是本应用,尝试检测用户常用浏览器"); - return self.detect_installed_browser().await; - } - } - - // 最后尝试检测已安装的浏览器,优先返回 Chrome - self.detect_installed_browser().await - } - - /// 将应用设置为默认浏览器 - async fn set_as_default_browser(&self) -> Result<()> { - tracing::info!("正在将应用设置为默认浏览器..."); - - // 使用 Swift 脚本 - let swift_code = r#" -import Foundation -import CoreServices - -let bundleId = "com.browser-interception.app" as CFString - -// 设置 HTTP handler -LSSetDefaultHandlerForURLScheme("http" as CFString, bundleId) - -// 设置 HTTPS handler -LSSetDefaultHandlerForURLScheme("https" as CFString, bundleId) - -print("OK") -"#; - - let output = Command::new("swift") - .args(["-e", swift_code]) - .output() - .map_err(|e| { - BrowserInterceptorError::PlatformError(format!("执行 Swift 脚本失败: {e}")) - })?; - - if output.status.success() { - tracing::info!("已通过 Launch Services API 设置默认浏览器"); - return Ok(()); - } - - // 提示用户手动设置 - let stderr = String::from_utf8_lossy(&output.stderr); - tracing::warn!("自动设置默认浏览器失败: {}", stderr); - tracing::info!("请手动在系统设置中将应用设置为默认浏览器"); - - // 打开系统设置 - let _ = Command::new("open") - .args(["x-apple.systempreferences:com.apple.preference.general"]) - .output(); - - Ok(()) - } - - /// 恢复原来的默认浏览器 - async fn restore_default_browser(&self) -> Result<()> { - // 获取要恢复的浏览器 - let browser_id = match &self.original_default_browser { - Some(id) if !id.contains("browser-interception") => id.clone(), - _ => { - // 如果没有记录原始浏览器,尝试检测已安装的浏览器 - self.detect_installed_browser() - .await - .unwrap_or_else(|| "com.google.Chrome".to_string()) - } - }; - - tracing::info!("正在恢复默认浏览器为: {}", browser_id); - - // 使用 Swift - let swift_code = format!( - r#" -import Foundation -import CoreServices - -let bundleId = "{browser_id}" as CFString -LSSetDefaultHandlerForURLScheme("http" as CFString, bundleId) -LSSetDefaultHandlerForURLScheme("https" as CFString, bundleId) -print("OK") -"# - ); - - let output = Command::new("swift").args(["-e", &swift_code]).output(); - - if let Ok(out) = output { - if out.status.success() { - tracing::info!("已通过 Swift 恢复默认浏览器为: {}", browser_id); - } else { - let stderr = String::from_utf8_lossy(&out.stderr); - tracing::warn!("Swift 恢复默认浏览器失败: {}", stderr); - } - } - - Ok(()) - } - - /// 检测已安装的浏览器(用于恢复时选择) - async fn detect_installed_browser(&self) -> Option { - // 按优先级排序的浏览器列表 - let browsers = [ - ("com.google.Chrome", "/Applications/Google Chrome.app"), - ("org.mozilla.firefox", "/Applications/Firefox.app"), - ("com.microsoft.edgemac", "/Applications/Microsoft Edge.app"), - ("com.brave.Browser", "/Applications/Brave Browser.app"), - ("com.apple.Safari", "/Applications/Safari.app"), - ]; - - for (bundle_id, app_path) in &browsers { - // 直接检查应用是否存在 - if std::path::Path::new(app_path).exists() { - tracing::info!("检测到已安装的浏览器: {} ({})", bundle_id, app_path); - return Some(bundle_id.to_string()); - } - } - - // Safari 总是存在 - Some("com.apple.Safari".to_string()) - } - - /// 停止拦截 - pub async fn stop(&mut self) -> Result<()> { - if !self.running { - return Ok(()); - } - - tracing::info!("正在停止 macOS 浏览器拦截器..."); - - // 清除全局 URL sender - if let Ok(mut global_sender) = GLOBAL_URL_SENDER.lock() { - *global_sender = None; - } - - // 恢复默认浏览器 - self.restore_default_browser().await?; - - self.running = false; - tracing::info!("macOS 浏览器拦截器已停止"); - Ok(()) - } - - /// 检查是否正在拦截 - pub fn is_running(&self) -> bool { - self.running - } - - /// 恢复系统默认设置 - pub async fn restore_system_defaults(&self) -> Result<()> { - self.restore_default_browser().await - } - - /// 临时禁用拦截 - pub async fn temporarily_disable(&mut self) -> Result<()> { - if !self.running { - return Ok(()); - } - - tracing::info!("正在临时禁用 macOS 浏览器拦截器..."); - - // 临时恢复默认浏览器 - self.restore_default_browser().await?; - - tracing::info!("macOS 浏览器拦截器已临时禁用"); - Ok(()) - } - - /// 重新启用拦截 - pub async fn re_enable(&mut self) -> Result<()> { - if !self.running { - return Err(BrowserInterceptorError::InterceptorError( - "拦截器未运行,无法重新启用".to_string(), - )); - } - - tracing::info!("正在重新启用 macOS 浏览器拦截器..."); - - // 重新设置为默认浏览器 - self.set_as_default_browser().await?; - - tracing::info!("macOS 浏览器拦截器已重新启用"); - Ok(()) - } -} - -impl Drop for MacOSInterceptor { - fn drop(&mut self) { - if self.running { - tracing::info!("macOS 浏览器拦截器资源已清理"); - } - } -} - -/// 处理从 deep-link 接收到的 URL(由 lib.rs 中的事件监听器调用) -pub fn handle_deep_link_url(url: String) { - tracing::info!("收到 deep-link URL: {}", url); - - // 检查是否是 http/https URL - if !url.starts_with("http://") && !url.starts_with("https://") { - tracing::debug!("忽略非 HTTP URL: {}", url); - return; - } - - // 尝试识别来源进程(macOS 上较难获取,使用默认值) - let source_process = detect_source_process(&url); - - let intercepted_url = InterceptedUrl::new(url, source_process); - - // 发送到处理器 - if let Ok(global_sender) = GLOBAL_URL_SENDER.lock() { - if let Some(ref sender) = *global_sender { - match sender.send(intercepted_url) { - Ok(_) => tracing::debug!("URL 已发送到处理器"), - Err(e) => tracing::error!("发送 URL 到处理器失败: {:?}", e), - } - } else { - tracing::warn!("拦截器未运行,忽略 URL"); - } - } else { - tracing::error!("无法获取 URL sender 锁"); - } -} - -/// 尝试检测 URL 的来源进程 -fn detect_source_process(url: &str) -> String { - // 基于 URL 特征推测来源 - if url.contains("kiro") || url.contains("amazon") || url.contains("aws") { - return "Kiro".to_string(); - } - if url.contains("cursor") || url.contains("anysphere") { - return "Cursor".to_string(); - } - if url.contains("vscode") || url.contains("microsoft") || url.contains("visualstudio") { - return "VSCode".to_string(); - } - if url.contains("claude") || url.contains("anthropic") { - return "Claude App".to_string(); - } - if url.contains("github") { - return "GitHub App".to_string(); - } - if url.contains("google") || url.contains("accounts.google") { - return "OAuth Request".to_string(); - } - - // 尝试获取前台应用 - if let Ok(output) = Command::new("osascript") - .args([ - "-e", - "tell application \"System Events\" to get the name of first process whose frontmost is true", - ]) - .output() - { - if output.status.success() { - let app_name = String::from_utf8_lossy(&output.stdout).trim().to_string(); - if !app_name.is_empty() && app_name != "浏览器拦截器" { - return app_name; - } - } - } - - "Unknown App".to_string() -} - -/// 检查拦截器是否有活跃的 URL sender -pub fn is_interceptor_active() -> bool { - GLOBAL_URL_SENDER - .lock() - .map(|sender| sender.is_some()) - .unwrap_or(false) -} diff --git a/src-tauri/src/browser_interceptor/platform/windows.rs b/src-tauri/src/browser_interceptor/platform/windows.rs deleted file mode 100644 index 6db64d404..000000000 --- a/src-tauri/src/browser_interceptor/platform/windows.rs +++ /dev/null @@ -1,323 +0,0 @@ -use crate::browser_interceptor::{BrowserInterceptorError, InterceptedUrl, Result}; -use chrono::Utc; -use std::sync::{Arc, Mutex}; -use std::thread; -use std::time::Duration; -use uuid::Uuid; - -#[cfg(windows)] -use std::ffi::OsStr; -#[cfg(windows)] -use std::iter::once; -#[cfg(windows)] -use std::os::windows::ffi::OsStrExt; - -/// Windows 平台的浏览器拦截器 -pub struct WindowsInterceptor { - running: bool, - original_browser: Option, - temp_exe_path: Option, - intercepted_urls_handler: Arc>>, - monitor_thread: Option>, -} - -#[cfg(windows)] -impl WindowsInterceptor { - pub fn new(url_handler: F) -> Self - where - F: Fn(InterceptedUrl) + Send + Sync + 'static, - { - Self { - running: false, - original_browser: None, - temp_exe_path: None, - intercepted_urls_handler: Arc::new(Mutex::new(Box::new(url_handler))), - monitor_thread: None, - } - } - - /// 启动拦截 - pub async fn start(&mut self) -> Result<()> { - if self.running { - return Err(BrowserInterceptorError::AlreadyRunning); - } - - // 备份当前默认浏览器设置 - self.backup_default_browser().await?; - - // 创建临时拦截程序 - self.create_interceptor_executable().await?; - - // 设置我们的程序为默认浏览器 - self.set_as_default_browser().await?; - - // 启动进程监控 - self.start_process_monitoring().await?; - - self.running = true; - tracing::info!("Windows 浏览器拦截器已启动"); - Ok(()) - } - - /// 停止拦截 - pub async fn stop(&mut self) -> Result<()> { - if !self.running { - return Ok(()); - } - - // 停止进程监控 - if let Some(handle) = self.monitor_thread.take() { - // 发送停止信号,等待线程结束 - handle.join().ok(); - } - - // 恢复原始默认浏览器 - self.restore_default_browser().await?; - - // 清理临时文件 - self.cleanup_temp_files().await?; - - self.running = false; - tracing::info!("Windows 浏览器拦截器已停止"); - Ok(()) - } - - /// 备份当前默认浏览器设置 - async fn backup_default_browser(&mut self) -> Result<()> { - // 简化实现 - tracing::info!("备份默认浏览器设置"); - Ok(()) - } - - /// 创建临时的拦截器可执行文件 - async fn create_interceptor_executable(&mut self) -> Result<()> { - // 创建一个简单的拦截器程序,用于接收 URL 参数 - let temp_dir = std::env::temp_dir(); - - // 创建拦截器脚本内容(批处理脚本) - let bat_content = format!( - r#"@echo off -echo URL被拦截: %1 >> "{}\browser_interception_urls.log" -"#, - temp_dir.to_string_lossy() - ); - - let bat_path = temp_dir.join("browser_interception.bat"); - std::fs::write(&bat_path, bat_content)?; - - self.temp_exe_path = Some(bat_path.to_string_lossy().to_string()); - - Ok(()) - } - - /// 设置我们的程序为默认浏览器 - async fn set_as_default_browser(&self) -> Result<()> { - if let Some(_exe_path) = &self.temp_exe_path { - tracing::info!("已设置拦截器为临时默认浏览器"); - } - - Ok(()) - } - - /// 启动进程监控 - async fn start_process_monitoring(&mut self) -> Result<()> { - let handler = Arc::clone(&self.intercepted_urls_handler); - let temp_dir = std::env::temp_dir(); - let log_file = temp_dir.join("browser_interception_urls.log"); - - let handle = thread::spawn(move || { - loop { - // 检查拦截日志文件 - if let Ok(content) = std::fs::read_to_string(&log_file) { - for line in content.lines() { - if line.starts_with("URL被拦截: ") { - let url = line.replace("URL被拦截: ", ""); - if !url.is_empty() && should_intercept_url(&url) { - let intercepted = InterceptedUrl { - id: Uuid::new_v4().to_string(), - url: url.clone(), - source_process: "Unknown".to_string(), - timestamp: Utc::now(), - copied: false, - opened_in_browser: false, - dismissed: false, - }; - - // 调用处理器 - if let Ok(handler_guard) = handler.lock() { - handler_guard(intercepted); - } - } - } - } - - // 清空日志文件避免重复处理 - std::fs::write(&log_file, "").ok(); - } - - thread::sleep(Duration::from_millis(1000)); - } - }); - - self.monitor_thread = Some(handle); - Ok(()) - } - - /// 恢复原始默认浏览器 - async fn restore_default_browser(&self) -> Result<()> { - if let Some(original_browser) = &self.original_browser { - tracing::info!("已恢复原始默认浏览器: {}", original_browser); - } - - Ok(()) - } - - /// 清理临时文件 - async fn cleanup_temp_files(&self) -> Result<()> { - if let Some(exe_path) = &self.temp_exe_path { - std::fs::remove_file(exe_path).ok(); - } - - let temp_dir = std::env::temp_dir(); - let log_file = temp_dir.join("browser_interception_urls.log"); - std::fs::remove_file(log_file).ok(); - - Ok(()) - } - - /// 检查是否正在拦截 - pub fn is_running(&self) -> bool { - self.running - } - - /// 恢复系统默认设置 - pub async fn restore_system_defaults(&self) -> Result<()> { - // 恢复默认浏览器设置 - self.restore_default_browser().await?; - tracing::info!("系统默认设置已恢复"); - Ok(()) - } - - /// 临时禁用拦截 - pub async fn temporarily_disable(&mut self) -> Result<()> { - if let Some(_original_browser) = &self.original_browser { - self.restore_default_browser().await?; - tracing::info!("拦截器已临时禁用"); - } - Ok(()) - } - - /// 重新启用拦截 - pub async fn re_enable(&mut self) -> Result<()> { - self.set_as_default_browser().await?; - tracing::info!("拦截器已重新启用"); - Ok(()) - } -} - -#[cfg(not(windows))] -impl WindowsInterceptor { - pub fn new(_url_handler: F) -> Self - where - F: Fn(InterceptedUrl) + Send + Sync + 'static, - { - Self { - running: false, - original_browser: None, - temp_exe_path: None, - intercepted_urls_handler: Arc::new(Mutex::new(Box::new(|_| {}))), - monitor_thread: None, - } - } - - pub async fn start(&mut self) -> Result<()> { - Err(BrowserInterceptorError::UnsupportedPlatform( - "Windows interceptor only supports Windows platform".to_string(), - )) - } - - pub async fn stop(&mut self) -> Result<()> { - Ok(()) - } - - pub fn is_running(&self) -> bool { - false - } - - pub async fn restore_system_defaults(&self) -> Result<()> { - Ok(()) - } - - pub async fn temporarily_disable(&mut self) -> Result<()> { - Ok(()) - } - - pub async fn re_enable(&mut self) -> Result<()> { - Ok(()) - } -} - -impl Drop for WindowsInterceptor { - fn drop(&mut self) { - if self.running { - // 在析构时恢复系统默认设置 - tokio::runtime::Handle::try_current().map(|handle| { - handle.block_on(async { - let _ = self.stop().await; - let _ = self.restore_system_defaults().await; - }) - }); - } - } -} - -/// 检查进程是否为目标应用 -pub fn is_target_process(process_name: &str) -> bool { - let target_processes = [ - "kiro", - "kiro.exe", - "cursor", - "cursor.exe", - "code", - "code.exe", - ]; - target_processes - .iter() - .any(|&target| process_name.to_lowercase().contains(&target.to_lowercase())) -} - -/// 检查 URL 是否匹配拦截模式 -pub fn should_intercept_url(url: &str) -> bool { - let patterns = [ - "https://auth.", - "https://accounts.google.com", - "https://github.com/login", - "https://login.microsoftonline.com", - "/oauth/", - "/auth/", - "localhost:8080/auth", // OAuth 回调地址 - ]; - - patterns.iter().any(|&pattern| url.contains(pattern)) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_is_target_process() { - assert!(is_target_process("kiro.exe")); - assert!(is_target_process("cursor")); - assert!(is_target_process("code.exe")); - assert!(!is_target_process("notepad.exe")); - } - - #[test] - fn test_should_intercept_url() { - assert!(should_intercept_url("https://accounts.google.com/oauth")); - assert!(should_intercept_url("https://github.com/login/oauth")); - assert!(should_intercept_url("localhost:8080/auth/callback")); - assert!(!should_intercept_url("https://example.com")); - } -} diff --git a/src-tauri/src/browser_interceptor/state_manager.rs b/src-tauri/src/browser_interceptor/state_manager.rs deleted file mode 100644 index 7935e1a5e..000000000 --- a/src-tauri/src/browser_interceptor/state_manager.rs +++ /dev/null @@ -1,303 +0,0 @@ -use crate::browser_interceptor::{BrowserInterceptorError, InterceptorState, Result}; -use chrono::Utc; -use std::sync::{Arc, RwLock}; -use tokio::time::{Duration, Instant}; - -/// 状态管理器,负责管理拦截器的状态和恢复机制 -pub struct StateManager { - state: Arc>, - original_system_state: Arc>>, - temporary_disable_timer: Arc>>, -} - -/// 系统原始状态备份 -#[derive(Debug, Clone)] -pub struct SystemState { - pub default_browser: Option, - pub registry_backup: std::collections::HashMap, - pub environment_backup: std::collections::HashMap, - pub timestamp: chrono::DateTime, -} - -impl StateManager { - pub fn new() -> Self { - Self { - state: Arc::new(RwLock::new(InterceptorState::default())), - original_system_state: Arc::new(RwLock::new(None)), - temporary_disable_timer: Arc::new(RwLock::new(None)), - } - } - - /// 获取当前状态 - pub fn get_state(&self) -> Result { - self.state - .read() - .map_err(|e| BrowserInterceptorError::StateError(format!("读取状态失败: {e}"))) - .map(|state| state.clone()) - } - - /// 启用拦截器 - pub async fn enable_interceptor(&self) -> Result<()> { - // 备份系统状态 - self.backup_system_state().await?; - - // 更新状态 - { - let mut state = self - .state - .write() - .map_err(|e| BrowserInterceptorError::StateError(format!("写入状态失败: {e}")))?; - - state.enabled = true; - state.can_restore = true; - state.last_activity = Some(Utc::now()); - } - - tracing::info!("浏览器拦截器已启用"); - Ok(()) - } - - /// 禁用拦截器 - pub async fn disable_interceptor(&self) -> Result<()> { - { - let mut state = self - .state - .write() - .map_err(|e| BrowserInterceptorError::StateError(format!("写入状态失败: {e}")))?; - - state.enabled = false; - state.active_hooks.clear(); - state.last_activity = Some(Utc::now()); - } - - tracing::info!("浏览器拦截器已禁用"); - Ok(()) - } - - /// 临时禁用拦截器 - pub async fn temporary_disable(&self, duration_seconds: u64) -> Result<()> { - self.disable_interceptor().await?; - - // 设置定时器 - { - let mut timer = self - .temporary_disable_timer - .write() - .map_err(|e| BrowserInterceptorError::StateError(format!("设置定时器失败: {e}")))?; - *timer = Some(Instant::now() + Duration::from_secs(duration_seconds)); - } - - // 启动后台任务来重新启用 - let state_manager = self.clone(); - tokio::spawn(async move { - tokio::time::sleep(Duration::from_secs(duration_seconds)).await; - if let Err(e) = state_manager.enable_interceptor().await { - tracing::error!("自动重新启用拦截器失败: {}", e); - } else { - tracing::info!("拦截器已自动重新启用"); - } - }); - - tracing::info!("拦截器已临时禁用 {} 秒", duration_seconds); - Ok(()) - } - - /// 恢复正常浏览器行为 - pub async fn restore_normal_behavior(&self) -> Result<()> { - // 先禁用拦截器 - self.disable_interceptor().await?; - - // 恢复系统状态 - self.restore_system_state().await?; - - { - let mut state = self - .state - .write() - .map_err(|e| BrowserInterceptorError::StateError(format!("写入状态失败: {e}")))?; - state.can_restore = false; - } - - tracing::info!("已恢复正常浏览器行为"); - Ok(()) - } - - /// 增加拦截计数 - pub fn increment_intercept_count(&self) -> Result<()> { - let mut state = self - .state - .write() - .map_err(|e| BrowserInterceptorError::StateError(format!("写入状态失败: {e}")))?; - - state.intercepted_count += 1; - state.last_activity = Some(Utc::now()); - - Ok(()) - } - - /// 添加活跃钩子 - pub fn add_active_hook(&self, hook_name: String) -> Result<()> { - let mut state = self - .state - .write() - .map_err(|e| BrowserInterceptorError::StateError(format!("写入状态失败: {e}")))?; - - if !state.active_hooks.contains(&hook_name) { - state.active_hooks.push(hook_name); - } - - Ok(()) - } - - /// 移除活跃钩子 - pub fn remove_active_hook(&self, hook_name: &str) -> Result<()> { - let mut state = self - .state - .write() - .map_err(|e| BrowserInterceptorError::StateError(format!("写入状态失败: {e}")))?; - - state.active_hooks.retain(|h| h != hook_name); - - Ok(()) - } - - /// 备份系统状态 - async fn backup_system_state(&self) -> Result<()> { - let system_state = SystemState { - default_browser: self.get_default_browser().await?, - registry_backup: self.backup_registry_keys().await?, - environment_backup: self.backup_environment_variables().await?, - timestamp: Utc::now(), - }; - - { - let mut backup = self.original_system_state.write().map_err(|e| { - BrowserInterceptorError::StateError(format!("备份系统状态失败: {e}")) - })?; - *backup = Some(system_state); - } - - tracing::info!("系统状态已备份"); - Ok(()) - } - - /// 恢复系统状态 - async fn restore_system_state(&self) -> Result<()> { - let backup = { - let backup_guard = self.original_system_state.read().map_err(|e| { - BrowserInterceptorError::StateError(format!("读取备份状态失败: {e}")) - })?; - backup_guard.clone() - }; - - if let Some(system_state) = backup { - // 恢复默认浏览器 - if let Some(default_browser) = &system_state.default_browser { - self.restore_default_browser(default_browser).await?; - } - - // 恢复注册表项 - self.restore_registry_keys(&system_state.registry_backup) - .await?; - - // 恢复环境变量 - self.restore_environment_variables(&system_state.environment_backup) - .await?; - - tracing::info!("系统状态已恢复到 {} 的备份", system_state.timestamp); - } else { - tracing::warn!("没有找到系统状态备份"); - } - - Ok(()) - } - - /// 获取默认浏览器(平台特定实现) - async fn get_default_browser(&self) -> Result> { - #[cfg(target_os = "macos")] - { - // macOS 实现 - Ok(None) - } - - #[cfg(target_os = "linux")] - { - // Linux 实现 - Ok(None) - } - - #[cfg(target_os = "windows")] - { - // Windows 实现 - 简化版本 - Ok(None) - } - - #[cfg(not(any(target_os = "windows", target_os = "macos", target_os = "linux")))] - { - Ok(None) - } - } - - /// 备份注册表项 - async fn backup_registry_keys(&self) -> Result> { - let backup = std::collections::HashMap::new(); - // 平台特定实现 - Ok(backup) - } - - /// 备份环境变量 - async fn backup_environment_variables( - &self, - ) -> Result> { - let mut backup = std::collections::HashMap::new(); - - // 备份可能影响浏览器启动的环境变量 - if let Ok(browser) = std::env::var("BROWSER") { - backup.insert("BROWSER".to_string(), browser); - } - - Ok(backup) - } - - /// 恢复默认浏览器 - async fn restore_default_browser(&self, _browser: &str) -> Result<()> { - // 平台特定实现 - Ok(()) - } - - /// 恢复注册表项 - async fn restore_registry_keys( - &self, - _backup: &std::collections::HashMap, - ) -> Result<()> { - // 平台特定实现 - Ok(()) - } - - /// 恢复环境变量 - async fn restore_environment_variables( - &self, - backup: &std::collections::HashMap, - ) -> Result<()> { - for (key, value) in backup { - std::env::set_var(key, value); - } - Ok(()) - } -} - -impl Clone for StateManager { - fn clone(&self) -> Self { - Self { - state: Arc::clone(&self.state), - original_system_state: Arc::clone(&self.original_system_state), - temporary_disable_timer: Arc::clone(&self.temporary_disable_timer), - } - } -} - -impl Default for StateManager { - fn default() -> Self { - Self::new() - } -} diff --git a/src-tauri/src/browser_interceptor/url_manager.rs b/src-tauri/src/browser_interceptor/url_manager.rs deleted file mode 100644 index c4ba9d7a6..000000000 --- a/src-tauri/src/browser_interceptor/url_manager.rs +++ /dev/null @@ -1,401 +0,0 @@ -use crate::browser_interceptor::{BrowserInterceptorError, InterceptedUrl, Result}; -use chrono::{DateTime, Utc}; -use serde::{Deserialize, Serialize}; -use std::collections::HashMap; -use std::fs; -use std::path::Path; -use std::sync::{Arc, RwLock}; - -/// URL 管理器,负责管理被拦截的 URL -pub struct UrlManager { - intercepted_urls: Arc>>, - history: Arc>>, - max_history_size: usize, - storage_path: Option, -} - -impl UrlManager { - pub fn new() -> Self { - Self { - intercepted_urls: Arc::new(RwLock::new(HashMap::new())), - history: Arc::new(RwLock::new(Vec::new())), - max_history_size: 1000, // 最多保存 1000 条历史记录 - storage_path: None, - } - } - - /// 创建带持久化存储的 URL 管理器 - pub fn with_storage>(storage_path: P) -> Result { - let storage_path = storage_path.as_ref().to_string_lossy().to_string(); - let mut manager = Self { - intercepted_urls: Arc::new(RwLock::new(HashMap::new())), - history: Arc::new(RwLock::new(Vec::new())), - max_history_size: 1000, - storage_path: Some(storage_path.clone()), - }; - - // 从文件加载历史记录 - manager.load_from_storage()?; - - Ok(manager) - } - - /// 添加被拦截的 URL - pub fn add_intercepted_url(&self, url: String, source_process: String) -> Result { - let intercepted_url = InterceptedUrl::new(url, source_process); - let id = intercepted_url.id.clone(); - let url_for_log = intercepted_url.url.clone(); - let process_for_log = intercepted_url.source_process.clone(); - - // 添加到当前拦截列表 - { - let mut urls = self.intercepted_urls.write().map_err(|e| { - BrowserInterceptorError::StateError(format!("添加拦截 URL 失败: {e}")) - })?; - urls.insert(id.clone(), intercepted_url.clone()); - } - - // 添加到历史记录 - { - let mut history = self.history.write().map_err(|e| { - BrowserInterceptorError::StateError(format!("添加历史记录失败: {e}")) - })?; - - history.push(intercepted_url); - - // 限制历史记录大小 - if history.len() > self.max_history_size { - history.remove(0); - } - } - - // 自动保存 - let _ = self.auto_save(); - - tracing::info!( - "已添加拦截 URL: {} (来源: {})", - url_for_log, - process_for_log - ); - Ok(id) - } - - /// 获取所有当前拦截的 URL - pub fn get_intercepted_urls(&self) -> Result> { - let urls = self - .intercepted_urls - .read() - .map_err(|e| BrowserInterceptorError::StateError(format!("读取拦截 URL 失败: {e}")))?; - - let mut result: Vec = urls.values().cloned().collect(); - result.sort_by(|a, b| b.timestamp.cmp(&a.timestamp)); // 按时间倒序排列 - - Ok(result) - } - - /// 获取指定 ID 的拦截 URL - pub fn get_intercepted_url(&self, id: &str) -> Result> { - let urls = self - .intercepted_urls - .read() - .map_err(|e| BrowserInterceptorError::StateError(format!("读取拦截 URL 失败: {e}")))?; - - Ok(urls.get(id).cloned()) - } - - /// 标记 URL 为已复制 - pub fn mark_as_copied(&self, id: &str) -> Result<()> { - let mut urls = self - .intercepted_urls - .write() - .map_err(|e| BrowserInterceptorError::StateError(format!("更新拦截 URL 失败: {e}")))?; - - if let Some(url) = urls.get_mut(id) { - url.copied = true; - tracing::info!("URL {} 已标记为已复制", id); - } - - Ok(()) - } - - /// 标记 URL 为已在浏览器中打开 - pub fn mark_as_opened(&self, id: &str) -> Result<()> { - let mut urls = self - .intercepted_urls - .write() - .map_err(|e| BrowserInterceptorError::StateError(format!("更新拦截 URL 失败: {e}")))?; - - if let Some(url) = urls.get_mut(id) { - url.opened_in_browser = true; - tracing::info!("URL {} 已标记为已在浏览器中打开", id); - } - - Ok(()) - } - - /// 忽略(移除)指定的 URL - pub fn dismiss_url(&self, id: &str) -> Result<()> { - let mut urls = self - .intercepted_urls - .write() - .map_err(|e| BrowserInterceptorError::StateError(format!("移除拦截 URL 失败: {e}")))?; - - if let Some(mut url) = urls.remove(id) { - url.dismissed = true; - - // 更新历史记录中的状态 - let mut history = self.history.write().map_err(|e| { - BrowserInterceptorError::StateError(format!("更新历史记录失败: {e}")) - })?; - - if let Some(history_url) = history.iter_mut().find(|u| u.id == id) { - history_url.dismissed = true; - } - - tracing::info!("URL {} 已被忽略", id); - } - - Ok(()) - } - - /// 清除所有当前拦截的 URL - pub fn clear_intercepted_urls(&self) -> Result<()> { - let mut urls = self - .intercepted_urls - .write() - .map_err(|e| BrowserInterceptorError::StateError(format!("清除拦截 URL 失败: {e}")))?; - - let count = urls.len(); - urls.clear(); - - tracing::info!("已清除 {} 个拦截的 URL", count); - Ok(()) - } - - /// 获取历史记录 - pub fn get_history(&self, limit: Option) -> Result> { - let history = self - .history - .read() - .map_err(|e| BrowserInterceptorError::StateError(format!("读取历史记录失败: {e}")))?; - - let mut result = history.clone(); - result.sort_by(|a, b| b.timestamp.cmp(&a.timestamp)); // 按时间倒序排列 - - if let Some(limit) = limit { - result.truncate(limit); - } - - Ok(result) - } - - /// 搜索历史记录 - pub fn search_history(&self, query: &str, limit: Option) -> Result> { - let history = self - .history - .read() - .map_err(|e| BrowserInterceptorError::StateError(format!("搜索历史记录失败: {e}")))?; - - let query_lower = query.to_lowercase(); - let mut result: Vec = history - .iter() - .filter(|url| { - url.url.to_lowercase().contains(&query_lower) - || url.source_process.to_lowercase().contains(&query_lower) - }) - .cloned() - .collect(); - - result.sort_by(|a, b| b.timestamp.cmp(&a.timestamp)); - - if let Some(limit) = limit { - result.truncate(limit); - } - - Ok(result) - } - - /// 获取统计信息 - pub fn get_statistics(&self) -> Result { - let urls = self - .intercepted_urls - .read() - .map_err(|e| BrowserInterceptorError::StateError(format!("读取拦截 URL 失败: {e}")))?; - - let history = self - .history - .read() - .map_err(|e| BrowserInterceptorError::StateError(format!("读取历史记录失败: {e}")))?; - - let current_count = urls.len(); - let total_intercepted = history.len(); - let copied_count = history.iter().filter(|u| u.copied).count(); - let opened_count = history.iter().filter(|u| u.opened_in_browser).count(); - let dismissed_count = history.iter().filter(|u| u.dismissed).count(); - - // 统计来源进程 - let mut process_stats = HashMap::new(); - for url in history.iter() { - *process_stats.entry(url.source_process.clone()).or_insert(0) += 1; - } - - Ok(UrlStatistics { - current_intercepted: current_count, - total_intercepted, - copied_count, - opened_count, - dismissed_count, - process_stats, - }) - } - - /// 保存到存储文件 - pub fn save_to_storage(&self) -> Result<()> { - if let Some(storage_path) = &self.storage_path { - let history = self.history.read().map_err(|e| { - BrowserInterceptorError::StateError(format!("读取历史记录失败: {e}")) - })?; - - let storage_data = UrlStorageData { - history: history.clone(), - max_history_size: self.max_history_size, - saved_at: Utc::now(), - }; - - let json_data = serde_json::to_string_pretty(&storage_data) - .map_err(|e| BrowserInterceptorError::StateError(format!("序列化数据失败: {e}")))?; - - // 确保目录存在 - if let Some(parent) = Path::new(storage_path).parent() { - fs::create_dir_all(parent).map_err(|e| { - BrowserInterceptorError::StateError(format!("创建目录失败: {e}")) - })?; - } - - fs::write(storage_path, json_data) - .map_err(|e| BrowserInterceptorError::StateError(format!("写入文件失败: {e}")))?; - - tracing::info!("已保存历史记录到: {}", storage_path); - } - - Ok(()) - } - - /// 从存储文件加载 - pub fn load_from_storage(&mut self) -> Result<()> { - if let Some(storage_path) = &self.storage_path { - if Path::new(storage_path).exists() { - let json_data = fs::read_to_string(storage_path).map_err(|e| { - BrowserInterceptorError::StateError(format!("读取文件失败: {e}")) - })?; - - let storage_data: UrlStorageData = - serde_json::from_str(&json_data).map_err(|e| { - BrowserInterceptorError::StateError(format!("反序列化数据失败: {e}")) - })?; - - { - let mut history = self.history.write().map_err(|e| { - BrowserInterceptorError::StateError(format!("写入历史记录失败: {e}")) - })?; - *history = storage_data.history; - } - - self.max_history_size = storage_data.max_history_size; - - tracing::info!("已从 {} 加载历史记录", storage_path); - } - } - - Ok(()) - } - - /// 自动保存(如果已配置存储路径) - fn auto_save(&self) -> Result<()> { - if self.storage_path.is_some() { - self.save_to_storage()?; - } - Ok(()) - } -} - -/// 存储数据结构 -#[derive(Debug, Clone, Serialize, Deserialize)] -struct UrlStorageData { - history: Vec, - max_history_size: usize, - saved_at: DateTime, -} - -/// URL 统计信息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct UrlStatistics { - pub current_intercepted: usize, - pub total_intercepted: usize, - pub copied_count: usize, - pub opened_count: usize, - pub dismissed_count: usize, - pub process_stats: HashMap, -} - -impl Default for UrlManager { - fn default() -> Self { - Self::new() - } -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_add_and_get_intercepted_url() { - let manager = UrlManager::new(); - - let id = manager - .add_intercepted_url( - "https://accounts.google.com/oauth/authorize".to_string(), - "kiro".to_string(), - ) - .unwrap(); - - let urls = manager.get_intercepted_urls().unwrap(); - assert_eq!(urls.len(), 1); - assert_eq!(urls[0].id, id); - assert_eq!(urls[0].url, "https://accounts.google.com/oauth/authorize"); - assert_eq!(urls[0].source_process, "kiro"); - } - - #[test] - fn test_mark_as_copied() { - let manager = UrlManager::new(); - - let id = manager - .add_intercepted_url("https://test.com".to_string(), "test".to_string()) - .unwrap(); - - manager.mark_as_copied(&id).unwrap(); - - let url = manager.get_intercepted_url(&id).unwrap().unwrap(); - assert!(url.copied); - } - - #[test] - fn test_dismiss_url() { - let manager = UrlManager::new(); - - let id = manager - .add_intercepted_url("https://test.com".to_string(), "test".to_string()) - .unwrap(); - - manager.dismiss_url(&id).unwrap(); - - let urls = manager.get_intercepted_urls().unwrap(); - assert_eq!(urls.len(), 0); - - // 但历史记录中应该还存在 - let history = manager.get_history(None).unwrap(); - assert_eq!(history.len(), 1); - assert!(history[0].dismissed); - } -} diff --git a/src-tauri/src/commands/browser_interceptor_cmd.rs b/src-tauri/src/commands/browser_interceptor_cmd.rs deleted file mode 100644 index 78a14cb7e..000000000 --- a/src-tauri/src/commands/browser_interceptor_cmd.rs +++ /dev/null @@ -1,247 +0,0 @@ -use crate::browser_interceptor::{ - BrowserInterceptor, BrowserInterceptorConfig, InterceptedUrl, InterceptorState, UrlStatistics, -}; -use once_cell::sync::Lazy; -use std::sync::Arc; -use tokio::sync::RwLock; - -/// 全局拦截器实例 -static INTERCEPTOR: Lazy>>> = - Lazy::new(|| Arc::new(RwLock::new(None))); - -/// 获取拦截器状态 -#[tauri::command] -pub async fn get_browser_interceptor_state() -> Result, String> { - let interceptor = INTERCEPTOR.read().await; - if let Some(ref int) = *interceptor { - int.get_state().await.map(Some).map_err(|e| e.to_string()) - } else { - Ok(None) - } -} - -/// 启动拦截器 -#[tauri::command] -pub async fn start_browser_interceptor(config: BrowserInterceptorConfig) -> Result { - let mut interceptor_guard = INTERCEPTOR.write().await; - - // 如果已有拦截器在运行,先停止 - if let Some(ref mut int) = *interceptor_guard { - int.stop().await.map_err(|e| e.to_string())?; - } - - // 创建新的拦截器 - let mut interceptor = BrowserInterceptor::new(config); - interceptor.start().await.map_err(|e| e.to_string())?; - - *interceptor_guard = Some(interceptor); - - Ok("拦截器已启动".to_string()) -} - -/// 停止拦截器 -#[tauri::command] -pub async fn stop_browser_interceptor() -> Result { - let mut interceptor_guard = INTERCEPTOR.write().await; - - if let Some(ref mut int) = *interceptor_guard { - int.stop().await.map_err(|e| e.to_string())?; - *interceptor_guard = None; - Ok("拦截器已停止".to_string()) - } else { - Ok("拦截器未运行".to_string()) - } -} - -/// 恢复正常浏览器行为 -#[tauri::command] -pub async fn restore_normal_browser_behavior() -> Result { - let mut interceptor_guard = INTERCEPTOR.write().await; - - if let Some(ref mut int) = *interceptor_guard { - int.restore_normal_behavior() - .await - .map_err(|e| e.to_string())?; - *interceptor_guard = None; - Ok("已恢复正常浏览器行为".to_string()) - } else { - Ok("拦截器未运行".to_string()) - } -} - -/// 临时禁用拦截器 -#[tauri::command] -pub async fn temporary_disable_interceptor(duration_seconds: u64) -> Result { - let mut interceptor_guard = INTERCEPTOR.write().await; - - if let Some(ref mut int) = *interceptor_guard { - int.temporary_disable(duration_seconds) - .await - .map_err(|e| e.to_string())?; - Ok(format!("拦截器已临时禁用 {duration_seconds} 秒")) - } else { - Err("拦截器未运行".to_string()) - } -} - -/// 获取拦截的 URL 列表 -#[tauri::command] -pub async fn get_intercepted_urls() -> Result, String> { - let interceptor = INTERCEPTOR.read().await; - - if let Some(ref int) = *interceptor { - int.get_intercepted_urls().await.map_err(|e| e.to_string()) - } else { - Ok(Vec::new()) - } -} - -/// 获取历史记录 -#[tauri::command] -pub async fn get_interceptor_history(limit: Option) -> Result, String> { - let interceptor = INTERCEPTOR.read().await; - - if let Some(ref int) = *interceptor { - int.get_history(limit).await.map_err(|e| e.to_string()) - } else { - Ok(Vec::new()) - } -} - -/// 复制 URL 到剪贴板 -#[tauri::command] -pub async fn copy_intercepted_url_to_clipboard(url_id: String) -> Result { - let interceptor = INTERCEPTOR.read().await; - - if let Some(ref int) = *interceptor { - int.copy_url_to_clipboard(&url_id) - .await - .map_err(|e| e.to_string())?; - Ok("URL 已复制到剪贴板".to_string()) - } else { - Err("拦截器未运行".to_string()) - } -} - -/// 在指纹浏览器中打开 URL -#[tauri::command] -pub async fn open_url_in_fingerprint_browser(url_id: String) -> Result { - let interceptor = INTERCEPTOR.read().await; - - if let Some(ref int) = *interceptor { - int.open_in_fingerprint_browser(&url_id) - .await - .map_err(|e| e.to_string())?; - Ok("URL 已在指纹浏览器中打开".to_string()) - } else { - Err("拦截器未运行".to_string()) - } -} - -/// 忽略 URL -#[tauri::command] -pub async fn dismiss_intercepted_url(url_id: String) -> Result { - let interceptor = INTERCEPTOR.read().await; - - if let Some(ref int) = *interceptor { - int.dismiss_url(&url_id).await.map_err(|e| e.to_string())?; - Ok("URL 已忽略".to_string()) - } else { - Err("拦截器未运行".to_string()) - } -} - -/// 更新配置 -#[tauri::command] -pub async fn update_browser_interceptor_config( - config: BrowserInterceptorConfig, -) -> Result { - let interceptor = INTERCEPTOR.read().await; - - if let Some(ref int) = *interceptor { - int.update_config(config).await.map_err(|e| e.to_string())?; - Ok("配置已更新".to_string()) - } else { - Err("拦截器未运行".to_string()) - } -} - -/// 获取默认配置 -#[tauri::command] -pub async fn get_default_browser_interceptor_config() -> Result { - Ok(BrowserInterceptorConfig::default()) -} - -/// 验证配置 -#[tauri::command] -pub async fn validate_browser_interceptor_config( - config: BrowserInterceptorConfig, -) -> Result { - config.validate()?; - Ok("配置验证通过".to_string()) -} - -/// 检查是否正在运行 -#[tauri::command] -pub async fn is_browser_interceptor_running() -> Result { - let interceptor = INTERCEPTOR.read().await; - Ok(interceptor.is_some()) -} - -/// 获取统计信息 -#[tauri::command] -pub async fn get_browser_interceptor_statistics() -> Result { - let interceptor = INTERCEPTOR.read().await; - - if let Some(ref _int) = *interceptor { - // 简化实现,返回默认统计 - Ok(UrlStatistics { - current_intercepted: 0, - total_intercepted: 0, - copied_count: 0, - opened_count: 0, - dismissed_count: 0, - process_stats: std::collections::HashMap::new(), - }) - } else { - Ok(UrlStatistics { - current_intercepted: 0, - total_intercepted: 0, - copied_count: 0, - opened_count: 0, - dismissed_count: 0, - process_stats: std::collections::HashMap::new(), - }) - } -} - -/// 显示通知 -#[tauri::command] -pub async fn show_notification( - title: String, - body: String, - _icon: Option, -) -> Result { - tracing::info!("[通知] {}: {}", title, body); - Ok("通知已显示".to_string()) -} - -/// 显示 URL 拦截通知 -#[tauri::command] -pub async fn show_url_intercept_notification( - url: String, - source_process: String, -) -> Result { - tracing::info!("[URL 拦截] 来自 {}: {}", source_process, url); - Ok("通知已显示".to_string()) -} - -/// 显示状态通知 -#[tauri::command] -pub async fn show_status_notification( - message: String, - notification_type: String, -) -> Result { - tracing::info!("[状态通知] [{}] {}", notification_type, message); - Ok("通知已显示".to_string()) -} diff --git a/src-tauri/src/commands/flow_monitor_cmd.rs b/src-tauri/src/commands/flow_monitor_cmd.rs deleted file mode 100644 index 9a8cb8029..000000000 --- a/src-tauri/src/commands/flow_monitor_cmd.rs +++ /dev/null @@ -1,3429 +0,0 @@ -//! Flow Monitor Tauri 命令 -//! -//! 提供 LLM Flow Monitor 的 Tauri 命令接口,用于前端访问 Flow 数据。 -//! -//! **Validates: Requirements 10.1-10.7** - -use serde::{Deserialize, Serialize}; -use std::sync::Arc; -use tauri::State; - -use crate::flow_monitor::{ - get_filter_help, BatchOperation, BatchOperations, BatchResult, DiffConfig, ExportFormat, - ExportOptions, FilterExpr, FilterParser, FlowAnnotations, FlowDiff, FlowDiffResult, - FlowExporter, FlowFilter, FlowMonitor, FlowQueryResult, FlowQueryService, FlowSearchResult, - FlowSortBy, FlowStats, LLMFlow, FILTER_HELP, -}; - -// ============================================================================ -// 状态封装 -// ============================================================================ - -/// FlowMonitor 状态封装 -pub struct FlowMonitorState(pub Arc); - -/// FlowQueryService 状态封装 -pub struct FlowQueryServiceState(pub Arc); - -// ============================================================================ -// 请求/响应类型 -// ============================================================================ - -/// 查询 Flow 请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct QueryFlowsRequest { - /// 过滤条件 - #[serde(default)] - pub filter: FlowFilter, - /// 排序字段 - #[serde(default)] - pub sort_by: FlowSortBy, - /// 是否降序 - #[serde(default = "default_true")] - pub sort_desc: bool, - /// 页码(从 1 开始) - #[serde(default = "default_page")] - pub page: usize, - /// 每页大小 - #[serde(default = "default_page_size")] - pub page_size: usize, -} - -fn default_true() -> bool { - true -} - -fn default_page() -> usize { - 1 -} - -fn default_page_size() -> usize { - 20 -} - -impl Default for QueryFlowsRequest { - fn default() -> Self { - Self { - filter: FlowFilter::default(), - sort_by: FlowSortBy::default(), - sort_desc: true, - page: 1, - page_size: 20, - } - } -} - -/// 搜索 Flow 请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct SearchFlowsRequest { - /// 搜索关键词 - pub query: String, - /// 最大返回数量 - #[serde(default = "default_search_limit")] - pub limit: usize, -} - -fn default_search_limit() -> usize { - 50 -} - -/// 导出 Flow 请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ExportFlowsRequest { - /// 导出格式 - pub format: ExportFormat, - /// 过滤条件 - #[serde(default)] - pub filter: Option, - /// 是否包含原始请求/响应体 - #[serde(default = "default_true")] - pub include_raw: bool, - /// 是否包含流式 chunks - #[serde(default)] - pub include_stream_chunks: bool, - /// 是否脱敏敏感数据 - #[serde(default)] - pub redact_sensitive: bool, - /// Flow ID 列表(如果指定,则只导出这些 Flow) - #[serde(default)] - pub flow_ids: Option>, -} - -/// 导出结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ExportFlowsResponse { - /// 导出的数据(JSON 字符串) - pub data: String, - /// 导出的 Flow 数量 - pub count: usize, - /// 导出格式 - pub format: ExportFormat, -} - -/// 更新标注请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct UpdateAnnotationsRequest { - /// Flow ID - pub flow_id: String, - /// 标注信息 - pub annotations: FlowAnnotations, -} - -/// 清理 Flow 请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CleanupFlowsRequest { - /// 清理类型 - pub cleanup_type: CleanupType, - /// 保留天数(清理此天数之前的数据)- 仅当 cleanup_type 为 ByTime 时使用 - pub retention_days: Option, - /// 保留小时数(清理此小时数之前的数据)- 仅当 cleanup_type 为 ByTime 时使用 - pub retention_hours: Option, - /// 保留的最大记录数 - 仅当 cleanup_type 为 ByCount 时使用 - pub max_records: Option, - /// 要清理的状态列表 - 仅当 cleanup_type 为 ByStatus 时使用 - pub target_states: Option>, - /// 要清理的 Provider 列表 - 仅当 cleanup_type 为 ByProvider 时使用 - pub target_providers: Option>, - /// 最大存储大小(字节)- 仅当 cleanup_type 为 BySize 时使用 - pub max_storage_bytes: Option, -} - -/// 清理类型 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub enum CleanupType { - /// 删除所有日志 - All, - /// 按时间清理(保留最近的数据) - ByTime, - /// 按数量清理(只保留最近N条记录) - ByCount, - /// 按状态清理(删除特定状态的Flow) - ByStatus, - /// 按Provider清理(删除特定Provider的日志) - ByProvider, - /// 按存储大小清理(当超过指定大小时清理最旧的数据) - BySize, -} - -/// 清理结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CleanupFlowsResponse { - /// 清理的 Flow 数量 - pub cleaned_count: usize, - /// 清理的文件数量 - pub cleaned_files: usize, - /// 释放的空间(字节) - pub freed_bytes: u64, -} - -// ============================================================================ -// Tauri 命令实现 -// ============================================================================ - -/// 查询 Flow 列表 -/// -/// **Validates: Requirements 10.1** -/// -/// # Arguments -/// * `request` - 查询请求参数 -/// * `query_service` - 查询服务状态 -/// -/// # Returns -/// * `Ok(FlowQueryResult)` - 成功时返回查询结果 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn query_flows( - request: QueryFlowsRequest, - query_service: State<'_, FlowQueryServiceState>, -) -> Result { - query_service - .0 - .query( - request.filter, - request.sort_by, - request.sort_desc, - request.page, - request.page_size, - ) - .await - .map_err(|e| format!("查询 Flow 失败: {e}")) -} - -/// 获取单个 Flow 详情 -/// -/// **Validates: Requirements 10.2** -/// -/// # Arguments -/// * `flow_id` - Flow ID -/// * `query_service` - 查询服务状态 -/// -/// # Returns -/// * `Ok(Some(LLMFlow))` - 成功时返回 Flow 详情 -/// * `Ok(None)` - Flow 不存在 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn get_flow_detail( - flow_id: String, - query_service: State<'_, FlowQueryServiceState>, -) -> Result, String> { - query_service - .0 - .get_flow(&flow_id) - .await - .map_err(|e| format!("获取 Flow 详情失败: {e}")) -} - -/// 全文搜索 Flow -/// -/// **Validates: Requirements 10.3** -/// -/// # Arguments -/// * `request` - 搜索请求参数 -/// * `query_service` - 查询服务状态 -/// -/// # Returns -/// * `Ok(Vec)` - 成功时返回搜索结果 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn search_flows( - request: SearchFlowsRequest, - query_service: State<'_, FlowQueryServiceState>, -) -> Result, String> { - query_service - .0 - .search(&request.query, request.limit) - .await - .map_err(|e| format!("搜索 Flow 失败: {e}")) -} - -/// 获取 Flow 统计信息 -/// -/// **Validates: Requirements 10.4** -/// -/// # Arguments -/// * `filter` - 过滤条件(可选) -/// * `query_service` - 查询服务状态 -/// -/// # Returns -/// * `Ok(FlowStats)` - 成功时返回统计信息 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn get_flow_stats( - filter: Option, - query_service: State<'_, FlowQueryServiceState>, -) -> Result { - let filter = filter.unwrap_or_default(); - Ok(query_service.0.get_stats(&filter).await) -} - -/// 导出 Flow -/// -/// **Validates: Requirements 10.5** -/// -/// # Arguments -/// * `request` - 导出请求参数 -/// * `query_service` - 查询服务状态 -/// -/// # Returns -/// * `Ok(ExportFlowsResponse)` - 成功时返回导出结果 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn export_flows( - request: ExportFlowsRequest, - query_service: State<'_, FlowQueryServiceState>, -) -> Result { - // 获取要导出的 Flow - let flows = if let Some(flow_ids) = request.flow_ids { - // 按 ID 列表获取 - let mut flows = Vec::new(); - for id in flow_ids { - if let Ok(Some(flow)) = query_service.0.get_flow(&id).await { - flows.push(flow); - } - } - flows - } else { - // 按过滤条件获取 - let filter = request.filter.unwrap_or_default(); - let result = query_service - .0 - .query(filter, FlowSortBy::CreatedAt, true, 1, 10000) - .await - .map_err(|e| format!("查询 Flow 失败: {e}"))?; - result.flows - }; - - let count = flows.len(); - - // 创建导出器 - let options = ExportOptions { - format: request.format, - filter: None, - include_raw: request.include_raw, - include_stream_chunks: request.include_stream_chunks, - redact_sensitive: request.redact_sensitive, - redaction_rules: Vec::new(), - compress: false, - }; - let exporter = FlowExporter::new(options); - - // 导出数据 - let data = match request.format { - ExportFormat::HAR => { - let har = exporter.export_har(&flows); - serde_json::to_string_pretty(&har).map_err(|e| format!("序列化 HAR 失败: {e}"))? - } - ExportFormat::JSON => { - let json = exporter.export_json(&flows); - serde_json::to_string_pretty(&json).map_err(|e| format!("序列化 JSON 失败: {e}"))? - } - ExportFormat::JSONL => exporter.export_jsonl(&flows), - ExportFormat::Markdown => exporter.export_markdown_multiple(&flows), - ExportFormat::CSV => exporter.export_csv(&flows), - }; - - Ok(ExportFlowsResponse { - data, - count, - format: request.format, - }) -} - -/// 更新 Flow 标注 -/// -/// **Validates: Requirements 10.6** -/// -/// # Arguments -/// * `request` - 更新标注请求参数 -/// * `monitor` - Flow 监控服务状态 -/// -/// # Returns -/// * `Ok(bool)` - 成功时返回是否更新成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn update_flow_annotations( - request: UpdateAnnotationsRequest, - monitor: State<'_, FlowMonitorState>, -) -> Result { - let updated = monitor - .0 - .update_annotations(&request.flow_id, request.annotations) - .await; - Ok(updated) -} - -/// 切换 Flow 收藏状态 -/// -/// **Validates: Requirements 10.6** -/// -/// # Arguments -/// * `flow_id` - Flow ID -/// * `monitor` - Flow 监控服务状态 -/// -/// # Returns -/// * `Ok(bool)` - 成功时返回是否更新成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn toggle_flow_starred( - flow_id: String, - monitor: State<'_, FlowMonitorState>, -) -> Result { - let updated = monitor.0.toggle_starred(&flow_id).await; - Ok(updated) -} - -/// 添加 Flow 评论 -/// -/// **Validates: Requirements 10.6** -/// -/// # Arguments -/// * `flow_id` - Flow ID -/// * `comment` - 评论内容 -/// * `monitor` - Flow 监控服务状态 -/// -/// # Returns -/// * `Ok(bool)` - 成功时返回是否更新成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn add_flow_comment( - flow_id: String, - comment: String, - monitor: State<'_, FlowMonitorState>, -) -> Result { - let updated = monitor.0.add_comment(&flow_id, comment).await; - Ok(updated) -} - -/// 添加 Flow 标签 -/// -/// **Validates: Requirements 10.6** -/// -/// # Arguments -/// * `flow_id` - Flow ID -/// * `tag` - 标签 -/// * `monitor` - Flow 监控服务状态 -/// -/// # Returns -/// * `Ok(bool)` - 成功时返回是否更新成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn add_flow_tag( - flow_id: String, - tag: String, - monitor: State<'_, FlowMonitorState>, -) -> Result { - let updated = monitor.0.add_tag(&flow_id, tag).await; - Ok(updated) -} - -/// 移除 Flow 标签 -/// -/// **Validates: Requirements 10.6** -/// -/// # Arguments -/// * `flow_id` - Flow ID -/// * `tag` - 标签 -/// * `monitor` - Flow 监控服务状态 -/// -/// # Returns -/// * `Ok(bool)` - 成功时返回是否更新成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn remove_flow_tag( - flow_id: String, - tag: String, - monitor: State<'_, FlowMonitorState>, -) -> Result { - let updated = monitor.0.remove_tag(&flow_id, &tag).await; - Ok(updated) -} - -/// 设置 Flow 标记 -/// -/// **Validates: Requirements 10.6** -/// -/// # Arguments -/// * `flow_id` - Flow ID -/// * `marker` - 标记(如 ⭐、🔴、🟢,None 表示清除标记) -/// * `monitor` - Flow 监控服务状态 -/// -/// # Returns -/// * `Ok(bool)` - 成功时返回是否更新成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn set_flow_marker( - flow_id: String, - marker: Option, - monitor: State<'_, FlowMonitorState>, -) -> Result { - let updated = monitor.0.set_marker(&flow_id, marker).await; - Ok(updated) -} - -/// 清理旧的 Flow 数据 -/// -/// **Validates: Requirements 10.7** -/// -/// # Arguments -/// * `request` - 清理请求参数 -/// * `monitor` - Flow 监控服务状态 -/// -/// # Returns -/// * `Ok(CleanupFlowsResponse)` - 成功时返回清理结果 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn cleanup_flows( - request: CleanupFlowsRequest, - monitor: State<'_, FlowMonitorState>, -) -> Result { - let mut cleaned_count = 0; - let mut cleaned_files = 0; - let mut freed_bytes = 0u64; - - // 根据清理类型执行不同的清理逻辑 - match request.cleanup_type { - CleanupType::All => { - // 清理所有数据 - 使用当前时间 + 1天,确保删除所有数据 - let before = chrono::Utc::now() + chrono::Duration::days(1); - - // 清理文件存储 - if let Some(file_store) = monitor.0.file_store() { - match file_store.cleanup(before) { - Ok(result) => { - cleaned_count = result.flows_deleted; - cleaned_files = result.files_deleted; - freed_bytes = result.bytes_freed; - } - Err(e) => { - tracing::error!("清理所有数据失败: {}", e); - return Err(format!("清理所有数据失败: {e}")); - } - } - } - - // 清理内存存储 - { - let memory_store = monitor.0.memory_store(); - let mut store = memory_store.write().await; - store.clear(); - tracing::info!("已清理内存存储"); - } - } - - CleanupType::ByTime => { - // 按时间清理 - let before = if let Some(hours) = request.retention_hours { - chrono::Utc::now() - chrono::Duration::hours(hours as i64) - } else if let Some(days) = request.retention_days { - chrono::Utc::now() - chrono::Duration::days(days as i64) - } else { - return Err("时间清理需要指定保留时间".to_string()); - }; - - // 清理文件存储 - if let Some(file_store) = monitor.0.file_store() { - match file_store.cleanup(before) { - Ok(result) => { - cleaned_count = result.flows_deleted; - cleaned_files = result.files_deleted; - freed_bytes = result.bytes_freed; - } - Err(e) => { - tracing::error!("按时间清理失败: {}", e); - return Err(format!("按时间清理失败: {e}")); - } - } - } - - // 清理内存存储 - { - let memory_store = monitor.0.memory_store(); - let mut store = memory_store.write().await; - let memory_cleaned = store.cleanup_before(before); - cleaned_count += memory_cleaned; - tracing::info!("已清理内存存储 {} 条记录", memory_cleaned); - } - } - - CleanupType::ByCount => { - // 按数量清理 - 使用retention清理,这是一个简化实现 - if let Some(file_store) = monitor.0.file_store() { - match file_store.cleanup_by_retention() { - Ok(result) => { - cleaned_count = result.flows_deleted; - cleaned_files = result.files_deleted; - freed_bytes = result.bytes_freed; - } - Err(e) => { - tracing::error!("按数量清理失败: {}", e); - return Err(format!("按数量清理失败: {e}")); - } - } - } - - // 清理内存存储 - 对于按数量清理,清空所有内存数据 - // 因为文件存储已按保留策略清理,内存应保持一致 - { - let memory_store = monitor.0.memory_store(); - let mut store = memory_store.write().await; - store.clear(); - tracing::info!("已清理内存存储"); - } - } - - CleanupType::ByStatus => { - // 按状态清理 - 暂时不支持,返回错误提示 - return Err("按状态清理功能暂未实现,请使用按时间清理".to_string()); - } - - CleanupType::ByProvider => { - // 按Provider清理 - 暂时不支持,返回错误提示 - return Err("按Provider清理功能暂未实现,请使用按时间清理".to_string()); - } - - CleanupType::BySize => { - // 按大小清理 - 暂时不支持,返回错误提示 - return Err("按大小清理功能暂未实现,请使用按时间清理".to_string()); - } - } - - Ok(CleanupFlowsResponse { - cleaned_count, - cleaned_files, - freed_bytes, - }) -} - -/// 获取最近的 Flow 列表 -/// -/// **Validates: Requirements 10.1** -/// -/// # Arguments -/// * `limit` - 最大返回数量 -/// * `query_service` - 查询服务状态 -/// -/// # Returns -/// * `Ok(Vec)` - 成功时返回 Flow 列表 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn get_recent_flows( - limit: Option, - query_service: State<'_, FlowQueryServiceState>, -) -> Result, String> { - let limit = limit.unwrap_or(20); - Ok(query_service.0.get_recent(limit).await) -} - -/// 获取 Flow Monitor 状态 -/// -/// **Validates: Requirements 10.1** -/// -/// # Arguments -/// * `monitor` - Flow 监控服务状态 -/// -/// # Returns -/// * `Ok(FlowMonitorStatus)` - 成功时返回监控状态 -/// * `Err(String)` - 失败时返回错误消息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct FlowMonitorStatus { - /// 是否启用 - pub enabled: bool, - /// 活跃 Flow 数量 - pub active_flow_count: usize, - /// 内存中的 Flow 数量 - pub memory_flow_count: usize, - /// 最大内存 Flow 数量 - pub max_memory_flows: usize, -} - -#[tauri::command] -pub async fn get_flow_monitor_status( - monitor: State<'_, FlowMonitorState>, -) -> Result { - let config = monitor.0.config().await; - Ok(FlowMonitorStatus { - enabled: monitor.0.is_enabled().await, - active_flow_count: monitor.0.active_flow_count().await, - memory_flow_count: monitor.0.memory_flow_count().await, - max_memory_flows: config.max_memory_flows, - }) -} - -/// 获取 Flow Monitor 状态(调试用) -/// -/// **Validates: Requirements 10.1** -/// -/// # Arguments -/// * `monitor` - Flow 监控服务状态 -/// * `query_service` - 查询服务状态 -/// -/// # Returns -/// * `Ok(FlowMonitorDebugInfo)` - 成功时返回调试信息 -/// * `Err(String)` - 失败时返回错误消息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct FlowMonitorDebugInfo { - /// 是否启用 - pub enabled: bool, - /// 活跃 Flow 数量 - pub active_flow_count: usize, - /// 内存中的 Flow 数量 - pub memory_flow_count: usize, - /// 最大内存 Flow 数量 - pub max_memory_flows: usize, - /// 内存中的 Flow ID 列表(最多显示10个) - pub memory_flow_ids: Vec, - /// 配置信息 - pub config_enabled: bool, -} - -#[tauri::command] -pub async fn get_flow_monitor_debug_info( - monitor: State<'_, FlowMonitorState>, - query_service: State<'_, FlowQueryServiceState>, -) -> Result { - let config = monitor.0.config().await; - let recent_flows = query_service.0.get_recent(10).await; - - Ok(FlowMonitorDebugInfo { - enabled: monitor.0.is_enabled().await, - active_flow_count: monitor.0.active_flow_count().await, - memory_flow_count: monitor.0.memory_flow_count().await, - max_memory_flows: config.max_memory_flows, - memory_flow_ids: recent_flows.into_iter().map(|f| f.id).collect(), - config_enabled: config.enabled, - }) -} - -/// 启用 Flow Monitor -/// -/// **Validates: Requirements 10.1** -/// -/// # Arguments -/// * `monitor` - Flow 监控服务状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn enable_flow_monitor(monitor: State<'_, FlowMonitorState>) -> Result<(), String> { - monitor.0.enable().await; - Ok(()) -} - -/// 创建测试 Flow 数据(仅用于调试) -/// -/// **Validates: Requirements 10.1** -/// -/// # Arguments -/// * `count` - 要创建的测试 Flow 数量 -/// * `monitor` - Flow 监控服务状态 -/// -/// # Returns -/// * `Ok(usize)` - 成功创建的 Flow 数量 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn create_test_flows( - count: Option, - monitor: State<'_, FlowMonitorState>, -) -> Result { - use crate::flow_monitor::{ - ClientInfo, FlowMetadata, LLMRequest, Message, MessageRole, ProviderType, - RequestParameters, RoutingInfo, - }; - use chrono::Utc; - - let count = count.unwrap_or(5); - let mut created = 0; - - for i in 0..count { - // 创建测试请求 - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - headers: std::collections::HashMap::new(), - body: serde_json::json!({ - "model": format!("gpt-4-test-{}", i), - "messages": [{"role": "user", "content": format!("测试消息 {}", i)}] - }), - messages: vec![Message { - role: MessageRole::User, - content: crate::flow_monitor::MessageContent::Text(format!("测试消息 {i}")), - tool_calls: None, - tool_result: None, - name: None, - }], - system_prompt: None, - tools: None, - model: format!("gpt-4-test-{i}"), - original_model: None, - parameters: RequestParameters { - temperature: Some(0.7), - top_p: Some(1.0), - max_tokens: Some(1000), - stop: None, - stream: false, - extra: std::collections::HashMap::new(), - }, - size_bytes: 100 + i * 10, - timestamp: Utc::now(), - }; - - // 创建测试元数据 - let metadata = FlowMetadata { - provider: ProviderType::OpenAI, - provider_id: Some("openai".to_string()), - credential_id: Some(format!("test-cred-{i}")), - credential_name: Some(format!("测试凭证 {i}")), - retry_count: 0, - client_info: ClientInfo { - ip: Some("127.0.0.1".to_string()), - user_agent: Some("test-agent".to_string()), - request_id: Some(format!("test-req-{i}")), - }, - routing_info: RoutingInfo { - target_url: Some("https://api.openai.com".to_string()), - route_rule: None, - load_balance_strategy: None, - }, - injected_params: None, - context_usage_percentage: Some(50.0), - }; - - // 启动 Flow - if let Some(flow_id) = monitor.0.start_flow(request, metadata).await { - // 模拟完成 Flow - let response = crate::flow_monitor::LLMResponse { - status_code: 200, - status_text: "OK".to_string(), - headers: std::collections::HashMap::new(), - body: serde_json::json!({ - "choices": [{"message": {"role": "assistant", "content": format!("测试响应 {}", i)}}] - }), - content: format!("测试响应 {i}"), - thinking: None, - tool_calls: Vec::new(), - usage: crate::flow_monitor::TokenUsage { - input_tokens: 10 + i as u32, - output_tokens: 20 + i as u32, - cache_read_tokens: None, - cache_write_tokens: None, - thinking_tokens: None, - total_tokens: 30 + i as u32 * 2, - }, - stop_reason: Some(crate::flow_monitor::StopReason::Stop), - size_bytes: 200 + i * 15, - timestamp_start: Utc::now(), - timestamp_end: Utc::now(), - stream_info: None, - }; - - monitor.0.complete_flow(&flow_id, Some(response)).await; - created += 1; - } - } - - Ok(created) -} - -/// 禁用 Flow Monitor -/// -/// **Validates: Requirements 10.1** -/// -/// # Arguments -/// * `monitor` - Flow 监控服务状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn disable_flow_monitor(monitor: State<'_, FlowMonitorState>) -> Result<(), String> { - monitor.0.disable().await; - Ok(()) -} - -// ============================================================================ -// 测试模块 -// ============================================================================ - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_query_flows_request_default() { - let request = QueryFlowsRequest::default(); - assert_eq!(request.page, 1); - assert_eq!(request.page_size, 20); - assert!(request.sort_desc); - } - - #[test] - fn test_search_flows_request_default_limit() { - let request = SearchFlowsRequest { - query: "test".to_string(), - limit: default_search_limit(), - }; - assert_eq!(request.limit, 50); - } - - #[test] - fn test_export_flows_request_serialization() { - let request = ExportFlowsRequest { - format: ExportFormat::JSON, - filter: None, - include_raw: true, - include_stream_chunks: false, - redact_sensitive: false, - flow_ids: None, - }; - - let json = serde_json::to_string(&request).unwrap(); - let deserialized: ExportFlowsRequest = serde_json::from_str(&json).unwrap(); - - assert_eq!(deserialized.format, ExportFormat::JSON); - assert!(deserialized.include_raw); - } -} - -// ============================================================================ -// 实时事件订阅命令 -// ============================================================================ - -use tauri::{AppHandle, Emitter}; - -/// 订阅 Flow 实时事件 -/// -/// 启动一个后台任务,将 Flow 事件通过 Tauri 事件系统推送到前端。 -/// 前端可以通过 `listen("flow-event", ...)` 来接收事件。 -/// -/// # Arguments -/// * `app` - Tauri AppHandle -/// * `monitor` - Flow 监控服务状态 -/// -/// # Returns -/// * `Ok(())` - 成功启动订阅 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn subscribe_flow_events( - app: AppHandle, - monitor: State<'_, FlowMonitorState>, -) -> Result<(), String> { - let mut receiver = monitor.0.subscribe(); - - // 启动后台任务来转发事件 - tokio::spawn(async move { - loop { - match receiver.recv().await { - Ok(event) => { - // 将事件发送到前端 - if let Err(e) = app.emit("flow-event", &event) { - tracing::warn!("发送 Flow 事件到前端失败: {}", e); - } - } - Err(tokio::sync::broadcast::error::RecvError::Lagged(n)) => { - tracing::warn!("Flow 事件接收器落后 {} 条消息", n); - } - Err(tokio::sync::broadcast::error::RecvError::Closed) => { - tracing::debug!("Flow 事件通道已关闭"); - break; - } - } - } - }); - - Ok(()) -} - -/// 获取所有可用的 Flow 标签 -/// -/// # Arguments -/// * `query_service` - 查询服务状态 -/// -/// # Returns -/// * `Ok(Vec)` - 成功时返回标签列表 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn get_all_flow_tags( - _query_service: State<'_, FlowQueryServiceState>, -) -> Result, String> { - // TODO: 实现从存储中获取所有标签 - // 目前返回空列表 - Ok(Vec::new()) -} - -// ============================================================================ -// 过滤表达式相关命令 -// ============================================================================ - -/// 过滤表达式解析结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ParseFilterResult { - /// 是否有效 - pub valid: bool, - /// 错误信息(如果无效) - pub error: Option, - /// 解析后的表达式(序列化为 JSON) - pub expr: Option, -} - -/// 过滤表达式帮助信息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct FilterHelpItem { - /// 语法 - pub syntax: String, - /// 描述 - pub description: String, -} - -/// 解析过滤表达式 -/// -/// **Validates: Requirements 1.1-1.17** -/// -/// 验证并解析过滤表达式字符串,返回解析结果。 -/// 如果表达式有效,返回解析后的 AST;如果无效,返回错误信息。 -/// -/// # Arguments -/// * `expression` - 过滤表达式字符串 -/// -/// # Returns -/// * `Ok(ParseFilterResult)` - 解析结果 -#[tauri::command] -pub async fn parse_filter(expression: String) -> Result { - match FilterParser::parse(&expression) { - Ok(expr) => Ok(ParseFilterResult { - valid: true, - error: None, - expr: Some(expr), - }), - Err(e) => Ok(ParseFilterResult { - valid: false, - error: Some(e.to_string()), - expr: None, - }), - } -} - -/// 验证过滤表达式 -/// -/// **Validates: Requirements 1.17** -/// -/// 仅验证过滤表达式语法是否正确,不返回解析后的 AST。 -/// -/// # Arguments -/// * `expression` - 过滤表达式字符串 -/// -/// # Returns -/// * `Ok(bool)` - 表达式是否有效 -/// * `Err(String)` - 验证过程中的错误 -#[tauri::command] -pub async fn validate_filter(expression: String) -> Result { - Ok(FilterParser::validate(&expression).is_ok()) -} - -/// 获取过滤表达式帮助信息 -/// -/// **Validates: Requirements 1.1-1.16** -/// -/// 返回所有支持的过滤表达式语法和描述。 -/// -/// # Returns -/// * `Ok(Vec)` - 帮助信息列表 -#[tauri::command] -pub async fn get_filter_help_items() -> Result, String> { - let items: Vec = FILTER_HELP - .iter() - .map(|(syntax, desc)| FilterHelpItem { - syntax: syntax.to_string(), - description: desc.to_string(), - }) - .collect(); - Ok(items) -} - -/// 获取过滤表达式帮助文本 -/// -/// **Validates: Requirements 1.1-1.16** -/// -/// 返回格式化的帮助文本,包含所有支持的过滤表达式语法和示例。 -/// -/// # Returns -/// * `Ok(String)` - 帮助文本 -#[tauri::command] -pub async fn get_filter_help_text() -> Result { - Ok(get_filter_help()) -} - -/// 使用过滤表达式查询 Flow 请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct QueryFlowsWithExpressionRequest { - /// 过滤表达式 - pub filter_expr: String, - /// 排序字段 - #[serde(default)] - pub sort_by: FlowSortBy, - /// 是否降序 - #[serde(default = "default_true")] - pub sort_desc: bool, - /// 页码(从 1 开始) - #[serde(default = "default_page")] - pub page: usize, - /// 每页大小 - #[serde(default = "default_page_size")] - pub page_size: usize, -} - -/// 使用过滤表达式查询 Flow -/// -/// **Validates: Requirements 1.1-1.16** -/// -/// 使用类似 mitmproxy 的过滤表达式语法查询 Flow。 -/// -/// # Arguments -/// * `request` - 查询请求参数 -/// * `query_service` - 查询服务状态 -/// -/// # Returns -/// * `Ok(FlowQueryResult)` - 成功时返回查询结果 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn query_flows_with_expression( - request: QueryFlowsWithExpressionRequest, - query_service: State<'_, FlowQueryServiceState>, -) -> Result { - query_service - .0 - .query_with_expression( - &request.filter_expr, - request.sort_by, - request.sort_desc, - request.page, - request.page_size, - ) - .await - .map_err(|e| format!("查询 Flow 失败: {e}")) -} - -// ============================================================================ -// 拦截器相关命令 -// ============================================================================ - -use crate::flow_monitor::{FlowInterceptor, InterceptConfig, InterceptedFlow, ModifiedData}; - -use crate::flow_monitor::{BatchReplayResult, FlowReplayer, ReplayConfig, ReplayResult}; - -/// 拦截器状态封装 -pub struct FlowInterceptorState(pub Arc); - -/// 重放器状态封装 -pub struct FlowReplayerState(pub Arc); - -/// 获取拦截器配置 -/// -/// **Validates: Requirements 2.7, 2.8** -/// -/// # Arguments -/// * `interceptor` - 拦截器状态 -/// -/// # Returns -/// * `Ok(InterceptConfig)` - 成功时返回拦截器配置 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn intercept_config_get( - interceptor: State<'_, FlowInterceptorState>, -) -> Result { - Ok(interceptor.0.config().await) -} - -/// 设置拦截器配置 -/// -/// **Validates: Requirements 2.7, 2.8** -/// -/// # Arguments -/// * `config` - 新的拦截器配置 -/// * `interceptor` - 拦截器状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn intercept_config_set( - config: InterceptConfig, - interceptor: State<'_, FlowInterceptorState>, -) -> Result<(), String> { - interceptor - .0 - .update_config(config) - .await - .map_err(|e| format!("设置拦截器配置失败: {e}")) -} - -/// 继续处理被拦截的 Flow -/// -/// **Validates: Requirements 2.3, 2.5** -/// -/// # Arguments -/// * `flow_id` - Flow ID -/// * `modified_request` - 修改后的请求(可选) -/// * `modified_response` - 修改后的响应(可选) -/// * `interceptor` - 拦截器状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn intercept_continue( - flow_id: String, - modified_request: Option, - modified_response: Option, - interceptor: State<'_, FlowInterceptorState>, -) -> Result<(), String> { - // 确定修改数据 - let modified = if let Some(req) = modified_request { - Some(ModifiedData::Request(req)) - } else { - modified_response.map(ModifiedData::Response) - }; - - interceptor - .0 - .continue_flow(&flow_id, modified) - .await - .map_err(|e| format!("继续处理 Flow 失败: {e}")) -} - -/// 取消被拦截的 Flow -/// -/// **Validates: Requirements 2.4** -/// -/// # Arguments -/// * `flow_id` - Flow ID -/// * `interceptor` - 拦截器状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn intercept_cancel( - flow_id: String, - interceptor: State<'_, FlowInterceptorState>, -) -> Result<(), String> { - interceptor - .0 - .cancel_flow(&flow_id) - .await - .map_err(|e| format!("取消 Flow 失败: {e}")) -} - -/// 获取被拦截的 Flow 详情 -/// -/// **Validates: Requirements 2.1** -/// -/// # Arguments -/// * `flow_id` - Flow ID -/// * `interceptor` - 拦截器状态 -/// -/// # Returns -/// * `Ok(Option)` - 成功时返回被拦截的 Flow 详情 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn intercept_get_flow( - flow_id: String, - interceptor: State<'_, FlowInterceptorState>, -) -> Result, String> { - Ok(interceptor.0.get_intercepted_flow(&flow_id).await) -} - -/// 获取所有被拦截的 Flow 列表 -/// -/// **Validates: Requirements 2.1** -/// -/// # Arguments -/// * `interceptor` - 拦截器状态 -/// -/// # Returns -/// * `Ok(Vec)` - 成功时返回被拦截的 Flow 列表 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn intercept_list_flows( - interceptor: State<'_, FlowInterceptorState>, -) -> Result, String> { - Ok(interceptor.0.list_intercepted_flows().await) -} - -/// 获取被拦截的 Flow 数量 -/// -/// **Validates: Requirements 2.1** -/// -/// # Arguments -/// * `interceptor` - 拦截器状态 -/// -/// # Returns -/// * `Ok(usize)` - 成功时返回被拦截的 Flow 数量 -#[tauri::command] -pub async fn intercept_count( - interceptor: State<'_, FlowInterceptorState>, -) -> Result { - Ok(interceptor.0.intercepted_count().await) -} - -/// 检查拦截是否启用 -/// -/// **Validates: Requirements 2.1** -/// -/// # Arguments -/// * `interceptor` - 拦截器状态 -/// -/// # Returns -/// * `Ok(bool)` - 成功时返回拦截是否启用 -#[tauri::command] -pub async fn intercept_is_enabled( - interceptor: State<'_, FlowInterceptorState>, -) -> Result { - Ok(interceptor.0.is_enabled().await) -} - -/// 启用拦截 -/// -/// **Validates: Requirements 2.1** -/// -/// # Arguments -/// * `interceptor` - 拦截器状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -#[tauri::command] -pub async fn intercept_enable(interceptor: State<'_, FlowInterceptorState>) -> Result<(), String> { - interceptor.0.enable().await; - Ok(()) -} - -/// 禁用拦截 -/// -/// **Validates: Requirements 2.1** -/// -/// # Arguments -/// * `interceptor` - 拦截器状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -#[tauri::command] -pub async fn intercept_disable(interceptor: State<'_, FlowInterceptorState>) -> Result<(), String> { - interceptor.0.disable().await; - Ok(()) -} - -/// 设置 Flow 为编辑状态 -/// -/// **Validates: Requirements 2.2** -/// -/// # Arguments -/// * `flow_id` - Flow ID -/// * `interceptor` - 拦截器状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn intercept_set_editing( - flow_id: String, - interceptor: State<'_, FlowInterceptorState>, -) -> Result<(), String> { - interceptor - .0 - .set_editing(&flow_id) - .await - .map_err(|e| format!("设置编辑状态失败: {e}")) -} - -/// 订阅拦截事件 -/// -/// **Validates: Requirements 2.1** -/// -/// 启动一个后台任务,将拦截事件通过 Tauri 事件系统推送到前端。 -/// 前端可以通过 `listen("intercept-event", ...)` 来接收事件。 -/// -/// # Arguments -/// * `app` - Tauri AppHandle -/// * `interceptor` - 拦截器状态 -/// -/// # Returns -/// * `Ok(())` - 成功启动订阅 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn subscribe_intercept_events( - app: AppHandle, - interceptor: State<'_, FlowInterceptorState>, -) -> Result<(), String> { - let mut receiver = interceptor.0.subscribe(); - - // 启动后台任务来转发事件 - tokio::spawn(async move { - loop { - match receiver.recv().await { - Ok(event) => { - // 将事件发送到前端 - if let Err(e) = app.emit("intercept-event", &event) { - tracing::warn!("发送拦截事件到前端失败: {}", e); - } - } - Err(tokio::sync::broadcast::error::RecvError::Lagged(n)) => { - tracing::warn!("拦截事件接收器落后 {} 条消息", n); - } - Err(tokio::sync::broadcast::error::RecvError::Closed) => { - tracing::debug!("拦截事件通道已关闭"); - break; - } - } - } - }); - - Ok(()) -} - -// ============================================================================ -// 重放器相关命令 -// ============================================================================ - -/// 重放 Flow 请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ReplayFlowRequest { - /// 要重放的 Flow ID - pub flow_id: String, - /// 重放配置 - #[serde(default)] - pub config: ReplayConfig, -} - -/// 批量重放 Flow 请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ReplayFlowsBatchRequest { - /// 要重放的 Flow ID 列表 - pub flow_ids: Vec, - /// 重放配置 - #[serde(default)] - pub config: ReplayConfig, -} - -/// 重放单个 Flow -/// -/// **Validates: Requirements 3.1, 3.3, 3.4** -/// -/// # Arguments -/// * `request` - 重放请求参数 -/// * `replayer` - 重放器状态 -/// -/// # Returns -/// * `Ok(ReplayResult)` - 成功时返回重放结果 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn replay_flow( - request: ReplayFlowRequest, - replayer: State<'_, FlowReplayerState>, -) -> Result { - replayer - .0 - .replay(&request.flow_id, request.config) - .await - .map_err(|e| format!("重放 Flow 失败: {e}")) -} - -/// 批量重放多个 Flow -/// -/// **Validates: Requirements 3.6, 3.7** -/// -/// # Arguments -/// * `request` - 批量重放请求参数 -/// * `replayer` - 重放器状态 -/// -/// # Returns -/// * `Ok(BatchReplayResult)` - 成功时返回批量重放结果 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn replay_flows_batch( - request: ReplayFlowsBatchRequest, - replayer: State<'_, FlowReplayerState>, -) -> Result { - Ok(replayer - .0 - .replay_batch(&request.flow_ids, request.config) - .await) -} - -// ============================================================================ -// 差异对比命令 -// ============================================================================ - -/// 差异对比请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct DiffFlowsRequest { - /// 左侧 Flow ID - pub left_flow_id: String, - /// 右侧 Flow ID - pub right_flow_id: String, - /// 差异配置 - #[serde(default)] - pub config: DiffConfig, -} - -/// 对比两个 Flow 的差异 -/// -/// **Validates: Requirements 4.1, 4.2** -/// -/// # Arguments -/// * `request` - 差异对比请求参数 -/// * `query_service` - 查询服务状态 -/// -/// # Returns -/// * `Ok(FlowDiffResult)` - 成功时返回差异结果 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn diff_flows( - request: DiffFlowsRequest, - query_service: State<'_, FlowQueryServiceState>, -) -> Result { - // 获取左侧 Flow - let left_flow = query_service - .0 - .get_flow(&request.left_flow_id) - .await - .map_err(|e| format!("获取左侧 Flow 失败: {e}"))? - .ok_or_else(|| format!("左侧 Flow 不存在: {}", request.left_flow_id))?; - - // 获取右侧 Flow - let right_flow = query_service - .0 - .get_flow(&request.right_flow_id) - .await - .map_err(|e| format!("获取右侧 Flow 失败: {e}"))? - .ok_or_else(|| format!("右侧 Flow 不存在: {}", request.right_flow_id))?; - - // 执行差异对比 - let result = FlowDiff::diff(&left_flow, &right_flow, &request.config); - - Ok(result) -} - -// ============================================================================ -// 重放器测试模块 -// ============================================================================ - -#[cfg(test)] -mod replayer_tests { - use super::*; - - #[test] - fn test_replay_flow_request_serialization() { - let request = ReplayFlowRequest { - flow_id: "test-flow-id".to_string(), - config: ReplayConfig::default(), - }; - - let json = serde_json::to_string(&request).unwrap(); - let deserialized: ReplayFlowRequest = serde_json::from_str(&json).unwrap(); - - assert_eq!(deserialized.flow_id, "test-flow-id"); - assert!(deserialized.config.credential_id.is_none()); - } - - #[test] - fn test_replay_flows_batch_request_serialization() { - let request = ReplayFlowsBatchRequest { - flow_ids: vec!["flow-1".to_string(), "flow-2".to_string()], - config: ReplayConfig { - credential_id: Some("cred-1".to_string()), - modify_request: None, - interval_ms: 500, - }, - }; - - let json = serde_json::to_string(&request).unwrap(); - let deserialized: ReplayFlowsBatchRequest = serde_json::from_str(&json).unwrap(); - - assert_eq!(deserialized.flow_ids.len(), 2); - assert_eq!(deserialized.config.interval_ms, 500); - assert_eq!( - deserialized.config.credential_id, - Some("cred-1".to_string()) - ); - } - - #[test] - fn test_replay_config_default() { - let config = ReplayConfig::default(); - assert!(config.credential_id.is_none()); - assert!(config.modify_request.is_none()); - assert_eq!(config.interval_ms, 1000); - } -} - -// ============================================================================ -// 差异对比测试模块 -// ============================================================================ - -#[cfg(test)] -mod diff_tests { - use super::*; - - #[test] - fn test_diff_flows_request_serialization() { - let request = DiffFlowsRequest { - left_flow_id: "flow-1".to_string(), - right_flow_id: "flow-2".to_string(), - config: DiffConfig::default(), - }; - - let json = serde_json::to_string(&request).unwrap(); - let deserialized: DiffFlowsRequest = serde_json::from_str(&json).unwrap(); - - assert_eq!(deserialized.left_flow_id, "flow-1"); - assert_eq!(deserialized.right_flow_id, "flow-2"); - assert!(deserialized.config.ignore_timestamps); - assert!(deserialized.config.ignore_ids); - } - - #[test] - fn test_diff_flows_request_with_custom_config() { - let request = DiffFlowsRequest { - left_flow_id: "flow-a".to_string(), - right_flow_id: "flow-b".to_string(), - config: DiffConfig { - ignore_fields: vec!["custom_field".to_string()], - ignore_timestamps: false, - ignore_ids: false, - }, - }; - - let json = serde_json::to_string(&request).unwrap(); - let deserialized: DiffFlowsRequest = serde_json::from_str(&json).unwrap(); - - assert_eq!(deserialized.config.ignore_fields.len(), 1); - assert!(!deserialized.config.ignore_timestamps); - assert!(!deserialized.config.ignore_ids); - } -} - -// ============================================================================ -// 会话管理命令 -// ============================================================================ - -use crate::flow_monitor::{AutoSessionConfig, FlowSession, SessionExportResult, SessionManager}; - -/// 会话管理器状态封装 -pub struct SessionManagerState(pub Arc); - -/// 创建会话请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CreateSessionRequest { - /// 会话名称 - pub name: String, - /// 会话描述(可选) - #[serde(default)] - pub description: Option, -} - -/// 更新会话请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct UpdateSessionRequest { - /// 会话 ID - pub session_id: String, - /// 新名称(可选) - #[serde(default)] - pub name: Option, - /// 新描述(可选,None 表示不更新,Some(None) 表示清除描述) - #[serde(default)] - pub description: Option>, -} - -/// 导出会话请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ExportSessionRequest { - /// 会话 ID - pub session_id: String, - /// 导出格式 - #[serde(default)] - pub format: ExportFormat, -} - -/// 创建新会话 -/// -/// **Validates: Requirements 5.1** -/// -/// # Arguments -/// * `request` - 创建会话请求参数 -/// * `session_manager` - 会话管理器状态 -/// -/// # Returns -/// * `Ok(FlowSession)` - 成功时返回新创建的会话 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn create_session( - request: CreateSessionRequest, - session_manager: State<'_, SessionManagerState>, -) -> Result { - session_manager - .0 - .create_session(&request.name, request.description.as_deref()) - .map_err(|e| format!("创建会话失败: {e}")) -} - -/// 获取会话详情 -/// -/// **Validates: Requirements 5.3** -/// -/// # Arguments -/// * `session_id` - 会话 ID -/// * `session_manager` - 会话管理器状态 -/// -/// # Returns -/// * `Ok(Option)` - 成功时返回会话详情 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn get_session( - session_id: String, - session_manager: State<'_, SessionManagerState>, -) -> Result, String> { - session_manager - .0 - .get_session(&session_id) - .map_err(|e| format!("获取会话失败: {e}")) -} - -/// 列出所有会话 -/// -/// **Validates: Requirements 5.3** -/// -/// # Arguments -/// * `include_archived` - 是否包含已归档的会话 -/// * `session_manager` - 会话管理器状态 -/// -/// # Returns -/// * `Ok(Vec)` - 成功时返回会话列表 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn list_sessions( - include_archived: Option, - session_manager: State<'_, SessionManagerState>, -) -> Result, String> { - session_manager - .0 - .list_sessions(include_archived.unwrap_or(false)) - .map_err(|e| format!("列出会话失败: {e}")) -} - -/// 添加 Flow 到会话 -/// -/// **Validates: Requirements 5.2** -/// -/// # Arguments -/// * `session_id` - 会话 ID -/// * `flow_id` - Flow ID -/// * `session_manager` - 会话管理器状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn add_flow_to_session( - session_id: String, - flow_id: String, - session_manager: State<'_, SessionManagerState>, -) -> Result<(), String> { - session_manager - .0 - .add_flow(&session_id, &flow_id) - .map_err(|e| format!("添加 Flow 到会话失败: {e}")) -} - -/// 从会话移除 Flow -/// -/// **Validates: Requirements 5.2** -/// -/// # Arguments -/// * `session_id` - 会话 ID -/// * `flow_id` - Flow ID -/// * `session_manager` - 会话管理器状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn remove_flow_from_session( - session_id: String, - flow_id: String, - session_manager: State<'_, SessionManagerState>, -) -> Result<(), String> { - session_manager - .0 - .remove_flow(&session_id, &flow_id) - .map_err(|e| format!("从会话移除 Flow 失败: {e}")) -} - -/// 更新会话信息 -/// -/// **Validates: Requirements 5.5** -/// -/// # Arguments -/// * `request` - 更新会话请求参数 -/// * `session_manager` - 会话管理器状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn update_session( - request: UpdateSessionRequest, - session_manager: State<'_, SessionManagerState>, -) -> Result<(), String> { - session_manager - .0 - .update_session( - &request.session_id, - request.name.as_deref(), - request.description.as_ref().map(|d| d.as_deref()), - ) - .map_err(|e| format!("更新会话失败: {e}")) -} - -/// 归档会话 -/// -/// **Validates: Requirements 5.7** -/// -/// # Arguments -/// * `session_id` - 会话 ID -/// * `session_manager` - 会话管理器状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn archive_session( - session_id: String, - session_manager: State<'_, SessionManagerState>, -) -> Result<(), String> { - session_manager - .0 - .archive_session(&session_id) - .map_err(|e| format!("归档会话失败: {e}")) -} - -/// 取消归档会话 -/// -/// # Arguments -/// * `session_id` - 会话 ID -/// * `session_manager` - 会话管理器状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn unarchive_session( - session_id: String, - session_manager: State<'_, SessionManagerState>, -) -> Result<(), String> { - session_manager - .0 - .unarchive_session(&session_id) - .map_err(|e| format!("取消归档会话失败: {e}")) -} - -/// 删除会话 -/// -/// **Validates: Requirements 5.7** -/// -/// # Arguments -/// * `session_id` - 会话 ID -/// * `session_manager` - 会话管理器状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn delete_session( - session_id: String, - session_manager: State<'_, SessionManagerState>, -) -> Result<(), String> { - session_manager - .0 - .delete_session(&session_id) - .map_err(|e| format!("删除会话失败: {e}")) -} - -/// 导出会话 -/// -/// **Validates: Requirements 5.6** -/// -/// # Arguments -/// * `request` - 导出会话请求参数 -/// * `session_manager` - 会话管理器状态 -/// * `query_service` - 查询服务状态 -/// -/// # Returns -/// * `Ok(SessionExportResult)` - 成功时返回导出结果 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn export_session( - request: ExportSessionRequest, - session_manager: State<'_, SessionManagerState>, - query_service: State<'_, FlowQueryServiceState>, -) -> Result { - // 获取会话中的 Flow ID - let flow_ids = session_manager - .0 - .get_session_flow_ids(&request.session_id) - .map_err(|e| format!("获取会话 Flow 列表失败: {e}"))?; - - // 获取所有 Flow - let mut flows = Vec::new(); - for flow_id in &flow_ids { - if let Ok(Some(flow)) = query_service.0.get_flow(flow_id).await { - flows.push(flow); - } - } - - // 导出会话 - session_manager - .0 - .export_session(&request.session_id, &flows, request.format) - .map_err(|e| format!("导出会话失败: {e}")) -} - -/// 获取会话中的 Flow 数量 -/// -/// # Arguments -/// * `session_id` - 会话 ID -/// * `session_manager` - 会话管理器状态 -/// -/// # Returns -/// * `Ok(usize)` - 成功时返回 Flow 数量 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn get_session_flow_count( - session_id: String, - session_manager: State<'_, SessionManagerState>, -) -> Result { - session_manager - .0 - .get_session_flow_count(&session_id) - .map_err(|e| format!("获取会话 Flow 数量失败: {e}")) -} - -/// 检查 Flow 是否在会话中 -/// -/// # Arguments -/// * `session_id` - 会话 ID -/// * `flow_id` - Flow ID -/// * `session_manager` - 会话管理器状态 -/// -/// # Returns -/// * `Ok(bool)` - 成功时返回是否在会话中 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn is_flow_in_session( - session_id: String, - flow_id: String, - session_manager: State<'_, SessionManagerState>, -) -> Result { - session_manager - .0 - .is_flow_in_session(&session_id, &flow_id) - .map_err(|e| format!("检查 Flow 是否在会话中失败: {e}")) -} - -/// 获取 Flow 所属的会话列表 -/// -/// # Arguments -/// * `flow_id` - Flow ID -/// * `session_manager` - 会话管理器状态 -/// -/// # Returns -/// * `Ok(Vec)` - 成功时返回会话 ID 列表 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn get_sessions_for_flow( - flow_id: String, - session_manager: State<'_, SessionManagerState>, -) -> Result, String> { - session_manager - .0 - .get_sessions_for_flow(&flow_id) - .map_err(|e| format!("获取 Flow 所属会话失败: {e}")) -} - -/// 获取自动会话检测配置 -/// -/// **Validates: Requirements 5.4** -/// -/// # Arguments -/// * `session_manager` - 会话管理器状态 -/// -/// # Returns -/// * `Ok(AutoSessionConfig)` - 成功时返回配置 -#[tauri::command] -pub async fn get_auto_session_config( - session_manager: State<'_, SessionManagerState>, -) -> Result { - Ok(session_manager.0.get_auto_config()) -} - -/// 设置自动会话检测配置 -/// -/// **Validates: Requirements 5.4** -/// -/// # Arguments -/// * `config` - 新配置 -/// * `session_manager` - 会话管理器状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -#[tauri::command] -pub async fn set_auto_session_config( - config: AutoSessionConfig, - session_manager: State<'_, SessionManagerState>, -) -> Result<(), String> { - session_manager.0.set_auto_config(config); - Ok(()) -} - -/// 注册活跃会话(用于自动检测) -/// -/// # Arguments -/// * `session_id` - 会话 ID -/// * `client_key` - 客户端标识(可选) -/// * `session_manager` - 会话管理器状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -#[tauri::command] -pub async fn register_active_session( - session_id: String, - client_key: Option, - session_manager: State<'_, SessionManagerState>, -) -> Result<(), String> { - session_manager - .0 - .register_active_session(&session_id, client_key.as_deref()); - Ok(()) -} - -// ============================================================================ -// 快速过滤器命令 -// ============================================================================ - -use crate::flow_monitor::{QuickFilter, QuickFilterManager, QuickFilterUpdate}; - -/// 快速过滤器管理器状态封装 -pub struct QuickFilterManagerState(pub Arc); - -/// 保存快速过滤器请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct SaveQuickFilterRequest { - /// 过滤器名称 - pub name: String, - /// 过滤表达式 - pub filter_expr: String, - /// 描述(可选) - #[serde(default)] - pub description: Option, - /// 分组(可选) - #[serde(default)] - pub group: Option, -} - -/// 更新快速过滤器请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct UpdateQuickFilterRequest { - /// 过滤器 ID - pub id: String, - /// 新名称(可选) - #[serde(default)] - pub name: Option, - /// 新描述(可选) - #[serde(default)] - pub description: Option>, - /// 新过滤表达式(可选) - #[serde(default)] - pub filter_expr: Option, - /// 新分组(可选) - #[serde(default)] - pub group: Option>, - /// 新排序顺序(可选) - #[serde(default)] - pub order: Option, -} - -/// 导入快速过滤器请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ImportQuickFiltersRequest { - /// JSON 格式的导入数据 - pub data: String, - /// 是否覆盖同名过滤器 - #[serde(default)] - pub overwrite: bool, -} - -/// 保存快速过滤器 -/// -/// **Validates: Requirements 6.1** -/// -/// # Arguments -/// * `request` - 保存快速过滤器请求参数 -/// * `quick_filter_manager` - 快速过滤器管理器状态 -/// -/// # Returns -/// * `Ok(QuickFilter)` - 成功时返回新创建的快速过滤器 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn save_quick_filter( - request: SaveQuickFilterRequest, - quick_filter_manager: State<'_, QuickFilterManagerState>, -) -> Result { - quick_filter_manager - .0 - .save( - &request.name, - &request.filter_expr, - request.description.as_deref(), - request.group.as_deref(), - ) - .map_err(|e| format!("保存快速过滤器失败: {e}")) -} - -/// 获取快速过滤器 -/// -/// # Arguments -/// * `id` - 过滤器 ID -/// * `quick_filter_manager` - 快速过滤器管理器状态 -/// -/// # Returns -/// * `Ok(Option)` - 成功时返回快速过滤器 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn get_quick_filter( - id: String, - quick_filter_manager: State<'_, QuickFilterManagerState>, -) -> Result, String> { - quick_filter_manager - .0 - .get(&id) - .map_err(|e| format!("获取快速过滤器失败: {e}")) -} - -/// 更新快速过滤器 -/// -/// **Validates: Requirements 6.4** -/// -/// # Arguments -/// * `request` - 更新快速过滤器请求参数 -/// * `quick_filter_manager` - 快速过滤器管理器状态 -/// -/// # Returns -/// * `Ok(QuickFilter)` - 成功时返回更新后的快速过滤器 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn update_quick_filter( - request: UpdateQuickFilterRequest, - quick_filter_manager: State<'_, QuickFilterManagerState>, -) -> Result { - let updates = QuickFilterUpdate { - name: request.name, - description: request.description, - filter_expr: request.filter_expr, - group: request.group, - order: request.order, - }; - - quick_filter_manager - .0 - .update(&request.id, updates) - .map_err(|e| format!("更新快速过滤器失败: {e}")) -} - -/// 删除快速过滤器 -/// -/// **Validates: Requirements 6.4** -/// -/// # Arguments -/// * `id` - 过滤器 ID -/// * `quick_filter_manager` - 快速过滤器管理器状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn delete_quick_filter( - id: String, - quick_filter_manager: State<'_, QuickFilterManagerState>, -) -> Result<(), String> { - quick_filter_manager - .0 - .delete(&id) - .map_err(|e| format!("删除快速过滤器失败: {e}")) -} - -/// 列出所有快速过滤器 -/// -/// **Validates: Requirements 6.2, 6.5** -/// -/// # Arguments -/// * `quick_filter_manager` - 快速过滤器管理器状态 -/// -/// # Returns -/// * `Ok(Vec)` - 成功时返回快速过滤器列表 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn list_quick_filters( - quick_filter_manager: State<'_, QuickFilterManagerState>, -) -> Result, String> { - quick_filter_manager - .0 - .list() - .map_err(|e| format!("列出快速过滤器失败: {e}")) -} - -/// 按分组列出快速过滤器 -/// -/// **Validates: Requirements 6.5** -/// -/// # Arguments -/// * `group` - 分组名称(可选,None 表示无分组的过滤器) -/// * `quick_filter_manager` - 快速过滤器管理器状态 -/// -/// # Returns -/// * `Ok(Vec)` - 成功时返回快速过滤器列表 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn list_quick_filters_by_group( - group: Option, - quick_filter_manager: State<'_, QuickFilterManagerState>, -) -> Result, String> { - quick_filter_manager - .0 - .list_by_group(group.as_deref()) - .map_err(|e| format!("按分组列出快速过滤器失败: {e}")) -} - -/// 列出所有分组 -/// -/// **Validates: Requirements 6.5** -/// -/// # Arguments -/// * `quick_filter_manager` - 快速过滤器管理器状态 -/// -/// # Returns -/// * `Ok(Vec)` - 成功时返回分组名称列表 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn list_quick_filter_groups( - quick_filter_manager: State<'_, QuickFilterManagerState>, -) -> Result, String> { - quick_filter_manager - .0 - .list_groups() - .map_err(|e| format!("列出快速过滤器分组失败: {e}")) -} - -/// 导出快速过滤器 -/// -/// **Validates: Requirements 6.7** -/// -/// # Arguments -/// * `include_presets` - 是否包含预设过滤器 -/// * `quick_filter_manager` - 快速过滤器管理器状态 -/// -/// # Returns -/// * `Ok(String)` - 成功时返回 JSON 格式的导出数据 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn export_quick_filters( - include_presets: Option, - quick_filter_manager: State<'_, QuickFilterManagerState>, -) -> Result { - quick_filter_manager - .0 - .export(include_presets.unwrap_or(false)) - .map_err(|e| format!("导出快速过滤器失败: {e}")) -} - -/// 导入快速过滤器 -/// -/// **Validates: Requirements 6.7** -/// -/// # Arguments -/// * `request` - 导入快速过滤器请求参数 -/// * `quick_filter_manager` - 快速过滤器管理器状态 -/// -/// # Returns -/// * `Ok(Vec)` - 成功时返回导入的快速过滤器列表 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn import_quick_filters( - request: ImportQuickFiltersRequest, - quick_filter_manager: State<'_, QuickFilterManagerState>, -) -> Result, String> { - quick_filter_manager - .0 - .import(&request.data, request.overwrite) - .map_err(|e| format!("导入快速过滤器失败: {e}")) -} - -/// 按名称查找快速过滤器 -/// -/// # Arguments -/// * `name` - 过滤器名称 -/// * `quick_filter_manager` - 快速过滤器管理器状态 -/// -/// # Returns -/// * `Ok(Option)` - 成功时返回快速过滤器 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn find_quick_filter_by_name( - name: String, - quick_filter_manager: State<'_, QuickFilterManagerState>, -) -> Result, String> { - quick_filter_manager - .0 - .find_by_name(&name) - .map_err(|e| format!("查找快速过滤器失败: {e}")) -} - -// ============================================================================ -// 代码导出命令 -// ============================================================================ - -use crate::flow_monitor::{CodeExporter, CodeFormat}; - -/// 代码导出请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ExportFlowAsCodeRequest { - /// Flow ID - pub flow_id: String, - /// 导出格式 - pub format: CodeFormat, -} - -/// 代码导出响应 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ExportFlowAsCodeResponse { - /// 导出的代码 - pub code: String, - /// 导出格式 - pub format: CodeFormat, -} - -/// 将 Flow 导出为代码 -/// -/// **Validates: Requirements 7.7, 7.8** -/// -/// 将指定的 Flow 导出为可执行的代码格式(curl、Python、TypeScript、JavaScript)。 -/// -/// # Arguments -/// * `request` - 代码导出请求参数 -/// * `query_service` - 查询服务状态 -/// -/// # Returns -/// * `Ok(ExportFlowAsCodeResponse)` - 成功时返回导出的代码 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn export_flow_as_code( - request: ExportFlowAsCodeRequest, - query_service: State<'_, FlowQueryServiceState>, -) -> Result { - // 获取 Flow - let flow = query_service - .0 - .get_flow(&request.flow_id) - .await - .map_err(|e| format!("获取 Flow 失败: {e}"))? - .ok_or_else(|| format!("Flow 不存在: {}", request.flow_id))?; - - // 导出为代码 - let code = CodeExporter::export(&flow, request.format); - - Ok(ExportFlowAsCodeResponse { - code, - format: request.format, - }) -} - -/// 批量导出 Flow 为代码 -/// -/// **Validates: Requirements 7.7, 7.8** -/// -/// 将多个 Flow 导出为可执行的代码格式。 -/// -/// # Arguments -/// * `flow_ids` - Flow ID 列表 -/// * `format` - 导出格式 -/// * `query_service` - 查询服务状态 -/// -/// # Returns -/// * `Ok(Vec)` - 成功时返回导出的代码列表 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn export_flows_as_code( - flow_ids: Vec, - format: CodeFormat, - query_service: State<'_, FlowQueryServiceState>, -) -> Result, String> { - let mut results = Vec::new(); - - for flow_id in flow_ids { - if let Ok(Some(flow)) = query_service.0.get_flow(&flow_id).await { - let code = CodeExporter::export(&flow, format); - results.push(ExportFlowAsCodeResponse { code, format }); - } - } - - Ok(results) -} - -/// 获取支持的代码导出格式 -/// -/// **Validates: Requirements 7.7, 7.8** -/// -/// 返回所有支持的代码导出格式列表。 -/// -/// # Returns -/// * `Ok(Vec)` - 成功时返回格式列表 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CodeFormatInfo { - /// 格式标识 - pub format: CodeFormat, - /// 格式名称 - pub name: String, - /// 格式描述 - pub description: String, -} - -#[tauri::command] -pub async fn get_code_export_formats() -> Result, String> { - Ok(vec![ - CodeFormatInfo { - format: CodeFormat::Curl, - name: "curl".to_string(), - description: "curl 命令行工具".to_string(), - }, - CodeFormatInfo { - format: CodeFormat::Python, - name: "Python".to_string(), - description: "Python requests 库".to_string(), - }, - CodeFormatInfo { - format: CodeFormat::TypeScript, - name: "TypeScript".to_string(), - description: "TypeScript fetch API".to_string(), - }, - CodeFormatInfo { - format: CodeFormat::JavaScript, - name: "JavaScript".to_string(), - description: "JavaScript fetch API".to_string(), - }, - ]) -} - -// ============================================================================ -// 书签管理命令 -// ============================================================================ - -use crate::flow_monitor::{BookmarkManager, FlowBookmark}; - -/// 书签管理器状态封装 -pub struct BookmarkManagerState(pub Arc); - -/// 添加书签请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct AddBookmarkRequest { - /// Flow ID - pub flow_id: String, - /// 书签名称(可选) - #[serde(default)] - pub name: Option, - /// 分组名称(可选) - #[serde(default)] - pub group: Option, -} - -/// 更新书签请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct UpdateBookmarkRequest { - /// 书签 ID - pub bookmark_id: String, - /// 新名称(可选) - #[serde(default)] - pub name: Option>, - /// 新分组(可选) - #[serde(default)] - pub group: Option>, -} - -/// 导入书签请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ImportBookmarksRequest { - /// JSON 格式的导入数据 - pub data: String, - /// 是否覆盖已存在的书签 - #[serde(default)] - pub overwrite: bool, -} - -/// 添加书签 -/// -/// **Validates: Requirements 8.1** -/// -/// # Arguments -/// * `request` - 添加书签请求参数 -/// * `bookmark_manager` - 书签管理器状态 -/// -/// # Returns -/// * `Ok(FlowBookmark)` - 成功时返回新创建的书签 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn add_bookmark( - request: AddBookmarkRequest, - bookmark_manager: State<'_, BookmarkManagerState>, -) -> Result { - bookmark_manager - .0 - .add( - &request.flow_id, - request.name.as_deref(), - request.group.as_deref(), - ) - .map_err(|e| format!("添加书签失败: {e}")) -} - -/// 获取书签 -/// -/// # Arguments -/// * `bookmark_id` - 书签 ID -/// * `bookmark_manager` - 书签管理器状态 -/// -/// # Returns -/// * `Ok(Option)` - 成功时返回书签 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn get_bookmark( - bookmark_id: String, - bookmark_manager: State<'_, BookmarkManagerState>, -) -> Result, String> { - bookmark_manager - .0 - .get(&bookmark_id) - .map_err(|e| format!("获取书签失败: {e}")) -} - -/// 根据 Flow ID 获取书签 -/// -/// # Arguments -/// * `flow_id` - Flow ID -/// * `bookmark_manager` - 书签管理器状态 -/// -/// # Returns -/// * `Ok(Option)` - 成功时返回书签 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn get_bookmark_by_flow_id( - flow_id: String, - bookmark_manager: State<'_, BookmarkManagerState>, -) -> Result, String> { - bookmark_manager - .0 - .get_by_flow_id(&flow_id) - .map_err(|e| format!("获取书签失败: {e}")) -} - -/// 移除书签 -/// -/// **Validates: Requirements 8.1** -/// -/// # Arguments -/// * `bookmark_id` - 书签 ID -/// * `bookmark_manager` - 书签管理器状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn remove_bookmark( - bookmark_id: String, - bookmark_manager: State<'_, BookmarkManagerState>, -) -> Result<(), String> { - bookmark_manager - .0 - .remove(&bookmark_id) - .map_err(|e| format!("移除书签失败: {e}")) -} - -/// 根据 Flow ID 移除书签 -/// -/// # Arguments -/// * `flow_id` - Flow ID -/// * `bookmark_manager` - 书签管理器状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn remove_bookmark_by_flow_id( - flow_id: String, - bookmark_manager: State<'_, BookmarkManagerState>, -) -> Result<(), String> { - bookmark_manager - .0 - .remove_by_flow_id(&flow_id) - .map_err(|e| format!("移除书签失败: {e}")) -} - -/// 更新书签 -/// -/// # Arguments -/// * `request` - 更新书签请求参数 -/// * `bookmark_manager` - 书签管理器状态 -/// -/// # Returns -/// * `Ok(FlowBookmark)` - 成功时返回更新后的书签 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn update_bookmark( - request: UpdateBookmarkRequest, - bookmark_manager: State<'_, BookmarkManagerState>, -) -> Result { - bookmark_manager - .0 - .update( - &request.bookmark_id, - request.name.as_ref().map(|n| n.as_deref()), - request.group.as_ref().map(|g| g.as_deref()), - ) - .map_err(|e| format!("更新书签失败: {e}")) -} - -/// 列出所有书签 -/// -/// **Validates: Requirements 8.3** -/// -/// # Arguments -/// * `group` - 分组名称(可选,None 表示所有书签) -/// * `bookmark_manager` - 书签管理器状态 -/// -/// # Returns -/// * `Ok(Vec)` - 成功时返回书签列表 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn list_bookmarks( - group: Option, - bookmark_manager: State<'_, BookmarkManagerState>, -) -> Result, String> { - bookmark_manager - .0 - .list(group.as_deref()) - .map_err(|e| format!("列出书签失败: {e}")) -} - -/// 列出所有书签分组 -/// -/// **Validates: Requirements 8.3** -/// -/// # Arguments -/// * `bookmark_manager` - 书签管理器状态 -/// -/// # Returns -/// * `Ok(Vec)` - 成功时返回分组名称列表 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn list_bookmark_groups( - bookmark_manager: State<'_, BookmarkManagerState>, -) -> Result, String> { - bookmark_manager - .0 - .list_groups() - .map_err(|e| format!("列出书签分组失败: {e}")) -} - -/// 检查 Flow 是否已添加书签 -/// -/// # Arguments -/// * `flow_id` - Flow ID -/// * `bookmark_manager` - 书签管理器状态 -/// -/// # Returns -/// * `Ok(bool)` - 成功时返回是否已添加书签 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn is_flow_bookmarked( - flow_id: String, - bookmark_manager: State<'_, BookmarkManagerState>, -) -> Result { - bookmark_manager - .0 - .is_bookmarked(&flow_id) - .map_err(|e| format!("检查书签状态失败: {e}")) -} - -/// 获取书签数量 -/// -/// # Arguments -/// * `bookmark_manager` - 书签管理器状态 -/// -/// # Returns -/// * `Ok(usize)` - 成功时返回书签数量 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn get_bookmark_count( - bookmark_manager: State<'_, BookmarkManagerState>, -) -> Result { - bookmark_manager - .0 - .count() - .map_err(|e| format!("获取书签数量失败: {e}")) -} - -/// 导出书签 -/// -/// **Validates: Requirements 8.6** -/// -/// # Arguments -/// * `bookmark_manager` - 书签管理器状态 -/// -/// # Returns -/// * `Ok(String)` - 成功时返回 JSON 格式的导出数据 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn export_bookmarks( - bookmark_manager: State<'_, BookmarkManagerState>, -) -> Result { - bookmark_manager - .0 - .export() - .map_err(|e| format!("导出书签失败: {e}")) -} - -/// 导入书签 -/// -/// **Validates: Requirements 8.6** -/// -/// # Arguments -/// * `request` - 导入书签请求参数 -/// * `bookmark_manager` - 书签管理器状态 -/// -/// # Returns -/// * `Ok(Vec)` - 成功时返回导入的书签列表 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn import_bookmarks( - request: ImportBookmarksRequest, - bookmark_manager: State<'_, BookmarkManagerState>, -) -> Result, String> { - bookmark_manager - .0 - .import(&request.data, request.overwrite) - .map_err(|e| format!("导入书签失败: {e}")) -} - -/// 切换书签状态 -/// -/// 如果 Flow 已添加书签则移除,否则添加书签。 -/// -/// **Validates: Requirements 8.1** -/// -/// # Arguments -/// * `flow_id` - Flow ID -/// * `name` - 书签名称(可选,仅在添加时使用) -/// * `group` - 分组名称(可选,仅在添加时使用) -/// * `bookmark_manager` - 书签管理器状态 -/// -/// # Returns -/// * `Ok(Option)` - 成功时返回书签(如果添加)或 None(如果移除) -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn toggle_bookmark( - flow_id: String, - name: Option, - group: Option, - bookmark_manager: State<'_, BookmarkManagerState>, -) -> Result, String> { - let is_bookmarked = bookmark_manager - .0 - .is_bookmarked(&flow_id) - .map_err(|e| format!("检查书签状态失败: {e}"))?; - - if is_bookmarked { - bookmark_manager - .0 - .remove_by_flow_id(&flow_id) - .map_err(|e| format!("移除书签失败: {e}"))?; - Ok(None) - } else { - let bookmark = bookmark_manager - .0 - .add(&flow_id, name.as_deref(), group.as_deref()) - .map_err(|e| format!("添加书签失败: {e}"))?; - Ok(Some(bookmark)) - } -} - -// ============================================================================ -// 增强统计相关命令 -// ============================================================================ - -use crate::flow_monitor::{ - Distribution, EnhancedStats, EnhancedStatsService, ReportFormat, StatsTimeRange, TrendData, -}; - -/// 增强统计服务状态封装 -pub struct EnhancedStatsServiceState(pub Arc); - -/// 获取增强统计请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct GetEnhancedStatsRequest { - /// 过滤条件 - #[serde(default)] - pub filter: FlowFilter, - /// 时间范围 - #[serde(default)] - pub time_range: StatsTimeRange, -} - -/// 获取请求趋势请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct GetRequestTrendRequest { - /// 过滤条件 - #[serde(default)] - pub filter: FlowFilter, - /// 时间范围 - #[serde(default)] - pub time_range: StatsTimeRange, - /// 时间间隔(如 "1h", "30m", "1d") - #[serde(default = "default_interval")] - pub interval: String, -} - -fn default_interval() -> String { - "1h".to_string() -} - -/// 获取延迟直方图请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct GetLatencyHistogramRequest { - /// 过滤条件 - #[serde(default)] - pub filter: FlowFilter, - /// 时间范围 - #[serde(default)] - pub time_range: StatsTimeRange, - /// 直方图桶边界(毫秒) - #[serde(default = "default_latency_buckets")] - pub buckets: Vec, -} - -fn default_latency_buckets() -> Vec { - vec![100, 500, 1000, 2000, 5000, 10000] -} - -/// 导出统计报告请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ExportStatsReportRequest { - /// 过滤条件 - #[serde(default)] - pub filter: FlowFilter, - /// 时间范围 - #[serde(default)] - pub time_range: StatsTimeRange, - /// 报告格式 - #[serde(default)] - pub format: ReportFormat, -} - -/// 获取增强统计 -/// -/// **Validates: Requirements 9.1-9.5** -/// -/// # Arguments -/// * `request` - 获取增强统计请求参数 -/// * `stats_service` - 增强统计服务状态 -/// -/// # Returns -/// * `Ok(EnhancedStats)` - 成功时返回增强统计结果 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn get_enhanced_stats( - request: GetEnhancedStatsRequest, - stats_service: State<'_, EnhancedStatsServiceState>, -) -> Result { - Ok(stats_service - .0 - .get_stats(&request.filter, &request.time_range) - .await) -} - -/// 获取请求趋势 -/// -/// **Validates: Requirements 9.1** -/// -/// # Arguments -/// * `request` - 获取请求趋势请求参数 -/// * `stats_service` - 增强统计服务状态 -/// -/// # Returns -/// * `Ok(TrendData)` - 成功时返回趋势数据 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn get_request_trend( - request: GetRequestTrendRequest, - stats_service: State<'_, EnhancedStatsServiceState>, -) -> Result { - Ok(stats_service - .0 - .get_request_trend(&request.filter, &request.time_range, &request.interval) - .await) -} - -/// 获取 Token 分布 -/// -/// **Validates: Requirements 9.2** -/// -/// # Arguments -/// * `request` - 获取增强统计请求参数(复用) -/// * `stats_service` - 增强统计服务状态 -/// -/// # Returns -/// * `Ok(Distribution)` - 成功时返回 Token 分布数据 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn get_token_distribution( - request: GetEnhancedStatsRequest, - stats_service: State<'_, EnhancedStatsServiceState>, -) -> Result { - Ok(stats_service - .0 - .get_token_distribution(&request.filter, &request.time_range) - .await) -} - -/// 获取延迟直方图 -/// -/// **Validates: Requirements 9.4** -/// -/// # Arguments -/// * `request` - 获取延迟直方图请求参数 -/// * `stats_service` - 增强统计服务状态 -/// -/// # Returns -/// * `Ok(Distribution)` - 成功时返回延迟直方图数据 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn get_latency_histogram( - request: GetLatencyHistogramRequest, - stats_service: State<'_, EnhancedStatsServiceState>, -) -> Result { - Ok(stats_service - .0 - .get_latency_histogram(&request.filter, &request.time_range, &request.buckets) - .await) -} - -/// 导出统计报告 -/// -/// **Validates: Requirements 9.7** -/// -/// # Arguments -/// * `request` - 导出统计报告请求参数 -/// * `stats_service` - 增强统计服务状态 -/// -/// # Returns -/// * `Ok(String)` - 成功时返回格式化的报告字符串 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn export_stats_report( - request: ExportStatsReportRequest, - stats_service: State<'_, EnhancedStatsServiceState>, -) -> Result { - Ok(stats_service - .0 - .export_report(&request.filter, &request.time_range, &request.format) - .await) -} -// ============================================================================ -// 批量操作状态封装 -// ============================================================================ - -/// BatchOperations 状态封装 -pub struct BatchOperationsState(pub Arc); - -// ============================================================================ -// 批量操作请求/响应类型 -// ============================================================================ - -/// 批量收藏 Flow 请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct BatchStarFlowsRequest { - /// Flow ID 列表 - pub flow_ids: Vec, -} - -/// 批量取消收藏 Flow 请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct BatchUnstarFlowsRequest { - /// Flow ID 列表 - pub flow_ids: Vec, -} - -/// 批量添加标签请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct BatchAddTagsRequest { - /// Flow ID 列表 - pub flow_ids: Vec, - /// 要添加的标签列表 - pub tags: Vec, -} - -/// 批量移除标签请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct BatchRemoveTagsRequest { - /// Flow ID 列表 - pub flow_ids: Vec, - /// 要移除的标签列表 - pub tags: Vec, -} - -/// 批量导出 Flow 请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct BatchExportFlowsRequest { - /// Flow ID 列表 - pub flow_ids: Vec, - /// 导出格式 - pub format: ExportFormat, -} - -/// 批量删除 Flow 请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct BatchDeleteFlowsRequest { - /// Flow ID 列表 - pub flow_ids: Vec, -} - -/// 批量添加到会话请求参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct BatchAddToSessionRequest { - /// Flow ID 列表 - pub flow_ids: Vec, - /// 会话 ID - pub session_id: String, -} - -// ============================================================================ -// 批量操作 Tauri 命令 -// ============================================================================ - -/// 批量收藏 Flow -/// -/// **Validates: Requirements 11.2** -/// -/// # Arguments -/// * `request` - 批量收藏请求参数 -/// * `batch_ops` - 批量操作服务状态 -/// -/// # Returns -/// * `Ok(BatchResult)` - 成功时返回批量操作结果 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn batch_star_flows( - request: BatchStarFlowsRequest, - batch_ops: State<'_, BatchOperationsState>, -) -> Result { - Ok(batch_ops - .0 - .execute(&request.flow_ids, BatchOperation::Star) - .await) -} - -/// 批量取消收藏 Flow -/// -/// **Validates: Requirements 11.2** -/// -/// # Arguments -/// * `request` - 批量取消收藏请求参数 -/// * `batch_ops` - 批量操作服务状态 -/// -/// # Returns -/// * `Ok(BatchResult)` - 成功时返回批量操作结果 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn batch_unstar_flows( - request: BatchUnstarFlowsRequest, - batch_ops: State<'_, BatchOperationsState>, -) -> Result { - Ok(batch_ops - .0 - .execute(&request.flow_ids, BatchOperation::Unstar) - .await) -} - -/// 批量添加标签 -/// -/// **Validates: Requirements 11.3** -/// -/// # Arguments -/// * `request` - 批量添加标签请求参数 -/// * `batch_ops` - 批量操作服务状态 -/// -/// # Returns -/// * `Ok(BatchResult)` - 成功时返回批量操作结果 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn batch_add_tags( - request: BatchAddTagsRequest, - batch_ops: State<'_, BatchOperationsState>, -) -> Result { - Ok(batch_ops - .0 - .execute( - &request.flow_ids, - BatchOperation::AddTags { tags: request.tags }, - ) - .await) -} - -/// 批量移除标签 -/// -/// **Validates: Requirements 11.4** -/// -/// # Arguments -/// * `request` - 批量移除标签请求参数 -/// * `batch_ops` - 批量操作服务状态 -/// -/// # Returns -/// * `Ok(BatchResult)` - 成功时返回批量操作结果 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn batch_remove_tags( - request: BatchRemoveTagsRequest, - batch_ops: State<'_, BatchOperationsState>, -) -> Result { - Ok(batch_ops - .0 - .execute( - &request.flow_ids, - BatchOperation::RemoveTags { tags: request.tags }, - ) - .await) -} - -/// 批量导出 Flow -/// -/// **Validates: Requirements 11.5** -/// -/// # Arguments -/// * `request` - 批量导出请求参数 -/// * `batch_ops` - 批量操作服务状态 -/// -/// # Returns -/// * `Ok(BatchResult)` - 成功时返回批量操作结果(包含导出数据) -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn batch_export_flows( - request: BatchExportFlowsRequest, - batch_ops: State<'_, BatchOperationsState>, -) -> Result { - Ok(batch_ops - .0 - .execute( - &request.flow_ids, - BatchOperation::Export { - format: request.format, - }, - ) - .await) -} - -/// 批量删除 Flow -/// -/// **Validates: Requirements 11.6** -/// -/// # Arguments -/// * `request` - 批量删除请求参数 -/// * `batch_ops` - 批量操作服务状态 -/// -/// # Returns -/// * `Ok(BatchResult)` - 成功时返回批量操作结果 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn batch_delete_flows( - request: BatchDeleteFlowsRequest, - batch_ops: State<'_, BatchOperationsState>, -) -> Result { - Ok(batch_ops - .0 - .execute(&request.flow_ids, BatchOperation::Delete) - .await) -} - -/// 批量添加到会话 -/// -/// **Validates: Requirements 11.2-11.6** -/// -/// # Arguments -/// * `request` - 批量添加到会话请求参数 -/// * `batch_ops` - 批量操作服务状态 -/// -/// # Returns -/// * `Ok(BatchResult)` - 成功时返回批量操作结果 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn batch_add_to_session( - request: BatchAddToSessionRequest, - batch_ops: State<'_, BatchOperationsState>, -) -> Result { - Ok(batch_ops - .0 - .execute( - &request.flow_ids, - BatchOperation::AddToSession { - session_id: request.session_id, - }, - ) - .await) -} - -// ============================================================================ -// 实时监控增强命令 -// ============================================================================ - -use crate::flow_monitor::ThresholdConfig; - -/// 阈值配置响应 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ThresholdConfigResponse { - /// 是否启用阈值检测 - pub enabled: bool, - /// 延迟阈值(毫秒) - pub latency_threshold_ms: u64, - /// Token 使用量阈值 - pub token_threshold: u32, - /// 输入 Token 阈值(可选) - pub input_token_threshold: Option, - /// 输出 Token 阈值(可选) - pub output_token_threshold: Option, -} - -impl From for ThresholdConfigResponse { - fn from(config: ThresholdConfig) -> Self { - Self { - enabled: config.enabled, - latency_threshold_ms: config.latency_threshold_ms, - token_threshold: config.token_threshold, - input_token_threshold: config.input_token_threshold, - output_token_threshold: config.output_token_threshold, - } - } -} - -/// 请求速率响应 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct RequestRateResponse { - /// 请求速率(每秒) - pub rate: f64, - /// 时间窗口内的请求数量 - pub count: usize, - /// 时间窗口(秒) - pub window_seconds: i64, -} - -/// 获取阈值配置 -/// -/// **Validates: Requirements 10.3, 10.4** -/// -/// # Arguments -/// * `monitor` - Flow 监控服务状态 -/// -/// # Returns -/// * `Ok(ThresholdConfigResponse)` - 成功时返回阈值配置 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn get_threshold_config( - monitor: State<'_, FlowMonitorState>, -) -> Result { - let config = monitor.0.threshold_config().await; - Ok(ThresholdConfigResponse::from(config)) -} - -/// 更新阈值配置 -/// -/// **Validates: Requirements 10.3, 10.4** -/// -/// # Arguments -/// * `config` - 新的阈值配置 -/// * `monitor` - Flow 监控服务状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn update_threshold_config( - config: ThresholdConfig, - monitor: State<'_, FlowMonitorState>, -) -> Result<(), String> { - monitor.0.update_threshold_config(config).await; - Ok(()) -} - -/// 获取请求速率 -/// -/// **Validates: Requirements 10.7** -/// -/// # Arguments -/// * `monitor` - Flow 监控服务状态 -/// -/// # Returns -/// * `Ok(RequestRateResponse)` - 成功时返回请求速率信息 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn get_request_rate( - monitor: State<'_, FlowMonitorState>, -) -> Result { - let rate = monitor.0.get_request_rate().await; - let count = monitor.0.get_request_count().await; - - Ok(RequestRateResponse { - rate, - count, - window_seconds: 60, // 默认 60 秒窗口 - }) -} - -/// 设置请求速率追踪器的时间窗口 -/// -/// **Validates: Requirements 10.7** -/// -/// # Arguments -/// * `window_seconds` - 时间窗口(秒) -/// * `monitor` - Flow 监控服务状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn set_rate_window( - window_seconds: i64, - monitor: State<'_, FlowMonitorState>, -) -> Result<(), String> { - if window_seconds <= 0 { - return Err("时间窗口必须大于 0".to_string()); - } - monitor.0.set_rate_window(window_seconds).await; - Ok(()) -} -// ============================================================================ -// 通知配置命令 -// ============================================================================ - -/* -/// 通知配置响应 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct NotificationConfigResponse { - /// 是否启用通知 - pub enabled: bool, - /// 新 Flow 通知配置 - pub new_flow: NotificationSettingsResponse, - /// 错误 Flow 通知配置 - pub error_flow: NotificationSettingsResponse, - /// 延迟警告通知配置 - pub latency_warning: NotificationSettingsResponse, - /// Token 警告通知配置 - pub token_warning: NotificationSettingsResponse, -} - -/// 通知设置响应 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct NotificationSettingsResponse { - /// 是否启用 - pub enabled: bool, - /// 是否显示桌面通知 - pub desktop: bool, - /// 是否播放声音 - pub sound: bool, - /// 声音文件路径(可选) - pub sound_file: Option, -} - -impl From for NotificationSettingsResponse { - fn from(settings: NotificationSettings) -> Self { - Self { - enabled: settings.enabled, - desktop: settings.desktop, - sound: settings.sound, - sound_file: settings.sound_file, - } - } -} - -impl From for NotificationConfigResponse { - fn from(config: NotificationConfig) -> Self { - Self { - enabled: config.enabled, - new_flow: NotificationSettingsResponse::from(config.new_flow), - error_flow: NotificationSettingsResponse::from(config.error_flow), - latency_warning: NotificationSettingsResponse::from(config.latency_warning), - token_warning: NotificationSettingsResponse::from(config.token_warning), - } - } -} - -/// 获取通知配置 -/// -/// **Validates: Requirements 10.1, 10.2** -/// -/// # Arguments -/// * `monitor` - Flow 监控服务状态 -/// -/// # Returns -/// * `Ok(NotificationConfigResponse)` - 成功时返回通知配置 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn get_notification_config( - monitor: State<'_, FlowMonitorState>, -) -> Result { - let config = monitor.0.notification_config().await; - Ok(NotificationConfigResponse::from(config)) -} - -/// 更新通知配置 -/// -/// **Validates: Requirements 10.1, 10.2** -/// -/// # Arguments -/// * `config` - 新的通知配置 -/// * `monitor` - Flow 监控服务状态 -/// -/// # Returns -/// * `Ok(())` - 成功 -/// * `Err(String)` - 失败时返回错误消息 -#[tauri::command] -pub async fn update_notification_config( - config: NotificationConfig, - monitor: State<'_, FlowMonitorState>, -) -> Result<(), String> { - monitor.0.update_notification_config(config).await; - Ok(()) -} -*/ diff --git a/src-tauri/src/commands/mod.rs b/src-tauri/src/commands/mod.rs index 48c8d0632..c70054af4 100644 --- a/src-tauri/src/commands/mod.rs +++ b/src-tauri/src/commands/mod.rs @@ -4,14 +4,13 @@ pub mod api_key_provider_cmd; pub mod asr_cmd; pub mod aster_agent_cmd; pub mod auto_fix_cmd; -pub mod browser_interceptor_cmd; pub mod config_cmd; pub mod connect_cmd; pub mod connection_cmd; pub mod content_cmd; pub mod context_memory; pub mod external_tools_cmd; -pub mod flow_monitor_cmd; + pub mod general_chat_cmd; pub mod injection_cmd; pub mod kiro_local; diff --git a/src-tauri/src/config/mod.rs b/src-tauri/src/config/mod.rs index ff9ac75e5..55a7085bb 100644 --- a/src-tauri/src/config/mod.rs +++ b/src-tauri/src/config/mod.rs @@ -1,71 +1,15 @@ //! 配置管理模块 //! -//! 提供 YAML 配置文件支持、热重载和配置导入导出功能 -//! 同时保持与旧版 JSON 配置的向后兼容性 +//! 核心配置类型、YAML 支持、热重载和导入导出功能已迁移到 proxycast-core crate。 +//! 本模块保留 observer(依赖 Tauri)和集成测试。 #![allow(unused_imports)] -mod export; -mod hot_reload; -mod import; -pub mod observer; -mod path_utils; -mod types; -mod yaml; +// 从 core crate 重新导出所有配置类型 +pub use proxycast_core::config::*; -pub use export::{ExportBundle, ExportOptions, ExportService, REDACTED_PLACEHOLDER}; -pub use hot_reload::{ - ConfigChangeEvent as FileChangeEvent, ConfigChangeKind, FileWatcher, HotReloadManager, - ReloadResult, -}; -pub use import::{ImportOptions, ImportService, ValidationResult}; -pub use path_utils::{collapse_tilde, contains_tilde, expand_tilde}; -pub use types::{ - generate_secure_api_key, - AmpConfig, - AmpModelMapping, - ApiKeyEntry, - AsrCredentialEntry, - // ASR 和语音输入相关类型 - AsrProviderType, - BaiduConfig, - Config, - CredentialEntry, - CredentialPoolConfig, - CustomProviderConfig, - EndpointProvidersConfig, - ExperimentalFeatures, - GeminiApiKeyEntry, - InjectionRuleConfig, - InjectionSettings, - LoggingConfig, - ModelInfo, - ModelsConfig, - NativeAgentConfig, - OpenAIAsrConfig, - ProviderConfig, - ProviderModelsConfig, - ProvidersConfig, - QuotaExceededConfig, - RemoteManagementConfig, - RetrySettings, - RoutingConfig, - ScreenshotChatConfig, - ServerConfig, - TlsConfig, - VertexApiKeyEntry, - VertexModelAlias, - VoiceInputConfig, - VoiceInstruction, - VoiceOutputConfig, - VoiceOutputMode, - VoiceProcessorConfig, - WhisperLocalConfig, - WhisperModelSize, - XunfeiConfig, - DEFAULT_API_KEY, -}; -pub use yaml::{load_config, save_config, ConfigError, ConfigManager, YamlService}; +// observer 模块保留在主 crate(依赖 Tauri) +pub mod observer; // 重新导出观察者模块的核心类型 pub use observer::{ diff --git a/src-tauri/src/config/tests.rs b/src-tauri/src/config/tests.rs index b72b69a68..f89e42162 100644 --- a/src-tauri/src/config/tests.rs +++ b/src-tauri/src/config/tests.rs @@ -2,12 +2,12 @@ //! //! 使用 proptest 进行属性测试 -use crate::config::types::{ContentCreatorConfig, NavigationConfig}; use crate::config::{ collapse_tilde, contains_tilde, expand_tilde, Config, ConfigManager, CustomProviderConfig, HotReloadManager, InjectionSettings, LoggingConfig, ProviderConfig, ProvidersConfig, ReloadResult, RetrySettings, RoutingConfig, ServerConfig, YamlService, }; +use crate::config::{ContentCreatorConfig, NavigationConfig}; use proptest::prelude::*; use std::io::Write; use tempfile::NamedTempFile; diff --git a/src-tauri/src/data/mod.rs b/src-tauri/src/data/mod.rs index 60a799247..802b30544 100644 --- a/src-tauri/src/data/mod.rs +++ b/src-tauri/src/data/mod.rs @@ -1,4 +1 @@ //! 静态数据模块 -//! -//! 模型数据现在从 aiclientproxy/models 仓库获取 -//! 本地硬编码数据已迁移到独立仓库: https://github.com/aiclientproxy/models diff --git a/src-tauri/src/flow_monitor/batch_ops.rs b/src-tauri/src/flow_monitor/batch_ops.rs deleted file mode 100644 index f2e221fc9..000000000 --- a/src-tauri/src/flow_monitor/batch_ops.rs +++ /dev/null @@ -1,570 +0,0 @@ -//! 批量操作服务 -//! -//! 该模块实现 Flow 批量操作功能,支持对多个 Flow 进行批量收藏、 -//! 添加标签、导出、删除等操作。 -//! -//! **Validates: Requirements 11.2-11.6** - -use serde::{Deserialize, Serialize}; -use std::sync::Arc; -use thiserror::Error; - -use super::exporter::{ExportFormat, ExportOptions, FlowExporter}; -use super::models::LLMFlow; -use super::monitor::FlowMonitor; -use super::session::SessionManager; - -/// 批量操作错误 -#[derive(Debug, Error)] -pub enum BatchOpsError { - #[error("Flow 不存在: {0}")] - FlowNotFound(String), - #[error("会话不存在: {0}")] - SessionNotFound(String), - #[error("导出错误: {0}")] - ExportError(String), - #[error("操作失败: {0}")] - OperationFailed(String), -} - -pub type Result = std::result::Result; - -/// 批量操作类型 -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] -#[serde(tag = "type", rename_all = "snake_case")] -pub enum BatchOperation { - Star, - Unstar, - AddTags { tags: Vec }, - RemoveTags { tags: Vec }, - Export { format: ExportFormat }, - Delete, - AddToSession { session_id: String }, -} - -/// 批量操作结果 -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] -pub struct BatchResult { - pub total: usize, - pub success: usize, - pub failed: usize, - pub errors: Vec<(String, String)>, - #[serde(skip_serializing_if = "Option::is_none")] - pub export_data: Option, -} - -impl BatchResult { - pub fn new(total: usize) -> Self { - Self { - total, - success: 0, - failed: 0, - errors: Vec::new(), - export_data: None, - } - } - pub fn record_success(&mut self) { - self.success += 1; - } - pub fn record_failure(&mut self, flow_id: impl Into, error: impl Into) { - self.failed += 1; - self.errors.push((flow_id.into(), error.into())); - } - pub fn is_all_success(&self) -> bool { - self.failed == 0 - } - pub fn is_all_failed(&self) -> bool { - self.success == 0 && self.total > 0 - } - pub fn is_partial_success(&self) -> bool { - self.success > 0 && self.failed > 0 - } -} - -/// 批量操作服务 -pub struct BatchOperations { - flow_monitor: Arc, - session_manager: Option>, -} - -impl BatchOperations { - pub fn new( - flow_monitor: Arc, - session_manager: Option>, - ) -> Self { - Self { - flow_monitor, - session_manager, - } - } - - pub async fn execute(&self, flow_ids: &[String], operation: BatchOperation) -> BatchResult { - self.execute_with_progress(flow_ids, operation, |_, _| {}) - .await - } - - pub async fn execute_with_progress( - &self, - flow_ids: &[String], - operation: BatchOperation, - progress: F, - ) -> BatchResult - where - F: Fn(usize, usize) + Send + Sync, - { - let mut result = BatchResult::new(flow_ids.len()); - match operation { - BatchOperation::Star => { - self.batch_star(flow_ids, true, &mut result, &progress) - .await - } - BatchOperation::Unstar => { - self.batch_star(flow_ids, false, &mut result, &progress) - .await - } - BatchOperation::AddTags { tags } => { - self.batch_add_tags(flow_ids, &tags, &mut result, &progress) - .await - } - BatchOperation::RemoveTags { tags } => { - self.batch_remove_tags(flow_ids, &tags, &mut result, &progress) - .await - } - BatchOperation::Export { format } => { - self.batch_export(flow_ids, format, &mut result, &progress) - .await - } - BatchOperation::Delete => self.batch_delete(flow_ids, &mut result, &progress).await, - BatchOperation::AddToSession { session_id } => { - self.batch_add_to_session(flow_ids, &session_id, &mut result, &progress) - .await - } - } - result - } - - async fn batch_star( - &self, - flow_ids: &[String], - starred: bool, - result: &mut BatchResult, - progress: &F, - ) where - F: Fn(usize, usize), - { - let total = flow_ids.len(); - for (i, flow_id) in flow_ids.iter().enumerate() { - progress(i + 1, total); - let memory_store = self.flow_monitor.memory_store(); - let store = memory_store.read().await; - let current_starred = store - .get(flow_id) - .and_then(|f| f.read().ok().map(|flow| flow.annotations.starred)); - drop(store); - match current_starred { - Some(current) if current != starred => { - if self.flow_monitor.toggle_starred(flow_id).await { - result.record_success(); - } else { - result.record_failure(flow_id, "更新收藏状态失败"); - } - } - Some(_) => { - result.record_success(); - } - None => { - result.record_failure(flow_id, "Flow 不存在"); - } - } - } - } - - async fn batch_add_tags( - &self, - flow_ids: &[String], - tags: &[String], - result: &mut BatchResult, - progress: &F, - ) where - F: Fn(usize, usize), - { - let total = flow_ids.len(); - for (i, flow_id) in flow_ids.iter().enumerate() { - progress(i + 1, total); - let memory_store = self.flow_monitor.memory_store(); - let store = memory_store.read().await; - let exists = store.get(flow_id).is_some(); - drop(store); - if !exists { - result.record_failure(flow_id, "Flow 不存在"); - continue; - } - let mut all_success = true; - for tag in tags { - if !self.flow_monitor.add_tag(flow_id, tag.clone()).await { - all_success = false; - break; - } - } - if all_success { - result.record_success(); - } else { - result.record_failure(flow_id, "添加标签失败"); - } - } - } - - async fn batch_remove_tags( - &self, - flow_ids: &[String], - tags: &[String], - result: &mut BatchResult, - progress: &F, - ) where - F: Fn(usize, usize), - { - let total = flow_ids.len(); - for (i, flow_id) in flow_ids.iter().enumerate() { - progress(i + 1, total); - let memory_store = self.flow_monitor.memory_store(); - let store = memory_store.read().await; - let exists = store.get(flow_id).is_some(); - drop(store); - if !exists { - result.record_failure(flow_id, "Flow 不存在"); - continue; - } - for tag in tags { - let _ = self.flow_monitor.remove_tag(flow_id, tag).await; - } - result.record_success(); - } - } - - async fn batch_export( - &self, - flow_ids: &[String], - format: ExportFormat, - result: &mut BatchResult, - progress: &F, - ) where - F: Fn(usize, usize), - { - let total = flow_ids.len(); - let mut flows: Vec = Vec::with_capacity(total); - for (i, flow_id) in flow_ids.iter().enumerate() { - progress(i + 1, total); - let memory_store = self.flow_monitor.memory_store(); - let store = memory_store.read().await; - if let Some(flow_lock) = store.get(flow_id) { - if let Ok(flow) = flow_lock.read() { - flows.push(flow.clone()); - result.record_success(); - } else { - result.record_failure(flow_id, "无法读取 Flow"); - } - } else { - result.record_failure(flow_id, "Flow 不存在"); - } - } - if !flows.is_empty() { - let options = ExportOptions { - format, - ..Default::default() - }; - let exporter = FlowExporter::new(options); - let export_result = exporter.export(&flows); - result.export_data = Some(export_result.to_string_pretty()); - } - } - - async fn batch_delete(&self, flow_ids: &[String], result: &mut BatchResult, progress: &F) - where - F: Fn(usize, usize), - { - let total = flow_ids.len(); - for (i, flow_id) in flow_ids.iter().enumerate() { - progress(i + 1, total); - let memory_store = self.flow_monitor.memory_store(); - let mut store = memory_store.write().await; - if store.remove(flow_id) { - result.record_success(); - } else { - result.record_failure(flow_id, "Flow 不存在或删除失败"); - } - } - } - - async fn batch_add_to_session( - &self, - flow_ids: &[String], - session_id: &str, - result: &mut BatchResult, - progress: &F, - ) where - F: Fn(usize, usize), - { - let total = flow_ids.len(); - let session_manager = match &self.session_manager { - Some(sm) => sm, - None => { - for flow_id in flow_ids { - result.record_failure(flow_id, "会话管理器不可用"); - } - return; - } - }; - match session_manager.get_session(session_id) { - Ok(Some(_)) => {} - Ok(None) => { - for flow_id in flow_ids { - result.record_failure(flow_id, format!("会话不存在: {session_id}")); - } - return; - } - Err(e) => { - for flow_id in flow_ids { - result.record_failure(flow_id, format!("查询会话失败: {e}")); - } - return; - } - } - for (i, flow_id) in flow_ids.iter().enumerate() { - progress(i + 1, total); - let memory_store = self.flow_monitor.memory_store(); - let store = memory_store.read().await; - let exists = store.get(flow_id).is_some(); - drop(store); - if !exists { - result.record_failure(flow_id, "Flow 不存在"); - continue; - } - match session_manager.add_flow(session_id, flow_id) { - Ok(_) => { - result.record_success(); - } - Err(e) => { - result.record_failure(flow_id, format!("添加到会话失败: {e}")); - } - } - } - } -} - -// ============================================================================ -// 属性测试 -// ============================================================================ - -#[cfg(test)] -mod property_tests { - use super::*; - use crate::flow_monitor::models::{FlowMetadata, FlowType, LLMRequest}; - use crate::flow_monitor::monitor::FlowMonitorConfig; - use proptest::prelude::*; - - fn create_test_flow_monitor() -> Arc { - let config = FlowMonitorConfig::default(); - Arc::new(FlowMonitor::new(config, None)) - } - - async fn create_test_flow(monitor: &FlowMonitor, flow_id: &str) -> String { - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model: "gpt-4".to_string(), - ..Default::default() - }; - let metadata = FlowMetadata::default(); - let mut flow = crate::flow_monitor::models::LLMFlow::new( - flow_id.to_string(), - FlowType::ChatCompletions, - request, - metadata, - ); - flow.state = crate::flow_monitor::models::FlowState::Completed; - let store = monitor.memory_store(); - store.write().await.add(flow); - flow_id.to_string() - } - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: flow-monitor-enhancement, Property 20: 批量操作正确性** - /// **Validates: Requirements 11.2-11.6** - /// - /// *对于任意* Flow 集合和批量操作,操作后所有 Flow 应该被正确更新。 - #[test] - fn prop_batch_star_correctness(flow_count in 1usize..10usize) { - let rt = tokio::runtime::Runtime::new().unwrap(); - rt.block_on(async { - let monitor = create_test_flow_monitor(); - let batch_ops = BatchOperations::new(monitor.clone(), None); - - // 创建测试 Flow - let mut flow_ids = Vec::new(); - for i in 0..flow_count { - let id = create_test_flow(&monitor, &format!("flow-{i}")).await; - flow_ids.push(id); - } - - // 执行批量收藏 - let result = batch_ops.execute(&flow_ids, BatchOperation::Star).await; - - // 验证结果 - prop_assert_eq!(result.total, flow_count); - prop_assert_eq!(result.success, flow_count); - prop_assert_eq!(result.failed, 0); - - // 验证所有 Flow 都被收藏 - let store = monitor.memory_store(); - let s = store.read().await; - for flow_id in &flow_ids { - if let Some(flow_lock) = s.get(flow_id) { - let flow = flow_lock.read().unwrap(); - prop_assert!(flow.annotations.starred, "Flow {} 应该被收藏", flow_id); - } - } - Ok(()) - })?; - } - - /// **Feature: flow-monitor-enhancement, Property 21: 批量操作原子性** - /// **Validates: Requirements 11.2-11.6** - /// - /// *对于任意* 批量操作,如果部分失败,成功的部分应该被正确应用,失败的部分应该被正确报告。 - #[test] - fn prop_batch_operation_atomicity( - valid_flow_count in 1usize..8usize, - invalid_flow_count in 1usize..5usize, - ) { - let rt = tokio::runtime::Runtime::new().unwrap(); - rt.block_on(async { - let monitor = create_test_flow_monitor(); - let batch_ops = BatchOperations::new(monitor.clone(), None); - - // 创建有效的 Flow - let mut valid_flow_ids = Vec::new(); - for i in 0..valid_flow_count { - let id = create_test_flow(&monitor, &format!("valid-flow-{i}")).await; - valid_flow_ids.push(id); - } - - // 创建无效的 Flow ID(不存在的) - let mut invalid_flow_ids = Vec::new(); - for i in 0..invalid_flow_count { - invalid_flow_ids.push(format!("invalid-flow-{i}")); - } - - // 混合有效和无效的 Flow ID - let mut all_flow_ids = valid_flow_ids.clone(); - all_flow_ids.extend(invalid_flow_ids.clone()); - - // 执行批量收藏操作 - let result = batch_ops.execute(&all_flow_ids, BatchOperation::Star).await; - - // 验证结果统计 - prop_assert_eq!(result.total, valid_flow_count + invalid_flow_count); - prop_assert_eq!(result.success, valid_flow_count); - prop_assert_eq!(result.failed, invalid_flow_count); - prop_assert_eq!(result.errors.len(), invalid_flow_count); - - // 验证成功的 Flow 被正确更新 - let store = monitor.memory_store(); - let s = store.read().await; - for flow_id in &valid_flow_ids { - if let Some(flow_lock) = s.get(flow_id) { - let flow = flow_lock.read().unwrap(); - prop_assert!(flow.annotations.starred, "有效的 Flow {} 应该被收藏", flow_id); - } - } - - // 验证失败的 Flow ID 被正确报告 - for invalid_id in &invalid_flow_ids { - let found_error = result.errors.iter().any(|(id, _)| id == invalid_id); - prop_assert!(found_error, "无效的 Flow ID {} 应该在错误列表中", invalid_id); - } - - // 验证部分成功状态 - prop_assert!(result.is_partial_success(), "应该是部分成功状态"); - prop_assert!(!result.is_all_success(), "不应该是全部成功"); - prop_assert!(!result.is_all_failed(), "不应该是全部失败"); - - Ok(()) - })?; - } - - /// **Feature: flow-monitor-enhancement, Property 21b: 批量标签操作原子性** - /// **Validates: Requirements 11.2-11.6** - /// - /// *对于任意* 批量标签操作,部分失败时应该正确处理成功和失败的情况。 - #[test] - fn prop_batch_tag_operation_atomicity( - valid_flow_count in 1usize..6usize, - invalid_flow_count in 1usize..4usize, - tag_count in 1usize..4usize, - ) { - let rt = tokio::runtime::Runtime::new().unwrap(); - rt.block_on(async { - let monitor = create_test_flow_monitor(); - let batch_ops = BatchOperations::new(monitor.clone(), None); - - // 创建有效的 Flow - let mut valid_flow_ids = Vec::new(); - for i in 0..valid_flow_count { - let id = create_test_flow(&monitor, &format!("valid-flow-{i}")).await; - valid_flow_ids.push(id); - } - - // 创建无效的 Flow ID - let mut invalid_flow_ids = Vec::new(); - for i in 0..invalid_flow_count { - invalid_flow_ids.push(format!("invalid-flow-{i}")); - } - - // 创建标签列表 - let tags: Vec = (0..tag_count).map(|i| format!("tag-{i}")).collect(); - - // 混合有效和无效的 Flow ID - let mut all_flow_ids = valid_flow_ids.clone(); - all_flow_ids.extend(invalid_flow_ids.clone()); - - // 执行批量添加标签操作 - let result = batch_ops.execute( - &all_flow_ids, - BatchOperation::AddTags { tags: tags.clone() } - ).await; - - // 验证结果统计 - prop_assert_eq!(result.total, valid_flow_count + invalid_flow_count); - prop_assert_eq!(result.success, valid_flow_count); - prop_assert_eq!(result.failed, invalid_flow_count); - - // 验证成功的 Flow 被正确添加标签 - let store = monitor.memory_store(); - let s = store.read().await; - for flow_id in &valid_flow_ids { - if let Some(flow_lock) = s.get(flow_id) { - let flow = flow_lock.read().unwrap(); - for tag in &tags { - prop_assert!( - flow.annotations.tags.contains(tag), - "有效的 Flow {} 应该包含标签 {}", - flow_id, - tag - ); - } - } - } - - // 验证失败的 Flow ID 被正确报告 - for invalid_id in &invalid_flow_ids { - let found_error = result.errors.iter().any(|(id, _)| id == invalid_id); - prop_assert!(found_error, "无效的 Flow ID {} 应该在错误列表中", invalid_id); - } - - Ok(()) - })?; - } - } -} diff --git a/src-tauri/src/flow_monitor/bookmark.rs b/src-tauri/src/flow_monitor/bookmark.rs deleted file mode 100644 index 801da8daf..000000000 --- a/src-tauri/src/flow_monitor/bookmark.rs +++ /dev/null @@ -1,1029 +0,0 @@ -//! 书签管理器 -//! -//! 该模块实现 Flow 书签功能,支持快速定位和导航到重要的 Flow。 -//! -//! **Validates: Requirements 8.1, 8.3, 8.6** - -use chrono::{DateTime, Utc}; -use rusqlite::{params, Connection, OptionalExtension}; -use serde::{Deserialize, Serialize}; -use std::path::PathBuf; -use std::sync::Mutex; -use thiserror::Error; -use uuid::Uuid; - -// ============================================================================ -// 错误类型 -// ============================================================================ - -/// 书签管理错误 -#[derive(Debug, Error)] -pub enum BookmarkError { - #[error("SQLite 错误: {0}")] - Sqlite(#[from] rusqlite::Error), - - #[error("书签不存在: {0}")] - BookmarkNotFound(String), - - #[error("Flow 不存在: {0}")] - FlowNotFound(String), - - #[error("JSON 序列化错误: {0}")] - Json(#[from] serde_json::Error), - - #[error("IO 错误: {0}")] - Io(#[from] std::io::Error), - - #[error("书签已存在: flow_id={0}")] - BookmarkAlreadyExists(String), -} - -pub type Result = std::result::Result; - -// ============================================================================ -// 数据结构 -// ============================================================================ - -/// Flow 书签 -/// -/// **Validates: Requirements 8.1, 8.3** -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] -pub struct FlowBookmark { - /// 唯一标识符 - pub id: String, - /// 关联的 Flow ID - pub flow_id: String, - /// 书签名称(可选) - #[serde(skip_serializing_if = "Option::is_none")] - pub name: Option, - /// 分组名称(可选) - #[serde(skip_serializing_if = "Option::is_none")] - pub group: Option, - /// 创建时间 - pub created_at: DateTime, -} - -impl FlowBookmark { - /// 创建新书签 - pub fn new(flow_id: impl Into, name: Option, group: Option) -> Self { - Self { - id: Uuid::new_v4().to_string(), - flow_id: flow_id.into(), - name, - group, - created_at: Utc::now(), - } - } -} - -/// 书签导出数据 -/// -/// **Validates: Requirements 8.6** -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct BookmarkExport { - /// 版本号 - pub version: String, - /// 导出时间 - pub exported_at: DateTime, - /// 书签列表 - pub bookmarks: Vec, -} - -impl BookmarkExport { - pub fn new(bookmarks: Vec) -> Self { - Self { - version: "1.0".to_string(), - exported_at: Utc::now(), - bookmarks, - } - } -} - -// ============================================================================ -// 书签管理器 -// ============================================================================ - -/// 书签管理器 -/// -/// **Validates: Requirements 8.1, 8.3, 8.6** -pub struct BookmarkManager { - /// SQLite 连接 - db: Mutex, -} - -impl BookmarkManager { - /// 创建新的书签管理器 - /// - /// # Arguments - /// * `db_path` - SQLite 数据库路径 - pub fn new(db_path: PathBuf) -> Result { - // 确保目录存在 - if let Some(parent) = db_path.parent() { - std::fs::create_dir_all(parent)?; - } - - let conn = Connection::open(&db_path)?; - Self::init_database(&conn)?; - - Ok(Self { - db: Mutex::new(conn), - }) - } - - /// 从现有连接创建书签管理器(用于测试) - pub fn from_connection(conn: Connection) -> Result { - Self::init_database(&conn)?; - - Ok(Self { - db: Mutex::new(conn), - }) - } - - /// 初始化数据库表 - fn init_database(conn: &Connection) -> Result<()> { - conn.execute_batch( - r#" - -- 书签表 - CREATE TABLE IF NOT EXISTS flow_bookmarks ( - id TEXT PRIMARY KEY, - flow_id TEXT NOT NULL, - name TEXT, - group_name TEXT, - created_at TEXT NOT NULL - ); - - CREATE INDEX IF NOT EXISTS idx_bookmarks_flow ON flow_bookmarks(flow_id); - CREATE INDEX IF NOT EXISTS idx_bookmarks_group ON flow_bookmarks(group_name); - CREATE INDEX IF NOT EXISTS idx_bookmarks_created ON flow_bookmarks(created_at); - "#, - )?; - - Ok(()) - } - - /// 添加书签 - /// - /// **Validates: Requirements 8.1** - /// - /// # Arguments - /// * `flow_id` - Flow ID - /// * `name` - 书签名称(可选) - /// * `group` - 分组名称(可选) - /// - /// # Returns - /// 新创建的书签 - pub fn add( - &self, - flow_id: impl Into, - name: Option<&str>, - group: Option<&str>, - ) -> Result { - let flow_id = flow_id.into(); - let bookmark = FlowBookmark::new(&flow_id, name.map(String::from), group.map(String::from)); - - let conn = self.db.lock().unwrap(); - - conn.execute( - r#" - INSERT INTO flow_bookmarks (id, flow_id, name, group_name, created_at) - VALUES (?1, ?2, ?3, ?4, ?5) - "#, - params![ - bookmark.id, - bookmark.flow_id, - bookmark.name, - bookmark.group, - bookmark.created_at.to_rfc3339(), - ], - )?; - - Ok(bookmark) - } - - /// 获取书签 - /// - /// # Arguments - /// * `bookmark_id` - 书签 ID - /// - /// # Returns - /// 书签信息(如果存在) - pub fn get(&self, bookmark_id: &str) -> Result> { - let conn = self.db.lock().unwrap(); - - let bookmark: Option<(String, String, Option, Option, String)> = conn - .query_row( - r#" - SELECT id, flow_id, name, group_name, created_at - FROM flow_bookmarks - WHERE id = ?1 - "#, - params![bookmark_id], - |row| { - Ok(( - row.get(0)?, - row.get(1)?, - row.get(2)?, - row.get(3)?, - row.get(4)?, - )) - }, - ) - .optional()?; - - match bookmark { - Some((id, flow_id, name, group, created_at)) => Ok(Some(FlowBookmark { - id, - flow_id, - name, - group, - created_at: DateTime::parse_from_rfc3339(&created_at) - .map(|dt| dt.with_timezone(&Utc)) - .unwrap_or_else(|_| Utc::now()), - })), - None => Ok(None), - } - } - - /// 根据 Flow ID 获取书签 - /// - /// # Arguments - /// * `flow_id` - Flow ID - /// - /// # Returns - /// 书签信息(如果存在) - pub fn get_by_flow_id(&self, flow_id: &str) -> Result> { - let conn = self.db.lock().unwrap(); - - let bookmark: Option<(String, String, Option, Option, String)> = conn - .query_row( - r#" - SELECT id, flow_id, name, group_name, created_at - FROM flow_bookmarks - WHERE flow_id = ?1 - "#, - params![flow_id], - |row| { - Ok(( - row.get(0)?, - row.get(1)?, - row.get(2)?, - row.get(3)?, - row.get(4)?, - )) - }, - ) - .optional()?; - - match bookmark { - Some((id, flow_id, name, group, created_at)) => Ok(Some(FlowBookmark { - id, - flow_id, - name, - group, - created_at: DateTime::parse_from_rfc3339(&created_at) - .map(|dt| dt.with_timezone(&Utc)) - .unwrap_or_else(|_| Utc::now()), - })), - None => Ok(None), - } - } - - /// 移除书签 - /// - /// **Validates: Requirements 8.1** - /// - /// # Arguments - /// * `bookmark_id` - 书签 ID - pub fn remove(&self, bookmark_id: &str) -> Result<()> { - let conn = self.db.lock().unwrap(); - - let rows_affected = conn.execute( - "DELETE FROM flow_bookmarks WHERE id = ?1", - params![bookmark_id], - )?; - - if rows_affected == 0 { - return Err(BookmarkError::BookmarkNotFound(bookmark_id.to_string())); - } - - Ok(()) - } - - /// 根据 Flow ID 移除书签 - /// - /// # Arguments - /// * `flow_id` - Flow ID - pub fn remove_by_flow_id(&self, flow_id: &str) -> Result<()> { - let conn = self.db.lock().unwrap(); - - conn.execute( - "DELETE FROM flow_bookmarks WHERE flow_id = ?1", - params![flow_id], - )?; - - Ok(()) - } - - /// 更新书签 - /// - /// # Arguments - /// * `bookmark_id` - 书签 ID - /// * `name` - 新名称(可选) - /// * `group` - 新分组(可选) - pub fn update( - &self, - bookmark_id: &str, - name: Option>, - group: Option>, - ) -> Result { - let conn = self.db.lock().unwrap(); - - // 检查书签是否存在 - let exists: bool = conn - .query_row( - "SELECT 1 FROM flow_bookmarks WHERE id = ?1", - params![bookmark_id], - |_| Ok(true), - ) - .optional()? - .unwrap_or(false); - - if !exists { - return Err(BookmarkError::BookmarkNotFound(bookmark_id.to_string())); - } - - // 更新名称 - if let Some(new_name) = name { - conn.execute( - "UPDATE flow_bookmarks SET name = ?1 WHERE id = ?2", - params![new_name, bookmark_id], - )?; - } - - // 更新分组 - if let Some(new_group) = group { - conn.execute( - "UPDATE flow_bookmarks SET group_name = ?1 WHERE id = ?2", - params![new_group, bookmark_id], - )?; - } - - drop(conn); - - // 返回更新后的书签 - self.get(bookmark_id)? - .ok_or_else(|| BookmarkError::BookmarkNotFound(bookmark_id.to_string())) - } - - /// 列出所有书签 - /// - /// **Validates: Requirements 8.3** - /// - /// # Arguments - /// * `group` - 分组名称(可选,None 表示所有书签) - /// - /// # Returns - /// 书签列表 - pub fn list(&self, group: Option<&str>) -> Result> { - let conn = self.db.lock().unwrap(); - - let mut stmt = if let Some(g) = group { - let mut stmt = conn.prepare( - r#" - SELECT id, flow_id, name, group_name, created_at - FROM flow_bookmarks - WHERE group_name = ?1 - ORDER BY created_at DESC - "#, - )?; - let bookmarks: Vec = stmt - .query_map(params![g], |row| { - Ok(FlowBookmark { - id: row.get(0)?, - flow_id: row.get(1)?, - name: row.get(2)?, - group: row.get(3)?, - created_at: DateTime::parse_from_rfc3339(&row.get::<_, String>(4)?) - .map(|dt| dt.with_timezone(&Utc)) - .unwrap_or_else(|_| Utc::now()), - }) - })? - .filter_map(|r| r.ok()) - .collect(); - return Ok(bookmarks); - } else { - conn.prepare( - r#" - SELECT id, flow_id, name, group_name, created_at - FROM flow_bookmarks - ORDER BY created_at DESC - "#, - )? - }; - - let bookmarks: Vec = stmt - .query_map([], |row| { - Ok(FlowBookmark { - id: row.get(0)?, - flow_id: row.get(1)?, - name: row.get(2)?, - group: row.get(3)?, - created_at: DateTime::parse_from_rfc3339(&row.get::<_, String>(4)?) - .map(|dt| dt.with_timezone(&Utc)) - .unwrap_or_else(|_| Utc::now()), - }) - })? - .filter_map(|r| r.ok()) - .collect(); - - Ok(bookmarks) - } - - /// 获取所有分组名称 - /// - /// **Validates: Requirements 8.3** - /// - /// # Returns - /// 分组名称列表 - pub fn list_groups(&self) -> Result> { - let conn = self.db.lock().unwrap(); - - let mut stmt = conn.prepare( - r#" - SELECT DISTINCT group_name - FROM flow_bookmarks - WHERE group_name IS NOT NULL - ORDER BY group_name ASC - "#, - )?; - - let groups: Vec = stmt - .query_map([], |row| row.get(0))? - .filter_map(|r| r.ok()) - .collect(); - - Ok(groups) - } - - /// 检查 Flow 是否已添加书签 - /// - /// # Arguments - /// * `flow_id` - Flow ID - /// - /// # Returns - /// 是否已添加书签 - pub fn is_bookmarked(&self, flow_id: &str) -> Result { - let conn = self.db.lock().unwrap(); - - let exists: bool = conn - .query_row( - "SELECT 1 FROM flow_bookmarks WHERE flow_id = ?1", - params![flow_id], - |_| Ok(true), - ) - .optional()? - .unwrap_or(false); - - Ok(exists) - } - - /// 获取书签数量 - pub fn count(&self) -> Result { - let conn = self.db.lock().unwrap(); - let count: i64 = - conn.query_row("SELECT COUNT(*) FROM flow_bookmarks", [], |row| row.get(0))?; - Ok(count as usize) - } - - /// 导出书签 - /// - /// **Validates: Requirements 8.6** - /// - /// # Returns - /// JSON 格式的导出数据 - pub fn export(&self) -> Result { - let bookmarks = self.list(None)?; - let export_data = BookmarkExport::new(bookmarks); - let json = serde_json::to_string_pretty(&export_data)?; - Ok(json) - } - - /// 导入书签 - /// - /// **Validates: Requirements 8.6** - /// - /// # Arguments - /// * `data` - JSON 格式的导入数据 - /// * `overwrite` - 是否覆盖已存在的书签(按 flow_id 判断) - /// - /// # Returns - /// 导入的书签列表 - pub fn import(&self, data: &str, overwrite: bool) -> Result> { - let export_data: BookmarkExport = serde_json::from_str(data)?; - - let mut imported = Vec::new(); - let conn = self.db.lock().unwrap(); - - for mut bookmark in export_data.bookmarks { - // 检查是否存在相同 flow_id 的书签 - let existing_id: Option = conn - .query_row( - "SELECT id FROM flow_bookmarks WHERE flow_id = ?1", - params![bookmark.flow_id], - |row| row.get(0), - ) - .optional()?; - - if let Some(existing) = existing_id { - if overwrite { - // 更新现有书签 - conn.execute( - r#" - UPDATE flow_bookmarks - SET name = ?1, group_name = ?2 - WHERE id = ?3 - "#, - params![bookmark.name, bookmark.group, existing], - )?; - bookmark.id = existing; - } else { - // 跳过已存在的书签 - continue; - } - } else { - // 生成新 ID - bookmark.id = Uuid::new_v4().to_string(); - bookmark.created_at = Utc::now(); - - conn.execute( - r#" - INSERT INTO flow_bookmarks (id, flow_id, name, group_name, created_at) - VALUES (?1, ?2, ?3, ?4, ?5) - "#, - params![ - bookmark.id, - bookmark.flow_id, - bookmark.name, - bookmark.group, - bookmark.created_at.to_rfc3339(), - ], - )?; - } - - imported.push(bookmark); - } - - Ok(imported) - } - - /// 清除所有书签(用于测试) - #[cfg(test)] - pub fn clear(&self) -> Result<()> { - let conn = self.db.lock().unwrap(); - conn.execute("DELETE FROM flow_bookmarks", [])?; - Ok(()) - } -} - -// ============================================================================ -// 测试模块 -// ============================================================================ - -#[cfg(test)] -mod tests { - use super::*; - - fn create_test_manager() -> BookmarkManager { - let conn = Connection::open_in_memory().unwrap(); - BookmarkManager::from_connection(conn).unwrap() - } - - #[test] - fn test_add_bookmark() { - let manager = create_test_manager(); - - let bookmark = manager - .add("flow-1", Some("Test Bookmark"), Some("Test Group")) - .unwrap(); - - assert!(!bookmark.id.is_empty()); - assert_eq!(bookmark.flow_id, "flow-1"); - assert_eq!(bookmark.name, Some("Test Bookmark".to_string())); - assert_eq!(bookmark.group, Some("Test Group".to_string())); - } - - #[test] - fn test_get_bookmark() { - let manager = create_test_manager(); - - let created = manager.add("flow-1", Some("Test"), None).unwrap(); - let retrieved = manager.get(&created.id).unwrap(); - - assert!(retrieved.is_some()); - let retrieved = retrieved.unwrap(); - assert_eq!(retrieved.id, created.id); - assert_eq!(retrieved.flow_id, "flow-1"); - assert_eq!(retrieved.name, Some("Test".to_string())); - } - - #[test] - fn test_get_by_flow_id() { - let manager = create_test_manager(); - - let created = manager.add("flow-1", Some("Test"), None).unwrap(); - let retrieved = manager.get_by_flow_id("flow-1").unwrap(); - - assert!(retrieved.is_some()); - let retrieved = retrieved.unwrap(); - assert_eq!(retrieved.id, created.id); - } - - #[test] - fn test_remove_bookmark() { - let manager = create_test_manager(); - - let bookmark = manager.add("flow-1", None, None).unwrap(); - manager.remove(&bookmark.id).unwrap(); - - let retrieved = manager.get(&bookmark.id).unwrap(); - assert!(retrieved.is_none()); - } - - #[test] - fn test_remove_by_flow_id() { - let manager = create_test_manager(); - - manager.add("flow-1", None, None).unwrap(); - manager.remove_by_flow_id("flow-1").unwrap(); - - let retrieved = manager.get_by_flow_id("flow-1").unwrap(); - assert!(retrieved.is_none()); - } - - #[test] - fn test_update_bookmark() { - let manager = create_test_manager(); - - let bookmark = manager.add("flow-1", Some("Original"), None).unwrap(); - - let updated = manager - .update(&bookmark.id, Some(Some("Updated")), Some(Some("New Group"))) - .unwrap(); - - assert_eq!(updated.name, Some("Updated".to_string())); - assert_eq!(updated.group, Some("New Group".to_string())); - } - - #[test] - fn test_list_bookmarks() { - let manager = create_test_manager(); - - manager.add("flow-1", None, None).unwrap(); - manager.add("flow-2", None, None).unwrap(); - manager.add("flow-3", None, None).unwrap(); - - let bookmarks = manager.list(None).unwrap(); - assert_eq!(bookmarks.len(), 3); - } - - #[test] - fn test_list_by_group() { - let manager = create_test_manager(); - - manager.add("flow-1", None, Some("Group A")).unwrap(); - manager.add("flow-2", None, Some("Group A")).unwrap(); - manager.add("flow-3", None, Some("Group B")).unwrap(); - - let group_a = manager.list(Some("Group A")).unwrap(); - assert_eq!(group_a.len(), 2); - - let group_b = manager.list(Some("Group B")).unwrap(); - assert_eq!(group_b.len(), 1); - } - - #[test] - fn test_list_groups() { - let manager = create_test_manager(); - - manager.add("flow-1", None, Some("Group A")).unwrap(); - manager.add("flow-2", None, Some("Group B")).unwrap(); - manager.add("flow-3", None, None).unwrap(); - - let groups = manager.list_groups().unwrap(); - assert_eq!(groups.len(), 2); - assert!(groups.contains(&"Group A".to_string())); - assert!(groups.contains(&"Group B".to_string())); - } - - #[test] - fn test_is_bookmarked() { - let manager = create_test_manager(); - - manager.add("flow-1", None, None).unwrap(); - - assert!(manager.is_bookmarked("flow-1").unwrap()); - assert!(!manager.is_bookmarked("flow-2").unwrap()); - } - - #[test] - fn test_count() { - let manager = create_test_manager(); - - assert_eq!(manager.count().unwrap(), 0); - - manager.add("flow-1", None, None).unwrap(); - manager.add("flow-2", None, None).unwrap(); - - assert_eq!(manager.count().unwrap(), 2); - } - - #[test] - fn test_bookmark_not_found() { - let manager = create_test_manager(); - - let result = manager.remove("non-existent"); - assert!(matches!(result, Err(BookmarkError::BookmarkNotFound(_)))); - } - - #[test] - fn test_export_import() { - let manager = create_test_manager(); - - manager - .add("flow-1", Some("Bookmark 1"), Some("Group")) - .unwrap(); - manager.add("flow-2", Some("Bookmark 2"), None).unwrap(); - - // 导出 - let exported = manager.export().unwrap(); - - // 创建新管理器并导入 - let manager2 = create_test_manager(); - let imported = manager2.import(&exported, false).unwrap(); - - assert_eq!(imported.len(), 2); - - // 验证导入的书签 - let bookmark1 = manager2.get_by_flow_id("flow-1").unwrap().unwrap(); - assert_eq!(bookmark1.name, Some("Bookmark 1".to_string())); - assert_eq!(bookmark1.group, Some("Group".to_string())); - } - - #[test] - fn test_import_overwrite() { - let manager = create_test_manager(); - - manager.add("flow-1", Some("Original"), None).unwrap(); - - // 创建导出数据 - let export_data = BookmarkExport::new(vec![FlowBookmark::new( - "flow-1", - Some("Updated".to_string()), - Some("New Group".to_string()), - )]); - let json = serde_json::to_string(&export_data).unwrap(); - - // 导入并覆盖 - manager.import(&json, true).unwrap(); - - let bookmark = manager.get_by_flow_id("flow-1").unwrap().unwrap(); - assert_eq!(bookmark.name, Some("Updated".to_string())); - assert_eq!(bookmark.group, Some("New Group".to_string())); - } - - #[test] - fn test_import_no_overwrite() { - let manager = create_test_manager(); - - manager.add("flow-1", Some("Original"), None).unwrap(); - - // 创建导出数据 - let export_data = BookmarkExport::new(vec![FlowBookmark::new( - "flow-1", - Some("Updated".to_string()), - None, - )]); - let json = serde_json::to_string(&export_data).unwrap(); - - // 导入但不覆盖 - let imported = manager.import(&json, false).unwrap(); - assert!(imported.is_empty()); - - let bookmark = manager.get_by_flow_id("flow-1").unwrap().unwrap(); - assert_eq!(bookmark.name, Some("Original".to_string())); - } -} - -// ============================================================================ -// 属性测试模块 -// ============================================================================ - -#[cfg(test)] -mod property_tests { - use super::*; - use proptest::prelude::*; - - // ======================================================================== - // 生成器 - // ======================================================================== - - /// 生成随机的 Flow ID - fn arb_flow_id() -> impl Strategy { - "[a-f0-9]{8}-[a-f0-9]{4}-[a-f0-9]{4}-[a-f0-9]{4}-[a-f0-9]{12}" - } - - /// 生成随机的书签名称 - fn arb_bookmark_name() -> impl Strategy> { - prop::option::of("[a-zA-Z0-9 _-]{1,50}") - } - - /// 生成随机的分组名称 - fn arb_group_name() -> impl Strategy> { - prop::option::of("[a-zA-Z0-9 _-]{1,30}") - } - - /// 生成随机的书签数据 - fn arb_bookmark_data() -> impl Strategy, Option)> { - (arb_flow_id(), arb_bookmark_name(), arb_group_name()) - } - - /// 生成多个书签数据 - fn arb_bookmarks( - max_len: usize, - ) -> impl Strategy, Option)>> { - prop::collection::vec(arb_bookmark_data(), 1..max_len) - } - - // ======================================================================== - // 属性测试 - // ======================================================================== - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: flow-monitor-enhancement, Property 15: 书签 Round-Trip** - /// **Validates: Requirements 8.1** - /// - /// *对于任意* 书签操作,添加后应该能够正确检索到该书签。 - #[test] - fn prop_bookmark_roundtrip( - (flow_id, name, group) in arb_bookmark_data() - ) { - let manager = create_test_manager(); - - // 添加书签 - let added = manager.add(&flow_id, name.as_deref(), group.as_deref()).unwrap(); - - // 通过 ID 检索 - let retrieved_by_id = manager.get(&added.id).unwrap().unwrap(); - prop_assert_eq!(&added.id, &retrieved_by_id.id); - prop_assert_eq!(&added.flow_id, &retrieved_by_id.flow_id); - prop_assert_eq!(&added.name, &retrieved_by_id.name); - prop_assert_eq!(&added.group, &retrieved_by_id.group); - - // 通过 Flow ID 检索 - let retrieved_by_flow = manager.get_by_flow_id(&flow_id).unwrap().unwrap(); - prop_assert_eq!(&added.id, &retrieved_by_flow.id); - prop_assert_eq!(&added.flow_id, &retrieved_by_flow.flow_id); - } - - /// **Feature: flow-monitor-enhancement, Property 16: 书签导入导出 Round-Trip** - /// **Validates: Requirements 8.6** - /// - /// *对于任意* 书签集合,导出后再导入应该得到等价的集合。 - #[test] - fn prop_bookmark_export_import_roundtrip( - bookmarks in arb_bookmarks(10) - ) { - let manager1 = create_test_manager(); - - // 添加所有书签(使用唯一的 flow_id) - let mut added_bookmarks = Vec::new(); - for (i, (flow_id, name, group)) in bookmarks.iter().enumerate() { - // 确保 flow_id 唯一 - let unique_flow_id = format!("{flow_id}_{i}"); - let bookmark = manager1.add(&unique_flow_id, name.as_deref(), group.as_deref()).unwrap(); - added_bookmarks.push(bookmark); - } - - // 导出 - let exported = manager1.export().unwrap(); - - // 创建新管理器并导入 - let manager2 = create_test_manager(); - let imported = manager2.import(&exported, false).unwrap(); - - // 验证导入数量 - prop_assert_eq!(imported.len(), added_bookmarks.len()); - - // 验证每个书签的内容 - for added in &added_bookmarks { - let found = manager2.get_by_flow_id(&added.flow_id).unwrap(); - prop_assert!(found.is_some(), "Bookmark for flow '{}' should be imported", added.flow_id); - - let found = found.unwrap(); - prop_assert_eq!(&added.flow_id, &found.flow_id); - prop_assert_eq!(&added.name, &found.name); - prop_assert_eq!(&added.group, &found.group); - } - } - - /// 书签删除后应该不存在 - #[test] - fn prop_bookmark_delete( - (flow_id, name, group) in arb_bookmark_data() - ) { - let manager = create_test_manager(); - - // 添加书签 - let bookmark = manager.add(&flow_id, name.as_deref(), group.as_deref()).unwrap(); - - // 删除书签 - manager.remove(&bookmark.id).unwrap(); - - // 验证不存在 - let found = manager.get(&bookmark.id).unwrap(); - prop_assert!(found.is_none()); - - let found_by_flow = manager.get_by_flow_id(&flow_id).unwrap(); - prop_assert!(found_by_flow.is_none()); - } - - /// 书签更新后应该保持一致性 - #[test] - fn prop_bookmark_update_consistency( - (flow_id, name, group) in arb_bookmark_data(), - (_, new_name, new_group) in arb_bookmark_data() - ) { - let manager = create_test_manager(); - - // 添加书签 - let original = manager.add(&flow_id, name.as_deref(), group.as_deref()).unwrap(); - - // 更新书签 - let updated = manager.update( - &original.id, - Some(new_name.as_deref()), - Some(new_group.as_deref()), - ).unwrap(); - - // 验证更新后的值 - prop_assert_eq!(updated.id, original.id); - prop_assert_eq!(updated.flow_id, original.flow_id); - prop_assert_eq!(updated.name, new_name); - prop_assert_eq!(updated.group, new_group); - } - - /// 列表应该包含所有添加的书签 - #[test] - fn prop_list_contains_all( - bookmarks in arb_bookmarks(5) - ) { - let manager = create_test_manager(); - - // 添加所有书签 - let mut added_ids = Vec::new(); - for (i, (flow_id, name, group)) in bookmarks.iter().enumerate() { - let unique_flow_id = format!("{flow_id}_{i}"); - let bookmark = manager.add(&unique_flow_id, name.as_deref(), group.as_deref()).unwrap(); - added_ids.push(bookmark.id); - } - - // 获取列表 - let list = manager.list(None).unwrap(); - - // 验证所有添加的书签都在列表中 - for id in &added_ids { - prop_assert!( - list.iter().any(|b| &b.id == id), - "Bookmark with id '{}' should be in list", - id - ); - } - } - - /// is_bookmarked 应该正确反映书签状态 - #[test] - fn prop_is_bookmarked_consistency( - (flow_id, name, group) in arb_bookmark_data() - ) { - let manager = create_test_manager(); - - // 初始状态:未添加书签 - prop_assert!(!manager.is_bookmarked(&flow_id).unwrap()); - - // 添加书签 - let bookmark = manager.add(&flow_id, name.as_deref(), group.as_deref()).unwrap(); - prop_assert!(manager.is_bookmarked(&flow_id).unwrap()); - - // 删除书签 - manager.remove(&bookmark.id).unwrap(); - prop_assert!(!manager.is_bookmarked(&flow_id).unwrap()); - } - } - - fn create_test_manager() -> BookmarkManager { - let conn = Connection::open_in_memory().unwrap(); - BookmarkManager::from_connection(conn).unwrap() - } -} diff --git a/src-tauri/src/flow_monitor/code_exporter.rs b/src-tauri/src/flow_monitor/code_exporter.rs deleted file mode 100644 index 16ea433a9..000000000 --- a/src-tauri/src/flow_monitor/code_exporter.rs +++ /dev/null @@ -1,1050 +0,0 @@ -//! 代码导出器 -//! -//! 提供将 LLM Flow 导出为可执行代码的功能,支持 curl、Python、TypeScript 等格式。 -//! -//! **Validates: Requirements 7.7, 7.8** - -use serde::{Deserialize, Serialize}; - -use super::models::{LLMFlow, LLMRequest}; - -// ============================================================================ -// 代码导出格式枚举 -// ============================================================================ - -/// 代码导出格式 -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] -#[serde(rename_all = "lowercase")] -#[derive(Default)] -pub enum CodeFormat { - /// curl 命令 - #[default] - Curl, - /// Python 代码 - Python, - /// TypeScript 代码 - TypeScript, - /// JavaScript 代码 - JavaScript, -} - -// ============================================================================ -// 代码导出器 -// ============================================================================ - -/// 代码导出器 -/// -/// 将 LLM Flow 导出为可执行的代码格式。 -pub struct CodeExporter; - -impl CodeExporter { - /// 导出为指定格式的代码 - /// - /// # Arguments - /// * `flow` - 要导出的 Flow - /// * `format` - 导出格式 - /// - /// # Returns - /// 导出的代码字符串 - pub fn export(flow: &LLMFlow, format: CodeFormat) -> String { - match format { - CodeFormat::Curl => Self::to_curl(flow), - CodeFormat::Python => Self::to_python(flow), - CodeFormat::TypeScript => Self::to_typescript(flow), - CodeFormat::JavaScript => Self::to_javascript(flow), - } - } - - /// 导出为 curl 命令 - /// - /// **Validates: Requirements 7.7** - /// - /// # Arguments - /// * `flow` - 要导出的 Flow - /// - /// # Returns - /// curl 命令字符串 - pub fn to_curl(flow: &LLMFlow) -> String { - Self::request_to_curl( - &flow.request, - flow.metadata.routing_info.target_url.as_deref(), - ) - } - - /// 将请求转换为 curl 命令 - pub fn request_to_curl(request: &LLMRequest, base_url: Option<&str>) -> String { - let mut parts = vec!["curl".to_string()]; - - // 添加方法(如果不是 GET) - if request.method != "GET" { - parts.push(format!("-X {}", request.method)); - } - - // 构建 URL - let url = if let Some(base) = base_url { - format!("{}{}", base.trim_end_matches('/'), request.path) - } else { - format!("http://localhost{}", request.path) - }; - parts.push(format!("'{url}'")); - - // 添加请求头 - for (key, value) in &request.headers { - // 跳过敏感头部或使用占位符 - let header_value = if key.to_lowercase() == "authorization" { - "$API_KEY".to_string() - } else if key.to_lowercase() == "x-api-key" { - "$API_KEY".to_string() - } else { - escape_shell_string(value) - }; - parts.push(format!("-H '{key}: {header_value}'")); - } - - // 确保有 Content-Type 头 - if !request - .headers - .keys() - .any(|k| k.to_lowercase() == "content-type") - { - parts.push("-H 'Content-Type: application/json'".to_string()); - } - - // 添加请求体 - if !request.body.is_null() { - let body_str = serde_json::to_string(&request.body).unwrap_or_default(); - parts.push(format!("-d '{}'", escape_shell_string(&body_str))); - } - - parts.join(" \\\n ") - } - - /// 导出为 Python 代码 - /// - /// **Validates: Requirements 7.8** - /// - /// # Arguments - /// * `flow` - 要导出的 Flow - /// - /// # Returns - /// Python 代码字符串 - pub fn to_python(flow: &LLMFlow) -> String { - Self::request_to_python( - &flow.request, - flow.metadata.routing_info.target_url.as_deref(), - ) - } - - /// 将请求转换为 Python 代码 - pub fn request_to_python(request: &LLMRequest, base_url: Option<&str>) -> String { - let mut code = String::new(); - - // 导入语句 - code.push_str("import requests\n"); - code.push_str("import json\n\n"); - - // URL - let url = if let Some(base) = base_url { - format!("{}{}", base.trim_end_matches('/'), request.path) - } else { - format!("http://localhost{}", request.path) - }; - code.push_str(&format!("url = \"{url}\"\n\n")); - - // 请求头 - code.push_str("headers = {\n"); - let mut has_content_type = false; - for (key, value) in &request.headers { - if key.to_lowercase() == "content-type" { - has_content_type = true; - } - let header_value = if key.to_lowercase() == "authorization" { - "os.environ.get('API_KEY', '')".to_string() - } else if key.to_lowercase() == "x-api-key" { - "os.environ.get('API_KEY', '')".to_string() - } else { - format!("\"{}\"", escape_python_string(value)) - }; - - if key.to_lowercase() == "authorization" || key.to_lowercase() == "x-api-key" { - code.push_str(&format!(" \"{key}\": {header_value},\n")); - } else { - code.push_str(&format!(" \"{key}\": {header_value},\n")); - } - } - if !has_content_type { - code.push_str(" \"Content-Type\": \"application/json\",\n"); - } - code.push_str("}\n\n"); - - // 请求体 - if !request.body.is_null() { - let body_str = serde_json::to_string_pretty(&request.body).unwrap_or_default(); - code.push_str(&format!("data = {body_str}\n\n")); - } else { - code.push_str("data = {}\n\n"); - } - - // 发送请求 - code.push_str(&format!( - "response = requests.{}(\n url,\n headers=headers,\n json=data\n)\n\n", - request.method.to_lowercase() - )); - - // 处理响应 - code.push_str("# 检查响应状态\n"); - code.push_str("response.raise_for_status()\n\n"); - code.push_str("# 解析响应\n"); - code.push_str("result = response.json()\n"); - code.push_str("print(json.dumps(result, indent=2, ensure_ascii=False))\n"); - - code - } - - /// 导出为 TypeScript 代码 - /// - /// **Validates: Requirements 7.8** - /// - /// # Arguments - /// * `flow` - 要导出的 Flow - /// - /// # Returns - /// TypeScript 代码字符串 - pub fn to_typescript(flow: &LLMFlow) -> String { - Self::request_to_typescript( - &flow.request, - flow.metadata.routing_info.target_url.as_deref(), - ) - } - - /// 将请求转换为 TypeScript 代码 - pub fn request_to_typescript(request: &LLMRequest, base_url: Option<&str>) -> String { - let mut code = String::new(); - - // URL - let url = if let Some(base) = base_url { - format!("{}{}", base.trim_end_matches('/'), request.path) - } else { - format!("http://localhost{}", request.path) - }; - - code.push_str("const url = '"); - code.push_str(&url); - code.push_str("';\n\n"); - - // 请求头 - code.push_str("const headers: Record = {\n"); - let mut has_content_type = false; - for (key, value) in &request.headers { - if key.to_lowercase() == "content-type" { - has_content_type = true; - } - let header_value = if key.to_lowercase() == "authorization" { - "process.env.API_KEY || ''".to_string() - } else if key.to_lowercase() == "x-api-key" { - "process.env.API_KEY || ''".to_string() - } else { - format!("'{}'", escape_js_string(value)) - }; - code.push_str(&format!(" '{key}': {header_value},\n")); - } - if !has_content_type { - code.push_str(" 'Content-Type': 'application/json',\n"); - } - code.push_str("};\n\n"); - - // 请求体 - if !request.body.is_null() { - let body_str = serde_json::to_string_pretty(&request.body).unwrap_or_default(); - code.push_str("const data = "); - code.push_str(&body_str); - code.push_str(";\n\n"); - } else { - code.push_str("const data = {};\n\n"); - } - - // 发送请求(使用 async/await) - code.push_str("async function makeRequest(): Promise {\n"); - code.push_str(" const response = await fetch(url, {\n"); - code.push_str(&format!(" method: '{}',\n", request.method)); - code.push_str(" headers,\n"); - code.push_str(" body: JSON.stringify(data),\n"); - code.push_str(" });\n\n"); - code.push_str(" if (!response.ok) {\n"); - code.push_str(" throw new Error(`HTTP error! status: ${response.status}`);\n"); - code.push_str(" }\n\n"); - code.push_str(" const result = await response.json();\n"); - code.push_str(" console.log(JSON.stringify(result, null, 2));\n"); - code.push_str("}\n\n"); - code.push_str("makeRequest().catch(console.error);\n"); - - code - } - - /// 导出为 JavaScript 代码 - /// - /// **Validates: Requirements 7.8** - /// - /// # Arguments - /// * `flow` - 要导出的 Flow - /// - /// # Returns - /// JavaScript 代码字符串 - pub fn to_javascript(flow: &LLMFlow) -> String { - Self::request_to_javascript( - &flow.request, - flow.metadata.routing_info.target_url.as_deref(), - ) - } - - /// 将请求转换为 JavaScript 代码 - pub fn request_to_javascript(request: &LLMRequest, base_url: Option<&str>) -> String { - let mut code = String::new(); - - // URL - let url = if let Some(base) = base_url { - format!("{}{}", base.trim_end_matches('/'), request.path) - } else { - format!("http://localhost{}", request.path) - }; - - code.push_str("const url = '"); - code.push_str(&url); - code.push_str("';\n\n"); - - // 请求头 - code.push_str("const headers = {\n"); - let mut has_content_type = false; - for (key, value) in &request.headers { - if key.to_lowercase() == "content-type" { - has_content_type = true; - } - let header_value = if key.to_lowercase() == "authorization" { - "process.env.API_KEY || ''".to_string() - } else if key.to_lowercase() == "x-api-key" { - "process.env.API_KEY || ''".to_string() - } else { - format!("'{}'", escape_js_string(value)) - }; - code.push_str(&format!(" '{key}': {header_value},\n")); - } - if !has_content_type { - code.push_str(" 'Content-Type': 'application/json',\n"); - } - code.push_str("};\n\n"); - - // 请求体 - if !request.body.is_null() { - let body_str = serde_json::to_string_pretty(&request.body).unwrap_or_default(); - code.push_str("const data = "); - code.push_str(&body_str); - code.push_str(";\n\n"); - } else { - code.push_str("const data = {};\n\n"); - } - - // 发送请求(使用 async/await) - code.push_str("async function makeRequest() {\n"); - code.push_str(" const response = await fetch(url, {\n"); - code.push_str(&format!(" method: '{}',\n", request.method)); - code.push_str(" headers,\n"); - code.push_str(" body: JSON.stringify(data),\n"); - code.push_str(" });\n\n"); - code.push_str(" if (!response.ok) {\n"); - code.push_str(" throw new Error(`HTTP error! status: ${response.status}`);\n"); - code.push_str(" }\n\n"); - code.push_str(" const result = await response.json();\n"); - code.push_str(" console.log(JSON.stringify(result, null, 2));\n"); - code.push_str("}\n\n"); - code.push_str("makeRequest().catch(console.error);\n"); - - code - } -} - -// ============================================================================ -// 辅助函数 -// ============================================================================ - -/// 转义 shell 字符串中的特殊字符 -fn escape_shell_string(s: &str) -> String { - s.replace('\\', "\\\\").replace('\'', "'\\''") -} - -/// 转义 Python 字符串中的特殊字符 -fn escape_python_string(s: &str) -> String { - s.replace('\\', "\\\\") - .replace('"', "\\\"") - .replace('\n', "\\n") - .replace('\r', "\\r") - .replace('\t', "\\t") -} - -/// 转义 JavaScript 字符串中的特殊字符 -fn escape_js_string(s: &str) -> String { - s.replace('\\', "\\\\") - .replace('\'', "\\'") - .replace('\n', "\\n") - .replace('\r', "\\r") - .replace('\t', "\\t") -} - -// ============================================================================ -// 测试模块 -// ============================================================================ - -#[cfg(test)] -mod tests { - use super::*; - use crate::flow_monitor::{ - FlowAnnotations, FlowMetadata, FlowState, FlowTimestamps, FlowType, Message, - MessageContent, MessageRole, RequestParameters, RoutingInfo, - }; - use crate::ProviderType; - use chrono::Utc; - use std::collections::HashMap; - - fn create_test_flow() -> LLMFlow { - let mut headers = HashMap::new(); - headers.insert("Content-Type".to_string(), "application/json".to_string()); - headers.insert( - "Authorization".to_string(), - "Bearer sk-test-key".to_string(), - ); - - let body = serde_json::json!({ - "model": "gpt-4", - "messages": [ - {"role": "user", "content": "Hello, world!"} - ], - "temperature": 0.7 - }); - - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - headers, - body, - messages: vec![Message { - role: MessageRole::User, - content: MessageContent::Text("Hello, world!".to_string()), - tool_calls: None, - tool_result: None, - name: None, - }], - system_prompt: None, - tools: None, - model: "gpt-4".to_string(), - original_model: None, - parameters: RequestParameters { - temperature: Some(0.7), - ..Default::default() - }, - size_bytes: 100, - timestamp: Utc::now(), - }; - - let mut metadata = FlowMetadata::default(); - metadata.provider = ProviderType::OpenAI; - metadata.routing_info = RoutingInfo { - target_url: Some("https://api.openai.com".to_string()), - route_rule: None, - load_balance_strategy: None, - }; - - LLMFlow { - id: "test-flow-id".to_string(), - flow_type: FlowType::ChatCompletions, - request, - response: None, - error: None, - metadata, - timestamps: FlowTimestamps::default(), - state: FlowState::Pending, - annotations: FlowAnnotations::default(), - } - } - - #[test] - fn test_to_curl() { - let flow = create_test_flow(); - let curl = CodeExporter::to_curl(&flow); - - // 验证 curl 命令包含必要的部分 - assert!(curl.contains("curl")); - assert!(curl.contains("-X POST")); - assert!(curl.contains("https://api.openai.com/v1/chat/completions")); - assert!(curl.contains("-H 'Content-Type: application/json'")); - assert!(curl.contains("-H 'Authorization: $API_KEY'")); - assert!(curl.contains("-d '")); - assert!(curl.contains("gpt-4")); - } - - #[test] - fn test_to_python() { - let flow = create_test_flow(); - let python = CodeExporter::to_python(&flow); - - // 验证 Python 代码包含必要的部分 - assert!(python.contains("import requests")); - assert!(python.contains("import json")); - assert!(python.contains("url = \"https://api.openai.com/v1/chat/completions\"")); - assert!(python.contains("headers = {")); - assert!(python.contains("\"Content-Type\": \"application/json\"")); - assert!(python.contains("data = {")); - assert!(python.contains("requests.post(")); - assert!(python.contains("response.raise_for_status()")); - assert!(python.contains("response.json()")); - } - - #[test] - fn test_to_typescript() { - let flow = create_test_flow(); - let typescript = CodeExporter::to_typescript(&flow); - - // 验证 TypeScript 代码包含必要的部分 - assert!(typescript.contains("const url = 'https://api.openai.com/v1/chat/completions'")); - assert!(typescript.contains("const headers: Record = {")); - assert!(typescript.contains("'Content-Type': 'application/json'")); - assert!(typescript.contains("const data = {")); - assert!(typescript.contains("async function makeRequest(): Promise")); - assert!(typescript.contains("await fetch(url")); - assert!(typescript.contains("method: 'POST'")); - assert!(typescript.contains("await response.json()")); - } - - #[test] - fn test_to_javascript() { - let flow = create_test_flow(); - let javascript = CodeExporter::to_javascript(&flow); - - // 验证 JavaScript 代码包含必要的部分 - assert!(javascript.contains("const url = 'https://api.openai.com/v1/chat/completions'")); - assert!(javascript.contains("const headers = {")); - assert!(javascript.contains("'Content-Type': 'application/json'")); - assert!(javascript.contains("const data = {")); - assert!(javascript.contains("async function makeRequest()")); - assert!(javascript.contains("await fetch(url")); - assert!(javascript.contains("method: 'POST'")); - assert!(javascript.contains("await response.json()")); - // TypeScript 和 JavaScript 的区别 - assert!(!javascript.contains(": Record")); - assert!(!javascript.contains(": Promise")); - } - - #[test] - fn test_export_with_format() { - let flow = create_test_flow(); - - let curl = CodeExporter::export(&flow, CodeFormat::Curl); - assert!(curl.contains("curl")); - - let python = CodeExporter::export(&flow, CodeFormat::Python); - assert!(python.contains("import requests")); - - let typescript = CodeExporter::export(&flow, CodeFormat::TypeScript); - assert!(typescript.contains("Record")); - - let javascript = CodeExporter::export(&flow, CodeFormat::JavaScript); - assert!(!javascript.contains("Record")); - } - - #[test] - fn test_escape_shell_string() { - assert_eq!(escape_shell_string("hello"), "hello"); - assert_eq!(escape_shell_string("it's"), "it'\\''s"); - assert_eq!(escape_shell_string("back\\slash"), "back\\\\slash"); - } - - #[test] - fn test_escape_python_string() { - assert_eq!(escape_python_string("hello"), "hello"); - assert_eq!(escape_python_string("say \"hi\""), "say \\\"hi\\\""); - assert_eq!(escape_python_string("line1\nline2"), "line1\\nline2"); - } - - #[test] - fn test_escape_js_string() { - assert_eq!(escape_js_string("hello"), "hello"); - assert_eq!(escape_js_string("it's"), "it\\'s"); - assert_eq!(escape_js_string("line1\nline2"), "line1\\nline2"); - } - - #[test] - fn test_curl_without_base_url() { - let mut flow = create_test_flow(); - flow.metadata.routing_info.target_url = None; - let curl = CodeExporter::to_curl(&flow); - - assert!(curl.contains("http://localhost/v1/chat/completions")); - } - - #[test] - fn test_api_key_placeholder() { - let flow = create_test_flow(); - - let curl = CodeExporter::to_curl(&flow); - assert!(curl.contains("$API_KEY")); - assert!(!curl.contains("sk-test-key")); - - let python = CodeExporter::to_python(&flow); - assert!(python.contains("os.environ.get('API_KEY'")); - assert!(!python.contains("sk-test-key")); - - let typescript = CodeExporter::to_typescript(&flow); - assert!(typescript.contains("process.env.API_KEY")); - assert!(!typescript.contains("sk-test-key")); - } -} - -// ============================================================================ -// 属性测试模块 -// ============================================================================ - -#[cfg(test)] -mod property_tests { - use super::*; - use crate::flow_monitor::{ - FlowAnnotations, FlowMetadata, FlowState, FlowTimestamps, FlowType, RequestParameters, - RoutingInfo, - }; - use crate::ProviderType; - use chrono::Utc; - use proptest::prelude::*; - use std::collections::HashMap; - - // ======================================================================== - // 生成器 - // ======================================================================== - - /// 生成随机的 HTTP 方法 - fn arb_http_method() -> impl Strategy { - prop_oneof![ - Just("GET".to_string()), - Just("POST".to_string()), - Just("PUT".to_string()), - Just("DELETE".to_string()), - Just("PATCH".to_string()), - ] - } - - /// 生成随机的 API 路径 - fn arb_api_path() -> impl Strategy { - prop_oneof![ - Just("/v1/chat/completions".to_string()), - Just("/v1/completions".to_string()), - Just("/v1/embeddings".to_string()), - Just("/v1/messages".to_string()), - Just("/api/generate".to_string()), - ] - } - - /// 生成随机的模型名称 - fn arb_model_name() -> impl Strategy { - prop_oneof![ - Just("gpt-4".to_string()), - Just("gpt-3.5-turbo".to_string()), - Just("claude-3-opus".to_string()), - Just("claude-3-sonnet".to_string()), - Just("gemini-pro".to_string()), - "[a-z]{3,10}-[0-9]{1,2}".prop_map(|s| s), - ] - } - - /// 生成随机的 URL - fn arb_base_url() -> impl Strategy> { - prop_oneof![ - Just(None), - Just(Some("https://api.openai.com".to_string())), - Just(Some("https://api.anthropic.com".to_string())), - Just(Some("http://localhost:8080".to_string())), - ] - } - - /// 生成随机的请求头 - fn arb_headers() -> impl Strategy> { - prop::collection::hash_map( - prop_oneof![ - Just("Content-Type".to_string()), - Just("Authorization".to_string()), - Just("X-Api-Key".to_string()), - Just("User-Agent".to_string()), - ], - "[a-zA-Z0-9-/]{5,30}", - 0..4, - ) - } - - /// 生成随机的请求体 - fn arb_request_body() -> impl Strategy { - prop_oneof![ - Just(serde_json::json!({})), - Just(serde_json::json!({"model": "gpt-4"})), - Just(serde_json::json!({ - "model": "gpt-4", - "messages": [{"role": "user", "content": "Hello"}] - })), - Just(serde_json::json!({ - "model": "claude-3", - "messages": [{"role": "user", "content": "Test"}], - "temperature": 0.7 - })), - ] - } - - /// 生成随机的 LLMRequest - fn arb_llm_request() -> impl Strategy { - ( - arb_http_method(), - arb_api_path(), - arb_model_name(), - arb_headers(), - arb_request_body(), - ) - .prop_map(|(method, path, model, headers, body)| LLMRequest { - method, - path, - headers, - body, - messages: vec![], - system_prompt: None, - tools: None, - model, - original_model: None, - parameters: RequestParameters::default(), - size_bytes: 0, - timestamp: Utc::now(), - }) - } - - /// 生成随机的 LLMFlow - fn arb_llm_flow() -> impl Strategy { - (arb_llm_request(), arb_base_url()).prop_map(|(request, base_url)| { - let mut metadata = FlowMetadata::default(); - metadata.provider = ProviderType::OpenAI; - metadata.routing_info = RoutingInfo { - target_url: base_url, - route_rule: None, - load_balance_strategy: None, - }; - - LLMFlow { - id: uuid::Uuid::new_v4().to_string(), - flow_type: FlowType::ChatCompletions, - request, - response: None, - error: None, - metadata, - timestamps: FlowTimestamps::default(), - state: FlowState::Pending, - annotations: FlowAnnotations::default(), - } - }) - } - - // ======================================================================== - // 属性测试 - // ======================================================================== - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: flow-monitor-enhancement, Property 13: curl 命令正确性** - /// **Validates: Requirements 7.7** - /// - /// *对于任意* 有效的 LLM Flow,生成的 curl 命令应该包含正确的 HTTP 方法、URL 和请求体。 - #[test] - fn prop_curl_command_correctness(flow in arb_llm_flow()) { - let curl = CodeExporter::to_curl(&flow); - - // 验证 curl 命令以 "curl" 开头 - prop_assert!(curl.starts_with("curl"), "curl 命令应该以 'curl' 开头"); - - // 验证包含正确的 HTTP 方法(如果不是 GET) - if flow.request.method != "GET" { - prop_assert!( - curl.contains(&format!("-X {}", flow.request.method)), - "curl 命令应该包含正确的 HTTP 方法: {}", - flow.request.method - ); - } - - // 验证包含 URL - let expected_path = &flow.request.path; - prop_assert!( - curl.contains(expected_path), - "curl 命令应该包含请求路径: {}", - expected_path - ); - - // 验证包含请求体(如果有) - if !flow.request.body.is_null() { - prop_assert!( - curl.contains("-d '"), - "curl 命令应该包含请求体" - ); - } - - // 验证敏感信息被替换 - prop_assert!( - !curl.contains("Bearer sk-") && !curl.contains("sk-ant-"), - "curl 命令不应该包含真实的 API 密钥" - ); - } - - /// **Feature: flow-monitor-enhancement, Property 13: curl 命令 URL 正确性** - /// **Validates: Requirements 7.7** - /// - /// *对于任意* 有效的 LLM Flow,生成的 curl 命令应该包含正确构建的 URL。 - #[test] - fn prop_curl_url_correctness(flow in arb_llm_flow()) { - let curl = CodeExporter::to_curl(&flow); - - // 构建预期的 URL - let expected_url = if let Some(ref base) = flow.metadata.routing_info.target_url { - format!("{}{}", base.trim_end_matches('/'), flow.request.path) - } else { - format!("http://localhost{}", flow.request.path) - }; - - prop_assert!( - curl.contains(&expected_url), - "curl 命令应该包含正确的 URL: {}, 实际: {}", - expected_url, - curl - ); - } - - /// **Feature: flow-monitor-enhancement, Property 14: Python 代码生成正确性** - /// **Validates: Requirements 7.8** - /// - /// *对于任意* 有效的 LLM Flow,生成的 Python 代码应该是语法正确的。 - #[test] - fn prop_python_code_correctness(flow in arb_llm_flow()) { - let python = CodeExporter::to_python(&flow); - - // 验证包含必要的导入语句 - prop_assert!( - python.contains("import requests"), - "Python 代码应该包含 'import requests'" - ); - prop_assert!( - python.contains("import json"), - "Python 代码应该包含 'import json'" - ); - - // 验证包含 URL 定义 - prop_assert!( - python.contains("url = \""), - "Python 代码应该包含 URL 定义" - ); - - // 验证包含 headers 定义 - prop_assert!( - python.contains("headers = {"), - "Python 代码应该包含 headers 定义" - ); - - // 验证包含 data 定义 - prop_assert!( - python.contains("data = "), - "Python 代码应该包含 data 定义" - ); - - // 验证包含 requests 调用 - prop_assert!( - python.contains(&format!("requests.{}(", flow.request.method.to_lowercase())), - "Python 代码应该包含正确的 requests 方法调用" - ); - - // 验证包含响应处理 - prop_assert!( - python.contains("response.raise_for_status()"), - "Python 代码应该包含错误处理" - ); - prop_assert!( - python.contains("response.json()"), - "Python 代码应该包含 JSON 解析" - ); - - // 验证敏感信息被替换 - prop_assert!( - !python.contains("Bearer sk-") && !python.contains("sk-ant-"), - "Python 代码不应该包含真实的 API 密钥" - ); - - // 验证基本的 Python 语法结构 - // 检查括号匹配 - let open_parens = python.matches('(').count(); - let close_parens = python.matches(')').count(); - prop_assert_eq!( - open_parens, close_parens, - "Python 代码的括号应该匹配" - ); - - // 检查花括号匹配 - let open_braces = python.matches('{').count(); - let close_braces = python.matches('}').count(); - prop_assert_eq!( - open_braces, close_braces, - "Python 代码的花括号应该匹配" - ); - } - - /// **Feature: flow-monitor-enhancement, Property 14: TypeScript 代码生成正确性** - /// **Validates: Requirements 7.8** - /// - /// *对于任意* 有效的 LLM Flow,生成的 TypeScript 代码应该是语法正确的。 - #[test] - fn prop_typescript_code_correctness(flow in arb_llm_flow()) { - let typescript = CodeExporter::to_typescript(&flow); - - // 验证包含 URL 定义 - prop_assert!( - typescript.contains("const url = '"), - "TypeScript 代码应该包含 URL 定义" - ); - - // 验证包含 headers 定义(带类型注解) - prop_assert!( - typescript.contains("const headers: Record = {"), - "TypeScript 代码应该包含带类型注解的 headers 定义" - ); - - // 验证包含 data 定义 - prop_assert!( - typescript.contains("const data = "), - "TypeScript 代码应该包含 data 定义" - ); - - // 验证包含 async 函数定义(带返回类型) - prop_assert!( - typescript.contains("async function makeRequest(): Promise"), - "TypeScript 代码应该包含带返回类型的 async 函数" - ); - - // 验证包含 fetch 调用 - prop_assert!( - typescript.contains("await fetch(url"), - "TypeScript 代码应该包含 fetch 调用" - ); - - // 验证包含正确的 HTTP 方法 - prop_assert!( - typescript.contains(&format!("method: '{}'", flow.request.method)), - "TypeScript 代码应该包含正确的 HTTP 方法" - ); - - // 验证包含错误处理 - prop_assert!( - typescript.contains("if (!response.ok)"), - "TypeScript 代码应该包含错误处理" - ); - - // 验证包含 JSON 解析 - prop_assert!( - typescript.contains("await response.json()"), - "TypeScript 代码应该包含 JSON 解析" - ); - - // 验证敏感信息被替换 - prop_assert!( - !typescript.contains("Bearer sk-") && !typescript.contains("sk-ant-"), - "TypeScript 代码不应该包含真实的 API 密钥" - ); - - // 验证基本的语法结构 - // 检查括号匹配 - let open_parens = typescript.matches('(').count(); - let close_parens = typescript.matches(')').count(); - prop_assert_eq!( - open_parens, close_parens, - "TypeScript 代码的括号应该匹配" - ); - - // 检查花括号匹配 - let open_braces = typescript.matches('{').count(); - let close_braces = typescript.matches('}').count(); - prop_assert_eq!( - open_braces, close_braces, - "TypeScript 代码的花括号应该匹配" - ); - } - - /// **Feature: flow-monitor-enhancement, Property 14: JavaScript 代码生成正确性** - /// **Validates: Requirements 7.8** - /// - /// *对于任意* 有效的 LLM Flow,生成的 JavaScript 代码应该是语法正确的,且不包含 TypeScript 类型注解。 - #[test] - fn prop_javascript_code_correctness(flow in arb_llm_flow()) { - let javascript = CodeExporter::to_javascript(&flow); - - // 验证包含 URL 定义 - prop_assert!( - javascript.contains("const url = '"), - "JavaScript 代码应该包含 URL 定义" - ); - - // 验证包含 headers 定义(不带类型注解) - prop_assert!( - javascript.contains("const headers = {"), - "JavaScript 代码应该包含 headers 定义" - ); - prop_assert!( - !javascript.contains("Record"), - "JavaScript 代码不应该包含 TypeScript 类型注解" - ); - - // 验证包含 data 定义 - prop_assert!( - javascript.contains("const data = "), - "JavaScript 代码应该包含 data 定义" - ); - - // 验证包含 async 函数定义(不带返回类型) - prop_assert!( - javascript.contains("async function makeRequest()"), - "JavaScript 代码应该包含 async 函数" - ); - prop_assert!( - !javascript.contains(": Promise"), - "JavaScript 代码不应该包含 TypeScript 返回类型" - ); - - // 验证包含 fetch 调用 - prop_assert!( - javascript.contains("await fetch(url"), - "JavaScript 代码应该包含 fetch 调用" - ); - - // 验证包含正确的 HTTP 方法 - prop_assert!( - javascript.contains(&format!("method: '{}'", flow.request.method)), - "JavaScript 代码应该包含正确的 HTTP 方法" - ); - - // 验证敏感信息被替换 - prop_assert!( - !javascript.contains("Bearer sk-") && !javascript.contains("sk-ant-"), - "JavaScript 代码不应该包含真实的 API 密钥" - ); - - // 验证基本的语法结构 - // 检查括号匹配 - let open_parens = javascript.matches('(').count(); - let close_parens = javascript.matches(')').count(); - prop_assert_eq!( - open_parens, close_parens, - "JavaScript 代码的括号应该匹配" - ); - - // 检查花括号匹配 - let open_braces = javascript.matches('{').count(); - let close_braces = javascript.matches('}').count(); - prop_assert_eq!( - open_braces, close_braces, - "JavaScript 代码的花括号应该匹配" - ); - } - } -} diff --git a/src-tauri/src/flow_monitor/diff.rs b/src-tauri/src/flow_monitor/diff.rs deleted file mode 100644 index 649ea61de..000000000 --- a/src-tauri/src/flow_monitor/diff.rs +++ /dev/null @@ -1,1593 +0,0 @@ -//! Flow 差异对比模块 -//! -//! 该模块实现两个 LLM Flow 之间的差异对比功能,支持请求、响应、元数据和 Token 使用量的对比。 -//! -//! # 主要功能 -//! -//! - 对比两个 Flow 的请求差异 -//! - 对比两个 Flow 的响应差异 -//! - 对比消息列表的差异 -//! - 计算 Token 使用量差异 -//! - 支持忽略动态字段(时间戳、ID 等) - -use serde::{Deserialize, Serialize}; -use serde_json::Value; - -use super::models::{LLMFlow, Message, MessageContent, TokenUsage}; - -// ============================================================================ -// 差异类型 -// ============================================================================ - -/// 差异类型 -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] -pub enum DiffType { - /// 新增 - Added, - /// 删除 - Removed, - /// 修改 - Modified, - /// 未变化 - #[default] - Unchanged, -} - -// ============================================================================ -// 差异项 -// ============================================================================ - -/// 差异项 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct DiffItem { - /// 字段路径 - pub path: String, - /// 差异类型 - pub diff_type: DiffType, - /// 左侧值(原始) - pub left_value: Option, - /// 右侧值(对比) - pub right_value: Option, -} - -impl DiffItem { - /// 创建新增差异项 - pub fn added(path: impl Into, value: Value) -> Self { - Self { - path: path.into(), - diff_type: DiffType::Added, - left_value: None, - right_value: Some(value), - } - } - - /// 创建删除差异项 - pub fn removed(path: impl Into, value: Value) -> Self { - Self { - path: path.into(), - diff_type: DiffType::Removed, - left_value: Some(value), - right_value: None, - } - } - - /// 创建修改差异项 - pub fn modified(path: impl Into, left: Value, right: Value) -> Self { - Self { - path: path.into(), - diff_type: DiffType::Modified, - left_value: Some(left), - right_value: Some(right), - } - } - - /// 创建未变化差异项 - pub fn unchanged(path: impl Into, value: Value) -> Self { - Self { - path: path.into(), - diff_type: DiffType::Unchanged, - left_value: Some(value.clone()), - right_value: Some(value), - } - } -} - -// ============================================================================ -// 差异配置 -// ============================================================================ - -/// 差异配置 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct DiffConfig { - /// 要忽略的字段列表 - pub ignore_fields: Vec, - /// 是否忽略时间戳 - pub ignore_timestamps: bool, - /// 是否忽略 ID - pub ignore_ids: bool, -} - -impl Default for DiffConfig { - fn default() -> Self { - Self { - ignore_fields: vec![], - ignore_timestamps: true, - ignore_ids: true, - } - } -} - -impl DiffConfig { - /// 创建新的配置 - pub fn new() -> Self { - Self::default() - } - - /// 设置忽略字段 - pub fn with_ignore_fields(mut self, fields: Vec) -> Self { - self.ignore_fields = fields; - self - } - - /// 设置是否忽略时间戳 - pub fn with_ignore_timestamps(mut self, ignore: bool) -> Self { - self.ignore_timestamps = ignore; - self - } - - /// 设置是否忽略 ID - pub fn with_ignore_ids(mut self, ignore: bool) -> Self { - self.ignore_ids = ignore; - self - } - - /// 检查字段是否应该被忽略 - pub fn should_ignore(&self, path: &str) -> bool { - // 检查自定义忽略字段 - if self.ignore_fields.iter().any(|f| path.contains(f)) { - return true; - } - - // 检查时间戳字段 - if self.ignore_timestamps { - let timestamp_fields = [ - "timestamp", - "created", - "updated", - "request_start", - "request_end", - "response_start", - "response_end", - "timestamp_start", - "timestamp_end", - "intercepted_at", - "added_at", - "created_at", - "updated_at", - ]; - if timestamp_fields.iter().any(|f| path.ends_with(f)) { - return true; - } - } - - // 检查 ID 字段 - if self.ignore_ids { - let id_fields = ["id", "flow_id", "request_id", "credential_id", "session_id"]; - if id_fields.iter().any(|f| path.ends_with(f)) { - return true; - } - } - - false - } -} - -// ============================================================================ -// Token 差异 -// ============================================================================ - -/// Token 差异 -#[derive(Debug, Clone, Serialize, Deserialize, Default)] -pub struct TokenDiff { - /// 输入 Token 差异 - pub input_diff: i64, - /// 输出 Token 差异 - pub output_diff: i64, - /// 总 Token 差异 - pub total_diff: i64, -} - -impl TokenDiff { - /// 计算两个 TokenUsage 之间的差异 - pub fn from_usage(left: &TokenUsage, right: &TokenUsage) -> Self { - Self { - input_diff: right.input_tokens as i64 - left.input_tokens as i64, - output_diff: right.output_tokens as i64 - left.output_tokens as i64, - total_diff: right.total_tokens as i64 - left.total_tokens as i64, - } - } - - /// 检查是否有差异 - pub fn has_diff(&self) -> bool { - self.input_diff != 0 || self.output_diff != 0 || self.total_diff != 0 - } -} - -// ============================================================================ -// 消息差异 -// ============================================================================ - -/// 消息差异项 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct MessageDiffItem { - /// 消息索引 - pub index: usize, - /// 差异类型 - pub diff_type: DiffType, - /// 左侧消息 - pub left_message: Option, - /// 右侧消息 - pub right_message: Option, - /// 内容差异详情 - pub content_diffs: Vec, -} - -// ============================================================================ -// Flow 差异结果 -// ============================================================================ - -/// Flow 差异结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct FlowDiffResult { - /// 左侧 Flow ID - pub left_flow_id: String, - /// 右侧 Flow ID - pub right_flow_id: String, - /// 请求差异 - pub request_diffs: Vec, - /// 响应差异 - pub response_diffs: Vec, - /// 元数据差异 - pub metadata_diffs: Vec, - /// 消息差异 - pub message_diffs: Vec, - /// Token 差异 - pub token_diff: TokenDiff, -} - -impl FlowDiffResult { - /// 检查是否有任何差异 - pub fn has_diff(&self) -> bool { - !self - .request_diffs - .iter() - .all(|d| d.diff_type == DiffType::Unchanged) - || !self - .response_diffs - .iter() - .all(|d| d.diff_type == DiffType::Unchanged) - || !self - .metadata_diffs - .iter() - .all(|d| d.diff_type == DiffType::Unchanged) - || !self - .message_diffs - .iter() - .all(|d| d.diff_type == DiffType::Unchanged) - || self.token_diff.has_diff() - } - - /// 获取所有有变化的差异项 - pub fn get_changed_items(&self) -> Vec<&DiffItem> { - let mut items = Vec::new(); - items.extend( - self.request_diffs - .iter() - .filter(|d| d.diff_type != DiffType::Unchanged), - ); - items.extend( - self.response_diffs - .iter() - .filter(|d| d.diff_type != DiffType::Unchanged), - ); - items.extend( - self.metadata_diffs - .iter() - .filter(|d| d.diff_type != DiffType::Unchanged), - ); - items - } -} - -// ============================================================================ -// FlowDiff 核心实现 -// ============================================================================ - -/// Flow 差异对比器 -pub struct FlowDiff; - -impl FlowDiff { - /// 对比两个 Flow - pub fn diff(left: &LLMFlow, right: &LLMFlow, config: &DiffConfig) -> FlowDiffResult { - let request_diffs = Self::diff_requests(&left.request, &right.request, config); - let response_diffs = - Self::diff_responses(left.response.as_ref(), right.response.as_ref(), config); - let metadata_diffs = Self::diff_metadata(&left.metadata, &right.metadata, config); - let message_diffs = Self::diff_messages(&left.request.messages, &right.request.messages); - let token_diff = Self::diff_tokens( - left.response.as_ref().map(|r| &r.usage), - right.response.as_ref().map(|r| &r.usage), - ); - - FlowDiffResult { - left_flow_id: left.id.clone(), - right_flow_id: right.id.clone(), - request_diffs, - response_diffs, - metadata_diffs, - message_diffs, - token_diff, - } - } - - /// 对比请求 - fn diff_requests( - left: &super::models::LLMRequest, - right: &super::models::LLMRequest, - config: &DiffConfig, - ) -> Vec { - let mut diffs = Vec::new(); - - // 对比模型 - if !config.should_ignore("request.model") && left.model != right.model { - diffs.push(DiffItem::modified( - "request.model", - Value::String(left.model.clone()), - Value::String(right.model.clone()), - )); - } - - // 对比方法 - if !config.should_ignore("request.method") && left.method != right.method { - diffs.push(DiffItem::modified( - "request.method", - Value::String(left.method.clone()), - Value::String(right.method.clone()), - )); - } - - // 对比路径 - if !config.should_ignore("request.path") && left.path != right.path { - diffs.push(DiffItem::modified( - "request.path", - Value::String(left.path.clone()), - Value::String(right.path.clone()), - )); - } - - // 对比系统提示词 - if !config.should_ignore("request.system_prompt") { - match (&left.system_prompt, &right.system_prompt) { - (Some(l), Some(r)) if l != r => { - diffs.push(DiffItem::modified( - "request.system_prompt", - Value::String(l.clone()), - Value::String(r.clone()), - )); - } - (Some(l), None) => { - diffs.push(DiffItem::removed( - "request.system_prompt", - Value::String(l.clone()), - )); - } - (None, Some(r)) => { - diffs.push(DiffItem::added( - "request.system_prompt", - Value::String(r.clone()), - )); - } - _ => {} - } - } - - // 对比参数 - if !config.should_ignore("request.parameters") { - Self::diff_parameters(&left.parameters, &right.parameters, &mut diffs, config); - } - - // 对比请求体 - if !config.should_ignore("request.body") { - let body_diffs = Self::diff_json(&left.body, &right.body, "request.body", config); - diffs.extend(body_diffs); - } - - diffs - } - - /// 对比请求参数 - fn diff_parameters( - left: &super::models::RequestParameters, - right: &super::models::RequestParameters, - diffs: &mut Vec, - config: &DiffConfig, - ) { - // 对比 temperature - if !config.should_ignore("request.parameters.temperature") { - match (left.temperature, right.temperature) { - (Some(l), Some(r)) if (l - r).abs() > f32::EPSILON => { - diffs.push(DiffItem::modified( - "request.parameters.temperature", - serde_json::json!(l), - serde_json::json!(r), - )); - } - (Some(l), None) => { - diffs.push(DiffItem::removed( - "request.parameters.temperature", - serde_json::json!(l), - )); - } - (None, Some(r)) => { - diffs.push(DiffItem::added( - "request.parameters.temperature", - serde_json::json!(r), - )); - } - _ => {} - } - } - - // 对比 top_p - if !config.should_ignore("request.parameters.top_p") { - match (left.top_p, right.top_p) { - (Some(l), Some(r)) if (l - r).abs() > f32::EPSILON => { - diffs.push(DiffItem::modified( - "request.parameters.top_p", - serde_json::json!(l), - serde_json::json!(r), - )); - } - (Some(l), None) => { - diffs.push(DiffItem::removed( - "request.parameters.top_p", - serde_json::json!(l), - )); - } - (None, Some(r)) => { - diffs.push(DiffItem::added( - "request.parameters.top_p", - serde_json::json!(r), - )); - } - _ => {} - } - } - - // 对比 max_tokens - if !config.should_ignore("request.parameters.max_tokens") { - match (left.max_tokens, right.max_tokens) { - (Some(l), Some(r)) if l != r => { - diffs.push(DiffItem::modified( - "request.parameters.max_tokens", - serde_json::json!(l), - serde_json::json!(r), - )); - } - (Some(l), None) => { - diffs.push(DiffItem::removed( - "request.parameters.max_tokens", - serde_json::json!(l), - )); - } - (None, Some(r)) => { - diffs.push(DiffItem::added( - "request.parameters.max_tokens", - serde_json::json!(r), - )); - } - _ => {} - } - } - - // 对比 stream - if !config.should_ignore("request.parameters.stream") && left.stream != right.stream { - diffs.push(DiffItem::modified( - "request.parameters.stream", - serde_json::json!(left.stream), - serde_json::json!(right.stream), - )); - } - } - - /// 对比响应 - fn diff_responses( - left: Option<&super::models::LLMResponse>, - right: Option<&super::models::LLMResponse>, - config: &DiffConfig, - ) -> Vec { - let mut diffs = Vec::new(); - - match (left, right) { - (Some(l), Some(r)) => { - // 对比状态码 - if !config.should_ignore("response.status_code") && l.status_code != r.status_code { - diffs.push(DiffItem::modified( - "response.status_code", - serde_json::json!(l.status_code), - serde_json::json!(r.status_code), - )); - } - - // 对比内容 - if !config.should_ignore("response.content") && l.content != r.content { - diffs.push(DiffItem::modified( - "response.content", - Value::String(l.content.clone()), - Value::String(r.content.clone()), - )); - } - - // 对比思维链 - if !config.should_ignore("response.thinking") { - match (&l.thinking, &r.thinking) { - (Some(lt), Some(rt)) if lt.text != rt.text => { - diffs.push(DiffItem::modified( - "response.thinking.text", - Value::String(lt.text.clone()), - Value::String(rt.text.clone()), - )); - } - (Some(lt), None) => { - diffs.push(DiffItem::removed( - "response.thinking", - serde_json::to_value(lt).unwrap_or(Value::Null), - )); - } - (None, Some(rt)) => { - diffs.push(DiffItem::added( - "response.thinking", - serde_json::to_value(rt).unwrap_or(Value::Null), - )); - } - _ => {} - } - } - - // 对比停止原因 - if !config.should_ignore("response.stop_reason") { - match (&l.stop_reason, &r.stop_reason) { - (Some(ls), Some(rs)) if ls != rs => { - diffs.push(DiffItem::modified( - "response.stop_reason", - serde_json::to_value(ls).unwrap_or(Value::Null), - serde_json::to_value(rs).unwrap_or(Value::Null), - )); - } - (Some(ls), None) => { - diffs.push(DiffItem::removed( - "response.stop_reason", - serde_json::to_value(ls).unwrap_or(Value::Null), - )); - } - (None, Some(rs)) => { - diffs.push(DiffItem::added( - "response.stop_reason", - serde_json::to_value(rs).unwrap_or(Value::Null), - )); - } - _ => {} - } - } - - // 对比工具调用数量 - if !config.should_ignore("response.tool_calls") - && l.tool_calls.len() != r.tool_calls.len() - { - diffs.push(DiffItem::modified( - "response.tool_calls.count", - serde_json::json!(l.tool_calls.len()), - serde_json::json!(r.tool_calls.len()), - )); - } - - // 对比响应体 - if !config.should_ignore("response.body") { - let body_diffs = Self::diff_json(&l.body, &r.body, "response.body", config); - diffs.extend(body_diffs); - } - } - (Some(l), None) => { - diffs.push(DiffItem::removed( - "response", - serde_json::to_value(l).unwrap_or(Value::Null), - )); - } - (None, Some(r)) => { - diffs.push(DiffItem::added( - "response", - serde_json::to_value(r).unwrap_or(Value::Null), - )); - } - (None, None) => {} - } - - diffs - } - - /// 对比元数据 - fn diff_metadata( - left: &super::models::FlowMetadata, - right: &super::models::FlowMetadata, - config: &DiffConfig, - ) -> Vec { - let mut diffs = Vec::new(); - - // 对比提供商 - if !config.should_ignore("metadata.provider") && left.provider != right.provider { - diffs.push(DiffItem::modified( - "metadata.provider", - serde_json::to_value(left.provider).unwrap_or(Value::Null), - serde_json::to_value(right.provider).unwrap_or(Value::Null), - )); - } - - // 对比凭证名称 - if !config.should_ignore("metadata.credential_name") { - match (&left.credential_name, &right.credential_name) { - (Some(l), Some(r)) if l != r => { - diffs.push(DiffItem::modified( - "metadata.credential_name", - Value::String(l.clone()), - Value::String(r.clone()), - )); - } - (Some(l), None) => { - diffs.push(DiffItem::removed( - "metadata.credential_name", - Value::String(l.clone()), - )); - } - (None, Some(r)) => { - diffs.push(DiffItem::added( - "metadata.credential_name", - Value::String(r.clone()), - )); - } - _ => {} - } - } - - // 对比重试次数 - if !config.should_ignore("metadata.retry_count") && left.retry_count != right.retry_count { - diffs.push(DiffItem::modified( - "metadata.retry_count", - serde_json::json!(left.retry_count), - serde_json::json!(right.retry_count), - )); - } - - diffs - } - - /// 对比消息列表 - pub fn diff_messages(left: &[Message], right: &[Message]) -> Vec { - let mut diffs = Vec::new(); - let max_len = left.len().max(right.len()); - - for i in 0..max_len { - match (left.get(i), right.get(i)) { - (Some(l), Some(r)) => { - let content_diffs = Self::diff_message_content(l, r, i); - let diff_type = if content_diffs.is_empty() { - DiffType::Unchanged - } else { - DiffType::Modified - }; - diffs.push(MessageDiffItem { - index: i, - diff_type, - left_message: Some(l.clone()), - right_message: Some(r.clone()), - content_diffs, - }); - } - (Some(l), None) => { - diffs.push(MessageDiffItem { - index: i, - diff_type: DiffType::Removed, - left_message: Some(l.clone()), - right_message: None, - content_diffs: vec![], - }); - } - (None, Some(r)) => { - diffs.push(MessageDiffItem { - index: i, - diff_type: DiffType::Added, - left_message: None, - right_message: Some(r.clone()), - content_diffs: vec![], - }); - } - (None, None) => {} - } - } - - diffs - } - - /// 对比单个消息的内容 - fn diff_message_content(left: &Message, right: &Message, index: usize) -> Vec { - let mut diffs = Vec::new(); - let prefix = format!("messages[{index}]"); - - // 对比角色 - if left.role != right.role { - diffs.push(DiffItem::modified( - format!("{prefix}.role"), - serde_json::to_value(&left.role).unwrap_or(Value::Null), - serde_json::to_value(&right.role).unwrap_or(Value::Null), - )); - } - - // 对比内容 - let left_text = Self::get_message_text(&left.content); - let right_text = Self::get_message_text(&right.content); - if left_text != right_text { - diffs.push(DiffItem::modified( - format!("{prefix}.content"), - Value::String(left_text), - Value::String(right_text), - )); - } - - // 对比名称 - match (&left.name, &right.name) { - (Some(l), Some(r)) if l != r => { - diffs.push(DiffItem::modified( - format!("{prefix}.name"), - Value::String(l.clone()), - Value::String(r.clone()), - )); - } - (Some(l), None) => { - diffs.push(DiffItem::removed( - format!("{prefix}.name"), - Value::String(l.clone()), - )); - } - (None, Some(r)) => { - diffs.push(DiffItem::added( - format!("{prefix}.name"), - Value::String(r.clone()), - )); - } - _ => {} - } - - // 对比工具调用 - match (&left.tool_calls, &right.tool_calls) { - (Some(l), Some(r)) if l.len() != r.len() => { - diffs.push(DiffItem::modified( - format!("{prefix}.tool_calls.count"), - serde_json::json!(l.len()), - serde_json::json!(r.len()), - )); - } - (Some(l), None) => { - diffs.push(DiffItem::removed( - format!("{prefix}.tool_calls"), - serde_json::to_value(l).unwrap_or(Value::Null), - )); - } - (None, Some(r)) => { - diffs.push(DiffItem::added( - format!("{prefix}.tool_calls"), - serde_json::to_value(r).unwrap_or(Value::Null), - )); - } - _ => {} - } - - diffs - } - - /// 获取消息文本内容 - fn get_message_text(content: &MessageContent) -> String { - match content { - MessageContent::Text(s) => s.clone(), - MessageContent::MultiModal(parts) => parts - .iter() - .filter_map(|p| { - if let super::models::ContentPart::Text { text } = p { - Some(text.as_str()) - } else { - None - } - }) - .collect::>() - .join("\n"), - } - } - - /// 对比 Token 使用量 - fn diff_tokens(left: Option<&TokenUsage>, right: Option<&TokenUsage>) -> TokenDiff { - match (left, right) { - (Some(l), Some(r)) => TokenDiff::from_usage(l, r), - (Some(l), None) => TokenDiff { - input_diff: -(l.input_tokens as i64), - output_diff: -(l.output_tokens as i64), - total_diff: -(l.total_tokens as i64), - }, - (None, Some(r)) => TokenDiff { - input_diff: r.input_tokens as i64, - output_diff: r.output_tokens as i64, - total_diff: r.total_tokens as i64, - }, - (None, None) => TokenDiff::default(), - } - } - - /// 对比两个 JSON 值 - pub fn diff_json( - left: &Value, - right: &Value, - path: &str, - config: &DiffConfig, - ) -> Vec { - if config.should_ignore(path) { - return vec![]; - } - - let mut diffs = Vec::new(); - - match (left, right) { - (Value::Object(l), Value::Object(r)) => { - // 收集所有键 - let mut all_keys: Vec<_> = l.keys().chain(r.keys()).collect(); - all_keys.sort(); - all_keys.dedup(); - - for key in all_keys { - let new_path = if path.is_empty() { - key.clone() - } else { - format!("{path}.{key}") - }; - - match (l.get(key), r.get(key)) { - (Some(lv), Some(rv)) => { - diffs.extend(Self::diff_json(lv, rv, &new_path, config)); - } - (Some(lv), None) => { - if !config.should_ignore(&new_path) { - diffs.push(DiffItem::removed(new_path, lv.clone())); - } - } - (None, Some(rv)) => { - if !config.should_ignore(&new_path) { - diffs.push(DiffItem::added(new_path, rv.clone())); - } - } - (None, None) => {} - } - } - } - (Value::Array(l), Value::Array(r)) => { - let max_len = l.len().max(r.len()); - for i in 0..max_len { - let new_path = format!("{path}[{i}]"); - match (l.get(i), r.get(i)) { - (Some(lv), Some(rv)) => { - diffs.extend(Self::diff_json(lv, rv, &new_path, config)); - } - (Some(lv), None) => { - if !config.should_ignore(&new_path) { - diffs.push(DiffItem::removed(new_path, lv.clone())); - } - } - (None, Some(rv)) => { - if !config.should_ignore(&new_path) { - diffs.push(DiffItem::added(new_path, rv.clone())); - } - } - (None, None) => {} - } - } - } - _ => { - if left != right && !config.should_ignore(path) { - diffs.push(DiffItem::modified(path, left.clone(), right.clone())); - } - } - } - - diffs - } -} - -// ============================================================================ -// 单元测试 -// ============================================================================ - -#[cfg(test)] -mod tests { - use super::*; - use crate::flow_monitor::models::{ - FlowMetadata, FlowType, LLMRequest, LLMResponse, Message, MessageRole, RequestParameters, - }; - use crate::ProviderType; - - /// 创建测试用的 Flow - fn create_test_flow(id: &str, model: &str, content: &str) -> LLMFlow { - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model: model.to_string(), - messages: vec![Message { - role: MessageRole::User, - content: MessageContent::Text(content.to_string()), - ..Default::default() - }], - parameters: RequestParameters::default(), - ..Default::default() - }; - - let metadata = FlowMetadata { - provider: ProviderType::OpenAI, - ..Default::default() - }; - - let mut flow = LLMFlow::new(id.to_string(), FlowType::ChatCompletions, request, metadata); - flow.response = Some(LLMResponse { - content: "Response content".to_string(), - usage: TokenUsage { - input_tokens: 100, - output_tokens: 50, - total_tokens: 150, - ..Default::default() - }, - ..Default::default() - }); - flow - } - - #[test] - fn test_diff_identical_flows() { - let flow1 = create_test_flow("id1", "gpt-4", "Hello"); - let flow2 = create_test_flow("id2", "gpt-4", "Hello"); - let config = DiffConfig::default(); - - let result = FlowDiff::diff(&flow1, &flow2, &config); - - // 由于 ID 被忽略,应该没有差异 - assert!(result.request_diffs.is_empty()); - assert!(result - .message_diffs - .iter() - .all(|d| d.diff_type == DiffType::Unchanged)); - } - - #[test] - fn test_diff_different_models() { - let flow1 = create_test_flow("id1", "gpt-4", "Hello"); - let flow2 = create_test_flow("id2", "gpt-3.5-turbo", "Hello"); - let config = DiffConfig::default(); - - let result = FlowDiff::diff(&flow1, &flow2, &config); - - let model_diff = result - .request_diffs - .iter() - .find(|d| d.path == "request.model"); - assert!(model_diff.is_some()); - assert_eq!(model_diff.unwrap().diff_type, DiffType::Modified); - } - - #[test] - fn test_diff_different_messages() { - let flow1 = create_test_flow("id1", "gpt-4", "Hello"); - let flow2 = create_test_flow("id2", "gpt-4", "World"); - let config = DiffConfig::default(); - - let result = FlowDiff::diff(&flow1, &flow2, &config); - - assert!(!result.message_diffs.is_empty()); - assert_eq!(result.message_diffs[0].diff_type, DiffType::Modified); - } - - #[test] - fn test_token_diff() { - let usage1 = TokenUsage { - input_tokens: 100, - output_tokens: 50, - total_tokens: 150, - ..Default::default() - }; - let usage2 = TokenUsage { - input_tokens: 120, - output_tokens: 60, - total_tokens: 180, - ..Default::default() - }; - - let diff = TokenDiff::from_usage(&usage1, &usage2); - - assert_eq!(diff.input_diff, 20); - assert_eq!(diff.output_diff, 10); - assert_eq!(diff.total_diff, 30); - assert!(diff.has_diff()); - } - - #[test] - fn test_diff_config_ignore_timestamps() { - let config = DiffConfig::default(); - assert!(config.should_ignore("timestamps.created")); - assert!(config.should_ignore("response.timestamp_start")); - assert!(!config.should_ignore("request.model")); - } - - #[test] - fn test_diff_config_ignore_ids() { - let config = DiffConfig::default(); - assert!(config.should_ignore("flow.id")); - assert!(config.should_ignore("metadata.credential_id")); - assert!(!config.should_ignore("request.model")); - } - - #[test] - fn test_diff_json_objects() { - let left = serde_json::json!({ - "a": 1, - "b": 2, - "c": 3 - }); - let right = serde_json::json!({ - "a": 1, - "b": 3, - "d": 4 - }); - let config = DiffConfig::new() - .with_ignore_timestamps(false) - .with_ignore_ids(false); - - let diffs = FlowDiff::diff_json(&left, &right, "root", &config); - - // b 被修改,c 被删除,d 被添加 - assert_eq!(diffs.len(), 3); - assert!(diffs - .iter() - .any(|d| d.path == "root.b" && d.diff_type == DiffType::Modified)); - assert!(diffs - .iter() - .any(|d| d.path == "root.c" && d.diff_type == DiffType::Removed)); - assert!(diffs - .iter() - .any(|d| d.path == "root.d" && d.diff_type == DiffType::Added)); - } - - #[test] - fn test_diff_messages_added() { - let left = vec![Message { - role: MessageRole::User, - content: MessageContent::Text("Hello".to_string()), - ..Default::default() - }]; - let right = vec![ - Message { - role: MessageRole::User, - content: MessageContent::Text("Hello".to_string()), - ..Default::default() - }, - Message { - role: MessageRole::Assistant, - content: MessageContent::Text("Hi there".to_string()), - ..Default::default() - }, - ]; - - let diffs = FlowDiff::diff_messages(&left, &right); - - assert_eq!(diffs.len(), 2); - assert_eq!(diffs[0].diff_type, DiffType::Unchanged); - assert_eq!(diffs[1].diff_type, DiffType::Added); - } - - #[test] - fn test_diff_messages_removed() { - let left = vec![ - Message { - role: MessageRole::User, - content: MessageContent::Text("Hello".to_string()), - ..Default::default() - }, - Message { - role: MessageRole::Assistant, - content: MessageContent::Text("Hi there".to_string()), - ..Default::default() - }, - ]; - let right = vec![Message { - role: MessageRole::User, - content: MessageContent::Text("Hello".to_string()), - ..Default::default() - }]; - - let diffs = FlowDiff::diff_messages(&left, &right); - - assert_eq!(diffs.len(), 2); - assert_eq!(diffs[0].diff_type, DiffType::Unchanged); - assert_eq!(diffs[1].diff_type, DiffType::Removed); - } -} - -// ============================================================================ -// 属性测试模块 -// ============================================================================ - -#[cfg(test)] -mod property_tests { - use super::*; - use crate::flow_monitor::models::{ - FlowMetadata, FlowType, LLMRequest, LLMResponse, Message, MessageRole, RequestParameters, - TokenUsage, - }; - use crate::ProviderType; - use chrono::Utc; - use proptest::prelude::*; - - // ======================================================================== - // 生成器 - // ======================================================================== - - /// 生成随机的 ProviderType - fn arb_provider_type() -> impl Strategy { - prop_oneof![ - Just(ProviderType::Kiro), - Just(ProviderType::Gemini), - Just(ProviderType::OpenAI), - Just(ProviderType::Claude), - ] - } - - /// 生成随机的 MessageRole - fn arb_message_role() -> impl Strategy { - prop_oneof![ - Just(MessageRole::System), - Just(MessageRole::User), - Just(MessageRole::Assistant), - ] - } - - /// 生成随机的 MessageContent - fn arb_message_content() -> impl Strategy { - "[a-zA-Z0-9 ]{1,100}".prop_map(MessageContent::Text) - } - - /// 生成随机的 Message - fn arb_message() -> impl Strategy { - (arb_message_role(), arb_message_content()).prop_map(|(role, content)| Message { - role, - content, - tool_calls: None, - tool_result: None, - name: None, - }) - } - - /// 生成随机的 RequestParameters - fn arb_request_parameters() -> impl Strategy { - ( - prop::option::of(0.0f32..2.0f32), - prop::option::of(0.0f32..1.0f32), - prop::option::of(1u32..4096u32), - any::(), - ) - .prop_map( - |(temperature, top_p, max_tokens, stream)| RequestParameters { - temperature, - top_p, - max_tokens, - stop: None, - stream, - extra: std::collections::HashMap::new(), - }, - ) - } - - /// 生成随机的 TokenUsage - fn arb_token_usage() -> impl Strategy { - (0u32..10000u32, 0u32..10000u32).prop_map(|(input, output)| TokenUsage { - input_tokens: input, - output_tokens: output, - total_tokens: input + output, - ..Default::default() - }) - } - - /// 生成随机的 LLMRequest - fn arb_llm_request() -> impl Strategy { - ( - "[a-z]{3,20}", // model - prop::collection::vec(arb_message(), 1..5), // messages - arb_request_parameters(), // parameters - prop::option::of("[a-zA-Z0-9 ]{10,50}"), // system_prompt - ) - .prop_map(|(model, messages, parameters, system_prompt)| LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - headers: std::collections::HashMap::new(), - body: serde_json::Value::Null, - messages, - system_prompt, - tools: None, - model, - original_model: None, - parameters, - size_bytes: 0, - timestamp: Utc::now(), - }) - } - - /// 生成随机的 LLMResponse - fn arb_llm_response() -> impl Strategy { - ( - "[a-zA-Z0-9 ]{10,200}", // content - arb_token_usage(), // usage - ) - .prop_map(|(content, usage)| LLMResponse { - status_code: 200, - status_text: "OK".to_string(), - headers: std::collections::HashMap::new(), - body: serde_json::Value::Null, - content, - thinking: None, - tool_calls: vec![], - usage, - stop_reason: None, - size_bytes: 0, - timestamp_start: Utc::now(), - timestamp_end: Utc::now(), - stream_info: None, - }) - } - - /// 生成随机的 FlowMetadata - fn arb_flow_metadata() -> impl Strategy { - arb_provider_type().prop_map(|provider| FlowMetadata { - provider, - provider_id: None, - credential_id: None, - credential_name: None, - retry_count: 0, - client_info: Default::default(), - routing_info: Default::default(), - injected_params: None, - context_usage_percentage: None, - }) - } - - /// 生成随机的 LLMFlow - fn arb_llm_flow() -> impl Strategy { - ( - "[a-f0-9]{8}", - arb_llm_request(), - arb_flow_metadata(), - prop::option::of(arb_llm_response()), - ) - .prop_map(|(id, request, metadata, response)| { - let mut flow = LLMFlow::new(id, FlowType::ChatCompletions, request, metadata); - flow.response = response; - flow - }) - } - - /// 生成随机的 DiffConfig - fn arb_diff_config() -> impl Strategy { - (any::(), any::()).prop_map(|(ignore_timestamps, ignore_ids)| DiffConfig { - ignore_fields: vec![], - ignore_timestamps, - ignore_ids, - }) - } - - // ======================================================================== - // 属性测试 - // ======================================================================== - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: flow-monitor-enhancement, Property 6: 差异计算正确性** - /// **Validates: Requirements 4.1, 4.2, 4.5, 4.6, 4.7** - /// - /// *对于任意* 两个 Flow,差异计算应该正确识别所有新增、删除和修改的字段, - /// 且忽略配置中指定的字段。 - #[test] - fn prop_diff_correctness( - flow1 in arb_llm_flow(), - flow2 in arb_llm_flow(), - config in arb_diff_config(), - ) { - let result = FlowDiff::diff(&flow1, &flow2, &config); - - // 验证 Flow ID 正确记录 - prop_assert_eq!(&result.left_flow_id, &flow1.id); - prop_assert_eq!(&result.right_flow_id, &flow2.id); - - // 验证模型差异检测 - if flow1.request.model != flow2.request.model { - let model_diff = result.request_diffs.iter().find(|d| d.path == "request.model"); - prop_assert!(model_diff.is_some(), "模型不同时应该检测到差异"); - prop_assert_eq!(model_diff.unwrap().diff_type, DiffType::Modified); - } - - // 验证消息数量差异检测 - let left_msg_count = flow1.request.messages.len(); - let right_msg_count = flow2.request.messages.len(); - prop_assert_eq!( - result.message_diffs.len(), - left_msg_count.max(right_msg_count), - "消息差异数量应该等于两个消息列表的最大长度" - ); - - // 验证 Token 差异计算 - if let (Some(r1), Some(r2)) = (&flow1.response, &flow2.response) { - let expected_input_diff = r2.usage.input_tokens as i64 - r1.usage.input_tokens as i64; - let expected_output_diff = r2.usage.output_tokens as i64 - r1.usage.output_tokens as i64; - prop_assert_eq!(result.token_diff.input_diff, expected_input_diff); - prop_assert_eq!(result.token_diff.output_diff, expected_output_diff); - } - - // 验证忽略字段配置生效 - for diff in &result.request_diffs { - prop_assert!( - !config.should_ignore(&diff.path), - "被忽略的字段不应该出现在差异结果中: {}", - diff.path - ); - } - for diff in &result.response_diffs { - prop_assert!( - !config.should_ignore(&diff.path), - "被忽略的字段不应该出现在差异结果中: {}", - diff.path - ); - } - for diff in &result.metadata_diffs { - prop_assert!( - !config.should_ignore(&diff.path), - "被忽略的字段不应该出现在差异结果中: {}", - diff.path - ); - } - } - - /// **Feature: flow-monitor-enhancement, Property 7: 差异计算对称性** - /// **Validates: Requirements 4.1, 4.2** - /// - /// *对于任意* 两个 Flow A 和 B,diff(A, B) 中的 "Added" 项应该对应 diff(B, A) 中的 "Removed" 项。 - #[test] - fn prop_diff_symmetry( - flow1 in arb_llm_flow(), - flow2 in arb_llm_flow(), - ) { - let config = DiffConfig::default(); - let result_ab = FlowDiff::diff(&flow1, &flow2, &config); - let result_ba = FlowDiff::diff(&flow2, &flow1, &config); - - // 验证请求差异对称性 - for diff_ab in &result_ab.request_diffs { - let corresponding = result_ba.request_diffs.iter().find(|d| d.path == diff_ab.path); - if let Some(diff_ba) = corresponding { - match diff_ab.diff_type { - DiffType::Added => { - prop_assert_eq!( - diff_ba.diff_type, - DiffType::Removed, - "A->B 的 Added 应该对应 B->A 的 Removed: {}", - diff_ab.path - ); - } - DiffType::Removed => { - prop_assert_eq!( - diff_ba.diff_type, - DiffType::Added, - "A->B 的 Removed 应该对应 B->A 的 Added: {}", - diff_ab.path - ); - } - DiffType::Modified => { - prop_assert_eq!( - diff_ba.diff_type, - DiffType::Modified, - "A->B 的 Modified 应该对应 B->A 的 Modified: {}", - diff_ab.path - ); - // 验证值交换 - prop_assert_eq!( - &diff_ab.left_value, - &diff_ba.right_value, - "Modified 差异的值应该交换" - ); - prop_assert_eq!( - &diff_ab.right_value, - &diff_ba.left_value, - "Modified 差异的值应该交换" - ); - } - DiffType::Unchanged => {} - } - } - } - - // 验证消息差异对称性 - for (i, diff_ab) in result_ab.message_diffs.iter().enumerate() { - if let Some(diff_ba) = result_ba.message_diffs.get(i) { - match diff_ab.diff_type { - DiffType::Added => { - prop_assert_eq!( - diff_ba.diff_type, - DiffType::Removed, - "消息 {} A->B 的 Added 应该对应 B->A 的 Removed", - i - ); - } - DiffType::Removed => { - prop_assert_eq!( - diff_ba.diff_type, - DiffType::Added, - "消息 {} A->B 的 Removed 应该对应 B->A 的 Added", - i - ); - } - _ => {} - } - } - } - - // 验证 Token 差异对称性 - prop_assert_eq!( - result_ab.token_diff.input_diff, - -result_ba.token_diff.input_diff, - "Token 输入差异应该相反" - ); - prop_assert_eq!( - result_ab.token_diff.output_diff, - -result_ba.token_diff.output_diff, - "Token 输出差异应该相反" - ); - prop_assert_eq!( - result_ab.token_diff.total_diff, - -result_ba.token_diff.total_diff, - "Token 总差异应该相反" - ); - } - - /// **Feature: flow-monitor-enhancement, Property 6b: 相同 Flow 无差异** - /// **Validates: Requirements 4.1, 4.2** - /// - /// *对于任意* Flow,与自身对比应该没有差异(除了被忽略的字段)。 - #[test] - fn prop_diff_self_no_changes( - flow in arb_llm_flow(), - ) { - let config = DiffConfig::default(); - let result = FlowDiff::diff(&flow, &flow, &config); - - // 验证请求差异为空或全部为 Unchanged - for diff in &result.request_diffs { - prop_assert_eq!( - diff.diff_type, - DiffType::Unchanged, - "自身对比不应该有请求差异: {}", - diff.path - ); - } - - // 验证响应差异为空或全部为 Unchanged - for diff in &result.response_diffs { - prop_assert_eq!( - diff.diff_type, - DiffType::Unchanged, - "自身对比不应该有响应差异: {}", - diff.path - ); - } - - // 验证消息差异全部为 Unchanged - for diff in &result.message_diffs { - prop_assert_eq!( - diff.diff_type, - DiffType::Unchanged, - "自身对比不应该有消息差异" - ); - } - - // 验证 Token 差异为零 - prop_assert_eq!(result.token_diff.input_diff, 0); - prop_assert_eq!(result.token_diff.output_diff, 0); - prop_assert_eq!(result.token_diff.total_diff, 0); - } - - /// **Feature: flow-monitor-enhancement, Property 6c: Token 差异计算正确性** - /// **Validates: Requirements 4.7** - /// - /// *对于任意* 两个 TokenUsage,差异计算应该正确。 - #[test] - fn prop_token_diff_correctness( - usage1 in arb_token_usage(), - usage2 in arb_token_usage(), - ) { - let diff = TokenDiff::from_usage(&usage1, &usage2); - - // 验证差异计算 - prop_assert_eq!( - diff.input_diff, - usage2.input_tokens as i64 - usage1.input_tokens as i64 - ); - prop_assert_eq!( - diff.output_diff, - usage2.output_tokens as i64 - usage1.output_tokens as i64 - ); - prop_assert_eq!( - diff.total_diff, - usage2.total_tokens as i64 - usage1.total_tokens as i64 - ); - - // 验证 has_diff 正确性 - let expected_has_diff = diff.input_diff != 0 || diff.output_diff != 0 || diff.total_diff != 0; - prop_assert_eq!(diff.has_diff(), expected_has_diff); - } - - /// **Feature: flow-monitor-enhancement, Property 6d: 消息差异计算正确性** - /// **Validates: Requirements 4.6** - /// - /// *对于任意* 两个消息列表,差异计算应该正确识别新增、删除和修改的消息。 - #[test] - fn prop_message_diff_correctness( - messages1 in prop::collection::vec(arb_message(), 0..5), - messages2 in prop::collection::vec(arb_message(), 0..5), - ) { - let diffs = FlowDiff::diff_messages(&messages1, &messages2); - - // 验证差异数量 - let expected_len = messages1.len().max(messages2.len()); - prop_assert_eq!(diffs.len(), expected_len); - - // 验证每个差异项 - for (i, diff) in diffs.iter().enumerate() { - prop_assert_eq!(diff.index, i); - - match (messages1.get(i), messages2.get(i)) { - (Some(_), Some(_)) => { - // 两边都有消息,应该是 Modified 或 Unchanged - prop_assert!( - diff.diff_type == DiffType::Modified || diff.diff_type == DiffType::Unchanged, - "两边都有消息时应该是 Modified 或 Unchanged" - ); - prop_assert!(diff.left_message.is_some()); - prop_assert!(diff.right_message.is_some()); - } - (Some(_), None) => { - // 只有左边有消息,应该是 Removed - prop_assert_eq!(diff.diff_type, DiffType::Removed); - prop_assert!(diff.left_message.is_some()); - prop_assert!(diff.right_message.is_none()); - } - (None, Some(_)) => { - // 只有右边有消息,应该是 Added - prop_assert_eq!(diff.diff_type, DiffType::Added); - prop_assert!(diff.left_message.is_none()); - prop_assert!(diff.right_message.is_some()); - } - (None, None) => { - // 不应该发生 - prop_assert!(false, "不应该有两边都没有消息的差异项"); - } - } - } - } - } -} diff --git a/src-tauri/src/flow_monitor/enhanced_stats.rs b/src-tauri/src/flow_monitor/enhanced_stats.rs deleted file mode 100644 index 8d12b65ae..000000000 --- a/src-tauri/src/flow_monitor/enhanced_stats.rs +++ /dev/null @@ -1,1010 +0,0 @@ -//! 增强统计服务 -//! -//! 该模块实现 LLM Flow 的增强统计功能,包括时间序列趋势、分布分析、直方图等。 -//! -//! **Validates: Requirements 9.1-9.7** - -use chrono::{DateTime, Duration, Utc}; -use serde::{Deserialize, Serialize}; -use std::collections::HashMap; -use std::sync::Arc; - -use super::memory_store::{FlowFilter, FlowMemoryStore, TimeRange}; -use super::models::{FlowState, LLMFlow}; -use tokio::sync::RwLock; - -// ============================================================================ -// 数据结构 -// ============================================================================ - -/// 时间序列数据点 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct TimeSeriesPoint { - /// 时间戳 - pub timestamp: DateTime, - /// 数值 - pub value: f64, -} - -/// 分布数据 -#[derive(Debug, Clone, Serialize, Deserialize, Default)] -pub struct Distribution { - /// 分布桶 (标签, 数量) - pub buckets: Vec<(String, u64)>, - /// 总数 - pub total: u64, -} - -/// 趋势数据 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct TrendData { - /// 数据点列表 - pub points: Vec, - /// 时间间隔 - pub interval: String, -} - -impl Default for TrendData { - fn default() -> Self { - Self { - points: Vec::new(), - interval: "1h".to_string(), - } - } -} - -/// 增强统计结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct EnhancedStats { - /// 请求趋势 - pub request_trend: TrendData, - /// 按模型的 Token 分布 - pub token_by_model: Distribution, - /// 按提供商的成功率 - pub success_by_provider: Vec<(String, f64)>, - /// 延迟直方图 - pub latency_histogram: Distribution, - /// 错误分布 - pub error_distribution: Distribution, - /// 请求速率(每秒) - pub request_rate: f64, - /// 时间范围 - pub time_range: StatsTimeRange, -} - -impl Default for EnhancedStats { - fn default() -> Self { - Self { - request_trend: TrendData::default(), - token_by_model: Distribution::default(), - success_by_provider: Vec::new(), - latency_histogram: Distribution::default(), - error_distribution: Distribution::default(), - request_rate: 0.0, - time_range: StatsTimeRange::default(), - } - } -} - -/// 统计时间范围 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct StatsTimeRange { - /// 开始时间 - pub start: DateTime, - /// 结束时间 - pub end: DateTime, -} - -impl Default for StatsTimeRange { - fn default() -> Self { - let now = Utc::now(); - Self { - start: now - Duration::hours(24), - end: now, - } - } -} - -/// 统计报告格式 -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] -#[serde(rename_all = "lowercase")] -#[derive(Default)] -pub enum ReportFormat { - /// JSON 格式 - #[default] - Json, - /// Markdown 格式 - Markdown, - /// CSV 格式 - Csv, -} - -// ============================================================================ -// 增强统计服务 -// ============================================================================ - -/// 增强统计服务 -/// -/// 提供更详细的统计分析功能,包括时间序列趋势、分布分析等。 -pub struct EnhancedStatsService { - /// 内存存储 - memory_store: Arc>, -} - -impl EnhancedStatsService { - /// 创建新的增强统计服务 - pub fn new(memory_store: Arc>) -> Self { - Self { memory_store } - } - - /// 获取增强统计 - /// - /// **Validates: Requirements 9.1-9.5** - /// - /// # Arguments - /// * `filter` - 过滤条件 - /// * `time_range` - 时间范围 - /// - /// # Returns - /// 增强统计结果 - pub async fn get_stats( - &self, - filter: &FlowFilter, - time_range: &StatsTimeRange, - ) -> EnhancedStats { - // 获取 Flow 数据 - let flows = self.get_flows_in_range(filter, time_range).await; - - if flows.is_empty() { - return EnhancedStats { - time_range: time_range.clone(), - ..Default::default() - }; - } - - // 计算各项统计 - let request_trend = self.calculate_request_trend(&flows, "1h"); - let token_by_model = self.calculate_token_distribution(&flows); - let success_by_provider = self.calculate_success_by_provider(&flows); - let latency_histogram = - self.calculate_latency_histogram(&flows, &default_latency_buckets()); - let error_distribution = self.calculate_error_distribution(&flows); - let request_rate = self.calculate_request_rate(&flows, time_range); - - EnhancedStats { - request_trend, - token_by_model, - success_by_provider, - latency_histogram, - error_distribution, - request_rate, - time_range: time_range.clone(), - } - } - - /// 获取请求趋势 - /// - /// **Validates: Requirements 9.1** - /// - /// # Arguments - /// * `filter` - 过滤条件 - /// * `time_range` - 时间范围 - /// * `interval` - 时间间隔(如 "1h", "30m", "1d") - /// - /// # Returns - /// 趋势数据 - pub async fn get_request_trend( - &self, - filter: &FlowFilter, - time_range: &StatsTimeRange, - interval: &str, - ) -> TrendData { - let flows = self.get_flows_in_range(filter, time_range).await; - self.calculate_request_trend(&flows, interval) - } - - /// 获取 Token 分布 - /// - /// **Validates: Requirements 9.2** - /// - /// # Arguments - /// * `filter` - 过滤条件 - /// * `time_range` - 时间范围 - /// - /// # Returns - /// Token 分布数据 - pub async fn get_token_distribution( - &self, - filter: &FlowFilter, - time_range: &StatsTimeRange, - ) -> Distribution { - let flows = self.get_flows_in_range(filter, time_range).await; - self.calculate_token_distribution(&flows) - } - - /// 获取延迟直方图 - /// - /// **Validates: Requirements 9.4** - /// - /// # Arguments - /// * `filter` - 过滤条件 - /// * `time_range` - 时间范围 - /// * `buckets` - 直方图桶边界(毫秒) - /// - /// # Returns - /// 延迟直方图数据 - pub async fn get_latency_histogram( - &self, - filter: &FlowFilter, - time_range: &StatsTimeRange, - buckets: &[u64], - ) -> Distribution { - let flows = self.get_flows_in_range(filter, time_range).await; - self.calculate_latency_histogram(&flows, buckets) - } - - /// 导出统计报告 - /// - /// **Validates: Requirements 9.7** - /// - /// # Arguments - /// * `filter` - 过滤条件 - /// * `time_range` - 时间范围 - /// * `format` - 报告格式 - /// - /// # Returns - /// 格式化的报告字符串 - pub async fn export_report( - &self, - filter: &FlowFilter, - time_range: &StatsTimeRange, - format: &ReportFormat, - ) -> String { - let stats = self.get_stats(filter, time_range).await; - - match format { - ReportFormat::Json => self.export_json(&stats), - ReportFormat::Markdown => self.export_markdown(&stats), - ReportFormat::Csv => self.export_csv(&stats), - } - } - - // ======================================================================== - // 内部方法 - // ======================================================================== - - /// 获取时间范围内的 Flow - async fn get_flows_in_range( - &self, - filter: &FlowFilter, - time_range: &StatsTimeRange, - ) -> Vec { - let store = self.memory_store.read().await; - - // 创建带时间范围的过滤器 - let mut filter_with_time = filter.clone(); - filter_with_time.time_range = Some(TimeRange { - start: Some(time_range.start), - end: Some(time_range.end), - }); - - store.query(&filter_with_time) - } - - /// 计算请求趋势 - fn calculate_request_trend(&self, flows: &[LLMFlow], interval: &str) -> TrendData { - if flows.is_empty() { - return TrendData { - points: Vec::new(), - interval: interval.to_string(), - }; - } - - // 解析时间间隔 - let interval_duration = parse_interval(interval); - - // 找到时间范围 - let min_time = flows - .iter() - .map(|f| f.timestamps.created) - .min() - .unwrap_or_else(Utc::now); - let max_time = flows - .iter() - .map(|f| f.timestamps.created) - .max() - .unwrap_or_else(Utc::now); - - // 按时间间隔分组计数 - let mut counts: HashMap = HashMap::new(); - - for flow in flows { - let bucket = (flow.timestamps.created.timestamp() / interval_duration.num_seconds()) - * interval_duration.num_seconds(); - *counts.entry(bucket).or_insert(0) += 1; - } - - // 生成完整的时间序列(包括零值点) - let mut points = Vec::new(); - let mut current = (min_time.timestamp() / interval_duration.num_seconds()) - * interval_duration.num_seconds(); - let end = max_time.timestamp(); - - while current <= end { - let count = counts.get(¤t).copied().unwrap_or(0); - if let Some(timestamp) = DateTime::from_timestamp(current, 0) { - points.push(TimeSeriesPoint { - timestamp: timestamp.with_timezone(&Utc), - value: count as f64, - }); - } - current += interval_duration.num_seconds(); - } - - TrendData { - points, - interval: interval.to_string(), - } - } - - /// 计算 Token 分布(按模型) - fn calculate_token_distribution(&self, flows: &[LLMFlow]) -> Distribution { - let mut model_tokens: HashMap = HashMap::new(); - let mut total: u64 = 0; - - for flow in flows { - if let Some(ref response) = flow.response { - let tokens = response.usage.total_tokens as u64; - *model_tokens.entry(flow.request.model.clone()).or_insert(0) += tokens; - total += tokens; - } - } - - // 按 Token 数量降序排序 - let mut buckets: Vec<(String, u64)> = model_tokens.into_iter().collect(); - buckets.sort_by(|a, b| b.1.cmp(&a.1)); - - Distribution { buckets, total } - } - - /// 计算按提供商的成功率 - fn calculate_success_by_provider(&self, flows: &[LLMFlow]) -> Vec<(String, f64)> { - let mut provider_stats: HashMap = HashMap::new(); - - for flow in flows { - let provider = format!("{:?}", flow.metadata.provider); - let entry = provider_stats.entry(provider).or_insert((0, 0)); - entry.0 += 1; // 总数 - if flow.state == FlowState::Completed { - entry.1 += 1; // 成功数 - } - } - - let mut result: Vec<(String, f64)> = provider_stats - .into_iter() - .map(|(provider, (total, success))| { - let rate = if total > 0 { - success as f64 / total as f64 - } else { - 0.0 - }; - (provider, rate) - }) - .collect(); - - // 按成功率降序排序 - result.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal)); - - result - } - - /// 计算延迟直方图 - fn calculate_latency_histogram(&self, flows: &[LLMFlow], buckets: &[u64]) -> Distribution { - let mut bucket_counts: Vec = vec![0; buckets.len() + 1]; - let mut total: u64 = 0; - - for flow in flows { - let latency = flow.timestamps.duration_ms; - total += 1; - - // 找到对应的桶 - let bucket_idx = buckets - .iter() - .position(|&b| latency < b) - .unwrap_or(buckets.len()); - bucket_counts[bucket_idx] += 1; - } - - // 生成桶标签 - let mut result_buckets = Vec::new(); - for (i, count) in bucket_counts.iter().enumerate() { - let label = if i == 0 { - format!("<{}ms", buckets.first().unwrap_or(&0)) - } else if i == buckets.len() { - format!(">={}ms", buckets.last().unwrap_or(&0)) - } else { - format!("{}-{}ms", buckets[i - 1], buckets[i]) - }; - result_buckets.push((label, *count)); - } - - Distribution { - buckets: result_buckets, - total, - } - } - - /// 计算错误分布 - fn calculate_error_distribution(&self, flows: &[LLMFlow]) -> Distribution { - let mut error_counts: HashMap = HashMap::new(); - let mut total: u64 = 0; - - for flow in flows { - if let Some(ref error) = flow.error { - let error_type = format!("{:?}", error.error_type); - *error_counts.entry(error_type).or_insert(0) += 1; - total += 1; - } - } - - // 按数量降序排序 - let mut buckets: Vec<(String, u64)> = error_counts.into_iter().collect(); - buckets.sort_by(|a, b| b.1.cmp(&a.1)); - - Distribution { buckets, total } - } - - /// 计算请求速率(每秒) - fn calculate_request_rate(&self, flows: &[LLMFlow], time_range: &StatsTimeRange) -> f64 { - if flows.is_empty() { - return 0.0; - } - - let duration_secs = (time_range.end - time_range.start).num_seconds() as f64; - if duration_secs <= 0.0 { - return 0.0; - } - - flows.len() as f64 / duration_secs - } - - /// 导出为 JSON 格式 - fn export_json(&self, stats: &EnhancedStats) -> String { - serde_json::to_string_pretty(stats).unwrap_or_else(|_| "{}".to_string()) - } - - /// 导出为 Markdown 格式 - fn export_markdown(&self, stats: &EnhancedStats) -> String { - let mut md = String::new(); - - md.push_str("# Flow 统计报告\n\n"); - md.push_str(&format!( - "**时间范围**: {} - {}\n\n", - stats.time_range.start.format("%Y-%m-%d %H:%M:%S"), - stats.time_range.end.format("%Y-%m-%d %H:%M:%S") - )); - md.push_str(&format!( - "**请求速率**: {:.2} 请求/秒\n\n", - stats.request_rate - )); - - // Token 分布 - md.push_str("## Token 分布(按模型)\n\n"); - md.push_str("| 模型 | Token 数 |\n"); - md.push_str("|------|----------|\n"); - for (model, tokens) in &stats.token_by_model.buckets { - md.push_str(&format!("| {model} | {tokens} |\n")); - } - md.push_str(&format!( - "| **总计** | **{}** |\n\n", - stats.token_by_model.total - )); - - // 成功率 - md.push_str("## 成功率(按提供商)\n\n"); - md.push_str("| 提供商 | 成功率 |\n"); - md.push_str("|--------|--------|\n"); - for (provider, rate) in &stats.success_by_provider { - md.push_str(&format!("| {} | {:.1}% |\n", provider, rate * 100.0)); - } - md.push('\n'); - - // 延迟直方图 - md.push_str("## 延迟分布\n\n"); - md.push_str("| 延迟范围 | 请求数 |\n"); - md.push_str("|----------|--------|\n"); - for (range, count) in &stats.latency_histogram.buckets { - md.push_str(&format!("| {range} | {count} |\n")); - } - md.push('\n'); - - // 错误分布 - if !stats.error_distribution.buckets.is_empty() { - md.push_str("## 错误分布\n\n"); - md.push_str("| 错误类型 | 数量 |\n"); - md.push_str("|----------|------|\n"); - for (error_type, count) in &stats.error_distribution.buckets { - md.push_str(&format!("| {error_type} | {count} |\n")); - } - md.push('\n'); - } - - md - } - - /// 导出为 CSV 格式 - fn export_csv(&self, stats: &EnhancedStats) -> String { - let mut csv = String::new(); - - // Token 分布 - csv.push_str("# Token Distribution by Model\n"); - csv.push_str("Model,Tokens\n"); - for (model, tokens) in &stats.token_by_model.buckets { - csv.push_str(&format!("{model},{tokens}\n")); - } - csv.push('\n'); - - // 成功率 - csv.push_str("# Success Rate by Provider\n"); - csv.push_str("Provider,SuccessRate\n"); - for (provider, rate) in &stats.success_by_provider { - csv.push_str(&format!("{provider},{rate:.4}\n")); - } - csv.push('\n'); - - // 延迟直方图 - csv.push_str("# Latency Histogram\n"); - csv.push_str("Range,Count\n"); - for (range, count) in &stats.latency_histogram.buckets { - csv.push_str(&format!("{range},{count}\n")); - } - csv.push('\n'); - - // 错误分布 - csv.push_str("# Error Distribution\n"); - csv.push_str("ErrorType,Count\n"); - for (error_type, count) in &stats.error_distribution.buckets { - csv.push_str(&format!("{error_type},{count}\n")); - } - - csv - } -} - -// ============================================================================ -// 辅助函数 -// ============================================================================ - -/// 解析时间间隔字符串 -fn parse_interval(interval: &str) -> Duration { - let interval = interval.trim().to_lowercase(); - - if let Some(num_str) = interval.strip_suffix('h') { - if let Ok(hours) = num_str.parse::() { - return Duration::hours(hours); - } - } else if let Some(num_str) = interval.strip_suffix('m') { - if let Ok(minutes) = num_str.parse::() { - return Duration::minutes(minutes); - } - } else if let Some(num_str) = interval.strip_suffix('d') { - if let Ok(days) = num_str.parse::() { - return Duration::days(days); - } - } else if let Some(num_str) = interval.strip_suffix('s') { - if let Ok(seconds) = num_str.parse::() { - return Duration::seconds(seconds); - } - } - - // 默认 1 小时 - Duration::hours(1) -} - -/// 默认延迟桶边界(毫秒) -fn default_latency_buckets() -> Vec { - vec![100, 500, 1000, 2000, 5000, 10000] -} - -// ============================================================================ -// 测试模块 -// ============================================================================ - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_parse_interval() { - assert_eq!(parse_interval("1h"), Duration::hours(1)); - assert_eq!(parse_interval("30m"), Duration::minutes(30)); - assert_eq!(parse_interval("1d"), Duration::days(1)); - assert_eq!(parse_interval("60s"), Duration::seconds(60)); - assert_eq!(parse_interval("invalid"), Duration::hours(1)); // 默认值 - } - - #[test] - fn test_default_latency_buckets() { - let buckets = default_latency_buckets(); - assert_eq!(buckets, vec![100, 500, 1000, 2000, 5000, 10000]); - } - - #[test] - fn test_distribution_default() { - let dist = Distribution::default(); - assert!(dist.buckets.is_empty()); - assert_eq!(dist.total, 0); - } - - #[test] - fn test_trend_data_default() { - let trend = TrendData::default(); - assert!(trend.points.is_empty()); - assert_eq!(trend.interval, "1h"); - } - - #[test] - fn test_enhanced_stats_default() { - let stats = EnhancedStats::default(); - assert!(stats.request_trend.points.is_empty()); - assert!(stats.token_by_model.buckets.is_empty()); - assert!(stats.success_by_provider.is_empty()); - assert_eq!(stats.request_rate, 0.0); - } - - #[test] - fn test_report_format_default() { - let format = ReportFormat::default(); - assert_eq!(format, ReportFormat::Json); - } - - #[test] - fn test_stats_time_range_default() { - let range = StatsTimeRange::default(); - assert!(range.start < range.end); - // 默认应该是 24 小时范围 - let diff = range.end - range.start; - assert_eq!(diff.num_hours(), 24); - } -} - -// ============================================================================ -// 属性测试模块 -// ============================================================================ - -#[cfg(test)] -mod property_tests { - use super::*; - use crate::flow_monitor::models::{ - FlowMetadata, FlowState, FlowType, LLMRequest, LLMResponse, Message, MessageContent, - MessageRole, RequestParameters, TokenUsage, - }; - use crate::ProviderType; - use proptest::prelude::*; - - // ======================================================================== - // 生成器 - // ======================================================================== - - /// 生成随机的 ProviderType - fn arb_provider_type() -> impl Strategy { - prop_oneof![ - Just(ProviderType::Kiro), - Just(ProviderType::Gemini), - Just(ProviderType::OpenAI), - Just(ProviderType::Claude), - ] - } - - /// 生成随机的 FlowState - fn arb_flow_state() -> impl Strategy { - prop_oneof![ - Just(FlowState::Pending), - Just(FlowState::Streaming), - Just(FlowState::Completed), - Just(FlowState::Failed), - Just(FlowState::Cancelled), - ] - } - - /// 生成随机的模型名称 - fn arb_model_name() -> impl Strategy { - prop_oneof![ - Just("gpt-4".to_string()), - Just("gpt-4-turbo".to_string()), - Just("gpt-3.5-turbo".to_string()), - Just("claude-3-opus".to_string()), - Just("claude-3-sonnet".to_string()), - Just("gemini-pro".to_string()), - ] - } - - /// 生成随机的 TokenUsage - fn arb_token_usage() -> impl Strategy { - (0u32..10000u32, 0u32..5000u32).prop_map(|(input, output)| TokenUsage { - input_tokens: input, - output_tokens: output, - total_tokens: input + output, - ..Default::default() - }) - } - - /// 生成随机的 LLMFlow - fn arb_llm_flow() -> impl Strategy { - ( - "[a-f0-9]{8}", - arb_model_name(), - arb_provider_type(), - arb_flow_state(), - 0u64..10000u64, // duration_ms - arb_token_usage(), - any::(), // has_response - ) - .prop_map( - |(id, model, provider, state, duration, usage, has_response)| { - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model, - messages: vec![Message { - role: MessageRole::User, - content: MessageContent::Text("test".to_string()), - ..Default::default() - }], - parameters: RequestParameters::default(), - ..Default::default() - }; - - let metadata = FlowMetadata { - provider, - ..Default::default() - }; - - let mut flow = LLMFlow::new(id, FlowType::ChatCompletions, request, metadata); - flow.state = state; - flow.timestamps.duration_ms = duration; - - if has_response { - flow.response = Some(LLMResponse { - usage, - ..Default::default() - }); - } - - flow - }, - ) - } - - /// 生成随机的 Flow 列表 - fn arb_flow_list() -> impl Strategy> { - prop::collection::vec(arb_llm_flow(), 0..50) - } - - // ======================================================================== - // 属性测试 - // ======================================================================== - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: flow-monitor-enhancement, Property 17: 统计计算正确性** - /// **Validates: Requirements 9.1-9.6** - /// - /// *对于任意* Flow 集合和时间范围,统计计算应该正确反映该范围内的数据。 - #[test] - fn prop_stats_calculation_correctness(flows in arb_flow_list()) { - // 创建一个临时的 EnhancedStatsService 实例来测试内部计算方法 - let service = EnhancedStatsService::new( - Arc::new(RwLock::new(FlowMemoryStore::new(1000))) - ); - - // 测试 Token 分布计算 - let token_dist = service.calculate_token_distribution(&flows); - - // 验证: Token 分布的总数应该等于所有 Flow 的 Token 总和 - let expected_total: u64 = flows - .iter() - .filter_map(|f| f.response.as_ref()) - .map(|r| r.usage.total_tokens as u64) - .sum(); - prop_assert_eq!( - token_dist.total, - expected_total, - "Token 分布总数应该等于所有 Flow 的 Token 总和" - ); - - // 验证: 每个模型的 Token 数应该正确 - let bucket_total: u64 = token_dist.buckets.iter().map(|(_, count)| *count).sum(); - prop_assert_eq!( - bucket_total, - expected_total, - "所有桶的 Token 数之和应该等于总数" - ); - - // 测试成功率计算 - let success_by_provider = service.calculate_success_by_provider(&flows); - - // 验证: 成功率应该在 0.0 到 1.0 之间 - for (_, rate) in &success_by_provider { - prop_assert!( - *rate >= 0.0 && *rate <= 1.0, - "成功率应该在 0.0 到 1.0 之间,实际值: {}", - rate - ); - } - - // 测试延迟直方图计算 - let buckets = vec![100, 500, 1000, 2000, 5000, 10000]; - let latency_hist = service.calculate_latency_histogram(&flows, &buckets); - - // 验证: 直方图总数应该等于 Flow 数量 - prop_assert_eq!( - latency_hist.total, - flows.len() as u64, - "延迟直方图总数应该等于 Flow 数量" - ); - - // 验证: 所有桶的数量之和应该等于总数 - let hist_bucket_total: u64 = latency_hist.buckets.iter().map(|(_, count)| *count).sum(); - prop_assert_eq!( - hist_bucket_total, - latency_hist.total, - "所有直方图桶的数量之和应该等于总数" - ); - - // 测试错误分布计算 - let error_dist = service.calculate_error_distribution(&flows); - - // 验证: 错误分布总数应该等于有错误的 Flow 数量 - let expected_error_count = flows.iter().filter(|f| f.error.is_some()).count() as u64; - prop_assert_eq!( - error_dist.total, - expected_error_count, - "错误分布总数应该等于有错误的 Flow 数量" - ); - } - - /// **Feature: flow-monitor-enhancement, Property 17b: 请求趋势计算正确性** - /// **Validates: Requirements 9.1** - /// - /// *对于任意* Flow 集合,请求趋势的数据点值之和应该等于 Flow 总数。 - #[test] - fn prop_request_trend_correctness(flows in arb_flow_list()) { - let service = EnhancedStatsService::new( - Arc::new(RwLock::new(FlowMemoryStore::new(1000))) - ); - - // 测试请求趋势计算 - let trend = service.calculate_request_trend(&flows, "1h"); - - // 验证: 趋势数据点的值之和应该等于 Flow 数量 - let trend_total: f64 = trend.points.iter().map(|p| p.value).sum(); - prop_assert_eq!( - trend_total as usize, - flows.len(), - "趋势数据点的值之和应该等于 Flow 数量" - ); - - // 验证: 间隔应该正确设置 - prop_assert_eq!( - trend.interval, - "1h", - "趋势间隔应该正确设置" - ); - } - - /// **Feature: flow-monitor-enhancement, Property 17c: 请求速率计算正确性** - /// **Validates: Requirements 9.1** - /// - /// *对于任意* Flow 集合和时间范围,请求速率应该正确计算。 - #[test] - fn prop_request_rate_correctness(flows in arb_flow_list()) { - let service = EnhancedStatsService::new( - Arc::new(RwLock::new(FlowMemoryStore::new(1000))) - ); - - let now = Utc::now(); - let time_range = StatsTimeRange { - start: now - Duration::hours(1), - end: now, - }; - - let rate = service.calculate_request_rate(&flows, &time_range); - - // 验证: 请求速率应该非负 - prop_assert!( - rate >= 0.0, - "请求速率应该非负,实际值: {}", - rate - ); - - // 验证: 如果有 Flow,速率应该大于 0 - if !flows.is_empty() { - prop_assert!( - rate > 0.0, - "如果有 Flow,请求速率应该大于 0" - ); - } - - // 验证: 速率计算正确(Flow 数量 / 时间范围秒数) - let duration_secs = (time_range.end - time_range.start).num_seconds() as f64; - let expected_rate = flows.len() as f64 / duration_secs; - prop_assert!( - (rate - expected_rate).abs() < 0.0001, - "请求速率计算应该正确,期望: {}, 实际: {}", - expected_rate, - rate - ); - } - - /// **Feature: flow-monitor-enhancement, Property 17d: 报告导出正确性** - /// **Validates: Requirements 9.7** - /// - /// *对于任意* 统计数据,导出的报告应该包含所有必要信息。 - #[test] - fn prop_report_export_correctness(flows in arb_flow_list()) { - let service = EnhancedStatsService::new( - Arc::new(RwLock::new(FlowMemoryStore::new(1000))) - ); - - let now = Utc::now(); - let time_range = StatsTimeRange { - start: now - Duration::hours(24), - end: now, - }; - - // 计算统计数据 - let token_dist = service.calculate_token_distribution(&flows); - let success_by_provider = service.calculate_success_by_provider(&flows); - let latency_hist = service.calculate_latency_histogram(&flows, &default_latency_buckets()); - let error_dist = service.calculate_error_distribution(&flows); - let request_rate = service.calculate_request_rate(&flows, &time_range); - - let stats = EnhancedStats { - request_trend: TrendData::default(), - token_by_model: token_dist, - success_by_provider, - latency_histogram: latency_hist, - error_distribution: error_dist, - request_rate, - time_range: time_range.clone(), - }; - - // 测试 JSON 导出 - let json_report = service.export_json(&stats); - prop_assert!( - !json_report.is_empty(), - "JSON 报告不应该为空" - ); - // 验证 JSON 可以解析 - let parsed: Result = serde_json::from_str(&json_report); - prop_assert!( - parsed.is_ok(), - "JSON 报告应该可以解析回 EnhancedStats" - ); - - // 测试 Markdown 导出 - let md_report = service.export_markdown(&stats); - prop_assert!( - !md_report.is_empty(), - "Markdown 报告不应该为空" - ); - prop_assert!( - md_report.contains("# Flow 统计报告"), - "Markdown 报告应该包含标题" - ); - - // 测试 CSV 导出 - let csv_report = service.export_csv(&stats); - prop_assert!( - !csv_report.is_empty(), - "CSV 报告不应该为空" - ); - prop_assert!( - csv_report.contains("Model,Tokens"), - "CSV 报告应该包含 Token 分布表头" - ); - } - } -} diff --git a/src-tauri/src/flow_monitor/exporter.rs b/src-tauri/src/flow_monitor/exporter.rs deleted file mode 100644 index 59671f06b..000000000 --- a/src-tauri/src/flow_monitor/exporter.rs +++ /dev/null @@ -1,2161 +0,0 @@ -//! LLM Flow 导出服务 -//! -//! 提供多种格式的 Flow 导出功能,包括 HAR、JSON、JSONL、Markdown 和 CSV。 -//! 支持敏感数据脱敏和导出前过滤。 - -use regex::Regex; -use serde::{Deserialize, Serialize}; - -use super::models::{ - FlowAnnotations, FlowError, LLMFlow, LLMRequest, LLMResponse, Message, MessageContent, - ThinkingContent, -}; -use super::FlowFilter; -#[cfg(test)] -use crate::ProviderType; - -// ============================================================================ -// 导出格式枚举 -// ============================================================================ - -/// 导出格式 -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] -#[serde(rename_all = "lowercase")] -#[derive(Default)] -pub enum ExportFormat { - /// HAR (HTTP Archive) 格式 - HAR, - /// JSON 格式 - #[default] - JSON, - /// JSONL (JSON Lines) 格式 - JSONL, - /// Markdown 格式 - Markdown, - /// CSV 格式(仅元数据) - CSV, -} - -// ============================================================================ -// 导出选项 -// ============================================================================ - -/// 导出选项 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ExportOptions { - /// 导出格式 - pub format: ExportFormat, - /// 过滤条件 - #[serde(default)] - pub filter: Option, - /// 是否包含原始请求/响应体 - #[serde(default = "default_true")] - pub include_raw: bool, - /// 是否包含流式 chunks - #[serde(default)] - pub include_stream_chunks: bool, - /// 是否脱敏敏感数据 - #[serde(default)] - pub redact_sensitive: bool, - /// 脱敏规则 - #[serde(default)] - pub redaction_rules: Vec, - /// 是否压缩输出 - #[serde(default)] - pub compress: bool, -} - -fn default_true() -> bool { - true -} - -impl Default for ExportOptions { - fn default() -> Self { - Self { - format: ExportFormat::JSON, - filter: None, - include_raw: true, - include_stream_chunks: false, - redact_sensitive: false, - redaction_rules: Vec::new(), - compress: false, - } - } -} - -// ============================================================================ -// 脱敏规则 -// ============================================================================ - -/// 脱敏规则 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct RedactionRule { - /// 规则名称 - pub name: String, - /// 匹配模式(正则表达式) - pub pattern: String, - /// 替换文本 - pub replacement: String, - /// 是否启用 - #[serde(default = "default_true")] - pub enabled: bool, -} - -impl RedactionRule { - /// 创建新的脱敏规则 - pub fn new( - name: impl Into, - pattern: impl Into, - replacement: impl Into, - ) -> Self { - Self { - name: name.into(), - pattern: pattern.into(), - replacement: replacement.into(), - enabled: true, - } - } -} - -/// 获取默认脱敏规则 -pub fn default_redaction_rules() -> Vec { - vec![ - // API 密钥模式 - RedactionRule::new( - "api_key", - r"(?i)(sk-[a-zA-Z0-9]{20,}|api[_-]?key[=:]\s*[a-zA-Z0-9_-]{20,})", - "[REDACTED_API_KEY]", - ), - // Bearer Token - RedactionRule::new( - "bearer_token", - r"(?i)bearer\s+[a-zA-Z0-9_.-]+", - "Bearer [REDACTED_TOKEN]", - ), - // 邮箱地址 - RedactionRule::new( - "email", - r"[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}", - "[REDACTED_EMAIL]", - ), - // 手机号(中国大陆) - RedactionRule::new("phone_cn", r"1[3-9]\d{9}", "[REDACTED_PHONE]"), - // 手机号(国际格式) - RedactionRule::new( - "phone_intl", - r"\+\d{1,3}[-.\s]?\d{1,4}[-.\s]?\d{1,4}[-.\s]?\d{1,9}", - "[REDACTED_PHONE]", - ), - // 信用卡号 - RedactionRule::new( - "credit_card", - r"\b\d{4}[-\s]?\d{4}[-\s]?\d{4}[-\s]?\d{4}\b", - "[REDACTED_CARD]", - ), - // 身份证号(中国大陆) - RedactionRule::new("id_card_cn", r"\b\d{17}[\dXx]\b", "[REDACTED_ID]"), - // AWS 密钥 - RedactionRule::new( - "aws_key", - r"(?i)(AKIA[0-9A-Z]{16}|aws[_-]?secret[_-]?access[_-]?key[=:]\s*[a-zA-Z0-9/+=]{40})", - "[REDACTED_AWS_KEY]", - ), - // OpenAI API Key - RedactionRule::new("openai_key", r"sk-[a-zA-Z0-9]{48}", "[REDACTED_OPENAI_KEY]"), - // Anthropic API Key - RedactionRule::new( - "anthropic_key", - r"sk-ant-[a-zA-Z0-9_-]{95}", - "[REDACTED_ANTHROPIC_KEY]", - ), - ] -} - -// ============================================================================ -// 脱敏器 -// ============================================================================ - -/// 敏感数据脱敏器 -pub struct Redactor { - rules: Vec<(String, Regex, String)>, -} - -impl Redactor { - /// 创建新的脱敏器 - pub fn new(rules: &[RedactionRule]) -> Self { - let compiled_rules: Vec<_> = rules - .iter() - .filter(|r| r.enabled) - .filter_map(|r| { - Regex::new(&r.pattern) - .ok() - .map(|regex| (r.name.clone(), regex, r.replacement.clone())) - }) - .collect(); - - Self { - rules: compiled_rules, - } - } - - /// 使用默认规则创建脱敏器 - pub fn with_defaults() -> Self { - Self::new(&default_redaction_rules()) - } - - /// 对文本应用脱敏 - pub fn redact(&self, text: &str) -> String { - let mut result = text.to_string(); - for (_, regex, replacement) in &self.rules { - result = regex.replace_all(&result, replacement.as_str()).to_string(); - } - result - } - - /// 对 JSON 值应用脱敏 - pub fn redact_json(&self, value: &serde_json::Value) -> serde_json::Value { - match value { - serde_json::Value::String(s) => serde_json::Value::String(self.redact(s)), - serde_json::Value::Array(arr) => { - serde_json::Value::Array(arr.iter().map(|v| self.redact_json(v)).collect()) - } - serde_json::Value::Object(obj) => { - let mut new_obj = serde_json::Map::new(); - for (k, v) in obj { - new_obj.insert(k.clone(), self.redact_json(v)); - } - serde_json::Value::Object(new_obj) - } - other => other.clone(), - } - } - - /// 对 Flow 应用脱敏 - pub fn redact_flow(&self, flow: &LLMFlow) -> LLMFlow { - let mut redacted = flow.clone(); - - // 脱敏请求 - redacted.request = self.redact_request(&flow.request); - - // 脱敏响应 - if let Some(ref response) = flow.response { - redacted.response = Some(self.redact_response(response)); - } - - // 脱敏错误信息 - if let Some(ref error) = flow.error { - redacted.error = Some(self.redact_error(error)); - } - - // 脱敏标注 - redacted.annotations = self.redact_annotations(&flow.annotations); - - redacted - } - - fn redact_request(&self, request: &LLMRequest) -> LLMRequest { - let mut redacted = request.clone(); - - // 脱敏请求头 - redacted.headers = request - .headers - .iter() - .map(|(k, v)| { - let redacted_value = if k.to_lowercase().contains("authorization") - || k.to_lowercase().contains("api-key") - || k.to_lowercase().contains("x-api-key") - { - "[REDACTED]".to_string() - } else { - self.redact(v) - }; - (k.clone(), redacted_value) - }) - .collect(); - - // 脱敏请求体 - redacted.body = self.redact_json(&request.body); - - // 脱敏消息 - redacted.messages = request - .messages - .iter() - .map(|m| self.redact_message(m)) - .collect(); - - // 脱敏系统提示词 - redacted.system_prompt = request.system_prompt.as_ref().map(|s| self.redact(s)); - - redacted - } - - fn redact_message(&self, message: &Message) -> Message { - let mut redacted = message.clone(); - - redacted.content = match &message.content { - MessageContent::Text(s) => MessageContent::Text(self.redact(s)), - MessageContent::MultiModal(parts) => MessageContent::MultiModal( - parts - .iter() - .map(|p| match p { - super::models::ContentPart::Text { text } => { - super::models::ContentPart::Text { - text: self.redact(text), - } - } - other => other.clone(), - }) - .collect(), - ), - }; - - redacted - } - - fn redact_response(&self, response: &LLMResponse) -> LLMResponse { - let mut redacted = response.clone(); - - // 脱敏响应头 - redacted.headers = response - .headers - .iter() - .map(|(k, v)| (k.clone(), self.redact(v))) - .collect(); - - // 脱敏响应体 - redacted.body = self.redact_json(&response.body); - - // 脱敏内容 - redacted.content = self.redact(&response.content); - - // 脱敏思维链 - if let Some(ref thinking) = response.thinking { - redacted.thinking = Some(ThinkingContent { - text: self.redact(&thinking.text), - tokens: thinking.tokens, - signature: thinking.signature.clone(), - }); - } - - redacted - } - - fn redact_error(&self, error: &FlowError) -> FlowError { - let mut redacted = error.clone(); - redacted.message = self.redact(&error.message); - redacted.raw_response = error.raw_response.as_ref().map(|s| self.redact(s)); - redacted - } - - fn redact_annotations(&self, annotations: &FlowAnnotations) -> FlowAnnotations { - let mut redacted = annotations.clone(); - redacted.comment = annotations.comment.as_ref().map(|s| self.redact(s)); - redacted - } -} - -// ============================================================================ -// HAR 格式结构 -// ============================================================================ - -/// HAR 存档 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct HarArchive { - pub log: HarLog, -} - -/// HAR 日志 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct HarLog { - pub version: String, - pub creator: HarCreator, - pub entries: Vec, - #[serde(skip_serializing_if = "Option::is_none")] - pub comment: Option, -} - -/// HAR 创建者信息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct HarCreator { - pub name: String, - pub version: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub comment: Option, -} - -/// HAR 条目 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct HarEntry { - pub started_date_time: String, - pub time: f64, - pub request: HarRequest, - pub response: HarResponse, - pub cache: HarCache, - pub timings: HarTimings, - #[serde(skip_serializing_if = "Option::is_none")] - pub server_ip_address: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub connection: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub comment: Option, - /// LLM 特定扩展 - #[serde(rename = "_llm", skip_serializing_if = "Option::is_none")] - pub llm_extension: Option, -} - -/// HAR 请求 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct HarRequest { - pub method: String, - pub url: String, - pub http_version: String, - pub cookies: Vec, - pub headers: Vec, - pub query_string: Vec, - #[serde(skip_serializing_if = "Option::is_none")] - pub post_data: Option, - pub headers_size: i64, - pub body_size: i64, - #[serde(skip_serializing_if = "Option::is_none")] - pub comment: Option, -} - -/// HAR 响应 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct HarResponse { - pub status: u16, - pub status_text: String, - pub http_version: String, - pub cookies: Vec, - pub headers: Vec, - pub content: HarContent, - pub redirect_url: String, - pub headers_size: i64, - pub body_size: i64, - #[serde(skip_serializing_if = "Option::is_none")] - pub comment: Option, -} - -/// HAR Cookie -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct HarCookie { - pub name: String, - pub value: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub path: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub domain: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub expires: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub http_only: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub secure: Option, -} - -/// HAR 请求头 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct HarHeader { - pub name: String, - pub value: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub comment: Option, -} - -/// HAR 查询参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct HarQueryParam { - pub name: String, - pub value: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub comment: Option, -} - -/// HAR POST 数据 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct HarPostData { - pub mime_type: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub params: Option>, - pub text: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub comment: Option, -} - -/// HAR 参数 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct HarParam { - pub name: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub value: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub file_name: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub content_type: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub comment: Option, -} - -/// HAR 内容 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct HarContent { - pub size: i64, - #[serde(skip_serializing_if = "Option::is_none")] - pub compression: Option, - pub mime_type: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub text: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub encoding: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub comment: Option, -} - -/// HAR 缓存 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct HarCache { - #[serde(skip_serializing_if = "Option::is_none")] - pub before_request: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub after_request: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub comment: Option, -} - -/// HAR 缓存状态 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct HarCacheState { - #[serde(skip_serializing_if = "Option::is_none")] - pub expires: Option, - pub last_access: String, - pub e_tag: String, - pub hit_count: i64, - #[serde(skip_serializing_if = "Option::is_none")] - pub comment: Option, -} - -/// HAR 时间 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct HarTimings { - pub blocked: f64, - pub dns: f64, - pub connect: f64, - pub send: f64, - pub wait: f64, - pub receive: f64, - pub ssl: f64, - #[serde(skip_serializing_if = "Option::is_none")] - pub comment: Option, -} - -/// LLM 特定扩展 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct HarLlmExtension { - /// Flow ID - pub flow_id: String, - /// 提供商 - pub provider: String, - /// 模型 - pub model: String, - /// Flow 类型 - pub flow_type: String, - /// Flow 状态 - pub state: String, - /// Token 使用 - #[serde(skip_serializing_if = "Option::is_none")] - pub tokens: Option, - /// 是否流式 - pub streaming: bool, - /// TTFB(毫秒) - #[serde(skip_serializing_if = "Option::is_none")] - pub ttfb_ms: Option, - /// 停止原因 - #[serde(skip_serializing_if = "Option::is_none")] - pub stop_reason: Option, - /// 是否有工具调用 - pub has_tool_calls: bool, - /// 是否有思维链 - pub has_thinking: bool, - /// 标注 - #[serde(skip_serializing_if = "Option::is_none")] - pub annotations: Option, -} - -/// LLM Token 信息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct HarLlmTokens { - pub input: u32, - pub output: u32, - pub total: u32, - #[serde(skip_serializing_if = "Option::is_none")] - pub cache_read: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub cache_write: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub thinking: Option, -} - -// ============================================================================ -// Flow 导出器 -// ============================================================================ - -/// Flow 导出器 -pub struct FlowExporter { - options: ExportOptions, - redactor: Option, -} - -impl FlowExporter { - /// 创建新的导出器 - pub fn new(options: ExportOptions) -> Self { - let redactor = if options.redact_sensitive { - let rules = if options.redaction_rules.is_empty() { - default_redaction_rules() - } else { - options.redaction_rules.clone() - }; - Some(Redactor::new(&rules)) - } else { - None - }; - - Self { options, redactor } - } - - /// 使用默认选项创建导出器 - pub fn with_defaults() -> Self { - Self::new(ExportOptions::default()) - } - - /// 预处理 Flow(应用脱敏等) - fn preprocess_flow(&self, flow: &LLMFlow) -> LLMFlow { - if let Some(ref redactor) = self.redactor { - redactor.redact_flow(flow) - } else { - flow.clone() - } - } - - /// 预处理多个 Flow - fn preprocess_flows(&self, flows: &[LLMFlow]) -> Vec { - flows.iter().map(|f| self.preprocess_flow(f)).collect() - } - - /// 导出为 HAR 格式 - pub fn export_har(&self, flows: &[LLMFlow]) -> HarArchive { - let processed = self.preprocess_flows(flows); - let entries: Vec = processed - .iter() - .map(|f| self.flow_to_har_entry(f)) - .collect(); - - HarArchive { - log: HarLog { - version: "1.2".to_string(), - creator: HarCreator { - name: "ProxyCast LLM Flow Monitor".to_string(), - version: env!("CARGO_PKG_VERSION").to_string(), - comment: Some("LLM API Flow Export".to_string()), - }, - entries, - comment: Some(format!("Exported {} flows", flows.len())), - }, - } - } - - /// 将 Flow 转换为 HAR Entry - fn flow_to_har_entry(&self, flow: &LLMFlow) -> HarEntry { - let request = &flow.request; - let response = flow.response.as_ref(); - - // 构建请求 URL - let base_url = flow - .metadata - .routing_info - .target_url - .clone() - .unwrap_or_else(|| "http://localhost".to_string()); - let url = format!("{}{}", base_url, request.path); - - // 构建请求头 - let headers: Vec = request - .headers - .iter() - .map(|(k, v)| HarHeader { - name: k.clone(), - value: v.clone(), - comment: None, - }) - .collect(); - - // 构建 POST 数据 - let post_data = if self.options.include_raw { - Some(HarPostData { - mime_type: "application/json".to_string(), - params: None, - text: serde_json::to_string(&request.body).unwrap_or_default(), - comment: None, - }) - } else { - None - }; - - // 构建响应 - let (har_response, _response_body_size) = if let Some(resp) = response { - let resp_headers: Vec = resp - .headers - .iter() - .map(|(k, v)| HarHeader { - name: k.clone(), - value: v.clone(), - comment: None, - }) - .collect(); - - let content_text = if self.options.include_raw { - Some(serde_json::to_string(&resp.body).unwrap_or_default()) - } else { - None - }; - - ( - HarResponse { - status: resp.status_code, - status_text: resp.status_text.clone(), - http_version: "HTTP/1.1".to_string(), - cookies: Vec::new(), - headers: resp_headers, - content: HarContent { - size: resp.size_bytes as i64, - compression: None, - mime_type: "application/json".to_string(), - text: content_text, - encoding: None, - comment: None, - }, - redirect_url: String::new(), - headers_size: -1, - body_size: resp.size_bytes as i64, - comment: None, - }, - resp.size_bytes as i64, - ) - } else { - ( - HarResponse { - status: 0, - status_text: "No Response".to_string(), - http_version: "HTTP/1.1".to_string(), - cookies: Vec::new(), - headers: Vec::new(), - content: HarContent { - size: 0, - compression: None, - mime_type: "application/json".to_string(), - text: None, - encoding: None, - comment: None, - }, - redirect_url: String::new(), - headers_size: -1, - body_size: 0, - comment: None, - }, - 0, - ) - }; - - // 构建 LLM 扩展 - let llm_extension = Some(HarLlmExtension { - flow_id: flow.id.clone(), - provider: format!("{:?}", flow.metadata.provider), - model: request.model.clone(), - flow_type: format!("{:?}", flow.flow_type), - state: format!("{:?}", flow.state), - tokens: response.map(|r| HarLlmTokens { - input: r.usage.input_tokens, - output: r.usage.output_tokens, - total: r.usage.total_tokens, - cache_read: r.usage.cache_read_tokens, - cache_write: r.usage.cache_write_tokens, - thinking: r.usage.thinking_tokens, - }), - streaming: request.parameters.stream, - ttfb_ms: flow.timestamps.ttfb_ms, - stop_reason: response.and_then(|r| r.stop_reason.as_ref().map(|s| format!("{s:?}"))), - has_tool_calls: response.map(|r| !r.tool_calls.is_empty()).unwrap_or(false), - has_thinking: response.map(|r| r.thinking.is_some()).unwrap_or(false), - annotations: if flow.annotations.starred - || flow.annotations.comment.is_some() - || !flow.annotations.tags.is_empty() - { - Some(flow.annotations.clone()) - } else { - None - }, - }); - - // 计算时间 - let ttfb = flow.timestamps.ttfb_ms.unwrap_or(0) as f64; - let total_time = flow.timestamps.duration_ms as f64; - - HarEntry { - started_date_time: flow.timestamps.request_start.to_rfc3339(), - time: total_time, - request: HarRequest { - method: request.method.clone(), - url, - http_version: "HTTP/1.1".to_string(), - cookies: Vec::new(), - headers, - query_string: Vec::new(), - post_data, - headers_size: -1, - body_size: request.size_bytes as i64, - comment: None, - }, - response: har_response, - cache: HarCache { - before_request: None, - after_request: None, - comment: None, - }, - timings: HarTimings { - blocked: -1.0, - dns: -1.0, - connect: -1.0, - send: 0.0, - wait: ttfb, - receive: total_time - ttfb, - ssl: -1.0, - comment: None, - }, - server_ip_address: None, - connection: None, - comment: flow.annotations.comment.clone(), - llm_extension, - } - } - - /// 导出为 JSON 格式 - pub fn export_json(&self, flows: &[LLMFlow]) -> serde_json::Value { - let processed = self.preprocess_flows(flows); - serde_json::to_value(&processed).unwrap_or(serde_json::Value::Array(Vec::new())) - } - - /// 导出为 JSONL 格式 - pub fn export_jsonl(&self, flows: &[LLMFlow]) -> String { - let processed = self.preprocess_flows(flows); - processed - .iter() - .filter_map(|f| serde_json::to_string(f).ok()) - .collect::>() - .join("\n") - } - - /// 导出单个 Flow 为 Markdown 格式 - pub fn export_markdown(&self, flow: &LLMFlow) -> String { - let processed = self.preprocess_flow(flow); - self.flow_to_markdown(&processed) - } - - /// 导出多个 Flow 为 Markdown 格式 - pub fn export_markdown_multiple(&self, flows: &[LLMFlow]) -> String { - let processed = self.preprocess_flows(flows); - processed - .iter() - .enumerate() - .map(|(i, f)| { - let md = self.flow_to_markdown(f); - if i > 0 { - format!("\n---\n\n{md}") - } else { - md - } - }) - .collect::>() - .join("") - } - - /// 将 Flow 转换为 Markdown - fn flow_to_markdown(&self, flow: &LLMFlow) -> String { - let mut md = String::new(); - - // 标题 - md.push_str(&format!("# LLM Flow: {}\n\n", flow.id)); - - // 元信息 - md.push_str("## 基本信息\n\n"); - md.push_str(&format!("- **Flow ID**: `{}`\n", flow.id)); - md.push_str(&format!("- **类型**: {:?}\n", flow.flow_type)); - md.push_str(&format!("- **状态**: {:?}\n", flow.state)); - md.push_str(&format!("- **提供商**: {:?}\n", flow.metadata.provider)); - md.push_str(&format!("- **模型**: {}\n", flow.request.model)); - md.push_str(&format!( - "- **创建时间**: {}\n", - flow.timestamps.created.format("%Y-%m-%d %H:%M:%S UTC") - )); - md.push_str(&format!("- **耗时**: {} ms\n", flow.timestamps.duration_ms)); - if let Some(ttfb) = flow.timestamps.ttfb_ms { - md.push_str(&format!("- **TTFB**: {ttfb} ms\n")); - } - md.push_str(&format!("- **流式**: {}\n", flow.request.parameters.stream)); - md.push('\n'); - - // Token 使用 - if let Some(ref response) = flow.response { - md.push_str("## Token 使用\n\n"); - md.push_str(&format!( - "- **输入 Token**: {}\n", - response.usage.input_tokens - )); - md.push_str(&format!( - "- **输出 Token**: {}\n", - response.usage.output_tokens - )); - md.push_str(&format!( - "- **总 Token**: {}\n", - response.usage.total_tokens - )); - if let Some(cache_read) = response.usage.cache_read_tokens { - md.push_str(&format!("- **缓存读取**: {cache_read}\n")); - } - if let Some(thinking) = response.usage.thinking_tokens { - md.push_str(&format!("- **思维链 Token**: {thinking}\n")); - } - md.push('\n'); - } - - // 请求 - md.push_str("## 请求\n\n"); - md.push_str(&format!( - "**{} {}**\n\n", - flow.request.method, flow.request.path - )); - - // 系统提示词 - if let Some(ref system) = flow.request.system_prompt { - md.push_str("### 系统提示词\n\n"); - md.push_str("```\n"); - md.push_str(system); - md.push_str("\n```\n\n"); - } - - // 消息 - if !flow.request.messages.is_empty() { - md.push_str("### 消息\n\n"); - for (i, msg) in flow.request.messages.iter().enumerate() { - md.push_str(&format!( - "#### {} {}\n\n", - i + 1, - format!("{:?}", msg.role).to_uppercase() - )); - let content = msg.content.get_all_text(); - if !content.is_empty() { - md.push_str("```\n"); - md.push_str(&content); - md.push_str("\n```\n\n"); - } - } - } - - // 响应 - if let Some(ref response) = flow.response { - md.push_str("## 响应\n\n"); - md.push_str(&format!( - "**状态**: {} {}\n\n", - response.status_code, response.status_text - )); - - // 思维链 - if let Some(ref thinking) = response.thinking { - md.push_str("### 思维链\n\n"); - md.push_str("
\n展开查看思维链内容\n\n"); - md.push_str("```\n"); - md.push_str(&thinking.text); - md.push_str("\n```\n\n"); - md.push_str("
\n\n"); - } - - // 内容 - if !response.content.is_empty() { - md.push_str("### 内容\n\n"); - md.push_str("```\n"); - md.push_str(&response.content); - md.push_str("\n```\n\n"); - } - - // 工具调用 - if !response.tool_calls.is_empty() { - md.push_str("### 工具调用\n\n"); - for (i, tc) in response.tool_calls.iter().enumerate() { - md.push_str(&format!("#### 工具调用 {}\n\n", i + 1)); - md.push_str(&format!("- **ID**: `{}`\n", tc.id)); - md.push_str(&format!("- **函数**: `{}`\n", tc.function.name)); - md.push_str("- **参数**:\n"); - md.push_str("```json\n"); - // 尝试格式化 JSON - if let Ok(parsed) = - serde_json::from_str::(&tc.function.arguments) - { - md.push_str( - &serde_json::to_string_pretty(&parsed) - .unwrap_or(tc.function.arguments.clone()), - ); - } else { - md.push_str(&tc.function.arguments); - } - md.push_str("\n```\n\n"); - } - } - - // 停止原因 - if let Some(ref stop_reason) = response.stop_reason { - md.push_str(&format!("**停止原因**: {stop_reason:?}\n\n")); - } - } - - // 错误 - if let Some(ref error) = flow.error { - md.push_str("## 错误\n\n"); - md.push_str(&format!("- **类型**: {:?}\n", error.error_type)); - md.push_str(&format!("- **消息**: {}\n", error.message)); - if let Some(code) = error.status_code { - md.push_str(&format!("- **状态码**: {code}\n")); - } - md.push_str(&format!("- **可重试**: {}\n", error.retryable)); - md.push('\n'); - } - - // 标注 - if flow.annotations.starred - || flow.annotations.comment.is_some() - || !flow.annotations.tags.is_empty() - { - md.push_str("## 标注\n\n"); - if flow.annotations.starred { - md.push_str("- ⭐ **已收藏**\n"); - } - if let Some(ref marker) = flow.annotations.marker { - md.push_str(&format!("- **标记**: {marker}\n")); - } - if !flow.annotations.tags.is_empty() { - md.push_str(&format!( - "- **标签**: {}\n", - flow.annotations.tags.join(", ") - )); - } - if let Some(ref comment) = flow.annotations.comment { - md.push_str(&format!("- **评论**: {comment}\n")); - } - md.push('\n'); - } - - md - } - - /// 导出为 CSV 格式(仅元数据) - pub fn export_csv(&self, flows: &[LLMFlow]) -> String { - let processed = self.preprocess_flows(flows); - let mut csv = String::new(); - - // CSV 头 - csv.push_str("id,created_at,provider,model,flow_type,state,method,path,"); - csv.push_str("status_code,duration_ms,ttfb_ms,input_tokens,output_tokens,total_tokens,"); - csv.push_str("streaming,has_error,has_tool_calls,has_thinking,starred,tags\n"); - - // 数据行 - for flow in &processed { - let response = flow.response.as_ref(); - let row = format!( - "{},{},{:?},{},{:?},{:?},{},{},{},{},{},{},{},{},{},{},{},{},{},{}\n", - escape_csv(&flow.id), - flow.timestamps.created.to_rfc3339(), - flow.metadata.provider, - escape_csv(&flow.request.model), - flow.flow_type, - flow.state, - escape_csv(&flow.request.method), - escape_csv(&flow.request.path), - response.map(|r| r.status_code).unwrap_or(0), - flow.timestamps.duration_ms, - flow.timestamps.ttfb_ms.unwrap_or(0), - response.map(|r| r.usage.input_tokens).unwrap_or(0), - response.map(|r| r.usage.output_tokens).unwrap_or(0), - response.map(|r| r.usage.total_tokens).unwrap_or(0), - flow.request.parameters.stream, - flow.error.is_some(), - response.map(|r| !r.tool_calls.is_empty()).unwrap_or(false), - response.map(|r| r.thinking.is_some()).unwrap_or(false), - flow.annotations.starred, - escape_csv(&flow.annotations.tags.join(";")) - ); - csv.push_str(&row); - } - - csv - } - - /// 根据选项导出 - pub fn export(&self, flows: &[LLMFlow]) -> ExportResult { - match self.options.format { - ExportFormat::HAR => { - let har = self.export_har(flows); - ExportResult::Har(har) - } - ExportFormat::JSON => { - let json = self.export_json(flows); - ExportResult::Json(json) - } - ExportFormat::JSONL => { - let jsonl = self.export_jsonl(flows); - ExportResult::Text(jsonl) - } - ExportFormat::Markdown => { - let md = self.export_markdown_multiple(flows); - ExportResult::Text(md) - } - ExportFormat::CSV => { - let csv = self.export_csv(flows); - ExportResult::Text(csv) - } - } - } -} - -/// CSV 字段转义 -fn escape_csv(s: &str) -> String { - if s.contains(',') || s.contains('"') || s.contains('\n') { - format!("\"{}\"", s.replace('"', "\"\"")) - } else { - s.to_string() - } -} - -/// 导出结果 -#[derive(Debug, Clone)] -pub enum ExportResult { - /// HAR 格式 - Har(HarArchive), - /// JSON 格式 - Json(serde_json::Value), - /// 文本格式(JSONL、Markdown、CSV) - Text(String), -} - -impl ExportResult { - /// 转换为字符串 - pub fn to_string_pretty(&self) -> String { - match self { - ExportResult::Har(har) => serde_json::to_string_pretty(har).unwrap_or_default(), - ExportResult::Json(json) => serde_json::to_string_pretty(json).unwrap_or_default(), - ExportResult::Text(text) => text.clone(), - } - } - - /// 转换为紧凑字符串 - pub fn to_string_compact(&self) -> String { - match self { - ExportResult::Har(har) => serde_json::to_string(har).unwrap_or_default(), - ExportResult::Json(json) => serde_json::to_string(json).unwrap_or_default(), - ExportResult::Text(text) => text.clone(), - } - } -} - -// ============================================================================ -// 测试模块 -// ============================================================================ - -#[cfg(test)] -mod tests { - use super::*; - use crate::flow_monitor::models::*; - use chrono::Utc; - use std::collections::HashMap; - - fn create_test_flow() -> LLMFlow { - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - headers: { - let mut h = HashMap::new(); - h.insert( - "Authorization".to_string(), - "Bearer sk-test123456789".to_string(), - ); - h.insert("Content-Type".to_string(), "application/json".to_string()); - h - }, - body: serde_json::json!({ - "model": "gpt-4", - "messages": [{"role": "user", "content": "Hello"}] - }), - messages: vec![Message { - role: MessageRole::User, - content: MessageContent::Text("Hello, my email is test@example.com".to_string()), - tool_calls: None, - tool_result: None, - name: None, - }], - system_prompt: Some("You are a helpful assistant.".to_string()), - tools: None, - model: "gpt-4".to_string(), - original_model: None, - parameters: RequestParameters { - temperature: Some(0.7), - stream: true, - ..Default::default() - }, - size_bytes: 256, - timestamp: Utc::now(), - }; - - let response = LLMResponse { - status_code: 200, - status_text: "OK".to_string(), - headers: HashMap::new(), - body: serde_json::json!({"choices": [{"message": {"content": "Hi there!"}}]}), - content: "Hi there!".to_string(), - thinking: None, - tool_calls: Vec::new(), - usage: TokenUsage { - input_tokens: 10, - output_tokens: 5, - total_tokens: 15, - ..Default::default() - }, - stop_reason: Some(StopReason::Stop), - size_bytes: 128, - timestamp_start: Utc::now(), - timestamp_end: Utc::now(), - stream_info: None, - }; - - let metadata = FlowMetadata { - provider: ProviderType::OpenAI, - credential_id: Some("cred-123".to_string()), - credential_name: Some("Test Credential".to_string()), - ..Default::default() - }; - - let mut flow = LLMFlow::new( - "test-flow-001".to_string(), - FlowType::ChatCompletions, - request, - metadata, - ); - flow.response = Some(response); - flow.state = FlowState::Completed; - flow.timestamps.duration_ms = 500; - flow.timestamps.ttfb_ms = Some(100); - - flow - } - - #[test] - fn test_export_format_default() { - assert_eq!(ExportFormat::default(), ExportFormat::JSON); - } - - #[test] - fn test_export_options_default() { - let options = ExportOptions::default(); - assert_eq!(options.format, ExportFormat::JSON); - assert!(options.include_raw); - assert!(!options.redact_sensitive); - } - - #[test] - fn test_redaction_rule_creation() { - let rule = RedactionRule::new("test", r"\d+", "[NUMBER]"); - assert_eq!(rule.name, "test"); - assert_eq!(rule.pattern, r"\d+"); - assert_eq!(rule.replacement, "[NUMBER]"); - assert!(rule.enabled); - } - - #[test] - fn test_default_redaction_rules() { - let rules = default_redaction_rules(); - assert!(!rules.is_empty()); - - // 验证包含常见规则 - let rule_names: Vec<_> = rules.iter().map(|r| r.name.as_str()).collect(); - assert!(rule_names.contains(&"api_key")); - assert!(rule_names.contains(&"email")); - assert!(rule_names.contains(&"phone_cn")); - } - - #[test] - fn test_redactor_email() { - let redactor = Redactor::with_defaults(); - let text = "Contact me at john@example.com for more info."; - let redacted = redactor.redact(text); - assert!(!redacted.contains("john@example.com")); - assert!(redacted.contains("[REDACTED_EMAIL]")); - } - - #[test] - fn test_redactor_phone() { - let redactor = Redactor::with_defaults(); - let text = "My phone is 13812345678"; - let redacted = redactor.redact(text); - assert!(!redacted.contains("13812345678")); - assert!(redacted.contains("[REDACTED_PHONE]")); - } - - #[test] - fn test_redactor_api_key() { - let redactor = Redactor::with_defaults(); - let text = "Use this key: sk-abcdefghijklmnopqrstuvwxyz123456"; - let redacted = redactor.redact(text); - assert!(!redacted.contains("sk-abcdefghijklmnopqrstuvwxyz123456")); - } - - #[test] - fn test_redactor_json() { - let redactor = Redactor::with_defaults(); - let json = serde_json::json!({ - "email": "test@example.com", - "nested": { - "phone": "13812345678" - } - }); - let redacted = redactor.redact_json(&json); - let redacted_str = serde_json::to_string(&redacted).unwrap(); - assert!(!redacted_str.contains("test@example.com")); - assert!(!redacted_str.contains("13812345678")); - } - - #[test] - fn test_export_json() { - let flow = create_test_flow(); - let exporter = FlowExporter::with_defaults(); - let json = exporter.export_json(&[flow]); - - assert!(json.is_array()); - let arr = json.as_array().unwrap(); - assert_eq!(arr.len(), 1); - } - - #[test] - fn test_export_jsonl() { - let flow = create_test_flow(); - let exporter = FlowExporter::with_defaults(); - let jsonl = exporter.export_jsonl(&[flow.clone(), flow]); - - let lines: Vec<_> = jsonl.lines().collect(); - assert_eq!(lines.len(), 2); - - // 验证每行都是有效的 JSON - for line in lines { - assert!(serde_json::from_str::(line).is_ok()); - } - } - - #[test] - fn test_export_har() { - let flow = create_test_flow(); - let exporter = FlowExporter::with_defaults(); - let har = exporter.export_har(&[flow]); - - assert_eq!(har.log.version, "1.2"); - assert_eq!(har.log.entries.len(), 1); - - let entry = &har.log.entries[0]; - assert_eq!(entry.request.method, "POST"); - assert!(entry.llm_extension.is_some()); - - let llm_ext = entry.llm_extension.as_ref().unwrap(); - assert_eq!(llm_ext.model, "gpt-4"); - assert!(llm_ext.streaming); - } - - #[test] - fn test_export_markdown() { - let flow = create_test_flow(); - let exporter = FlowExporter::with_defaults(); - let md = exporter.export_markdown(&flow); - - assert!(md.contains("# LLM Flow:")); - assert!(md.contains("test-flow-001")); - assert!(md.contains("gpt-4")); - assert!(md.contains("## 请求")); - assert!(md.contains("## 响应")); - } - - #[test] - fn test_export_csv() { - let flow = create_test_flow(); - let exporter = FlowExporter::with_defaults(); - let csv = exporter.export_csv(&[flow]); - - let lines: Vec<_> = csv.lines().collect(); - assert_eq!(lines.len(), 2); // header + 1 data row - - // 验证头部 - assert!(lines[0].contains("id,created_at,provider")); - - // 验证数据行 - assert!(lines[1].contains("test-flow-001")); - } - - #[test] - fn test_export_with_redaction() { - let flow = create_test_flow(); - let options = ExportOptions { - format: ExportFormat::JSON, - redact_sensitive: true, - ..Default::default() - }; - let exporter = FlowExporter::new(options); - let json = exporter.export_json(&[flow]); - - let json_str = serde_json::to_string(&json).unwrap(); - // 验证敏感数据已被脱敏 - assert!(!json_str.contains("test@example.com")); - } - - #[test] - fn test_export_result_to_string() { - let flow = create_test_flow(); - let exporter = FlowExporter::with_defaults(); - let result = exporter.export(&[flow]); - - let pretty = result.to_string_pretty(); - let compact = result.to_string_compact(); - - assert!(!pretty.is_empty()); - assert!(!compact.is_empty()); - // Pretty 格式应该比 compact 更长(有缩进) - assert!(pretty.len() >= compact.len()); - } - - #[test] - fn test_escape_csv() { - assert_eq!(escape_csv("simple"), "simple"); - assert_eq!(escape_csv("with,comma"), "\"with,comma\""); - assert_eq!(escape_csv("with\"quote"), "\"with\"\"quote\""); - assert_eq!(escape_csv("with\nnewline"), "\"with\nnewline\""); - } - - #[test] - fn test_har_llm_extension() { - let flow = create_test_flow(); - let exporter = FlowExporter::with_defaults(); - let har = exporter.export_har(&[flow]); - - let entry = &har.log.entries[0]; - let llm_ext = entry.llm_extension.as_ref().unwrap(); - - assert_eq!(llm_ext.flow_id, "test-flow-001"); - assert!(llm_ext.tokens.is_some()); - - let tokens = llm_ext.tokens.as_ref().unwrap(); - assert_eq!(tokens.input, 10); - assert_eq!(tokens.output, 5); - assert_eq!(tokens.total, 15); - } -} - -// ============================================================================ -// 属性测试模块 -// ============================================================================ - -#[cfg(test)] -mod property_tests { - use super::*; - use crate::flow_monitor::models::*; - use chrono::Utc; - use proptest::prelude::*; - use std::collections::HashMap; - - // ======================================================================== - // 生成器 - // ======================================================================== - - /// 生成随机的 ProviderType - fn arb_provider_type() -> impl Strategy { - prop_oneof![ - Just(ProviderType::Kiro), - Just(ProviderType::Gemini), - Just(ProviderType::OpenAI), - Just(ProviderType::Claude), - Just(ProviderType::Antigravity), - ] - } - - /// 生成随机的 FlowType - fn arb_flow_type() -> impl Strategy { - prop_oneof![ - Just(FlowType::ChatCompletions), - Just(FlowType::AnthropicMessages), - Just(FlowType::GeminiGenerateContent), - Just(FlowType::Embeddings), - ] - } - - /// 生成随机的 MessageRole - fn arb_message_role() -> impl Strategy { - prop_oneof![ - Just(MessageRole::System), - Just(MessageRole::User), - Just(MessageRole::Assistant), - ] - } - - /// 生成随机的文本内容(不包含敏感数据) - fn arb_safe_text() -> impl Strategy { - "[a-zA-Z0-9 ,.!?]{0,100}" - } - - /// 生成随机的 MessageContent - fn arb_message_content() -> impl Strategy { - arb_safe_text().prop_map(MessageContent::Text) - } - - /// 生成随机的 Message - fn arb_message() -> impl Strategy { - (arb_message_role(), arb_message_content()).prop_map(|(role, content)| Message { - role, - content, - tool_calls: None, - tool_result: None, - name: None, - }) - } - - /// 生成随机的 RequestParameters - fn arb_request_parameters() -> impl Strategy { - ( - prop::option::of(0.0f32..2.0f32), - prop::option::of(0.0f32..1.0f32), - prop::option::of(1u32..4096u32), - any::(), - ) - .prop_map( - |(temperature, top_p, max_tokens, stream)| RequestParameters { - temperature, - top_p, - max_tokens, - stop: None, - stream, - extra: HashMap::new(), - }, - ) - } - - /// 生成随机的 LLMRequest - fn arb_llm_request() -> impl Strategy { - ( - "[a-z0-9-]{3,20}", // model - prop::collection::vec(arb_message(), 0..3), // messages - arb_request_parameters(), // parameters - prop::option::of(arb_safe_text()), // system_prompt - ) - .prop_map(|(model, messages, parameters, system_prompt)| LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - headers: HashMap::new(), - body: serde_json::Value::Null, - messages, - system_prompt, - tools: None, - model, - original_model: None, - parameters, - size_bytes: 0, - timestamp: Utc::now(), - }) - } - - /// 生成随机的 TokenUsage - fn arb_token_usage() -> impl Strategy { - (0u32..10000u32, 0u32..10000u32).prop_map(|(input, output)| TokenUsage { - input_tokens: input, - output_tokens: output, - total_tokens: input + output, - cache_read_tokens: None, - cache_write_tokens: None, - thinking_tokens: None, - }) - } - - /// 生成随机的 LLMResponse - fn arb_llm_response() -> impl Strategy { - (arb_safe_text(), arb_token_usage()).prop_map(|(content, usage)| LLMResponse { - status_code: 200, - status_text: "OK".to_string(), - headers: HashMap::new(), - body: serde_json::Value::Null, - content, - thinking: None, - tool_calls: Vec::new(), - usage, - stop_reason: Some(StopReason::Stop), - size_bytes: 0, - timestamp_start: Utc::now(), - timestamp_end: Utc::now(), - stream_info: None, - }) - } - - /// 生成随机的 FlowMetadata - fn arb_flow_metadata() -> impl Strategy { - arb_provider_type().prop_map(|provider| FlowMetadata { - provider, - provider_id: None, - credential_id: None, - credential_name: None, - retry_count: 0, - client_info: ClientInfo::default(), - routing_info: RoutingInfo::default(), - injected_params: None, - context_usage_percentage: None, - }) - } - - /// 生成随机的 Flow ID - fn arb_flow_id() -> impl Strategy { - "[a-f0-9]{8}-[a-f0-9]{4}-[a-f0-9]{4}-[a-f0-9]{4}-[a-f0-9]{12}" - } - - /// 生成随机的 LLMFlow - fn arb_llm_flow() -> impl Strategy { - ( - arb_flow_id(), - arb_flow_type(), - arb_llm_request(), - arb_flow_metadata(), - prop::option::of(arb_llm_response()), - ) - .prop_map(|(id, flow_type, request, metadata, response)| { - let mut flow = LLMFlow::new(id, flow_type, request, metadata); - flow.response = response; - if flow.response.is_some() { - flow.state = FlowState::Completed; - } - flow.timestamps.duration_ms = 100; - flow - }) - } - - // ======================================================================== - // 属性测试 - // ======================================================================== - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: llm-flow-monitor, Property 8: 导出 Round-Trip** - /// **Validates: Requirements 5.2** - /// - /// *对于任意* 有效的 LLM_Flow,导出为 JSON 格式后再解析, - /// 解析的 Flow 应该与原始 Flow 等价。 - #[test] - fn prop_export_json_roundtrip(flow in arb_llm_flow()) { - let exporter = FlowExporter::with_defaults(); - - // 导出为 JSON - let json = exporter.export_json(&[flow.clone()]); - - // 验证是数组 - prop_assert!(json.is_array(), "导出结果应该是 JSON 数组"); - - let arr = json.as_array().unwrap(); - prop_assert_eq!(arr.len(), 1, "数组应该包含一个元素"); - - // 反序列化 - let deserialized: LLMFlow = serde_json::from_value(arr[0].clone()) - .expect("应该能够反序列化"); - - // 验证关键字段一致 - prop_assert_eq!(&flow.id, &deserialized.id, "ID 应该在往返后保持一致"); - prop_assert_eq!(flow.state, deserialized.state, "状态应该在往返后保持一致"); - prop_assert_eq!(flow.flow_type, deserialized.flow_type, "FlowType 应该在往返后保持一致"); - prop_assert_eq!(&flow.request.model, &deserialized.request.model, "模型应该在往返后保持一致"); - prop_assert_eq!(&flow.request.method, &deserialized.request.method, "方法应该在往返后保持一致"); - prop_assert_eq!(flow.metadata.provider, deserialized.metadata.provider, "Provider 应该在往返后保持一致"); - - // 验证响应 - prop_assert_eq!(flow.response.is_some(), deserialized.response.is_some(), "响应存在性应该一致"); - if let (Some(ref orig), Some(ref deser)) = (&flow.response, &deserialized.response) { - prop_assert_eq!(orig.status_code, deser.status_code, "状态码应该一致"); - prop_assert_eq!(&orig.content, &deser.content, "内容应该一致"); - prop_assert_eq!(orig.usage.input_tokens, deser.usage.input_tokens, "输入 Token 应该一致"); - prop_assert_eq!(orig.usage.output_tokens, deser.usage.output_tokens, "输出 Token 应该一致"); - } - } - - /// **Feature: llm-flow-monitor, Property 8b: JSONL 导出 Round-Trip** - /// **Validates: Requirements 5.3** - /// - /// *对于任意* 有效的 LLM_Flow 列表,导出为 JSONL 格式后再解析, - /// 每行都应该能够正确反序列化为 LLMFlow。 - #[test] - fn prop_export_jsonl_roundtrip( - flows in prop::collection::vec(arb_llm_flow(), 1..5) - ) { - let exporter = FlowExporter::with_defaults(); - - // 导出为 JSONL - let jsonl = exporter.export_jsonl(&flows); - - // 验证行数 - let lines: Vec<_> = jsonl.lines().collect(); - prop_assert_eq!(lines.len(), flows.len(), "JSONL 行数应该等于 Flow 数量"); - - // 验证每行都能反序列化 - for (i, line) in lines.iter().enumerate() { - let deserialized: LLMFlow = serde_json::from_str(line) - .unwrap_or_else(|_| panic!("第 {i} 行应该能够反序列化")); - - prop_assert_eq!( - &flows[i].id, &deserialized.id, - "第 {} 个 Flow 的 ID 应该一致", i - ); - } - } - - /// **Feature: llm-flow-monitor, Property 8c: HAR 导出结构正确性** - /// **Validates: Requirements 5.1, 5.7** - /// - /// *对于任意* 有效的 LLM_Flow 列表,导出为 HAR 格式后, - /// HAR 结构应该符合规范,且包含 LLM 特定扩展。 - #[test] - fn prop_export_har_structure( - flows in prop::collection::vec(arb_llm_flow(), 1..5) - ) { - let exporter = FlowExporter::with_defaults(); - - // 导出为 HAR - let har = exporter.export_har(&flows); - - // 验证 HAR 结构 - prop_assert_eq!(har.log.version, "1.2", "HAR 版本应该是 1.2"); - prop_assert_eq!(har.log.entries.len(), flows.len(), "HAR 条目数应该等于 Flow 数量"); - - // 验证每个条目 - for (i, entry) in har.log.entries.iter().enumerate() { - // 验证请求 - prop_assert_eq!(&entry.request.method, &flows[i].request.method, "请求方法应该一致"); - - // 验证 LLM 扩展存在 - prop_assert!(entry.llm_extension.is_some(), "应该包含 LLM 扩展"); - - let llm_ext = entry.llm_extension.as_ref().unwrap(); - prop_assert_eq!(&llm_ext.flow_id, &flows[i].id, "Flow ID 应该一致"); - prop_assert_eq!(&llm_ext.model, &flows[i].request.model, "模型应该一致"); - prop_assert_eq!(llm_ext.streaming, flows[i].request.parameters.stream, "流式标志应该一致"); - } - } - - /// **Feature: llm-flow-monitor, Property 8d: CSV 导出包含所有 Flow** - /// **Validates: Requirements 5.5** - /// - /// *对于任意* 有效的 LLM_Flow 列表,导出为 CSV 格式后, - /// CSV 应该包含头部和所有 Flow 的数据行。 - #[test] - fn prop_export_csv_completeness( - flows in prop::collection::vec(arb_llm_flow(), 1..5) - ) { - let exporter = FlowExporter::with_defaults(); - - // 导出为 CSV - let csv = exporter.export_csv(&flows); - - // 验证行数(头部 + 数据行) - let lines: Vec<_> = csv.lines().collect(); - prop_assert_eq!(lines.len(), flows.len() + 1, "CSV 行数应该等于 Flow 数量 + 1(头部)"); - - // 验证头部 - prop_assert!(lines[0].contains("id"), "头部应该包含 id 列"); - prop_assert!(lines[0].contains("provider"), "头部应该包含 provider 列"); - prop_assert!(lines[0].contains("model"), "头部应该包含 model 列"); - - // 验证每个数据行包含 Flow ID - for (i, flow) in flows.iter().enumerate() { - prop_assert!( - lines[i + 1].contains(&flow.id), - "第 {} 行应该包含 Flow ID", i - ); - } - } - - /// **Feature: llm-flow-monitor, Property 8e: Markdown 导出包含关键信息** - /// **Validates: Requirements 5.4** - /// - /// *对于任意* 有效的 LLM_Flow,导出为 Markdown 格式后, - /// 应该包含 Flow 的关键信息。 - #[test] - fn prop_export_markdown_content(flow in arb_llm_flow()) { - let exporter = FlowExporter::with_defaults(); - - // 导出为 Markdown - let md = exporter.export_markdown(&flow); - - // 验证包含关键信息 - prop_assert!(md.contains(&flow.id), "Markdown 应该包含 Flow ID"); - prop_assert!(md.contains(&flow.request.model), "Markdown 应该包含模型名称"); - prop_assert!(md.contains("## 请求"), "Markdown 应该包含请求部分"); - - // 如果有响应,验证包含响应部分 - if flow.response.is_some() { - prop_assert!(md.contains("## 响应"), "Markdown 应该包含响应部分"); - } - } - } -} - -// ============================================================================ -// 脱敏属性测试模块 -// ============================================================================ - -#[cfg(test)] -mod redaction_property_tests { - use super::*; - use crate::flow_monitor::models::*; - use chrono::Utc; - use proptest::prelude::*; - use std::collections::HashMap; - - // ======================================================================== - // 敏感数据生成器 - // ======================================================================== - - /// 生成随机邮箱地址 - fn arb_email() -> impl Strategy { - ( - "[a-z]{3,10}", - "[a-z]{3,10}", - prop_oneof!["com", "org", "net", "io"], - ) - .prop_map(|(user, domain, tld)| format!("{user}@{domain}.{tld}")) - } - - /// 生成随机中国手机号 - fn arb_phone_cn() -> impl Strategy { - ( - prop_oneof![Just("13"), Just("15"), Just("18"), Just("19")], - "[0-9]{9}", - ) - .prop_map(|(prefix, suffix)| format!("{prefix}{suffix}")) - } - - /// 生成随机 API 密钥 - fn arb_api_key() -> impl Strategy { - "[a-zA-Z0-9]{20,40}".prop_map(|s| format!("sk-{s}")) - } - - /// 生成随机 Bearer Token - fn arb_bearer_token() -> impl Strategy { - "[a-zA-Z0-9_.-]{20,50}".prop_map(|s| format!("Bearer {s}")) - } - - /// 生成包含敏感数据的文本 - fn arb_text_with_sensitive_data() -> impl Strategy)> { - prop_oneof![ - // 包含邮箱 - arb_email().prop_map(|email| { - let text = format!("Contact me at {email} for more info."); - (text, vec![email]) - }), - // 包含手机号 - arb_phone_cn().prop_map(|phone| { - let text = format!("My phone number is {phone}."); - (text, vec![phone]) - }), - // 包含 API 密钥 - arb_api_key().prop_map(|key| { - let text = format!("Use this API key: {key}"); - (text, vec![key]) - }), - // 包含 Bearer Token - arb_bearer_token().prop_map(|token| { - let text = format!("Authorization: {token}"); - (text, vec![token]) - }), - // 包含多种敏感数据 - (arb_email(), arb_phone_cn()).prop_map(|(email, phone)| { - let text = format!("Email: {email}, Phone: {phone}"); - (text, vec![email, phone]) - }), - ] - } - - /// 生成随机的 ProviderType - fn arb_provider_type() -> impl Strategy { - prop_oneof![ - Just(ProviderType::Kiro), - Just(ProviderType::OpenAI), - Just(ProviderType::Claude), - ] - } - - /// 生成随机的 FlowType - fn arb_flow_type() -> impl Strategy { - prop_oneof![ - Just(FlowType::ChatCompletions), - Just(FlowType::AnthropicMessages), - ] - } - - /// 生成包含敏感数据的 LLMFlow - fn arb_flow_with_sensitive_data() -> impl Strategy)> { - ( - "[a-f0-9]{8}-[a-f0-9]{4}-[a-f0-9]{4}-[a-f0-9]{4}-[a-f0-9]{12}", - arb_flow_type(), - arb_provider_type(), - arb_text_with_sensitive_data(), - arb_text_with_sensitive_data(), - ) - .prop_map( - |( - id, - flow_type, - provider, - (req_content, req_sensitive), - (resp_content, resp_sensitive), - )| { - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - headers: HashMap::new(), - body: serde_json::Value::Null, - messages: vec![Message { - role: MessageRole::User, - content: MessageContent::Text(req_content), - tool_calls: None, - tool_result: None, - name: None, - }], - system_prompt: None, - tools: None, - model: "gpt-4".to_string(), - original_model: None, - parameters: RequestParameters::default(), - size_bytes: 0, - timestamp: Utc::now(), - }; - - let response = LLMResponse { - status_code: 200, - status_text: "OK".to_string(), - headers: HashMap::new(), - body: serde_json::Value::Null, - content: resp_content, - thinking: None, - tool_calls: Vec::new(), - usage: TokenUsage::default(), - stop_reason: Some(StopReason::Stop), - size_bytes: 0, - timestamp_start: Utc::now(), - timestamp_end: Utc::now(), - stream_info: None, - }; - - let metadata = FlowMetadata { - provider, - provider_id: None, - credential_id: None, - credential_name: None, - retry_count: 0, - client_info: ClientInfo::default(), - routing_info: RoutingInfo::default(), - injected_params: None, - context_usage_percentage: None, - }; - - let mut flow = LLMFlow::new(id, flow_type, request, metadata); - flow.response = Some(response); - flow.state = FlowState::Completed; - - // 合并所有敏感数据 - let mut all_sensitive = req_sensitive; - all_sensitive.extend(resp_sensitive); - - (flow, all_sensitive) - }, - ) - } - - // ======================================================================== - // 属性测试 - // ======================================================================== - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: llm-flow-monitor, Property 11: 脱敏正确性** - /// **Validates: Requirements 8.1, 8.2, 8.3** - /// - /// *对于任意* 包含敏感数据(API 密钥、邮箱、手机号)的 Flow, - /// 应用脱敏规则后,输出不应该包含原始敏感数据。 - #[test] - fn prop_redaction_removes_sensitive_data( - (flow, sensitive_data) in arb_flow_with_sensitive_data() - ) { - let redactor = Redactor::with_defaults(); - - // 应用脱敏 - let redacted_flow = redactor.redact_flow(&flow); - - // 序列化为 JSON 以便检查 - let redacted_json = serde_json::to_string(&redacted_flow) - .expect("应该能够序列化"); - - // 验证所有敏感数据都已被脱敏 - for sensitive in &sensitive_data { - prop_assert!( - !redacted_json.contains(sensitive), - "脱敏后的 JSON 不应该包含敏感数据: {}", - sensitive - ); - } - } - - /// **Feature: llm-flow-monitor, Property 11b: 脱敏后导出不包含敏感数据** - /// **Validates: Requirements 8.1, 8.2, 8.3** - /// - /// *对于任意* 包含敏感数据的 Flow,使用启用脱敏的导出器导出后, - /// 导出结果不应该包含原始敏感数据。 - #[test] - fn prop_export_with_redaction_removes_sensitive_data( - (flow, sensitive_data) in arb_flow_with_sensitive_data() - ) { - let options = ExportOptions { - format: ExportFormat::JSON, - redact_sensitive: true, - ..Default::default() - }; - let exporter = FlowExporter::new(options); - - // 导出 - let json = exporter.export_json(&[flow]); - let json_str = serde_json::to_string(&json).expect("应该能够序列化"); - - // 验证所有敏感数据都已被脱敏 - for sensitive in &sensitive_data { - prop_assert!( - !json_str.contains(sensitive), - "导出的 JSON 不应该包含敏感数据: {}", - sensitive - ); - } - } - - /// **Feature: llm-flow-monitor, Property 11c: 脱敏保留非敏感数据** - /// **Validates: Requirements 8.1, 8.2, 8.3** - /// - /// *对于任意* Flow,脱敏后应该保留非敏感的关键字段。 - #[test] - fn prop_redaction_preserves_non_sensitive_data( - (flow, _) in arb_flow_with_sensitive_data() - ) { - let redactor = Redactor::with_defaults(); - - // 应用脱敏 - let redacted_flow = redactor.redact_flow(&flow); - - // 验证关键字段保持不变 - prop_assert_eq!(&flow.id, &redacted_flow.id, "Flow ID 应该保持不变"); - prop_assert_eq!(flow.state, redacted_flow.state, "状态应该保持不变"); - prop_assert_eq!(flow.flow_type, redacted_flow.flow_type, "FlowType 应该保持不变"); - prop_assert_eq!(&flow.request.model, &redacted_flow.request.model, "模型应该保持不变"); - prop_assert_eq!(&flow.request.method, &redacted_flow.request.method, "方法应该保持不变"); - prop_assert_eq!(flow.metadata.provider, redacted_flow.metadata.provider, "Provider 应该保持不变"); - - // 验证响应存在性 - prop_assert_eq!( - flow.response.is_some(), - redacted_flow.response.is_some(), - "响应存在性应该保持不变" - ); - - // 验证 Token 使用量保持不变 - if let (Some(ref orig), Some(ref redacted)) = (&flow.response, &redacted_flow.response) { - prop_assert_eq!( - orig.usage.input_tokens, - redacted.usage.input_tokens, - "输入 Token 应该保持不变" - ); - prop_assert_eq!( - orig.usage.output_tokens, - redacted.usage.output_tokens, - "输出 Token 应该保持不变" - ); - } - } - - /// **Feature: llm-flow-monitor, Property 11d: 邮箱脱敏** - /// **Validates: Requirements 8.2** - /// - /// *对于任意* 包含邮箱的文本,脱敏后不应该包含原始邮箱。 - #[test] - fn prop_redact_email(email in arb_email()) { - let redactor = Redactor::with_defaults(); - let text = format!("Contact: {email}"); - - let redacted = redactor.redact(&text); - - prop_assert!( - !redacted.contains(&email), - "脱敏后不应该包含邮箱: {}", - email - ); - prop_assert!( - redacted.contains("[REDACTED_EMAIL]"), - "脱敏后应该包含占位符" - ); - } - - /// **Feature: llm-flow-monitor, Property 11e: 手机号脱敏** - /// **Validates: Requirements 8.2** - /// - /// *对于任意* 包含中国手机号的文本,脱敏后不应该包含原始手机号。 - #[test] - fn prop_redact_phone(phone in arb_phone_cn()) { - let redactor = Redactor::with_defaults(); - let text = format!("Phone: {phone}"); - - let redacted = redactor.redact(&text); - - prop_assert!( - !redacted.contains(&phone), - "脱敏后不应该包含手机号: {}", - phone - ); - prop_assert!( - redacted.contains("[REDACTED_PHONE]"), - "脱敏后应该包含占位符" - ); - } - - /// **Feature: llm-flow-monitor, Property 11f: API 密钥脱敏** - /// **Validates: Requirements 8.1** - /// - /// *对于任意* 包含 API 密钥的文本,脱敏后不应该包含原始密钥。 - #[test] - fn prop_redact_api_key(key in arb_api_key()) { - let redactor = Redactor::with_defaults(); - let text = format!("API Key: {key}"); - - let redacted = redactor.redact(&text); - - prop_assert!( - !redacted.contains(&key), - "脱敏后不应该包含 API 密钥: {}", - key - ); - } - } -} diff --git a/src-tauri/src/flow_monitor/file_store.rs b/src-tauri/src/flow_monitor/file_store.rs deleted file mode 100644 index f3895d431..000000000 --- a/src-tauri/src/flow_monitor/file_store.rs +++ /dev/null @@ -1,1407 +0,0 @@ -//! Flow 文件存储 -//! -//! 该模块实现 LLM Flow 的文件持久化存储,支持 JSONL 格式写入、 -//! SQLite 索引、文件轮转和自动清理功能。 - -use chrono::{DateTime, NaiveDate, Utc}; -use rusqlite::{params, Connection, OptionalExtension}; -use serde::{Deserialize, Serialize}; -use std::fs::{self, File, OpenOptions}; -use std::io::{BufRead, BufReader, BufWriter, Seek, SeekFrom, Write}; -use std::path::{Path, PathBuf}; -use std::sync::Mutex; -use thiserror::Error; - -use super::memory_store::FlowFilter; -use super::models::LLMFlow; - -// ============================================================================ -// 错误类型 -// ============================================================================ - -/// 文件存储错误 -#[derive(Debug, Error)] -pub enum FileStoreError { - #[error("IO 错误: {0}")] - Io(#[from] std::io::Error), - - #[error("JSON 序列化错误: {0}")] - Json(#[from] serde_json::Error), - - #[error("SQLite 错误: {0}")] - Sqlite(#[from] rusqlite::Error), - - #[error("存储目录不存在: {0}")] - DirectoryNotFound(PathBuf), - - #[error("Flow 不存在: {0}")] - FlowNotFound(String), - - #[error("文件轮转失败: {0}")] - RotationFailed(String), -} - -pub type Result = std::result::Result; - -// ============================================================================ -// 配置结构 -// ============================================================================ - -/// 文件轮转配置 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct RotationConfig { - /// 是否按日期轮转 - pub rotate_daily: bool, - /// 单个文件最大大小(字节) - pub max_file_size: u64, - /// 保留天数 - pub retention_days: u32, - /// 是否压缩旧文件 - pub compress_old: bool, -} - -impl Default for RotationConfig { - fn default() -> Self { - Self { - rotate_daily: true, - max_file_size: 100 * 1024 * 1024, // 100MB - retention_days: 7, - compress_old: false, // 暂不实现压缩 - } - } -} - -/// 清理结果 -#[derive(Debug, Clone, Default)] -pub struct CleanupResult { - /// 删除的文件数 - pub files_deleted: usize, - /// 删除的 Flow 数 - pub flows_deleted: usize, - /// 释放的空间(字节) - pub bytes_freed: u64, -} - -// ============================================================================ -// 索引记录 -// ============================================================================ - -/// Flow 索引记录(存储在 SQLite 中) -#[derive(Debug, Clone)] -pub struct FlowIndexRecord { - pub id: String, - pub created_at: DateTime, - pub provider: String, - pub model: String, - pub status: String, - pub duration_ms: Option, - pub input_tokens: Option, - pub output_tokens: Option, - pub has_error: bool, - pub has_tool_calls: bool, - pub has_thinking: bool, - pub file_path: String, - pub file_offset: i64, - pub content_preview: Option, - pub request_preview: Option, -} - -/// FTS 搜索结果 -#[derive(Debug, Clone)] -pub struct FtsSearchResult { - /// Flow ID - pub id: String, - /// 创建时间(RFC3339 格式字符串) - pub created_at: String, - /// 模型名称 - pub model: String, - /// 提供商 - pub provider: String, - /// 匹配的内容片段 - pub snippet: String, -} - -impl FlowIndexRecord { - /// 从 LLMFlow 创建索引记录 - pub fn from_flow(flow: &LLMFlow, file_path: &str, file_offset: i64) -> Self { - let content_preview = flow - .response - .as_ref() - .map(|r| r.content.chars().take(200).collect::()); - - let request_preview = flow - .request - .system_prompt - .as_ref() - .map(|s| s.chars().take(200).collect::()) - .or_else(|| { - flow.request.messages.first().map(|m| { - m.content - .get_all_text() - .chars() - .take(200) - .collect::() - }) - }); - - Self { - id: flow.id.clone(), - created_at: flow.timestamps.created, - provider: format!("{:?}", flow.metadata.provider), - model: flow.request.model.clone(), - status: format!("{:?}", flow.state), - duration_ms: Some(flow.timestamps.duration_ms as i64), - input_tokens: flow.response.as_ref().map(|r| r.usage.input_tokens as i32), - output_tokens: flow.response.as_ref().map(|r| r.usage.output_tokens as i32), - has_error: flow.error.is_some(), - has_tool_calls: flow - .response - .as_ref() - .is_some_and(|r| !r.tool_calls.is_empty()), - has_thinking: flow.response.as_ref().is_some_and(|r| r.thinking.is_some()), - file_path: file_path.to_string(), - file_offset, - content_preview, - request_preview, - } - } -} - -// ============================================================================ -// 文件写入器 -// ============================================================================ - -/// JSONL 文件写入器 -struct FlowWriter { - file: BufWriter, - path: PathBuf, - current_offset: u64, - current_size: u64, -} - -impl FlowWriter { - /// 创建新的写入器 - fn new(path: PathBuf) -> Result { - let file = OpenOptions::new().create(true).append(true).open(&path)?; - - let current_size = file.metadata()?.len(); - let current_offset = current_size; - - Ok(Self { - file: BufWriter::new(file), - path, - current_offset, - current_size, - }) - } - - /// 写入 Flow 并返回偏移量 - fn write(&mut self, flow: &LLMFlow) -> Result { - let offset = self.current_offset; - let json = serde_json::to_string(flow)?; - let line = format!("{json}\n"); - let bytes = line.as_bytes(); - - self.file.write_all(bytes)?; - self.file.flush()?; - - self.current_offset += bytes.len() as u64; - self.current_size += bytes.len() as u64; - - Ok(offset) - } - - /// 获取当前文件大小 - fn size(&self) -> u64 { - self.current_size - } - - /// 获取文件路径 - fn path(&self) -> &Path { - &self.path - } -} - -// ============================================================================ -// Flow 文件存储 -// ============================================================================ - -/// Flow 文件存储 -/// -/// 使用 JSONL 格式存储 Flow,SQLite 索引支持快速查询。 -pub struct FlowFileStore { - /// 存储目录 - base_dir: PathBuf, - /// 当前写入器 - current_writer: Mutex>, - /// 当前日期(用于日期轮转) - current_date: Mutex, - /// 当前文件序号 - current_file_index: Mutex, - /// 轮转配置 - rotation_config: RotationConfig, - /// SQLite 连接 - index_db: Mutex, -} - -impl FlowFileStore { - /// 创建新的文件存储 - /// - /// # 参数 - /// - `base_dir`: 存储目录 - /// - `config`: 轮转配置 - pub fn new(base_dir: PathBuf, config: RotationConfig) -> Result { - // 创建存储目录 - fs::create_dir_all(&base_dir)?; - - // 创建全局索引数据库 - let db_path = base_dir.join("global_index.sqlite"); - let conn = Connection::open(&db_path)?; - - // 初始化数据库表 - Self::init_database(&conn)?; - - let today = Utc::now().date_naive(); - - Ok(Self { - base_dir, - current_writer: Mutex::new(None), - current_date: Mutex::new(today), - current_file_index: Mutex::new(1), - rotation_config: config, - index_db: Mutex::new(conn), - }) - } - - /// 初始化数据库表 - fn init_database(conn: &Connection) -> Result<()> { - conn.execute_batch( - r#" - -- 全局索引表 - CREATE TABLE IF NOT EXISTS flow_index ( - id TEXT PRIMARY KEY, - created_at TEXT NOT NULL, - provider TEXT NOT NULL, - model TEXT NOT NULL, - status TEXT NOT NULL, - duration_ms INTEGER, - input_tokens INTEGER, - output_tokens INTEGER, - has_error INTEGER DEFAULT 0, - has_tool_calls INTEGER DEFAULT 0, - has_thinking INTEGER DEFAULT 0, - file_path TEXT NOT NULL, - file_offset INTEGER NOT NULL, - content_preview TEXT, - request_preview TEXT - ); - - CREATE INDEX IF NOT EXISTS idx_created_at ON flow_index(created_at); - CREATE INDEX IF NOT EXISTS idx_provider ON flow_index(provider); - CREATE INDEX IF NOT EXISTS idx_model ON flow_index(model); - CREATE INDEX IF NOT EXISTS idx_status ON flow_index(status); - - -- 标注表 - CREATE TABLE IF NOT EXISTS flow_annotations ( - flow_id TEXT PRIMARY KEY, - starred INTEGER DEFAULT 0, - marker TEXT, - comment TEXT, - updated_at TEXT NOT NULL, - FOREIGN KEY (flow_id) REFERENCES flow_index(id) - ); - - -- 标签表 - CREATE TABLE IF NOT EXISTS flow_tags ( - flow_id TEXT NOT NULL, - tag TEXT NOT NULL, - PRIMARY KEY (flow_id, tag), - FOREIGN KEY (flow_id) REFERENCES flow_index(id) - ); - - CREATE INDEX IF NOT EXISTS idx_tags ON flow_tags(tag); - - -- 全文搜索表(FTS5) - -- 注意:这是一个独立的 FTS5 表,不使用 content= 选项 - -- 数据通过 INSERT 语句直接插入 - CREATE VIRTUAL TABLE IF NOT EXISTS flow_fts USING fts5( - id, - content_text, - request_text, - model - ); - "#, - )?; - - Ok(()) - } - - /// 获取存储目录 - pub fn base_dir(&self) -> &Path { - &self.base_dir - } - - /// 获取轮转配置 - pub fn rotation_config(&self) -> &RotationConfig { - &self.rotation_config - } - - /// 写入 Flow 到文件 - /// - /// # 参数 - /// - `flow`: 要写入的 Flow - pub fn write(&self, flow: &LLMFlow) -> Result<()> { - // 检查是否需要轮转 - self.check_rotation()?; - - // 获取或创建写入器 - let mut writer_guard = self.current_writer.lock().unwrap(); - if writer_guard.is_none() { - *writer_guard = Some(self.create_writer()?); - } - - let writer = writer_guard.as_mut().unwrap(); - - // 写入 Flow - let offset = writer.write(flow)?; - let file_path = writer.path().to_string_lossy().to_string(); - - // 更新索引 - self.update_index(flow, &file_path, offset as i64)?; - - // 检查文件大小是否需要轮转 - if writer.size() >= self.rotation_config.max_file_size { - drop(writer_guard); - self.rotate()?; - } - - Ok(()) - } - - /// 创建新的写入器 - fn create_writer(&self) -> Result { - let date = *self.current_date.lock().unwrap(); - let index = *self.current_file_index.lock().unwrap(); - - // 创建日期目录 - let date_dir = self.base_dir.join(date.format("%Y-%m-%d").to_string()); - fs::create_dir_all(&date_dir)?; - - // 创建文件路径 - let file_name = format!("flows_{index:03}.jsonl"); - let file_path = date_dir.join(file_name); - - FlowWriter::new(file_path) - } - - /// 检查是否需要日期轮转 - fn check_rotation(&self) -> Result<()> { - if !self.rotation_config.rotate_daily { - return Ok(()); - } - - let today = Utc::now().date_naive(); - let mut current_date = self.current_date.lock().unwrap(); - - if *current_date != today { - // 日期变化,需要轮转 - *current_date = today; - *self.current_file_index.lock().unwrap() = 1; - *self.current_writer.lock().unwrap() = None; - } - - Ok(()) - } - - /// 轮转到新文件 - pub fn rotate(&self) -> Result<()> { - // 关闭当前写入器 - *self.current_writer.lock().unwrap() = None; - - // 增加文件序号 - let mut index = self.current_file_index.lock().unwrap(); - *index += 1; - - Ok(()) - } - - /// 更新索引 - fn update_index(&self, flow: &LLMFlow, file_path: &str, file_offset: i64) -> Result<()> { - let record = FlowIndexRecord::from_flow(flow, file_path, file_offset); - let conn = self.index_db.lock().unwrap(); - - conn.execute( - r#" - INSERT OR REPLACE INTO flow_index ( - id, created_at, provider, model, status, - duration_ms, input_tokens, output_tokens, - has_error, has_tool_calls, has_thinking, - file_path, file_offset, content_preview, request_preview - ) VALUES ( - ?1, ?2, ?3, ?4, ?5, - ?6, ?7, ?8, - ?9, ?10, ?11, - ?12, ?13, ?14, ?15 - ) - "#, - params![ - record.id, - record.created_at.to_rfc3339(), - record.provider, - record.model, - record.status, - record.duration_ms, - record.input_tokens, - record.output_tokens, - record.has_error as i32, - record.has_tool_calls as i32, - record.has_thinking as i32, - record.file_path, - record.file_offset, - record.content_preview, - record.request_preview, - ], - )?; - - // 更新标注 - if flow.annotations.starred - || flow.annotations.marker.is_some() - || flow.annotations.comment.is_some() - { - conn.execute( - r#" - INSERT OR REPLACE INTO flow_annotations ( - flow_id, starred, marker, comment, updated_at - ) VALUES (?1, ?2, ?3, ?4, ?5) - "#, - params![ - flow.id, - flow.annotations.starred as i32, - flow.annotations.marker, - flow.annotations.comment, - Utc::now().to_rfc3339(), - ], - )?; - } - - // 更新标签 - if !flow.annotations.tags.is_empty() { - // 先删除旧标签 - conn.execute("DELETE FROM flow_tags WHERE flow_id = ?1", params![flow.id])?; - - // 插入新标签 - for tag in &flow.annotations.tags { - conn.execute( - "INSERT INTO flow_tags (flow_id, tag) VALUES (?1, ?2)", - params![flow.id, tag], - )?; - } - } - - // 更新 FTS5 索引 - let content_text = flow - .response - .as_ref() - .map_or(String::new(), |r| r.content.clone()); - let request_text = Self::get_request_text_for_fts(flow); - - // 先删除旧的 FTS 记录 - conn.execute("DELETE FROM flow_fts WHERE id = ?1", params![flow.id])?; - - // 插入新的 FTS 记录 - conn.execute( - "INSERT INTO flow_fts (id, content_text, request_text, model) VALUES (?1, ?2, ?3, ?4)", - params![flow.id, content_text, request_text, flow.request.model], - )?; - - Ok(()) - } - - /// 获取请求文本(用于 FTS 索引) - fn get_request_text_for_fts(flow: &LLMFlow) -> String { - let mut text = String::new(); - - // 添加系统提示词 - if let Some(ref system) = flow.request.system_prompt { - text.push_str(system); - text.push('\n'); - } - - // 添加消息内容 - for msg in &flow.request.messages { - text.push_str(&msg.content.get_all_text()); - text.push('\n'); - } - - text - } - - /// 根据 ID 获取 Flow - pub fn get(&self, id: &str) -> Result> { - let conn = self.index_db.lock().unwrap(); - - let result: Option<(String, i64)> = conn - .query_row( - "SELECT file_path, file_offset FROM flow_index WHERE id = ?1", - params![id], - |row| Ok((row.get(0)?, row.get(1)?)), - ) - .optional()?; - - match result { - Some((file_path, file_offset)) => self.read_flow_from_file(&file_path, file_offset), - None => Ok(None), - } - } - - /// 从文件读取 Flow - fn read_flow_from_file(&self, file_path: &str, file_offset: i64) -> Result> { - let path = Path::new(file_path); - if !path.exists() { - return Ok(None); - } - - let file = File::open(path)?; - let mut reader = BufReader::new(file); - - // 跳转到指定偏移量 - reader.seek(SeekFrom::Start(file_offset as u64))?; - - // 读取一行 - let mut line = String::new(); - reader.read_line(&mut line)?; - - if line.is_empty() { - return Ok(None); - } - - let flow: LLMFlow = serde_json::from_str(&line)?; - Ok(Some(flow)) - } - - /// 查询 Flow(从索引) - pub fn query(&self, filter: &FlowFilter, limit: usize, offset: usize) -> Result> { - // 先获取所有文件位置信息 - let file_locations = self.query_index(filter, limit, offset)?; - - // 读取 Flow - let mut flows = Vec::new(); - for (file_path, file_offset) in file_locations { - if let Some(flow) = self.read_flow_from_file(&file_path, file_offset)? { - // 再次用内存过滤器验证(处理复杂条件) - if filter.matches(&flow) { - flows.push(flow); - } - } - } - - Ok(flows) - } - - /// 从索引查询文件位置 - fn query_index( - &self, - filter: &FlowFilter, - limit: usize, - offset: usize, - ) -> Result> { - let conn = self.index_db.lock().unwrap(); - - // 构建查询条件 - let mut conditions: Vec = Vec::new(); - let mut params_vec: Vec> = Vec::new(); - - // 时间范围 - if let Some(ref time_range) = filter.time_range { - if let Some(start) = time_range.start { - conditions.push("created_at >= ?".to_string()); - params_vec.push(Box::new(start.to_rfc3339())); - } - if let Some(end) = time_range.end { - conditions.push("created_at <= ?".to_string()); - params_vec.push(Box::new(end.to_rfc3339())); - } - } - - // 提供商过滤 - if let Some(ref providers) = filter.providers { - if !providers.is_empty() { - let placeholders: Vec = providers.iter().map(|_| "?".to_string()).collect(); - conditions.push(format!("provider IN ({})", placeholders.join(", "))); - for p in providers { - params_vec.push(Box::new(format!("{p:?}"))); - } - } - } - - // 状态过滤 - if let Some(ref states) = filter.states { - if !states.is_empty() { - let placeholders: Vec = states.iter().map(|_| "?".to_string()).collect(); - conditions.push(format!("status IN ({})", placeholders.join(", "))); - for s in states { - params_vec.push(Box::new(format!("{s:?}"))); - } - } - } - - // 错误过滤 - if let Some(has_error) = filter.has_error { - conditions.push("has_error = ?".to_string()); - params_vec.push(Box::new(has_error as i32)); - } - - // 工具调用过滤 - if let Some(has_tool_calls) = filter.has_tool_calls { - conditions.push("has_tool_calls = ?".to_string()); - params_vec.push(Box::new(has_tool_calls as i32)); - } - - // 思维链过滤 - if let Some(has_thinking) = filter.has_thinking { - conditions.push("has_thinking = ?".to_string()); - params_vec.push(Box::new(has_thinking as i32)); - } - - // 构建 SQL - let where_clause = if conditions.is_empty() { - String::new() - } else { - format!("WHERE {}", conditions.join(" AND ")) - }; - - let sql = format!( - "SELECT file_path, file_offset FROM flow_index {where_clause} ORDER BY created_at DESC LIMIT ? OFFSET ?" - ); - - params_vec.push(Box::new(limit as i64)); - params_vec.push(Box::new(offset as i64)); - - // 执行查询 - let params_refs: Vec<&dyn rusqlite::ToSql> = - params_vec.iter().map(|p| p.as_ref()).collect(); - let mut stmt = conn.prepare(&sql)?; - let rows = stmt.query_map(params_refs.as_slice(), |row| { - Ok((row.get::<_, String>(0)?, row.get::<_, i64>(1)?)) - })?; - - let mut results = Vec::new(); - for row in rows { - results.push(row?); - } - - Ok(results) - } - - /// 获取索引中的 Flow 数量 - pub fn count(&self) -> Result { - let conn = self.index_db.lock().unwrap(); - let count: i64 = conn.query_row("SELECT COUNT(*) FROM flow_index", [], |row| row.get(0))?; - Ok(count as usize) - } - - /// 全文搜索 - /// - /// 使用 SQLite FTS5 进行全文搜索 - /// - /// # 参数 - /// - `query`: 搜索关键词 - /// - `limit`: 最大返回数量 - /// - /// # 返回 - /// 匹配的 Flow ID、创建时间、模型、提供商和匹配片段 - pub fn search(&self, query: &str, limit: usize) -> Result> { - let conn = self.index_db.lock().unwrap(); - - // 转义特殊字符并构建 FTS5 查询 - let escaped_query = Self::escape_fts_query(query); - - let sql = r#" - SELECT - f.id, - f.created_at, - f.model, - f.provider, - snippet(flow_fts, 1, '', '', '...', 32) as snippet - FROM flow_fts - JOIN flow_index f ON flow_fts.id = f.id - WHERE flow_fts MATCH ?1 - ORDER BY rank - LIMIT ?2 - "#; - - let mut stmt = conn.prepare(sql)?; - let rows = stmt.query_map(params![escaped_query, limit as i64], |row| { - Ok(FtsSearchResult { - id: row.get(0)?, - created_at: row.get(1)?, - model: row.get(2)?, - provider: row.get(3)?, - snippet: row.get(4)?, - }) - })?; - - let mut results = Vec::new(); - for row in rows { - results.push(row?); - } - - Ok(results) - } - - /// 转义 FTS5 查询中的特殊字符 - fn escape_fts_query(query: &str) -> String { - // FTS5 特殊字符: " * - ^ : ( ) - // 对于简单搜索,我们使用双引号包裹整个查询 - format!("\"{}\"", query.replace('"', "\"\"")) - } - - /// 更新 Flow 标注 - /// - /// # 参数 - /// - `flow_id`: Flow ID - /// - `annotations`: 新的标注信息 - pub fn update_annotations( - &self, - flow_id: &str, - annotations: &crate::flow_monitor::models::FlowAnnotations, - ) -> Result<()> { - let conn = self.index_db.lock().unwrap(); - - // 更新或插入标注 - conn.execute( - r#" - INSERT OR REPLACE INTO flow_annotations ( - flow_id, starred, marker, comment, updated_at - ) VALUES (?1, ?2, ?3, ?4, ?5) - "#, - params![ - flow_id, - annotations.starred as i32, - annotations.marker, - annotations.comment, - Utc::now().to_rfc3339(), - ], - )?; - - // 更新标签 - // 先删除旧标签 - conn.execute("DELETE FROM flow_tags WHERE flow_id = ?1", params![flow_id])?; - - // 插入新标签 - for tag in &annotations.tags { - conn.execute( - "INSERT INTO flow_tags (flow_id, tag) VALUES (?1, ?2)", - params![flow_id, tag], - )?; - } - - Ok(()) - } - - /// 清理过期数据 - /// - /// # 参数 - /// - `before`: 清理此时间之前的数据 - pub fn cleanup(&self, before: DateTime) -> Result { - let mut result = CleanupResult::default(); - - // 获取要删除的文件列表和执行删除操作 - let file_paths = { - let conn = self.index_db.lock().unwrap(); - - // 获取要删除的文件列表 - let mut stmt = - conn.prepare("SELECT DISTINCT file_path FROM flow_index WHERE created_at < ?1")?; - - let file_paths: Vec = stmt - .query_map(params![before.to_rfc3339()], |row| row.get(0))? - .filter_map(|r| r.ok()) - .collect(); - - // 统计要删除的 Flow 数量 - let count: i64 = conn.query_row( - "SELECT COUNT(*) FROM flow_index WHERE created_at < ?1", - params![before.to_rfc3339()], - |row| row.get(0), - )?; - result.flows_deleted = count as usize; - - // 删除索引记录 - conn.execute( - "DELETE FROM flow_annotations WHERE flow_id IN (SELECT id FROM flow_index WHERE created_at < ?1)", - params![before.to_rfc3339()], - )?; - - conn.execute( - "DELETE FROM flow_tags WHERE flow_id IN (SELECT id FROM flow_index WHERE created_at < ?1)", - params![before.to_rfc3339()], - )?; - - conn.execute( - "DELETE FROM flow_index WHERE created_at < ?1", - params![before.to_rfc3339()], - )?; - - file_paths - }; // conn 在这里被释放 - - // 删除文件 - for file_path in file_paths { - let path = Path::new(&file_path); - if path.exists() { - if let Ok(metadata) = fs::metadata(path) { - result.bytes_freed += metadata.len(); - } - if fs::remove_file(path).is_ok() { - result.files_deleted += 1; - } - } - } - - // 清理空目录 - self.cleanup_empty_dirs()?; - - Ok(result) - } - - /// 清理空目录 - fn cleanup_empty_dirs(&self) -> Result<()> { - if let Ok(entries) = fs::read_dir(&self.base_dir) { - for entry in entries.flatten() { - let path = entry.path(); - if path.is_dir() { - // 检查目录是否为空(除了 .sqlite 文件) - if let Ok(mut dir_entries) = fs::read_dir(&path) { - let has_jsonl = dir_entries.any(|e| { - e.ok() - .map(|e| e.path().extension().is_some_and(|ext| ext == "jsonl")) - .unwrap_or(false) - }); - - if !has_jsonl { - // 删除目录中的所有文件 - if let Ok(files) = fs::read_dir(&path) { - for file in files.flatten() { - let _ = fs::remove_file(file.path()); - } - } - let _ = fs::remove_dir(&path); - } - } - } - } - } - - Ok(()) - } - - /// 根据保留天数清理 - pub fn cleanup_by_retention(&self) -> Result { - let retention_days = self.rotation_config.retention_days; - let before = Utc::now() - chrono::Duration::days(retention_days as i64); - self.cleanup(before) - } -} - -// ============================================================================ -// 测试模块 -// ============================================================================ - -#[cfg(test)] -mod tests { - use super::*; - use crate::flow_monitor::models::{FlowMetadata, FlowType, LLMRequest, RequestParameters}; - use crate::ProviderType; - use tempfile::TempDir; - - /// 创建测试用的 Flow - fn create_test_flow(id: &str, model: &str, provider: ProviderType) -> LLMFlow { - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model: model.to_string(), - parameters: RequestParameters { - stream: false, - ..Default::default() - }, - ..Default::default() - }; - - let metadata = FlowMetadata { - provider, - ..Default::default() - }; - - LLMFlow::new(id.to_string(), FlowType::ChatCompletions, request, metadata) - } - - #[test] - fn test_file_store_creation() { - let temp_dir = TempDir::new().unwrap(); - let store = FlowFileStore::new(temp_dir.path().to_path_buf(), RotationConfig::default()); - - assert!(store.is_ok()); - let store = store.unwrap(); - assert!(store.base_dir().exists()); - } - - #[test] - fn test_file_store_write_and_get() { - let temp_dir = TempDir::new().unwrap(); - let store = - FlowFileStore::new(temp_dir.path().to_path_buf(), RotationConfig::default()).unwrap(); - - let flow = create_test_flow("test-1", "gpt-4", ProviderType::OpenAI); - store.write(&flow).unwrap(); - - // 验证可以读取 - let retrieved = store.get("test-1").unwrap(); - assert!(retrieved.is_some()); - - let retrieved = retrieved.unwrap(); - assert_eq!(retrieved.id, "test-1"); - assert_eq!(retrieved.request.model, "gpt-4"); - } - - #[test] - fn test_file_store_multiple_writes() { - let temp_dir = TempDir::new().unwrap(); - let store = - FlowFileStore::new(temp_dir.path().to_path_buf(), RotationConfig::default()).unwrap(); - - // 写入多个 Flow - for i in 0..10 { - let flow = create_test_flow(&format!("flow-{i}"), "gpt-4", ProviderType::OpenAI); - store.write(&flow).unwrap(); - } - - // 验证数量 - assert_eq!(store.count().unwrap(), 10); - - // 验证可以读取每个 - for i in 0..10 { - let retrieved = store.get(&format!("flow-{i}")).unwrap(); - assert!(retrieved.is_some()); - } - } - - #[test] - fn test_file_store_query() { - let temp_dir = TempDir::new().unwrap(); - let store = - FlowFileStore::new(temp_dir.path().to_path_buf(), RotationConfig::default()).unwrap(); - - // 写入不同提供商的 Flow - store - .write(&create_test_flow("flow-1", "gpt-4", ProviderType::OpenAI)) - .unwrap(); - store - .write(&create_test_flow( - "flow-2", - "claude-3", - ProviderType::Claude, - )) - .unwrap(); - store - .write(&create_test_flow( - "flow-3", - "gpt-4-turbo", - ProviderType::OpenAI, - )) - .unwrap(); - - // 查询所有 - let filter = FlowFilter::default(); - let results = store.query(&filter, 100, 0).unwrap(); - assert_eq!(results.len(), 3); - - // 按提供商过滤 - let filter = FlowFilter { - providers: Some(vec![ProviderType::OpenAI]), - ..Default::default() - }; - let results = store.query(&filter, 100, 0).unwrap(); - assert_eq!(results.len(), 2); - } - - #[test] - fn test_file_store_rotation() { - let temp_dir = TempDir::new().unwrap(); - let config = RotationConfig { - max_file_size: 100, // 很小的文件大小,强制轮转 - ..Default::default() - }; - let store = FlowFileStore::new(temp_dir.path().to_path_buf(), config).unwrap(); - - // 写入多个 Flow,应该触发轮转 - for i in 0..5 { - let flow = create_test_flow(&format!("flow-{i}"), "gpt-4", ProviderType::OpenAI); - store.write(&flow).unwrap(); - } - - // 验证所有 Flow 都可以读取 - for i in 0..5 { - let retrieved = store.get(&format!("flow-{i}")).unwrap(); - assert!(retrieved.is_some()); - } - } - - #[test] - fn test_file_store_cleanup() { - let temp_dir = TempDir::new().unwrap(); - let store = - FlowFileStore::new(temp_dir.path().to_path_buf(), RotationConfig::default()).unwrap(); - - // 写入一些 Flow - for i in 0..5 { - let flow = create_test_flow(&format!("flow-{i}"), "gpt-4", ProviderType::OpenAI); - store.write(&flow).unwrap(); - } - - assert_eq!(store.count().unwrap(), 5); - - // 清理未来时间之前的数据(应该清理所有) - let future = Utc::now() + chrono::Duration::days(1); - let result = store.cleanup(future).unwrap(); - - assert_eq!(result.flows_deleted, 5); - assert_eq!(store.count().unwrap(), 0); - } - - #[test] - fn test_index_record_from_flow() { - let flow = create_test_flow("test-1", "gpt-4", ProviderType::OpenAI); - let record = FlowIndexRecord::from_flow(&flow, "/path/to/file.jsonl", 0); - - assert_eq!(record.id, "test-1"); - assert_eq!(record.model, "gpt-4"); - assert_eq!(record.provider, "OpenAI"); - assert_eq!(record.status, "Pending"); - assert!(!record.has_error); - assert!(!record.has_tool_calls); - assert!(!record.has_thinking); - } -} - -// ============================================================================ -// 属性测试模块 -// ============================================================================ - -#[cfg(test)] -mod property_tests { - use super::*; - use crate::flow_monitor::models::{ - FlowAnnotations, FlowMetadata, FlowType, LLMRequest, LLMResponse, Message, MessageContent, - MessageRole, RequestParameters, TokenUsage, - }; - use crate::ProviderType; - use proptest::prelude::*; - use tempfile::TempDir; - - // ======================================================================== - // 生成器 - // ======================================================================== - - /// 生成随机的 ProviderType - fn arb_provider_type() -> impl Strategy { - prop_oneof![ - Just(ProviderType::Kiro), - Just(ProviderType::Gemini), - Just(ProviderType::OpenAI), - Just(ProviderType::Claude), - Just(ProviderType::Antigravity), - ] - } - - /// 生成随机的 FlowType - fn arb_flow_type() -> impl Strategy { - prop_oneof![ - Just(FlowType::ChatCompletions), - Just(FlowType::AnthropicMessages), - Just(FlowType::GeminiGenerateContent), - Just(FlowType::Embeddings), - ] - } - - /// 生成随机的 Flow ID - fn arb_flow_id() -> impl Strategy { - "[a-f0-9]{8}-[a-f0-9]{4}-[a-f0-9]{4}-[a-f0-9]{4}-[a-f0-9]{12}" - } - - /// 生成随机的模型名称 - fn arb_model_name() -> impl Strategy { - prop_oneof![ - Just("gpt-4".to_string()), - Just("gpt-4-turbo".to_string()), - Just("gpt-3.5-turbo".to_string()), - Just("claude-3-opus".to_string()), - Just("claude-3-sonnet".to_string()), - Just("gemini-pro".to_string()), - ] - } - - /// 生成随机的 MessageContent - fn arb_message_content() -> impl Strategy { - "[a-zA-Z0-9 ]{1,100}".prop_map(MessageContent::Text) - } - - /// 生成随机的 Message - fn arb_message() -> impl Strategy { - ( - prop_oneof![ - Just(MessageRole::System), - Just(MessageRole::User), - Just(MessageRole::Assistant), - ], - arb_message_content(), - ) - .prop_map(|(role, content)| Message { - role, - content, - tool_calls: None, - tool_result: None, - name: None, - }) - } - - /// 生成随机的 LLMRequest - fn arb_llm_request() -> impl Strategy { - ( - arb_model_name(), - prop::collection::vec(arb_message(), 0..3), - prop::option::of("[a-zA-Z0-9 ]{10,50}"), - any::(), - ) - .prop_map(|(model, messages, system_prompt, stream)| LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model, - messages, - system_prompt, - parameters: RequestParameters { - stream, - temperature: Some(0.7), - max_tokens: Some(1000), - ..Default::default() - }, - ..Default::default() - }) - } - - /// 生成随机的 FlowMetadata - fn arb_flow_metadata() -> impl Strategy { - arb_provider_type().prop_map(|provider| FlowMetadata { - provider, - ..Default::default() - }) - } - - /// 生成随机的 LLMResponse - fn arb_llm_response() -> impl Strategy> { - prop::option::of( - ("[a-zA-Z0-9 ]{10,200}", 0u32..1000u32, 0u32..500u32).prop_map( - |(content, input_tokens, output_tokens)| LLMResponse { - status_code: 200, - status_text: "OK".to_string(), - content, - usage: TokenUsage { - input_tokens, - output_tokens, - total_tokens: input_tokens + output_tokens, - ..Default::default() - }, - ..Default::default() - }, - ), - ) - } - - /// 生成随机的 FlowAnnotations - fn arb_flow_annotations() -> impl Strategy { - ( - any::(), - prop::option::of("[a-zA-Z0-9 ]{5,20}"), - prop::collection::vec("[a-z]{3,10}", 0..3), - ) - .prop_map(|(starred, comment, tags)| FlowAnnotations { - starred, - comment, - tags, - marker: None, - }) - } - - /// 生成随机的 LLMFlow - fn arb_llm_flow() -> impl Strategy { - ( - arb_flow_id(), - arb_flow_type(), - arb_llm_request(), - arb_flow_metadata(), - arb_llm_response(), - arb_flow_annotations(), - ) - .prop_map( - |(id, flow_type, request, metadata, response, annotations)| { - let mut flow = LLMFlow::new(id, flow_type, request, metadata); - flow.response = response; - flow.annotations = annotations; - flow - }, - ) - } - - // ======================================================================== - // 属性测试 - // ======================================================================== - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: llm-flow-monitor, Property 3: 存储 Round-Trip** - /// **Validates: Requirements 3.3, 3.5** - /// - /// *对于任意* 有效的 LLMFlow,存储到 Flow_Store 后再读取, - /// 读取的 Flow 应该与原始 Flow 等价。 - #[test] - fn prop_file_store_roundtrip( - flow in arb_llm_flow(), - ) { - let temp_dir = TempDir::new().unwrap(); - let store = FlowFileStore::new( - temp_dir.path().to_path_buf(), - RotationConfig::default(), - ).unwrap(); - - let original_id = flow.id.clone(); - let original_model = flow.request.model.clone(); - let original_provider = flow.metadata.provider; - let original_state = flow.state.clone(); - let original_content = flow.response.as_ref().map(|r| r.content.clone()); - let original_starred = flow.annotations.starred; - - // 写入 - store.write(&flow).unwrap(); - - // 读取 - let retrieved = store.get(&original_id).unwrap(); - prop_assert!(retrieved.is_some(), "Flow 应该能够被读取"); - - let retrieved = retrieved.unwrap(); - - // 验证关键字段一致 - prop_assert_eq!(&retrieved.id, &original_id, "ID 应该一致"); - prop_assert_eq!(&retrieved.request.model, &original_model, "模型应该一致"); - prop_assert_eq!(&retrieved.metadata.provider, &original_provider, "Provider 应该一致"); - prop_assert_eq!(&retrieved.state, &original_state, "状态应该一致"); - prop_assert_eq!( - retrieved.response.as_ref().map(|r| r.content.clone()), - original_content, - "响应内容应该一致" - ); - prop_assert_eq!(retrieved.annotations.starred, original_starred, "收藏状态应该一致"); - } - } - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: llm-flow-monitor, Property 3b: 多 Flow 存储 Round-Trip** - /// **Validates: Requirements 3.3, 3.5** - /// - /// *对于任意* 多个有效的 LLMFlow,存储后都应该能够正确读取。 - #[test] - fn prop_file_store_multiple_roundtrip( - flow_count in 1usize..=20usize, - ) { - let temp_dir = TempDir::new().unwrap(); - let store = FlowFileStore::new( - temp_dir.path().to_path_buf(), - RotationConfig::default(), - ).unwrap(); - - // 创建并写入多个 Flow - let mut original_flows = Vec::new(); - for i in 0..flow_count { - let id = format!("flow-{i:04}"); - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model: "gpt-4".to_string(), - ..Default::default() - }; - let metadata = FlowMetadata { - provider: ProviderType::OpenAI, - ..Default::default() - }; - let flow = LLMFlow::new(id, FlowType::ChatCompletions, request, metadata); - store.write(&flow).unwrap(); - original_flows.push(flow); - } - - // 验证所有 Flow 都可以读取 - for original in &original_flows { - let retrieved = store.get(&original.id).unwrap(); - prop_assert!(retrieved.is_some(), "Flow {} 应该能够被读取", original.id); - - let retrieved = retrieved.unwrap(); - prop_assert_eq!(&retrieved.id, &original.id, "ID 应该一致"); - prop_assert_eq!(&retrieved.request.model, &original.request.model, "模型应该一致"); - } - - // 验证索引数量正确 - prop_assert_eq!( - store.count().unwrap(), - flow_count, - "索引中的 Flow 数量应该正确" - ); - } - } - - proptest! { - #![proptest_config(ProptestConfig::with_cases(50))] - - /// **Feature: llm-flow-monitor, Property 3c: 文件轮转后 Round-Trip** - /// **Validates: Requirements 3.3, 3.4** - /// - /// *对于任意* Flow 序列,即使触发文件轮转,所有 Flow 都应该能够正确读取。 - #[test] - fn prop_file_store_rotation_roundtrip( - flow_count in 5usize..=15usize, - ) { - let temp_dir = TempDir::new().unwrap(); - // 使用很小的文件大小强制轮转 - let config = RotationConfig { - max_file_size: 500, // 500 字节,强制频繁轮转 - ..Default::default() - }; - let store = FlowFileStore::new(temp_dir.path().to_path_buf(), config).unwrap(); - - // 创建并写入多个 Flow - let mut original_ids = Vec::new(); - for i in 0..flow_count { - let id = format!("rotation-flow-{i:04}"); - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model: "gpt-4".to_string(), - ..Default::default() - }; - let metadata = FlowMetadata { - provider: ProviderType::OpenAI, - ..Default::default() - }; - let flow = LLMFlow::new(id.clone(), FlowType::ChatCompletions, request, metadata); - store.write(&flow).unwrap(); - original_ids.push(id); - } - - // 验证所有 Flow 都可以读取(即使跨多个文件) - for id in &original_ids { - let retrieved = store.get(id).unwrap(); - prop_assert!(retrieved.is_some(), "Flow {} 应该能够被读取(即使在轮转后)", id); - } - } - } -} diff --git a/src-tauri/src/flow_monitor/filter_parser.rs b/src-tauri/src/flow_monitor/filter_parser.rs deleted file mode 100644 index 645f37b3d..000000000 --- a/src-tauri/src/flow_monitor/filter_parser.rs +++ /dev/null @@ -1,1977 +0,0 @@ -//! 过滤表达式解析器 -//! -//! 该模块实现类似 mitmproxy 的过滤表达式语法,支持组合条件过滤 Flow。 -//! -//! # 支持的过滤器 -//! -//! - `~m `: 模型名称匹配 -//! - `~p `: 提供商匹配 -//! - `~s `: 状态匹配 (pending/streaming/completed/failed) -//! - `~e`: 有错误 - -#![allow(dead_code)] -//! - `~t`: 有工具调用 -//! - `~k`: 有思维链 -//! - `~starred`: 已收藏 -//! - `~tag `: 包含标签 -//! - `~b `: 请求或响应内容匹配 -//! - `~bq `: 请求内容匹配 -//! - `~bs `: 响应内容匹配 -//! - `~tokens `: Token 数量比较 -//! - `~latency `: 延迟比较 (支持 s/ms 后缀) -//! - `&`: AND 逻辑 -//! - `|`: OR 逻辑 -//! - `!`: NOT 逻辑 -//! - `()`: 分组 - -use regex::Regex; -use serde::{Deserialize, Serialize}; -use std::fmt; -use thiserror::Error; - -use super::models::{FlowState, LLMFlow, MessageContent}; - -// ============================================================================ -// 错误类型 -// ============================================================================ - -/// 过滤表达式解析错误 -#[derive(Debug, Clone, Error, PartialEq, Eq, Serialize, Deserialize)] -pub enum FilterParseError { - /// 意外的字符 - #[error("意外的字符 '{0}' 在位置 {1}")] - UnexpectedChar(char, usize), - - /// 意外的 Token - #[error("意外的 Token '{0}' 在位置 {1}")] - UnexpectedToken(String, usize), - - /// 意外的输入结束 - #[error("意外的输入结束")] - UnexpectedEof, - - /// 未知的过滤器类型 - #[error("未知的过滤器类型 '{0}'")] - UnknownFilter(String), - - /// 缺少参数 - #[error("过滤器 '{0}' 缺少参数")] - MissingArgument(String), - - /// 无效的比较运算符 - #[error("无效的比较运算符 '{0}'")] - InvalidComparisonOp(String), - - /// 无效的数值 - #[error("无效的数值 '{0}'")] - InvalidNumber(String), - - /// 无效的状态值 - #[error("无效的状态值 '{0}',有效值: pending, streaming, completed, failed, cancelled")] - InvalidState(String), - - /// 无效的正则表达式 - #[error("无效的正则表达式: {0}")] - InvalidRegex(String), - - /// 括号不匹配 - #[error("括号不匹配")] - UnmatchedParen, - - /// 空表达式 - #[error("空表达式")] - EmptyExpression, -} - -// ============================================================================ -// 比较运算符 -// ============================================================================ - -/// 比较运算符 -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -pub enum ComparisonOp { - /// 大于 - Gt, - /// 大于等于 - Gte, - /// 小于 - Lt, - /// 小于等于 - Lte, - /// 等于 - Eq, -} - -impl fmt::Display for ComparisonOp { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - match self { - ComparisonOp::Gt => write!(f, ">"), - ComparisonOp::Gte => write!(f, ">="), - ComparisonOp::Lt => write!(f, "<"), - ComparisonOp::Lte => write!(f, "<="), - ComparisonOp::Eq => write!(f, "="), - } - } -} - -/// 数值比较 -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -pub struct Comparison { - pub op: ComparisonOp, - pub value: i64, -} - -impl Comparison { - /// 执行比较 - pub fn compare(&self, actual: i64) -> bool { - match self.op { - ComparisonOp::Gt => actual > self.value, - ComparisonOp::Gte => actual >= self.value, - ComparisonOp::Lt => actual < self.value, - ComparisonOp::Lte => actual <= self.value, - ComparisonOp::Eq => actual == self.value, - } - } -} - -impl fmt::Display for Comparison { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - write!(f, "{}{}", self.op, self.value) - } -} - -// ============================================================================ -// Token 类型 -// ============================================================================ - -/// 过滤表达式 Token -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] -pub enum FilterToken { - // 基础过滤器 - /// 模型名称匹配 (~m ) - Model(String), - /// 提供商匹配 (~p ) - Provider(String), - /// 状态匹配 (~s ) - State(FlowState), - /// 有错误 (~e) - HasError, - /// 有工具调用 (~t) - HasToolCalls, - /// 有思维链 (~k) - HasThinking, - /// 已收藏 (~starred) - Starred, - /// 包含标签 (~tag ) - Tag(String), - - // 内容搜索 - /// 请求或响应内容匹配 (~b ) - Body(String), - /// 请求内容匹配 (~bq ) - BodyRequest(String), - /// 响应内容匹配 (~bs ) - BodyResponse(String), - - // 数值比较 - /// Token 数量比较 (~tokens ) - Tokens(Comparison), - /// 延迟比较 (~latency ) - Latency(Comparison), - - // 逻辑运算 - /// AND 逻辑 (&) - And, - /// OR 逻辑 (|) - Or, - /// NOT 逻辑 (!) - Not, - - // 分组 - /// 左括号 ( - LeftParen, - /// 右括号 ) - RightParen, -} - -impl fmt::Display for FilterToken { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - match self { - FilterToken::Model(s) => write!(f, "~m {s}"), - FilterToken::Provider(s) => write!(f, "~p {s}"), - FilterToken::State(s) => write!(f, "~s {}", state_to_string(s)), - FilterToken::HasError => write!(f, "~e"), - FilterToken::HasToolCalls => write!(f, "~t"), - FilterToken::HasThinking => write!(f, "~k"), - FilterToken::Starred => write!(f, "~starred"), - FilterToken::Tag(s) => write!(f, "~tag {s}"), - FilterToken::Body(s) => write!(f, "~b {s}"), - FilterToken::BodyRequest(s) => write!(f, "~bq {s}"), - FilterToken::BodyResponse(s) => write!(f, "~bs {s}"), - FilterToken::Tokens(c) => write!(f, "~tokens {c}"), - FilterToken::Latency(c) => write!(f, "~latency {c}"), - FilterToken::And => write!(f, "&"), - FilterToken::Or => write!(f, "|"), - FilterToken::Not => write!(f, "!"), - FilterToken::LeftParen => write!(f, "("), - FilterToken::RightParen => write!(f, ")"), - } - } -} - -/// 将 FlowState 转换为字符串 -fn state_to_string(state: &FlowState) -> &'static str { - match state { - FlowState::Pending => "pending", - FlowState::Streaming => "streaming", - FlowState::Completed => "completed", - FlowState::Failed => "failed", - FlowState::Cancelled => "cancelled", - } -} - -/// 从字符串解析 FlowState -fn parse_state(s: &str) -> Result { - match s.to_lowercase().as_str() { - "pending" => Ok(FlowState::Pending), - "streaming" => Ok(FlowState::Streaming), - "completed" => Ok(FlowState::Completed), - "failed" => Ok(FlowState::Failed), - "cancelled" => Ok(FlowState::Cancelled), - _ => Err(FilterParseError::InvalidState(s.to_string())), - } -} - -// ============================================================================ -// AST 表达式 -// ============================================================================ - -/// 过滤表达式 AST -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] -pub enum FilterExpr { - /// 单个 Token - Token(FilterToken), - /// AND 表达式 - And(Box, Box), - /// OR 表达式 - Or(Box, Box), - /// NOT 表达式 - Not(Box), -} - -impl fmt::Display for FilterExpr { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - match self { - FilterExpr::Token(t) => write!(f, "{t}"), - FilterExpr::And(left, right) => write!(f, "({left} & {right})"), - FilterExpr::Or(left, right) => write!(f, "({left} | {right})"), - FilterExpr::Not(expr) => write!(f, "!{expr}"), - } - } -} - -// ============================================================================ -// 词法分析器 (Lexer) -// ============================================================================ - -/// 词法分析器 -struct Lexer<'a> { - input: &'a str, - chars: std::iter::Peekable>, - pos: usize, -} - -impl<'a> Lexer<'a> { - fn new(input: &'a str) -> Self { - Self { - input, - chars: input.char_indices().peekable(), - pos: 0, - } - } - - /// 跳过空白字符 - fn skip_whitespace(&mut self) { - while let Some(&(_, c)) = self.chars.peek() { - if c.is_whitespace() { - self.chars.next(); - } else { - break; - } - } - } - - /// 读取一个单词(字母数字和下划线、连字符) - fn read_word(&mut self) -> String { - let mut word = String::new(); - while let Some(&(_, c)) = self.chars.peek() { - if c.is_alphanumeric() || c == '_' || c == '-' || c == '.' || c == '*' { - word.push(c); - self.chars.next(); - } else { - break; - } - } - word - } - - /// 读取带引号的字符串 - fn read_quoted_string(&mut self, quote: char) -> Result { - let mut s = String::new(); - // 跳过开始引号 - self.chars.next(); - - while let Some((pos, c)) = self.chars.next() { - if c == quote { - return Ok(s); - } else if c == '\\' { - // 转义字符 - if let Some((_, next_c)) = self.chars.next() { - s.push(next_c); - } else { - return Err(FilterParseError::UnexpectedEof); - } - } else { - s.push(c); - } - self.pos = pos; - } - Err(FilterParseError::UnexpectedEof) - } - - /// 读取参数(可能带引号或不带引号) - fn read_argument(&mut self) -> Result { - self.skip_whitespace(); - - if let Some(&(_, c)) = self.chars.peek() { - if c == '"' || c == '\'' { - return self.read_quoted_string(c); - } - } - - let word = self.read_word(); - if word.is_empty() { - return Err(FilterParseError::UnexpectedEof); - } - Ok(word) - } - - /// 解析比较运算符和数值 - fn parse_comparison(&mut self, filter_name: &str) -> Result { - self.skip_whitespace(); - - // 读取运算符 - let op = match self.chars.peek() { - Some(&(_, '>')) => { - self.chars.next(); - if let Some(&(_, '=')) = self.chars.peek() { - self.chars.next(); - ComparisonOp::Gte - } else { - ComparisonOp::Gt - } - } - Some(&(_, '<')) => { - self.chars.next(); - if let Some(&(_, '=')) = self.chars.peek() { - self.chars.next(); - ComparisonOp::Lte - } else { - ComparisonOp::Lt - } - } - Some(&(_, '=')) => { - self.chars.next(); - ComparisonOp::Eq - } - Some(&(_pos, c)) => { - return Err(FilterParseError::InvalidComparisonOp(c.to_string())); - } - None => { - return Err(FilterParseError::MissingArgument(filter_name.to_string())); - } - }; - - self.skip_whitespace(); - - // 读取数值(可能带单位) - let value_str = self.read_word(); - if value_str.is_empty() { - return Err(FilterParseError::MissingArgument(filter_name.to_string())); - } - - let value = self.parse_value_with_unit(&value_str, filter_name)?; - - Ok(Comparison { op, value }) - } - - /// 解析带单位的数值 - fn parse_value_with_unit(&self, s: &str, filter_name: &str) -> Result { - let s = s.to_lowercase(); - - // 检查是否有单位后缀 - if filter_name == "latency" { - if let Some(num_str) = s.strip_suffix("ms") { - return num_str - .parse::() - .map_err(|_| FilterParseError::InvalidNumber(s.clone())); - } else if let Some(num_str) = s.strip_suffix('s') { - return num_str - .parse::() - .map(|n| n * 1000) - .map_err(|_| FilterParseError::InvalidNumber(s.clone())); - } - } - - // 尝试直接解析为数字 - s.parse::() - .map_err(|_| FilterParseError::InvalidNumber(s)) - } - - /// 解析过滤器 Token - fn parse_filter(&mut self) -> Result { - self.skip_whitespace(); - - // 读取过滤器名称 - let filter_name = self.read_word(); - - match filter_name.as_str() { - "m" => { - let pattern = self.read_argument()?; - Ok(FilterToken::Model(pattern)) - } - "p" => { - let provider = self.read_argument()?; - Ok(FilterToken::Provider(provider)) - } - "s" => { - let state_str = self.read_argument()?; - let state = parse_state(&state_str)?; - Ok(FilterToken::State(state)) - } - "e" => Ok(FilterToken::HasError), - "t" => Ok(FilterToken::HasToolCalls), - "k" => Ok(FilterToken::HasThinking), - "starred" => Ok(FilterToken::Starred), - "tag" => { - let tag = self.read_argument()?; - Ok(FilterToken::Tag(tag)) - } - "b" => { - let pattern = self.read_argument()?; - // 验证正则表达式 - Regex::new(&pattern).map_err(|e| FilterParseError::InvalidRegex(e.to_string()))?; - Ok(FilterToken::Body(pattern)) - } - "bq" => { - let pattern = self.read_argument()?; - Regex::new(&pattern).map_err(|e| FilterParseError::InvalidRegex(e.to_string()))?; - Ok(FilterToken::BodyRequest(pattern)) - } - "bs" => { - let pattern = self.read_argument()?; - Regex::new(&pattern).map_err(|e| FilterParseError::InvalidRegex(e.to_string()))?; - Ok(FilterToken::BodyResponse(pattern)) - } - "tokens" => { - let comparison = self.parse_comparison("tokens")?; - Ok(FilterToken::Tokens(comparison)) - } - "latency" => { - let comparison = self.parse_comparison("latency")?; - Ok(FilterToken::Latency(comparison)) - } - _ => Err(FilterParseError::UnknownFilter(filter_name)), - } - } - - /// 获取下一个 Token - fn next_token(&mut self) -> Result, FilterParseError> { - self.skip_whitespace(); - - match self.chars.peek() { - None => Ok(None), - Some(&(pos, c)) => { - self.pos = pos; - match c { - '~' => { - self.chars.next(); - let token = self.parse_filter()?; - Ok(Some(token)) - } - '&' => { - self.chars.next(); - Ok(Some(FilterToken::And)) - } - '|' => { - self.chars.next(); - Ok(Some(FilterToken::Or)) - } - '!' => { - self.chars.next(); - Ok(Some(FilterToken::Not)) - } - '(' => { - self.chars.next(); - Ok(Some(FilterToken::LeftParen)) - } - ')' => { - self.chars.next(); - Ok(Some(FilterToken::RightParen)) - } - _ => Err(FilterParseError::UnexpectedChar(c, pos)), - } - } - } - } - - /// 词法分析,返回所有 Token - fn tokenize(&mut self) -> Result, FilterParseError> { - let mut tokens = Vec::new(); - while let Some(token) = self.next_token()? { - tokens.push(token); - } - Ok(tokens) - } -} - -// ============================================================================ -// 语法分析器 (Parser) -// ============================================================================ - -/// 语法分析器 -struct Parser { - tokens: Vec, - pos: usize, -} - -impl Parser { - fn new(tokens: Vec) -> Self { - Self { tokens, pos: 0 } - } - - /// 查看当前 Token - fn peek(&self) -> Option<&FilterToken> { - self.tokens.get(self.pos) - } - - /// 消费当前 Token - fn advance(&mut self) -> Option { - if self.pos < self.tokens.len() { - let token = self.tokens[self.pos].clone(); - self.pos += 1; - Some(token) - } else { - None - } - } - - /// 检查当前 Token 是否匹配 - fn check(&self, token: &FilterToken) -> bool { - self.peek() - .is_some_and(|t| std::mem::discriminant(t) == std::mem::discriminant(token)) - } - - /// 解析表达式 - fn parse_expr(&mut self) -> Result { - self.parse_or() - } - - /// 解析 OR 表达式 - fn parse_or(&mut self) -> Result { - let mut left = self.parse_and()?; - - while self.check(&FilterToken::Or) { - self.advance(); // 消费 | - let right = self.parse_and()?; - left = FilterExpr::Or(Box::new(left), Box::new(right)); - } - - Ok(left) - } - - /// 解析 AND 表达式 - fn parse_and(&mut self) -> Result { - let mut left = self.parse_unary()?; - - while self.check(&FilterToken::And) { - self.advance(); // 消费 & - let right = self.parse_unary()?; - left = FilterExpr::And(Box::new(left), Box::new(right)); - } - - Ok(left) - } - - /// 解析一元表达式 (NOT) - fn parse_unary(&mut self) -> Result { - if self.check(&FilterToken::Not) { - self.advance(); // 消费 ! - let expr = self.parse_unary()?; - return Ok(FilterExpr::Not(Box::new(expr))); - } - - self.parse_primary() - } - - /// 解析基本表达式 - fn parse_primary(&mut self) -> Result { - match self.peek() { - Some(FilterToken::LeftParen) => { - self.advance(); // 消费 ( - let expr = self.parse_expr()?; - - // 期望 ) - match self.peek() { - Some(FilterToken::RightParen) => { - self.advance(); - Ok(expr) - } - _ => Err(FilterParseError::UnmatchedParen), - } - } - Some(token) => { - // 检查是否是过滤器 Token - match token { - FilterToken::And | FilterToken::Or | FilterToken::RightParen => Err( - FilterParseError::UnexpectedToken(format!("{token}"), self.pos), - ), - _ => { - let token = self.advance().unwrap(); - Ok(FilterExpr::Token(token)) - } - } - } - None => Err(FilterParseError::UnexpectedEof), - } - } -} - -// ============================================================================ -// FilterParser 公共接口 -// ============================================================================ - -/// 过滤表达式解析器 -pub struct FilterParser; - -impl FilterParser { - /// 解析过滤表达式字符串 - pub fn parse(input: &str) -> Result { - let input = input.trim(); - if input.is_empty() { - return Err(FilterParseError::EmptyExpression); - } - - let mut lexer = Lexer::new(input); - let tokens = lexer.tokenize()?; - - if tokens.is_empty() { - return Err(FilterParseError::EmptyExpression); - } - - let mut parser = Parser::new(tokens); - let expr = parser.parse_expr()?; - - // 检查是否还有未消费的 Token - if parser.peek().is_some() { - return Err(FilterParseError::UnexpectedToken( - format!("{}", parser.peek().unwrap()), - parser.pos, - )); - } - - Ok(expr) - } - - /// 验证表达式语法 - pub fn validate(input: &str) -> Result<(), FilterParseError> { - Self::parse(input)?; - Ok(()) - } - - /// 将 FilterExpr 编译为可执行的过滤函数 - pub fn compile(expr: &FilterExpr) -> Box bool + Send + Sync> { - let expr = expr.clone(); - Box::new(move |flow| Self::evaluate(&expr, flow)) - } - - /// 评估表达式 - fn evaluate(expr: &FilterExpr, flow: &LLMFlow) -> bool { - match expr { - FilterExpr::Token(token) => Self::evaluate_token(token, flow), - FilterExpr::And(left, right) => { - Self::evaluate(left, flow) && Self::evaluate(right, flow) - } - FilterExpr::Or(left, right) => { - Self::evaluate(left, flow) || Self::evaluate(right, flow) - } - FilterExpr::Not(inner) => !Self::evaluate(inner, flow), - } - } - - /// 评估单个 Token - fn evaluate_token(token: &FilterToken, flow: &LLMFlow) -> bool { - match token { - FilterToken::Model(pattern) => Self::match_pattern(pattern, &flow.request.model), - FilterToken::Provider(provider) => { - let flow_provider = format!("{:?}", flow.metadata.provider).to_lowercase(); - flow_provider.contains(&provider.to_lowercase()) - } - FilterToken::State(state) => flow.state == *state, - FilterToken::HasError => flow.error.is_some(), - FilterToken::HasToolCalls => flow - .response - .as_ref() - .is_some_and(|r| !r.tool_calls.is_empty()), - FilterToken::HasThinking => { - flow.response.as_ref().is_some_and(|r| r.thinking.is_some()) - } - FilterToken::Starred => flow.annotations.starred, - FilterToken::Tag(tag) => flow - .annotations - .tags - .iter() - .any(|t| t.to_lowercase() == tag.to_lowercase()), - FilterToken::Body(pattern) => { - let request_text = Self::get_request_text(flow); - let response_text = flow - .response - .as_ref() - .map_or(String::new(), |r| r.content.clone()); - let combined = format!("{request_text}\n{response_text}"); - - if let Ok(re) = Regex::new(pattern) { - re.is_match(&combined) - } else { - combined.to_lowercase().contains(&pattern.to_lowercase()) - } - } - FilterToken::BodyRequest(pattern) => { - let request_text = Self::get_request_text(flow); - - if let Ok(re) = Regex::new(pattern) { - re.is_match(&request_text) - } else { - request_text - .to_lowercase() - .contains(&pattern.to_lowercase()) - } - } - FilterToken::BodyResponse(pattern) => { - let response_text = flow - .response - .as_ref() - .map_or(String::new(), |r| r.content.clone()); - - if let Ok(re) = Regex::new(pattern) { - re.is_match(&response_text) - } else { - response_text - .to_lowercase() - .contains(&pattern.to_lowercase()) - } - } - FilterToken::Tokens(comparison) => { - let total_tokens = flow - .response - .as_ref() - .map_or(0, |r| r.usage.total_tokens as i64); - comparison.compare(total_tokens) - } - FilterToken::Latency(comparison) => { - comparison.compare(flow.timestamps.duration_ms as i64) - } - // 逻辑运算符和括号不应该在这里出现 - FilterToken::And - | FilterToken::Or - | FilterToken::Not - | FilterToken::LeftParen - | FilterToken::RightParen => false, - } - } - - /// 模式匹配(支持 * 通配符) - fn match_pattern(pattern: &str, text: &str) -> bool { - if pattern == "*" { - return true; - } - - let pattern_lower = pattern.to_lowercase(); - let text_lower = text.to_lowercase(); - - if pattern.contains('*') { - // 通配符匹配 - let parts: Vec<&str> = pattern_lower.split('*').collect(); - let mut pos = 0; - - for (i, part) in parts.iter().enumerate() { - if part.is_empty() { - continue; - } - - if let Some(found_pos) = text_lower[pos..].find(part) { - // 第一个部分必须从开头匹配(如果模式不以 * 开头) - if i == 0 && found_pos != 0 && !pattern_lower.starts_with('*') { - return false; - } - pos += found_pos + part.len(); - } else { - return false; - } - } - - // 最后一个部分必须匹配到结尾(如果模式不以 * 结尾) - if !pattern_lower.ends_with('*') && pos != text_lower.len() { - return false; - } - - true - } else { - // 不含通配符时,检查是否包含该模式 - text_lower.contains(&pattern_lower) - } - } - - /// 获取请求文本(用于搜索) - fn get_request_text(flow: &LLMFlow) -> String { - let mut text = String::new(); - - if let Some(ref system) = flow.request.system_prompt { - text.push_str(system); - text.push('\n'); - } - - for msg in &flow.request.messages { - match &msg.content { - MessageContent::Text(s) => { - text.push_str(s); - text.push('\n'); - } - MessageContent::MultiModal(parts) => { - for part in parts { - if let super::models::ContentPart::Text { text: t } = part { - text.push_str(t); - text.push('\n'); - } - } - } - } - } - - text - } -} - -// ============================================================================ -// 帮助信息 -// ============================================================================ - -/// 过滤表达式帮助信息 -pub const FILTER_HELP: &[(&str, &str)] = &[ - ("~m ", "模型名称匹配(支持 * 通配符)"), - ("~p ", "提供商匹配"), - ( - "~s ", - "状态匹配 (pending/streaming/completed/failed/cancelled)", - ), - ("~e", "有错误"), - ("~t", "有工具调用"), - ("~k", "有思维链"), - ("~starred", "已收藏"), - ("~tag ", "包含标签"), - ("~b ", "请求或响应内容匹配(正则表达式)"), - ("~bq ", "请求内容匹配(正则表达式)"), - ("~bs ", "响应内容匹配(正则表达式)"), - ("~tokens ", "Token 数量比较 (>, >=, <, <=, =)"), - ("~latency ", "延迟比较 (支持 s/ms 后缀)"), - ("&", "AND 逻辑"), - ("|", "OR 逻辑"), - ("!", "NOT 逻辑"), - ("()", "分组"), -]; - -/// 获取帮助文本 -pub fn get_filter_help() -> String { - let mut help = String::from("过滤表达式语法:\n\n"); - for (syntax, desc) in FILTER_HELP { - help.push_str(&format!(" {syntax:<20} {desc}\n")); - } - help.push_str("\n示例:\n"); - help.push_str(" ~m claude 模型名称包含 'claude'\n"); - help.push_str(" ~p kiro & ~m claude 提供商为 kiro 且模型包含 claude\n"); - help.push_str(" ~e | ~latency >5s 有错误或延迟超过 5 秒\n"); - help.push_str(" !~e 没有错误\n"); - help.push_str(" (~p kiro | ~p gemini) & ~tokens >1000\n"); - help -} - -// ============================================================================ -// 单元测试 -// ============================================================================ - -#[cfg(test)] -mod tests { - use super::*; - use crate::flow_monitor::models::{ - FlowMetadata, FlowType, LLMRequest, LLMResponse, RequestParameters, TokenUsage, - }; - use crate::ProviderType; - - /// 创建测试用的 Flow - fn create_test_flow(model: &str, provider: ProviderType) -> LLMFlow { - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model: model.to_string(), - parameters: RequestParameters::default(), - ..Default::default() - }; - - let metadata = FlowMetadata { - provider, - ..Default::default() - }; - - LLMFlow::new( - "test-id".to_string(), - FlowType::ChatCompletions, - request, - metadata, - ) - } - - #[test] - fn test_parse_model_filter() { - let expr = FilterParser::parse("~m claude").unwrap(); - assert!(matches!(expr, FilterExpr::Token(FilterToken::Model(s)) if s == "claude")); - } - - #[test] - fn test_parse_provider_filter() { - let expr = FilterParser::parse("~p kiro").unwrap(); - assert!(matches!(expr, FilterExpr::Token(FilterToken::Provider(s)) if s == "kiro")); - } - - #[test] - fn test_parse_state_filter() { - let expr = FilterParser::parse("~s completed").unwrap(); - assert!(matches!( - expr, - FilterExpr::Token(FilterToken::State(FlowState::Completed)) - )); - } - - #[test] - fn test_parse_has_error_filter() { - let expr = FilterParser::parse("~e").unwrap(); - assert!(matches!(expr, FilterExpr::Token(FilterToken::HasError))); - } - - #[test] - fn test_parse_has_tool_calls_filter() { - let expr = FilterParser::parse("~t").unwrap(); - assert!(matches!(expr, FilterExpr::Token(FilterToken::HasToolCalls))); - } - - #[test] - fn test_parse_has_thinking_filter() { - let expr = FilterParser::parse("~k").unwrap(); - assert!(matches!(expr, FilterExpr::Token(FilterToken::HasThinking))); - } - - #[test] - fn test_parse_starred_filter() { - let expr = FilterParser::parse("~starred").unwrap(); - assert!(matches!(expr, FilterExpr::Token(FilterToken::Starred))); - } - - #[test] - fn test_parse_tag_filter() { - let expr = FilterParser::parse("~tag important").unwrap(); - assert!(matches!(expr, FilterExpr::Token(FilterToken::Tag(s)) if s == "important")); - } - - #[test] - fn test_parse_body_filter() { - let expr = FilterParser::parse("~b hello").unwrap(); - assert!(matches!(expr, FilterExpr::Token(FilterToken::Body(s)) if s == "hello")); - } - - #[test] - fn test_parse_body_request_filter() { - let expr = FilterParser::parse("~bq request").unwrap(); - assert!(matches!(expr, FilterExpr::Token(FilterToken::BodyRequest(s)) if s == "request")); - } - - #[test] - fn test_parse_body_response_filter() { - let expr = FilterParser::parse("~bs response").unwrap(); - assert!(matches!(expr, FilterExpr::Token(FilterToken::BodyResponse(s)) if s == "response")); - } - - #[test] - fn test_parse_tokens_filter() { - let expr = FilterParser::parse("~tokens >1000").unwrap(); - if let FilterExpr::Token(FilterToken::Tokens(c)) = expr { - assert_eq!(c.op, ComparisonOp::Gt); - assert_eq!(c.value, 1000); - } else { - panic!("Expected Tokens filter"); - } - } - - #[test] - fn test_parse_latency_filter_seconds() { - let expr = FilterParser::parse("~latency >5s").unwrap(); - if let FilterExpr::Token(FilterToken::Latency(c)) = expr { - assert_eq!(c.op, ComparisonOp::Gt); - assert_eq!(c.value, 5000); // 5 seconds = 5000 ms - } else { - panic!("Expected Latency filter"); - } - } - - #[test] - fn test_parse_latency_filter_milliseconds() { - let expr = FilterParser::parse("~latency >=500ms").unwrap(); - if let FilterExpr::Token(FilterToken::Latency(c)) = expr { - assert_eq!(c.op, ComparisonOp::Gte); - assert_eq!(c.value, 500); - } else { - panic!("Expected Latency filter"); - } - } - - #[test] - fn test_parse_and_expression() { - let expr = FilterParser::parse("~p kiro & ~m claude").unwrap(); - assert!(matches!(expr, FilterExpr::And(_, _))); - } - - #[test] - fn test_parse_or_expression() { - let expr = FilterParser::parse("~p kiro | ~p gemini").unwrap(); - assert!(matches!(expr, FilterExpr::Or(_, _))); - } - - #[test] - fn test_parse_not_expression() { - let expr = FilterParser::parse("!~e").unwrap(); - assert!(matches!(expr, FilterExpr::Not(_))); - } - - #[test] - fn test_parse_grouped_expression() { - let expr = FilterParser::parse("(~p kiro | ~p gemini) & ~m claude").unwrap(); - assert!(matches!(expr, FilterExpr::And(_, _))); - } - - #[test] - fn test_parse_complex_expression() { - let expr = FilterParser::parse("~p kiro & ~m claude & !~e").unwrap(); - // Should parse as ((~p kiro & ~m claude) & !~e) - assert!(matches!(expr, FilterExpr::And(_, _))); - } - - #[test] - fn test_parse_error_unknown_filter() { - let result = FilterParser::parse("~unknown"); - assert!(matches!(result, Err(FilterParseError::UnknownFilter(_)))); - } - - #[test] - fn test_parse_error_invalid_state() { - let result = FilterParser::parse("~s invalid"); - assert!(matches!(result, Err(FilterParseError::InvalidState(_)))); - } - - #[test] - fn test_parse_error_empty_expression() { - let result = FilterParser::parse(""); - assert!(matches!(result, Err(FilterParseError::EmptyExpression))); - } - - #[test] - fn test_parse_error_unmatched_paren() { - let result = FilterParser::parse("(~m claude"); - assert!(matches!(result, Err(FilterParseError::UnmatchedParen))); - } - - #[test] - fn test_evaluate_model_filter() { - let flow = create_test_flow("claude-3-opus", ProviderType::Kiro); - let expr = FilterParser::parse("~m claude").unwrap(); - let filter = FilterParser::compile(&expr); - assert!(filter(&flow)); - - let expr = FilterParser::parse("~m gpt").unwrap(); - let filter = FilterParser::compile(&expr); - assert!(!filter(&flow)); - } - - #[test] - fn test_evaluate_provider_filter() { - let flow = create_test_flow("claude-3", ProviderType::Kiro); - let expr = FilterParser::parse("~p kiro").unwrap(); - let filter = FilterParser::compile(&expr); - assert!(filter(&flow)); - - let expr = FilterParser::parse("~p openai").unwrap(); - let filter = FilterParser::compile(&expr); - assert!(!filter(&flow)); - } - - #[test] - fn test_evaluate_state_filter() { - let mut flow = create_test_flow("claude-3", ProviderType::Kiro); - flow.state = FlowState::Completed; - - let expr = FilterParser::parse("~s completed").unwrap(); - let filter = FilterParser::compile(&expr); - assert!(filter(&flow)); - - let expr = FilterParser::parse("~s pending").unwrap(); - let filter = FilterParser::compile(&expr); - assert!(!filter(&flow)); - } - - #[test] - fn test_evaluate_starred_filter() { - let mut flow = create_test_flow("claude-3", ProviderType::Kiro); - - let expr = FilterParser::parse("~starred").unwrap(); - let filter = FilterParser::compile(&expr); - assert!(!filter(&flow)); - - flow.annotations.starred = true; - assert!(filter(&flow)); - } - - #[test] - fn test_evaluate_tag_filter() { - let mut flow = create_test_flow("claude-3", ProviderType::Kiro); - flow.annotations.tags = vec!["important".to_string(), "test".to_string()]; - - let expr = FilterParser::parse("~tag important").unwrap(); - let filter = FilterParser::compile(&expr); - assert!(filter(&flow)); - - let expr = FilterParser::parse("~tag missing").unwrap(); - let filter = FilterParser::compile(&expr); - assert!(!filter(&flow)); - } - - #[test] - fn test_evaluate_tokens_filter() { - let mut flow = create_test_flow("claude-3", ProviderType::Kiro); - flow.response = Some(LLMResponse { - usage: TokenUsage { - input_tokens: 500, - output_tokens: 600, - total_tokens: 1100, - ..Default::default() - }, - ..Default::default() - }); - - let expr = FilterParser::parse("~tokens >1000").unwrap(); - let filter = FilterParser::compile(&expr); - assert!(filter(&flow)); - - let expr = FilterParser::parse("~tokens <1000").unwrap(); - let filter = FilterParser::compile(&expr); - assert!(!filter(&flow)); - } - - #[test] - fn test_evaluate_latency_filter() { - let mut flow = create_test_flow("claude-3", ProviderType::Kiro); - flow.timestamps.duration_ms = 6000; // 6 seconds - - let expr = FilterParser::parse("~latency >5s").unwrap(); - let filter = FilterParser::compile(&expr); - assert!(filter(&flow)); - - let expr = FilterParser::parse("~latency <5000ms").unwrap(); - let filter = FilterParser::compile(&expr); - assert!(!filter(&flow)); - } - - #[test] - fn test_evaluate_and_expression() { - let flow = create_test_flow("claude-3-opus", ProviderType::Kiro); - - let expr = FilterParser::parse("~p kiro & ~m claude").unwrap(); - let filter = FilterParser::compile(&expr); - assert!(filter(&flow)); - - let expr = FilterParser::parse("~p openai & ~m claude").unwrap(); - let filter = FilterParser::compile(&expr); - assert!(!filter(&flow)); - } - - #[test] - fn test_evaluate_or_expression() { - let flow = create_test_flow("claude-3", ProviderType::Kiro); - - let expr = FilterParser::parse("~p kiro | ~p gemini").unwrap(); - let filter = FilterParser::compile(&expr); - assert!(filter(&flow)); - - let expr = FilterParser::parse("~p openai | ~p gemini").unwrap(); - let filter = FilterParser::compile(&expr); - assert!(!filter(&flow)); - } - - #[test] - fn test_evaluate_not_expression() { - let flow = create_test_flow("claude-3", ProviderType::Kiro); - - let expr = FilterParser::parse("!~e").unwrap(); - let filter = FilterParser::compile(&expr); - assert!(filter(&flow)); // No error, so !~e is true - } - - #[test] - fn test_display_filter_expr() { - let expr = FilterParser::parse("~p kiro & ~m claude").unwrap(); - let display = format!("{expr}"); - assert!(display.contains("~p kiro")); - assert!(display.contains("~m claude")); - } - - #[test] - fn test_round_trip_simple() { - let original = "~m claude"; - let expr = FilterParser::parse(original).unwrap(); - let display = format!("{expr}"); - let reparsed = FilterParser::parse(&display).unwrap(); - assert_eq!(format!("{expr}"), format!("{}", reparsed)); - } -} - -// ============================================================================ -// 属性测试 -// ============================================================================ - -#[cfg(test)] -mod property_tests { - use super::*; - use crate::flow_monitor::models::{ - FlowError, FlowErrorType, FlowMetadata, FlowType, FunctionCall, LLMRequest, LLMResponse, - RequestParameters, ThinkingContent, TokenUsage, ToolCall, - }; - use crate::ProviderType; - use proptest::prelude::*; - - // ======================================================================== - // 生成器 - // ======================================================================== - - /// 生成随机的 ProviderType - fn arb_provider_type() -> impl Strategy { - prop_oneof![ - Just(ProviderType::Kiro), - Just(ProviderType::Gemini), - Just(ProviderType::OpenAI), - Just(ProviderType::Claude), - Just(ProviderType::Antigravity), - ] - } - - /// 生成随机的 FlowState - fn arb_flow_state() -> impl Strategy { - prop_oneof![ - Just(FlowState::Pending), - Just(FlowState::Streaming), - Just(FlowState::Completed), - Just(FlowState::Failed), - Just(FlowState::Cancelled), - ] - } - - /// 生成随机的模型名称 - fn arb_model_name() -> impl Strategy { - prop_oneof![ - Just("gpt-4".to_string()), - Just("gpt-4-turbo".to_string()), - Just("gpt-3.5-turbo".to_string()), - Just("claude-3-opus".to_string()), - Just("claude-3-sonnet".to_string()), - Just("gemini-pro".to_string()), - Just("qwen-max".to_string()), - ] - } - - /// 生成随机的标签 - fn arb_tags() -> impl Strategy> { - prop::collection::vec("[a-z]{3,10}", 0..5) - } - - /// 生成随机的 LLMFlow - fn arb_llm_flow() -> impl Strategy { - ( - "[a-f0-9]{8}", - arb_model_name(), - arb_provider_type(), - arb_flow_state(), - any::(), // starred - arb_tags(), // tags - any::(), // has_error - any::(), // has_tool_calls - any::(), // has_thinking - 0u32..50000u32, // total_tokens - 0u64..30000u64, // duration_ms - ) - .prop_map( - |( - id, - model, - provider, - state, - starred, - tags, - has_error, - has_tool_calls, - has_thinking, - total_tokens, - duration_ms, - )| { - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model, - parameters: RequestParameters::default(), - ..Default::default() - }; - - let metadata = FlowMetadata { - provider, - ..Default::default() - }; - - let mut flow = LLMFlow::new(id, FlowType::ChatCompletions, request, metadata); - flow.state = state; - flow.annotations.starred = starred; - flow.annotations.tags = tags; - flow.timestamps.duration_ms = duration_ms; - - // 设置错误 - if has_error { - flow.error = Some(FlowError::new(FlowErrorType::ServerError, "Test error")); - } - - // 设置响应 - let mut response = LLMResponse { - usage: TokenUsage { - input_tokens: total_tokens / 2, - output_tokens: total_tokens / 2, - total_tokens, - ..Default::default() - }, - ..Default::default() - }; - - // 设置工具调用 - if has_tool_calls { - response.tool_calls = vec![ToolCall { - id: "call_1".to_string(), - tool_type: "function".to_string(), - function: FunctionCall { - name: "test_function".to_string(), - arguments: "{}".to_string(), - }, - }]; - } - - // 设置思维链 - if has_thinking { - response.thinking = Some(ThinkingContent { - text: "Thinking...".to_string(), - tokens: Some(100), - signature: None, - }); - } - - flow.response = Some(response); - flow - }, - ) - } - - // ======================================================================== - // Property 1: 过滤表达式正确性 - // ======================================================================== - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: flow-monitor-enhancement, Property 1: 过滤表达式正确性** - /// **Validates: Requirements 1.1-1.16** - /// - /// *对于任意* 有效的过滤表达式和 Flow 集合,解析并执行过滤后, - /// 返回的所有 Flow 都应该满足该表达式定义的条件。 - #[test] - fn prop_filter_model_correctness( - flow in arb_llm_flow(), - ) { - // 测试模型过滤器正确性 - let model = flow.request.model.clone(); - let expr_str = format!("~m {model}"); - let expr = FilterParser::parse(&expr_str).unwrap(); - let filter = FilterParser::compile(&expr); - - // 使用完整模型名称过滤应该匹配 - prop_assert!( - filter(&flow), - "模型过滤器 '{}' 应该匹配模型 '{}'", - expr_str, - model - ); - } - - #[test] - fn prop_filter_provider_correctness( - flow in arb_llm_flow(), - ) { - // 测试提供商过滤器正确性 - let provider_str = format!("{:?}", flow.metadata.provider).to_lowercase(); - let expr_str = format!("~p {provider_str}"); - let expr = FilterParser::parse(&expr_str).unwrap(); - let filter = FilterParser::compile(&expr); - - prop_assert!( - filter(&flow), - "提供商过滤器 '{}' 应该匹配提供商 '{:?}'", - expr_str, - flow.metadata.provider - ); - } - - #[test] - fn prop_filter_state_correctness( - flow in arb_llm_flow(), - ) { - // 测试状态过滤器正确性 - let state_str = state_to_string(&flow.state); - let expr_str = format!("~s {state_str}"); - let expr = FilterParser::parse(&expr_str).unwrap(); - let filter = FilterParser::compile(&expr); - - prop_assert!( - filter(&flow), - "状态过滤器 '{}' 应该匹配状态 '{:?}'", - expr_str, - flow.state - ); - } - - #[test] - fn prop_filter_error_correctness( - flow in arb_llm_flow(), - ) { - // 测试错误过滤器正确性 - let expr = FilterParser::parse("~e").unwrap(); - let filter = FilterParser::compile(&expr); - let result = filter(&flow); - - prop_assert_eq!( - result, - flow.error.is_some(), - "错误过滤器结果应该与 flow.error.is_some() 一致" - ); - } - - #[test] - fn prop_filter_tool_calls_correctness( - flow in arb_llm_flow(), - ) { - // 测试工具调用过滤器正确性 - let expr = FilterParser::parse("~t").unwrap(); - let filter = FilterParser::compile(&expr); - let result = filter(&flow); - - let has_tool_calls = flow - .response - .as_ref() - .is_some_and(|r| !r.tool_calls.is_empty()); - - prop_assert_eq!( - result, - has_tool_calls, - "工具调用过滤器结果应该与实际工具调用状态一致" - ); - } - - #[test] - fn prop_filter_thinking_correctness( - flow in arb_llm_flow(), - ) { - // 测试思维链过滤器正确性 - let expr = FilterParser::parse("~k").unwrap(); - let filter = FilterParser::compile(&expr); - let result = filter(&flow); - - let has_thinking = flow - .response - .as_ref() - .is_some_and(|r| r.thinking.is_some()); - - prop_assert_eq!( - result, - has_thinking, - "思维链过滤器结果应该与实际思维链状态一致" - ); - } - - #[test] - fn prop_filter_starred_correctness( - flow in arb_llm_flow(), - ) { - // 测试收藏过滤器正确性 - let expr = FilterParser::parse("~starred").unwrap(); - let filter = FilterParser::compile(&expr); - let result = filter(&flow); - - prop_assert_eq!( - result, - flow.annotations.starred, - "收藏过滤器结果应该与 flow.annotations.starred 一致" - ); - } - - #[test] - fn prop_filter_tokens_correctness( - flow in arb_llm_flow(), - threshold in 0i64..50000i64, - ) { - // 测试 Token 数量过滤器正确性 - let total_tokens = flow - .response - .as_ref() - .map_or(0, |r| r.usage.total_tokens as i64); - - // 测试大于 - let expr_str = format!("~tokens >{threshold}"); - let expr = FilterParser::parse(&expr_str).unwrap(); - let filter = FilterParser::compile(&expr); - let result = filter(&flow); - - prop_assert_eq!( - result, - total_tokens > threshold, - "Token 过滤器 '{}' 结果应该正确 (actual: {}, threshold: {})", - expr_str, - total_tokens, - threshold - ); - } - - #[test] - fn prop_filter_latency_correctness( - flow in arb_llm_flow(), - threshold in 0i64..30000i64, - ) { - // 测试延迟过滤器正确性 - let duration_ms = flow.timestamps.duration_ms as i64; - - // 测试大于 - let expr_str = format!("~latency >{threshold}ms"); - let expr = FilterParser::parse(&expr_str).unwrap(); - let filter = FilterParser::compile(&expr); - let result = filter(&flow); - - prop_assert_eq!( - result, - duration_ms > threshold, - "延迟过滤器 '{}' 结果应该正确 (actual: {}, threshold: {})", - expr_str, - duration_ms, - threshold - ); - } - - #[test] - fn prop_filter_and_correctness( - flow in arb_llm_flow(), - ) { - // 测试 AND 逻辑正确性 - let model = flow.request.model.clone(); - let provider_str = format!("{:?}", flow.metadata.provider).to_lowercase(); - - let expr_str = format!("~m {model} & ~p {provider_str}"); - let expr = FilterParser::parse(&expr_str).unwrap(); - let filter = FilterParser::compile(&expr); - - // 两个条件都应该满足 - prop_assert!( - filter(&flow), - "AND 表达式 '{}' 应该匹配", - expr_str - ); - } - - #[test] - fn prop_filter_or_correctness( - flow in arb_llm_flow(), - ) { - // 测试 OR 逻辑正确性 - let model = flow.request.model.clone(); - - // 使用一个匹配的条件和一个不匹配的条件 - let expr_str = format!("~m {model} | ~m nonexistent-model-xyz"); - let expr = FilterParser::parse(&expr_str).unwrap(); - let filter = FilterParser::compile(&expr); - - // 至少一个条件满足 - prop_assert!( - filter(&flow), - "OR 表达式 '{}' 应该匹配", - expr_str - ); - } - - #[test] - fn prop_filter_not_correctness( - flow in arb_llm_flow(), - ) { - // 测试 NOT 逻辑正确性 - let expr = FilterParser::parse("~e").unwrap(); - let filter_e = FilterParser::compile(&expr); - let result_e = filter_e(&flow); - - let expr_not = FilterParser::parse("!~e").unwrap(); - let filter_not_e = FilterParser::compile(&expr_not); - let result_not_e = filter_not_e(&flow); - - prop_assert_eq!( - result_not_e, - !result_e, - "NOT 表达式结果应该是原表达式的取反" - ); - } - } - - // ======================================================================== - // Property 2: 过滤表达式 Round-Trip - // ======================================================================== - - /// 生成随机的比较运算符 - fn arb_comparison_op() -> impl Strategy { - prop_oneof![ - Just(ComparisonOp::Gt), - Just(ComparisonOp::Gte), - Just(ComparisonOp::Lt), - Just(ComparisonOp::Lte), - Just(ComparisonOp::Eq), - ] - } - - /// 生成随机的 Comparison - fn arb_comparison() -> impl Strategy { - (arb_comparison_op(), 0i64..100000i64).prop_map(|(op, value)| Comparison { op, value }) - } - - /// 生成随机的简单 FilterToken(不包括逻辑运算符和括号) - fn arb_simple_filter_token() -> impl Strategy { - prop_oneof![ - arb_model_name().prop_map(FilterToken::Model), - prop_oneof![ - Just("kiro".to_string()), - Just("openai".to_string()), - Just("claude".to_string()), - Just("gemini".to_string()), - ] - .prop_map(FilterToken::Provider), - arb_flow_state().prop_map(FilterToken::State), - Just(FilterToken::HasError), - Just(FilterToken::HasToolCalls), - Just(FilterToken::HasThinking), - Just(FilterToken::Starred), - "[a-z]{3,8}".prop_map(FilterToken::Tag), - arb_comparison().prop_map(FilterToken::Tokens), - arb_comparison().prop_map(FilterToken::Latency), - ] - } - - /// 生成随机的 FilterExpr - fn arb_filter_expr() -> impl Strategy { - arb_simple_filter_token() - .prop_map(FilterExpr::Token) - .prop_recursive(3, 10, 5, |inner| { - prop_oneof![ - inner.clone().prop_map(|e| FilterExpr::Not(Box::new(e))), - (inner.clone(), inner.clone()) - .prop_map(|(l, r)| FilterExpr::And(Box::new(l), Box::new(r))), - (inner.clone(), inner) - .prop_map(|(l, r)| FilterExpr::Or(Box::new(l), Box::new(r))), - ] - }) - } - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: flow-monitor-enhancement, Property 2: 过滤表达式 Round-Trip** - /// **Validates: Requirements 1.1-1.16** - /// - /// *对于任意* 有效的过滤表达式,解析为 AST 后再序列化回字符串, - /// 重新解析应该产生语义等价的 AST(对同一 Flow 产生相同的过滤结果)。 - #[test] - fn prop_filter_expr_round_trip( - expr in arb_filter_expr(), - flow in arb_llm_flow(), - ) { - // 序列化为字符串 - let expr_str = format!("{expr}"); - - // 重新解析 - let reparsed = FilterParser::parse(&expr_str); - prop_assert!( - reparsed.is_ok(), - "序列化后的表达式 '{}' 应该能够重新解析", - expr_str - ); - - let reparsed_expr = reparsed.unwrap(); - - // 编译两个表达式 - let filter1 = FilterParser::compile(&expr); - let filter2 = FilterParser::compile(&reparsed_expr); - - // 对同一 Flow 应该产生相同的结果 - let result1 = filter1(&flow); - let result2 = filter2(&flow); - - prop_assert_eq!( - result1, - result2, - "原始表达式和重新解析的表达式对同一 Flow 应该产生相同的结果\n原始: {}\n重新解析: {}", - format!("{}", expr), - format!("{}", reparsed_expr) - ); - } - - /// 测试简单表达式的 Round-Trip - #[test] - fn prop_simple_filter_round_trip( - token in arb_simple_filter_token(), - flow in arb_llm_flow(), - ) { - let expr = FilterExpr::Token(token); - let expr_str = format!("{expr}"); - - // 重新解析 - let reparsed = FilterParser::parse(&expr_str); - prop_assert!( - reparsed.is_ok(), - "简单表达式 '{}' 应该能够重新解析", - expr_str - ); - - let reparsed_expr = reparsed.unwrap(); - - // 编译并比较结果 - let filter1 = FilterParser::compile(&expr); - let filter2 = FilterParser::compile(&reparsed_expr); - - prop_assert_eq!( - filter1(&flow), - filter2(&flow), - "简单表达式 Round-Trip 应该保持语义一致" - ); - } - } - - // ======================================================================== - // Property 3: 过滤表达式错误处理 - // ======================================================================== - - /// 生成无效的过滤器名称 - fn arb_invalid_filter_name() -> impl Strategy { - prop_oneof![ - Just("unknown".to_string()), - Just("invalid".to_string()), - Just("xyz".to_string()), - Just("foo".to_string()), - Just("bar".to_string()), - "[a-z]{5,10}".prop_filter("Filter out valid names", |s| { - ![ - "m", "p", "s", "e", "t", "k", "b", "bq", "bs", "starred", "tag", "tokens", - "latency", - ] - .contains(&s.as_str()) - }), - ] - } - - /// 生成无效的状态值 - fn arb_invalid_state() -> impl Strategy { - prop_oneof![ - Just("invalid".to_string()), - Just("unknown".to_string()), - Just("running".to_string()), - Just("stopped".to_string()), - "[a-z]{5,10}".prop_filter("Filter out valid states", |s| { - !["pending", "streaming", "completed", "failed", "cancelled"] - .contains(&s.to_lowercase().as_str()) - }), - ] - } - - /// 生成无效的比较运算符 - fn arb_invalid_comparison_op() -> impl Strategy { - prop_oneof![ - Just("==".to_string()), - Just("!=".to_string()), - Just("<>".to_string()), - Just("~".to_string()), - Just("@".to_string()), - ] - } - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: flow-monitor-enhancement, Property 3: 过滤表达式错误处理** - /// **Validates: Requirements 1.17** - /// - /// *对于任意* 无效的过滤表达式,解析器应该返回错误而不是 panic, - /// 且错误信息应该包含有用的诊断信息。 - #[test] - fn prop_invalid_filter_returns_error( - filter_name in arb_invalid_filter_name(), - ) { - let expr_str = format!("~{filter_name}"); - let result = FilterParser::parse(&expr_str); - - // 应该返回错误 - prop_assert!( - result.is_err(), - "无效的过滤器 '{}' 应该返回错误", - expr_str - ); - - // 错误应该是 UnknownFilter - if let Err(e) = result { - prop_assert!( - matches!(e, FilterParseError::UnknownFilter(_)), - "错误类型应该是 UnknownFilter,实际是: {:?}", - e - ); - } - } - - #[test] - fn prop_invalid_state_returns_error( - state in arb_invalid_state(), - ) { - let expr_str = format!("~s {state}"); - let result = FilterParser::parse(&expr_str); - - // 应该返回错误 - prop_assert!( - result.is_err(), - "无效的状态 '{}' 应该返回错误", - expr_str - ); - - // 错误应该是 InvalidState - if let Err(e) = result { - prop_assert!( - matches!(e, FilterParseError::InvalidState(_)), - "错误类型应该是 InvalidState,实际是: {:?}", - e - ); - } - } - - #[test] - fn prop_invalid_comparison_op_returns_error( - op in arb_invalid_comparison_op(), - ) { - let expr_str = format!("~tokens {op}100"); - let result = FilterParser::parse(&expr_str); - - // 应该返回错误 - prop_assert!( - result.is_err(), - "无效的比较运算符 '{}' 应该返回错误", - expr_str - ); - } - - #[test] - fn prop_unmatched_paren_returns_error( - depth in 1usize..5usize, - ) { - // 生成不匹配的括号 - let open_parens: String = (0..depth).map(|_| '(').collect(); - let expr_str = format!("{open_parens}~e"); - let result = FilterParser::parse(&expr_str); - - // 应该返回错误 - prop_assert!( - result.is_err(), - "不匹配的括号 '{}' 应该返回错误", - expr_str - ); - } - - #[test] - fn prop_empty_expression_returns_error( - spaces in " {0,10}", - ) { - let result = FilterParser::parse(&spaces); - - // 应该返回错误 - prop_assert!( - result.is_err(), - "空表达式应该返回错误" - ); - - // 错误应该是 EmptyExpression - if let Err(e) = result { - prop_assert!( - matches!(e, FilterParseError::EmptyExpression), - "错误类型应该是 EmptyExpression,实际是: {:?}", - e - ); - } - } - - #[test] - fn prop_missing_argument_returns_error( - filter in prop_oneof![ - Just("m"), - Just("p"), - Just("s"), - Just("tag"), - Just("b"), - Just("bq"), - Just("bs"), - ], - ) { - // 缺少参数的过滤器 - let expr_str = format!("~{filter}"); - let result = FilterParser::parse(&expr_str); - - // 应该返回错误(缺少参数) - prop_assert!( - result.is_err(), - "缺少参数的过滤器 '{}' 应该返回错误", - expr_str - ); - } - - /// 测试解析器不会 panic - #[test] - fn prop_parser_never_panics( - input in "[ -~]{0,50}", - ) { - // 尝试解析任意输入,不应该 panic - let _ = FilterParser::parse(&input); - // 如果没有 panic,测试通过 - } - } -} diff --git a/src-tauri/src/flow_monitor/interceptor.rs b/src-tauri/src/flow_monitor/interceptor.rs deleted file mode 100644 index 27d1620ee..000000000 --- a/src-tauri/src/flow_monitor/interceptor.rs +++ /dev/null @@ -1,1497 +0,0 @@ -//! Flow 拦截器 -//! -//! 该模块实现 LLM Flow 的拦截功能,允许用户暂停、查看和修改请求/响应。 -//! -//! # 功能 -//! -//! - 根据过滤表达式拦截匹配的 Flow -//! - 支持拦截请求、响应或两者 -//! - 支持超时自动处理 -//! - 实时事件广播 - -use chrono::{DateTime, Utc}; -use serde::{Deserialize, Serialize}; -use std::collections::HashMap; -use std::sync::Arc; -use tokio::sync::{broadcast, oneshot, RwLock}; -use tokio::time::{timeout, Duration}; - -use super::filter_parser::FilterParser; -use super::models::{LLMFlow, LLMRequest, LLMResponse}; - -// ============================================================================ -// 配置结构 -// ============================================================================ - -/// 超时动作 -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -#[serde(rename_all = "lowercase")] -#[derive(Default)] -pub enum TimeoutAction { - /// 超时后继续处理 - #[default] - Continue, - /// 超时后取消请求 - Cancel, -} - -/// 拦截配置 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct InterceptConfig { - /// 是否启用拦截 - #[serde(default)] - pub enabled: bool, - /// 过滤表达式(可选,为空时拦截所有) - #[serde(skip_serializing_if = "Option::is_none")] - pub filter_expr: Option, - /// 是否拦截请求 - #[serde(default = "default_intercept_request")] - pub intercept_request: bool, - /// 是否拦截响应 - #[serde(default)] - pub intercept_response: bool, - /// 超时时间(毫秒) - #[serde(default = "default_timeout_ms")] - pub timeout_ms: u64, - /// 超时动作 - #[serde(default)] - pub timeout_action: TimeoutAction, -} - -fn default_intercept_request() -> bool { - true -} - -fn default_timeout_ms() -> u64 { - 30000 // 30 秒 -} - -impl Default for InterceptConfig { - fn default() -> Self { - Self { - enabled: false, - filter_expr: None, - intercept_request: default_intercept_request(), - intercept_response: false, - timeout_ms: default_timeout_ms(), - timeout_action: TimeoutAction::default(), - } - } -} - -// ============================================================================ -// 拦截状态和类型 -// ============================================================================ - -/// 拦截类型 -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -#[serde(rename_all = "lowercase")] -pub enum InterceptType { - /// 拦截请求 - Request, - /// 拦截响应 - Response, -} - -/// 拦截状态 -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -#[serde(rename_all = "lowercase")] -#[derive(Default)] -pub enum InterceptState { - /// 等待用户操作 - #[default] - Pending, - /// 用户正在编辑 - Editing, - /// 已继续处理 - Continued, - /// 已取消 - Cancelled, - /// 已超时 - TimedOut, -} - -// ============================================================================ -// 被拦截的 Flow -// ============================================================================ - -/// 被拦截的 Flow -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct InterceptedFlow { - /// Flow ID - pub flow_id: String, - /// 拦截状态 - pub state: InterceptState, - /// 拦截类型 - pub intercept_type: InterceptType, - /// 原始请求(如果拦截请求) - #[serde(skip_serializing_if = "Option::is_none")] - pub original_request: Option, - /// 修改后的请求 - #[serde(skip_serializing_if = "Option::is_none")] - pub modified_request: Option, - /// 原始响应(如果拦截响应) - #[serde(skip_serializing_if = "Option::is_none")] - pub original_response: Option, - /// 修改后的响应 - #[serde(skip_serializing_if = "Option::is_none")] - pub modified_response: Option, - /// 拦截时间 - pub intercepted_at: DateTime, -} - -impl InterceptedFlow { - /// 创建新的拦截请求 - pub fn new_request(flow_id: String, request: LLMRequest) -> Self { - Self { - flow_id, - state: InterceptState::Pending, - intercept_type: InterceptType::Request, - original_request: Some(request), - modified_request: None, - original_response: None, - modified_response: None, - intercepted_at: Utc::now(), - } - } - - /// 创建新的拦截响应 - pub fn new_response(flow_id: String, response: LLMResponse) -> Self { - Self { - flow_id, - state: InterceptState::Pending, - intercept_type: InterceptType::Response, - original_request: None, - modified_request: None, - original_response: Some(response), - modified_response: None, - intercepted_at: Utc::now(), - } - } -} - -// ============================================================================ -// 修改数据 -// ============================================================================ - -/// 修改后的数据 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(untagged)] -pub enum ModifiedData { - /// 修改后的请求 - Request(LLMRequest), - /// 修改后的响应 - Response(LLMResponse), -} - -// ============================================================================ -// 拦截事件 -// ============================================================================ - -/// 拦截事件 -#[derive(Debug, Clone, Serialize)] -#[serde(tag = "type")] -pub enum InterceptEvent { - /// Flow 被拦截 - FlowIntercepted { - /// 被拦截的 Flow 信息 - flow: InterceptedFlow, - }, - /// Flow 继续处理 - FlowContinued { - /// Flow ID - flow_id: String, - /// 是否有修改 - modified: bool, - }, - /// Flow 被取消 - FlowCancelled { - /// Flow ID - flow_id: String, - }, - /// Flow 超时 - FlowTimedOut { - /// Flow ID - flow_id: String, - /// 超时动作 - action: TimeoutAction, - }, - /// 配置已更新 - ConfigUpdated { - /// 新配置 - config: InterceptConfig, - }, -} - -// ============================================================================ -// 拦截动作 -// ============================================================================ - -/// 用户拦截动作 -#[derive(Debug, Clone)] -pub enum InterceptAction { - /// 继续处理(可能带有修改) - Continue(Option), - /// 取消请求 - Cancel, - /// 超时 - Timeout(TimeoutAction), -} - -// ============================================================================ -// 等待中的拦截 -// ============================================================================ - -/// 等待中的拦截 -struct PendingIntercept { - /// 被拦截的 Flow 信息 - flow: InterceptedFlow, - /// 动作发送器 - action_sender: Option>, -} - -// ============================================================================ -// 拦截器错误 -// ============================================================================ - -/// 拦截器错误 -#[derive(Debug, Clone, thiserror::Error, Serialize, Deserialize)] -pub enum InterceptorError { - /// Flow 不存在 - #[error("Flow '{0}' 不存在或未被拦截")] - FlowNotFound(String), - /// 无效的过滤表达式 - #[error("无效的过滤表达式: {0}")] - InvalidFilterExpr(String), - /// 操作已完成 - #[error("Flow '{0}' 的拦截操作已完成")] - AlreadyCompleted(String), - /// 内部错误 - #[error("内部错误: {0}")] - Internal(String), -} - -// ============================================================================ -// Flow 拦截器 -// ============================================================================ - -/// Flow 拦截器 -/// -/// 负责拦截和管理 LLM Flow 的核心服务。 -pub struct FlowInterceptor { - /// 拦截配置 - config: RwLock, - /// 编译后的过滤器 - filter: RwLock bool + Send + Sync>>>, - /// 等待中的拦截 - pending_intercepts: RwLock>, - /// 事件发送器 - event_sender: broadcast::Sender, -} - -impl FlowInterceptor { - /// 创建新的拦截器 - pub fn new(config: InterceptConfig) -> Self { - let (event_sender, _) = broadcast::channel(100); - let filter = Self::compile_filter(&config.filter_expr); - - Self { - config: RwLock::new(config), - filter: RwLock::new(filter), - pending_intercepts: RwLock::new(HashMap::new()), - event_sender, - } - } - - /// 编译过滤表达式 - fn compile_filter( - filter_expr: &Option, - ) -> Option bool + Send + Sync>> { - filter_expr.as_ref().and_then(|expr| { - FilterParser::parse(expr).ok().map(|parsed| { - let filter = FilterParser::compile(&parsed); - Arc::new(move |flow: &LLMFlow| filter(flow)) - as Arc bool + Send + Sync> - }) - }) - } - - /// 获取当前配置 - pub async fn config(&self) -> InterceptConfig { - self.config.read().await.clone() - } - - /// 更新配置 - pub async fn update_config(&self, config: InterceptConfig) -> Result<(), InterceptorError> { - // 验证过滤表达式 - if let Some(ref expr) = config.filter_expr { - FilterParser::parse(expr) - .map_err(|e| InterceptorError::InvalidFilterExpr(e.to_string()))?; - } - - // 编译新的过滤器 - let new_filter = Self::compile_filter(&config.filter_expr); - - // 更新配置和过滤器 - { - let mut current_config = self.config.write().await; - *current_config = config.clone(); - } - { - let mut current_filter = self.filter.write().await; - *current_filter = new_filter; - } - - // 发送配置更新事件 - let _ = self - .event_sender - .send(InterceptEvent::ConfigUpdated { config }); - - Ok(()) - } - - /// 订阅拦截事件 - pub fn subscribe(&self) -> broadcast::Receiver { - self.event_sender.subscribe() - } - - /// 检查是否应该拦截 - pub async fn should_intercept(&self, flow: &LLMFlow, intercept_type: &InterceptType) -> bool { - let config = self.config.read().await; - - // 检查是否启用 - if !config.enabled { - return false; - } - - // 检查拦截类型 - match intercept_type { - InterceptType::Request => { - if !config.intercept_request { - return false; - } - } - InterceptType::Response => { - if !config.intercept_response { - return false; - } - } - } - - // 检查过滤器 - let filter = self.filter.read().await; - if let Some(ref f) = *filter { - f(flow) - } else { - // 没有过滤器时,拦截所有 - true - } - } - - /// 拦截请求 - pub async fn intercept_request(&self, flow_id: &str, request: LLMRequest) -> InterceptedFlow { - let intercepted = InterceptedFlow::new_request(flow_id.to_string(), request); - self.add_pending_intercept(intercepted.clone()).await; - - // 发送拦截事件 - let _ = self.event_sender.send(InterceptEvent::FlowIntercepted { - flow: intercepted.clone(), - }); - - intercepted - } - - /// 拦截响应 - pub async fn intercept_response( - &self, - flow_id: &str, - response: LLMResponse, - ) -> InterceptedFlow { - let intercepted = InterceptedFlow::new_response(flow_id.to_string(), response); - self.add_pending_intercept(intercepted.clone()).await; - - // 发送拦截事件 - let _ = self.event_sender.send(InterceptEvent::FlowIntercepted { - flow: intercepted.clone(), - }); - - intercepted - } - - /// 添加等待中的拦截 - async fn add_pending_intercept(&self, flow: InterceptedFlow) { - let mut pending = self.pending_intercepts.write().await; - pending.insert( - flow.flow_id.clone(), - PendingIntercept { - flow, - action_sender: None, - }, - ); - } - - /// 继续处理 Flow - pub async fn continue_flow( - &self, - flow_id: &str, - modified: Option, - ) -> Result<(), InterceptorError> { - let mut pending = self.pending_intercepts.write().await; - - if let Some(mut intercept) = pending.remove(flow_id) { - // 更新状态 - intercept.flow.state = InterceptState::Continued; - - // 更新修改后的数据 - if let Some(ref data) = modified { - match data { - ModifiedData::Request(req) => { - intercept.flow.modified_request = Some(req.clone()); - } - ModifiedData::Response(resp) => { - intercept.flow.modified_response = Some(resp.clone()); - } - } - } - - // 发送动作 - if let Some(sender) = intercept.action_sender { - let _ = sender.send(InterceptAction::Continue(modified.clone())); - } - - // 发送事件 - let _ = self.event_sender.send(InterceptEvent::FlowContinued { - flow_id: flow_id.to_string(), - modified: modified.is_some(), - }); - - Ok(()) - } else { - Err(InterceptorError::FlowNotFound(flow_id.to_string())) - } - } - - /// 取消 Flow - pub async fn cancel_flow(&self, flow_id: &str) -> Result<(), InterceptorError> { - let mut pending = self.pending_intercepts.write().await; - - if let Some(mut intercept) = pending.remove(flow_id) { - // 更新状态 - intercept.flow.state = InterceptState::Cancelled; - - // 发送动作 - if let Some(sender) = intercept.action_sender { - let _ = sender.send(InterceptAction::Cancel); - } - - // 发送事件 - let _ = self.event_sender.send(InterceptEvent::FlowCancelled { - flow_id: flow_id.to_string(), - }); - - Ok(()) - } else { - Err(InterceptorError::FlowNotFound(flow_id.to_string())) - } - } - - /// 等待用户操作 - /// - /// 此方法会阻塞直到用户执行操作或超时。 - pub async fn wait_for_action(&self, flow_id: &str) -> InterceptAction { - let config = self.config.read().await.clone(); - let timeout_ms = config.timeout_ms; - let timeout_action = config.timeout_action.clone(); - drop(config); - - // 创建 oneshot channel - let (tx, rx) = oneshot::channel(); - - // 设置 action_sender - { - let mut pending = self.pending_intercepts.write().await; - if let Some(intercept) = pending.get_mut(flow_id) { - intercept.action_sender = Some(tx); - } else { - // Flow 不存在,返回取消 - return InterceptAction::Cancel; - } - } - - // 等待动作或超时 - let result = timeout(Duration::from_millis(timeout_ms), rx).await; - - match result { - Ok(Ok(action)) => action, - Ok(Err(_)) => { - // Channel 被关闭,视为取消 - InterceptAction::Cancel - } - Err(_) => { - // 超时 - self.handle_timeout(flow_id, &timeout_action).await; - InterceptAction::Timeout(timeout_action) - } - } - } - - /// 处理超时 - async fn handle_timeout(&self, flow_id: &str, timeout_action: &TimeoutAction) { - let mut pending = self.pending_intercepts.write().await; - - if let Some(mut intercept) = pending.remove(flow_id) { - intercept.flow.state = InterceptState::TimedOut; - - // 发送超时事件 - let _ = self.event_sender.send(InterceptEvent::FlowTimedOut { - flow_id: flow_id.to_string(), - action: timeout_action.clone(), - }); - } - } - - /// 获取被拦截的 Flow - pub async fn get_intercepted_flow(&self, flow_id: &str) -> Option { - let pending = self.pending_intercepts.read().await; - pending.get(flow_id).map(|p| p.flow.clone()) - } - - /// 获取所有被拦截的 Flow - pub async fn list_intercepted_flows(&self) -> Vec { - let pending = self.pending_intercepts.read().await; - pending.values().map(|p| p.flow.clone()).collect() - } - - /// 获取被拦截的 Flow 数量 - pub async fn intercepted_count(&self) -> usize { - self.pending_intercepts.read().await.len() - } - - /// 检查拦截是否启用 - pub async fn is_enabled(&self) -> bool { - self.config.read().await.enabled - } - - /// 启用拦截 - pub async fn enable(&self) { - let mut config = self.config.write().await; - config.enabled = true; - } - - /// 禁用拦截 - pub async fn disable(&self) { - let mut config = self.config.write().await; - config.enabled = false; - } - - /// 设置编辑状态 - pub async fn set_editing(&self, flow_id: &str) -> Result<(), InterceptorError> { - let mut pending = self.pending_intercepts.write().await; - - if let Some(intercept) = pending.get_mut(flow_id) { - intercept.flow.state = InterceptState::Editing; - Ok(()) - } else { - Err(InterceptorError::FlowNotFound(flow_id.to_string())) - } - } -} - -impl Default for FlowInterceptor { - fn default() -> Self { - Self::new(InterceptConfig::default()) - } -} - -// ============================================================================ -// 单元测试 -// ============================================================================ - -#[cfg(test)] -mod tests { - use super::*; - use crate::flow_monitor::models::{ - FlowMetadata, FlowType, LLMRequest, Message, MessageContent, MessageRole, - RequestParameters, TokenUsage, - }; - use crate::ProviderType; - use std::collections::HashMap; - - /// 创建测试用的 LLMRequest - fn create_test_request(model: &str) -> LLMRequest { - LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - headers: HashMap::new(), - body: serde_json::Value::Null, - messages: vec![Message { - role: MessageRole::User, - content: MessageContent::Text("Hello".to_string()), - tool_calls: None, - tool_result: None, - name: None, - }], - system_prompt: None, - tools: None, - model: model.to_string(), - original_model: None, - parameters: RequestParameters::default(), - size_bytes: 0, - timestamp: Utc::now(), - } - } - - /// 创建测试用的 LLMResponse - fn create_test_response() -> LLMResponse { - LLMResponse { - status_code: 200, - status_text: "OK".to_string(), - headers: HashMap::new(), - body: serde_json::Value::Null, - content: "Hello, world!".to_string(), - thinking: None, - tool_calls: Vec::new(), - usage: TokenUsage::default(), - stop_reason: None, - size_bytes: 0, - timestamp_start: Utc::now(), - timestamp_end: Utc::now(), - stream_info: None, - } - } - - /// 创建测试用的 LLMFlow - fn create_test_flow(model: &str, provider: ProviderType) -> LLMFlow { - let request = create_test_request(model); - let metadata = FlowMetadata { - provider, - ..Default::default() - }; - LLMFlow::new( - "test-flow-id".to_string(), - FlowType::ChatCompletions, - request, - metadata, - ) - } - - #[tokio::test] - async fn test_interceptor_creation() { - let config = InterceptConfig::default(); - let interceptor = FlowInterceptor::new(config); - - assert!(!interceptor.is_enabled().await); - assert_eq!(interceptor.intercepted_count().await, 0); - } - - #[tokio::test] - async fn test_interceptor_enable_disable() { - let interceptor = FlowInterceptor::default(); - - assert!(!interceptor.is_enabled().await); - - interceptor.enable().await; - assert!(interceptor.is_enabled().await); - - interceptor.disable().await; - assert!(!interceptor.is_enabled().await); - } - - #[tokio::test] - async fn test_should_intercept_disabled() { - let interceptor = FlowInterceptor::default(); - let flow = create_test_flow("gpt-4", ProviderType::OpenAI); - - // 禁用时不应该拦截 - assert!( - !interceptor - .should_intercept(&flow, &InterceptType::Request) - .await - ); - } - - #[tokio::test] - async fn test_should_intercept_enabled_no_filter() { - let config = InterceptConfig { - enabled: true, - intercept_request: true, - ..Default::default() - }; - let interceptor = FlowInterceptor::new(config); - let flow = create_test_flow("gpt-4", ProviderType::OpenAI); - - // 启用且无过滤器时应该拦截所有 - assert!( - interceptor - .should_intercept(&flow, &InterceptType::Request) - .await - ); - } - - #[tokio::test] - async fn test_should_intercept_with_filter() { - let config = InterceptConfig { - enabled: true, - filter_expr: Some("~m claude".to_string()), - intercept_request: true, - ..Default::default() - }; - let interceptor = FlowInterceptor::new(config); - - let flow_claude = create_test_flow("claude-3-opus", ProviderType::Claude); - let flow_gpt = create_test_flow("gpt-4", ProviderType::OpenAI); - - // 应该拦截 claude 模型 - assert!( - interceptor - .should_intercept(&flow_claude, &InterceptType::Request) - .await - ); - // 不应该拦截 gpt 模型 - assert!( - !interceptor - .should_intercept(&flow_gpt, &InterceptType::Request) - .await - ); - } - - #[tokio::test] - async fn test_should_intercept_request_only() { - let config = InterceptConfig { - enabled: true, - intercept_request: true, - intercept_response: false, - ..Default::default() - }; - let interceptor = FlowInterceptor::new(config); - let flow = create_test_flow("gpt-4", ProviderType::OpenAI); - - assert!( - interceptor - .should_intercept(&flow, &InterceptType::Request) - .await - ); - assert!( - !interceptor - .should_intercept(&flow, &InterceptType::Response) - .await - ); - } - - #[tokio::test] - async fn test_should_intercept_response_only() { - let config = InterceptConfig { - enabled: true, - intercept_request: false, - intercept_response: true, - ..Default::default() - }; - let interceptor = FlowInterceptor::new(config); - let flow = create_test_flow("gpt-4", ProviderType::OpenAI); - - assert!( - !interceptor - .should_intercept(&flow, &InterceptType::Request) - .await - ); - assert!( - interceptor - .should_intercept(&flow, &InterceptType::Response) - .await - ); - } - - #[tokio::test] - async fn test_intercept_request() { - let interceptor = FlowInterceptor::default(); - let request = create_test_request("gpt-4"); - - let intercepted = interceptor - .intercept_request("flow-1", request.clone()) - .await; - - assert_eq!(intercepted.flow_id, "flow-1"); - assert_eq!(intercepted.state, InterceptState::Pending); - assert_eq!(intercepted.intercept_type, InterceptType::Request); - assert!(intercepted.original_request.is_some()); - assert!(intercepted.modified_request.is_none()); - assert_eq!(interceptor.intercepted_count().await, 1); - } - - #[tokio::test] - async fn test_intercept_response() { - let interceptor = FlowInterceptor::default(); - let response = create_test_response(); - - let intercepted = interceptor - .intercept_response("flow-1", response.clone()) - .await; - - assert_eq!(intercepted.flow_id, "flow-1"); - assert_eq!(intercepted.state, InterceptState::Pending); - assert_eq!(intercepted.intercept_type, InterceptType::Response); - assert!(intercepted.original_response.is_some()); - assert!(intercepted.modified_response.is_none()); - assert_eq!(interceptor.intercepted_count().await, 1); - } - - #[tokio::test] - async fn test_continue_flow() { - let interceptor = FlowInterceptor::default(); - let request = create_test_request("gpt-4"); - - interceptor.intercept_request("flow-1", request).await; - - // 继续处理 - let result = interceptor.continue_flow("flow-1", None).await; - assert!(result.is_ok()); - assert_eq!(interceptor.intercepted_count().await, 0); - } - - #[tokio::test] - async fn test_continue_flow_with_modification() { - let interceptor = FlowInterceptor::default(); - let request = create_test_request("gpt-4"); - - interceptor.intercept_request("flow-1", request).await; - - // 修改请求并继续 - let modified_request = create_test_request("gpt-4-turbo"); - let result = interceptor - .continue_flow("flow-1", Some(ModifiedData::Request(modified_request))) - .await; - assert!(result.is_ok()); - } - - #[tokio::test] - async fn test_cancel_flow() { - let interceptor = FlowInterceptor::default(); - let request = create_test_request("gpt-4"); - - interceptor.intercept_request("flow-1", request).await; - - // 取消 - let result = interceptor.cancel_flow("flow-1").await; - assert!(result.is_ok()); - assert_eq!(interceptor.intercepted_count().await, 0); - } - - #[tokio::test] - async fn test_continue_nonexistent_flow() { - let interceptor = FlowInterceptor::default(); - - let result = interceptor.continue_flow("nonexistent", None).await; - assert!(matches!(result, Err(InterceptorError::FlowNotFound(_)))); - } - - #[tokio::test] - async fn test_cancel_nonexistent_flow() { - let interceptor = FlowInterceptor::default(); - - let result = interceptor.cancel_flow("nonexistent").await; - assert!(matches!(result, Err(InterceptorError::FlowNotFound(_)))); - } - - #[tokio::test] - async fn test_update_config() { - let interceptor = FlowInterceptor::default(); - - let new_config = InterceptConfig { - enabled: true, - filter_expr: Some("~m claude".to_string()), - intercept_request: true, - intercept_response: true, - timeout_ms: 60000, - timeout_action: TimeoutAction::Cancel, - }; - - let result = interceptor.update_config(new_config.clone()).await; - assert!(result.is_ok()); - - let config = interceptor.config().await; - assert!(config.enabled); - assert_eq!(config.filter_expr, Some("~m claude".to_string())); - assert_eq!(config.timeout_ms, 60000); - assert_eq!(config.timeout_action, TimeoutAction::Cancel); - } - - #[tokio::test] - async fn test_update_config_invalid_filter() { - let interceptor = FlowInterceptor::default(); - - let new_config = InterceptConfig { - enabled: true, - filter_expr: Some("~invalid".to_string()), - ..Default::default() - }; - - let result = interceptor.update_config(new_config).await; - assert!(matches!( - result, - Err(InterceptorError::InvalidFilterExpr(_)) - )); - } - - #[tokio::test] - async fn test_get_intercepted_flow() { - let interceptor = FlowInterceptor::default(); - let request = create_test_request("gpt-4"); - - interceptor.intercept_request("flow-1", request).await; - - let flow = interceptor.get_intercepted_flow("flow-1").await; - assert!(flow.is_some()); - assert_eq!(flow.unwrap().flow_id, "flow-1"); - - let nonexistent = interceptor.get_intercepted_flow("nonexistent").await; - assert!(nonexistent.is_none()); - } - - #[tokio::test] - async fn test_list_intercepted_flows() { - let interceptor = FlowInterceptor::default(); - - interceptor - .intercept_request("flow-1", create_test_request("gpt-4")) - .await; - interceptor - .intercept_request("flow-2", create_test_request("claude-3")) - .await; - - let flows = interceptor.list_intercepted_flows().await; - assert_eq!(flows.len(), 2); - } - - #[tokio::test] - async fn test_set_editing() { - let interceptor = FlowInterceptor::default(); - let request = create_test_request("gpt-4"); - - interceptor.intercept_request("flow-1", request).await; - - let result = interceptor.set_editing("flow-1").await; - assert!(result.is_ok()); - - let flow = interceptor.get_intercepted_flow("flow-1").await.unwrap(); - assert_eq!(flow.state, InterceptState::Editing); - } - - #[tokio::test] - async fn test_event_subscription() { - let interceptor = FlowInterceptor::default(); - let mut receiver = interceptor.subscribe(); - - let request = create_test_request("gpt-4"); - interceptor.intercept_request("flow-1", request).await; - - // 应该收到 FlowIntercepted 事件 - let event = receiver.try_recv(); - assert!(event.is_ok()); - if let InterceptEvent::FlowIntercepted { flow } = event.unwrap() { - assert_eq!(flow.flow_id, "flow-1"); - } else { - panic!("Expected FlowIntercepted event"); - } - } -} - -// ============================================================================ -// 属性测试 -// ============================================================================ - -#[cfg(test)] -mod property_tests { - #![allow(dead_code)] - use super::*; - use crate::flow_monitor::models::{ - FlowError, FlowErrorType, FlowMetadata, FlowType, FunctionCall, LLMRequest, LLMResponse, - RequestParameters, ThinkingContent, TokenUsage, ToolCall, - }; - use crate::ProviderType; - use proptest::prelude::*; - use tokio::runtime::Runtime; - - // ======================================================================== - // 生成器 - // ======================================================================== - - /// 生成随机的 ProviderType - fn arb_provider_type() -> impl Strategy { - prop_oneof![ - Just(ProviderType::Kiro), - Just(ProviderType::Gemini), - Just(ProviderType::OpenAI), - Just(ProviderType::Claude), - Just(ProviderType::Antigravity), - ] - } - - /// 生成随机的 FlowState - fn arb_flow_state() -> impl Strategy { - prop_oneof![ - Just(crate::flow_monitor::models::FlowState::Pending), - Just(crate::flow_monitor::models::FlowState::Streaming), - Just(crate::flow_monitor::models::FlowState::Completed), - Just(crate::flow_monitor::models::FlowState::Failed), - Just(crate::flow_monitor::models::FlowState::Cancelled), - ] - } - - /// 生成随机的模型名称 - fn arb_model_name() -> impl Strategy { - prop_oneof![ - Just("gpt-4".to_string()), - Just("gpt-4-turbo".to_string()), - Just("gpt-3.5-turbo".to_string()), - Just("claude-3-opus".to_string()), - Just("claude-3-sonnet".to_string()), - Just("gemini-pro".to_string()), - Just("qwen-max".to_string()), - ] - } - - /// 生成随机的标签 - fn arb_tags() -> impl Strategy> { - prop::collection::vec("[a-z]{3,10}", 0..5) - } - - /// 生成随机的 LLMFlow - fn arb_llm_flow() -> impl Strategy { - ( - "[a-f0-9]{8}", - arb_model_name(), - arb_provider_type(), - arb_flow_state(), - any::(), // starred - arb_tags(), // tags - any::(), // has_error - any::(), // has_tool_calls - any::(), // has_thinking - 0u32..50000u32, // total_tokens - 0u64..30000u64, // duration_ms - ) - .prop_map( - |( - id, - model, - provider, - state, - starred, - tags, - has_error, - has_tool_calls, - has_thinking, - total_tokens, - duration_ms, - )| { - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model, - parameters: RequestParameters::default(), - ..Default::default() - }; - - let metadata = FlowMetadata { - provider, - ..Default::default() - }; - - let mut flow = LLMFlow::new(id, FlowType::ChatCompletions, request, metadata); - flow.state = state; - flow.annotations.starred = starred; - flow.annotations.tags = tags; - flow.timestamps.duration_ms = duration_ms; - - if has_error { - flow.error = Some(FlowError::new(FlowErrorType::ServerError, "Test error")); - } - - let mut response = LLMResponse { - usage: TokenUsage { - input_tokens: total_tokens / 2, - output_tokens: total_tokens / 2, - total_tokens, - ..Default::default() - }, - ..Default::default() - }; - - if has_tool_calls { - response.tool_calls = vec![ToolCall { - id: "call_1".to_string(), - tool_type: "function".to_string(), - function: FunctionCall { - name: "test_function".to_string(), - arguments: "{}".to_string(), - }, - }]; - } - - if has_thinking { - response.thinking = Some(ThinkingContent { - text: "Thinking...".to_string(), - tokens: Some(100), - signature: None, - }); - } - - flow.response = Some(response); - flow - }, - ) - } - - /// 生成随机的 InterceptType - fn arb_intercept_type() -> impl Strategy { - prop_oneof![Just(InterceptType::Request), Just(InterceptType::Response),] - } - - /// 生成随机的过滤表达式 - fn arb_filter_expr() -> impl Strategy { - prop_oneof![ - arb_model_name().prop_map(|m| format!("~m {m}")), - prop_oneof![ - Just("kiro".to_string()), - Just("openai".to_string()), - Just("claude".to_string()), - Just("gemini".to_string()), - ] - .prop_map(|p| format!("~p {p}")), - Just("~e".to_string()), - Just("~t".to_string()), - Just("~k".to_string()), - Just("~starred".to_string()), - (0i64..50000i64).prop_map(|n| format!("~tokens >{n}")), - (0i64..30000i64).prop_map(|n| format!("~latency >{n}ms")), - ] - } - - /// 生成随机的 InterceptConfig - fn arb_intercept_config() -> impl Strategy { - ( - any::(), // enabled - prop::option::of(arb_filter_expr()), // filter_expr - any::(), // intercept_request - any::(), // intercept_response - 1000u64..60000u64, // timeout_ms - prop_oneof![Just(TimeoutAction::Continue), Just(TimeoutAction::Cancel),], - ) - .prop_map( - |( - enabled, - filter_expr, - intercept_request, - intercept_response, - timeout_ms, - timeout_action, - )| { - InterceptConfig { - enabled, - filter_expr, - intercept_request, - intercept_response, - timeout_ms, - timeout_action, - } - }, - ) - } - - // ======================================================================== - // Property 4: 拦截规则匹配正确性 - // ======================================================================== - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: flow-monitor-enhancement, Property 4: 拦截规则匹配正确性** - /// **Validates: Requirements 2.1, 2.7** - /// - /// *对于任意* 拦截配置和 Flow,拦截器的 should_intercept 方法应该正确判断是否需要拦截。 - #[test] - fn prop_intercept_disabled_never_intercepts( - flow in arb_llm_flow(), - intercept_type in arb_intercept_type(), - ) { - let rt = Runtime::new().unwrap(); - rt.block_on(async { - // 禁用拦截时,永远不应该拦截 - let config = InterceptConfig { - enabled: false, - ..Default::default() - }; - let interceptor = FlowInterceptor::new(config); - - let should_intercept = interceptor.should_intercept(&flow, &intercept_type).await; - prop_assert!( - !should_intercept, - "禁用拦截时不应该拦截任何 Flow" - ); - Ok(()) - })?; - } - - #[test] - fn prop_intercept_type_respected( - flow in arb_llm_flow(), - ) { - let rt = Runtime::new().unwrap(); - rt.block_on(async { - // 只拦截请求 - let config_request_only = InterceptConfig { - enabled: true, - intercept_request: true, - intercept_response: false, - ..Default::default() - }; - let interceptor = FlowInterceptor::new(config_request_only); - - let should_intercept_request = interceptor.should_intercept(&flow, &InterceptType::Request).await; - let should_intercept_response = interceptor.should_intercept(&flow, &InterceptType::Response).await; - - prop_assert!( - should_intercept_request, - "配置为拦截请求时应该拦截请求" - ); - prop_assert!( - !should_intercept_response, - "配置为不拦截响应时不应该拦截响应" - ); - - // 只拦截响应 - let config_response_only = InterceptConfig { - enabled: true, - intercept_request: false, - intercept_response: true, - ..Default::default() - }; - let interceptor2 = FlowInterceptor::new(config_response_only); - - let should_intercept_request2 = interceptor2.should_intercept(&flow, &InterceptType::Request).await; - let should_intercept_response2 = interceptor2.should_intercept(&flow, &InterceptType::Response).await; - - prop_assert!( - !should_intercept_request2, - "配置为不拦截请求时不应该拦截请求" - ); - prop_assert!( - should_intercept_response2, - "配置为拦截响应时应该拦截响应" - ); - - Ok(()) - })?; - } - - #[test] - fn prop_filter_model_intercept_correctness( - flow in arb_llm_flow(), - ) { - let rt = Runtime::new().unwrap(); - rt.block_on(async { - let model = flow.request.model.clone(); - - // 使用模型过滤器 - let config = InterceptConfig { - enabled: true, - filter_expr: Some(format!("~m {model}")), - intercept_request: true, - ..Default::default() - }; - let interceptor = FlowInterceptor::new(config); - - let should_intercept = interceptor.should_intercept(&flow, &InterceptType::Request).await; - - prop_assert!( - should_intercept, - "使用模型 '{}' 的过滤器应该拦截模型为 '{}' 的 Flow", - model, - flow.request.model - ); - - Ok(()) - })?; - } - - #[test] - fn prop_filter_provider_intercept_correctness( - flow in arb_llm_flow(), - ) { - let rt = Runtime::new().unwrap(); - rt.block_on(async { - let provider_str = format!("{:?}", flow.metadata.provider).to_lowercase(); - - // 使用提供商过滤器 - let config = InterceptConfig { - enabled: true, - filter_expr: Some(format!("~p {provider_str}")), - intercept_request: true, - ..Default::default() - }; - let interceptor = FlowInterceptor::new(config); - - let should_intercept = interceptor.should_intercept(&flow, &InterceptType::Request).await; - - prop_assert!( - should_intercept, - "使用提供商 '{}' 的过滤器应该拦截提供商为 '{:?}' 的 Flow", - provider_str, - flow.metadata.provider - ); - - Ok(()) - })?; - } - - #[test] - fn prop_filter_error_intercept_correctness( - flow in arb_llm_flow(), - ) { - let rt = Runtime::new().unwrap(); - rt.block_on(async { - // 使用错误过滤器 - let config = InterceptConfig { - enabled: true, - filter_expr: Some("~e".to_string()), - intercept_request: true, - ..Default::default() - }; - let interceptor = FlowInterceptor::new(config); - - let should_intercept = interceptor.should_intercept(&flow, &InterceptType::Request).await; - let has_error = flow.error.is_some(); - - prop_assert_eq!( - should_intercept, - has_error, - "错误过滤器的拦截结果应该与 flow.error.is_some() 一致" - ); - - Ok(()) - })?; - } - - #[test] - fn prop_filter_starred_intercept_correctness( - flow in arb_llm_flow(), - ) { - let rt = Runtime::new().unwrap(); - rt.block_on(async { - // 使用收藏过滤器 - let config = InterceptConfig { - enabled: true, - filter_expr: Some("~starred".to_string()), - intercept_request: true, - ..Default::default() - }; - let interceptor = FlowInterceptor::new(config); - - let should_intercept = interceptor.should_intercept(&flow, &InterceptType::Request).await; - - prop_assert_eq!( - should_intercept, - flow.annotations.starred, - "收藏过滤器的拦截结果应该与 flow.annotations.starred 一致" - ); - - Ok(()) - })?; - } - - #[test] - fn prop_filter_tokens_intercept_correctness( - flow in arb_llm_flow(), - threshold in 0i64..50000i64, - ) { - let rt = Runtime::new().unwrap(); - rt.block_on(async { - let total_tokens = flow - .response - .as_ref() - .map_or(0, |r| r.usage.total_tokens as i64); - - // 使用 Token 过滤器 - let config = InterceptConfig { - enabled: true, - filter_expr: Some(format!("~tokens >{threshold}")), - intercept_request: true, - ..Default::default() - }; - let interceptor = FlowInterceptor::new(config); - - let should_intercept = interceptor.should_intercept(&flow, &InterceptType::Request).await; - - prop_assert_eq!( - should_intercept, - total_tokens > threshold, - "Token 过滤器的拦截结果应该正确 (actual: {}, threshold: {})", - total_tokens, - threshold - ); - - Ok(()) - })?; - } - - #[test] - fn prop_filter_latency_intercept_correctness( - flow in arb_llm_flow(), - threshold in 0i64..30000i64, - ) { - let rt = Runtime::new().unwrap(); - rt.block_on(async { - let duration_ms = flow.timestamps.duration_ms as i64; - - // 使用延迟过滤器 - let config = InterceptConfig { - enabled: true, - filter_expr: Some(format!("~latency >{threshold}ms")), - intercept_request: true, - ..Default::default() - }; - let interceptor = FlowInterceptor::new(config); - - let should_intercept = interceptor.should_intercept(&flow, &InterceptType::Request).await; - - prop_assert_eq!( - should_intercept, - duration_ms > threshold, - "延迟过滤器的拦截结果应该正确 (actual: {}, threshold: {})", - duration_ms, - threshold - ); - - Ok(()) - })?; - } - - #[test] - fn prop_no_filter_intercepts_all( - flow in arb_llm_flow(), - intercept_type in arb_intercept_type(), - ) { - let rt = Runtime::new().unwrap(); - rt.block_on(async { - // 启用但无过滤器时应该拦截所有 - let config = InterceptConfig { - enabled: true, - filter_expr: None, - intercept_request: true, - intercept_response: true, - ..Default::default() - }; - let interceptor = FlowInterceptor::new(config); - - let should_intercept = interceptor.should_intercept(&flow, &intercept_type).await; - - prop_assert!( - should_intercept, - "无过滤器时应该拦截所有 Flow" - ); - - Ok(()) - })?; - } - } -} diff --git a/src-tauri/src/flow_monitor/memory_store.rs b/src-tauri/src/flow_monitor/memory_store.rs deleted file mode 100644 index 7b3438553..000000000 --- a/src-tauri/src/flow_monitor/memory_store.rs +++ /dev/null @@ -1,1252 +0,0 @@ -//! Flow 内存存储 -//! -//! 该模块实现 LLM Flow 的内存缓存存储,支持 LRU 驱逐策略。 -//! 提供快速的 Flow 访问和查询功能。 - -use chrono::{DateTime, Utc}; -use serde::{Deserialize, Serialize}; -use std::collections::{HashMap, VecDeque}; -use std::sync::{Arc, RwLock}; - -use super::models::{FlowState, FlowType, LLMFlow}; -use crate::ProviderType; - -// ============================================================================ -// 过滤器结构 -// ============================================================================ - -/// 时间范围 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct TimeRange { - /// 开始时间 - pub start: Option>, - /// 结束时间 - pub end: Option>, -} - -impl TimeRange { - /// 创建新的时间范围 - pub fn new(start: Option>, end: Option>) -> Self { - Self { start, end } - } - - /// 检查时间是否在范围内 - pub fn contains(&self, time: &DateTime) -> bool { - let after_start = self.start.is_none_or(|s| time >= &s); - let before_end = self.end.is_none_or(|e| time <= &e); - after_start && before_end - } -} - -/// Token 范围 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct TokenRange { - /// 最小 Token 数 - pub min: Option, - /// 最大 Token 数 - pub max: Option, -} - -impl TokenRange { - /// 检查 Token 数是否在范围内 - pub fn contains(&self, tokens: u32) -> bool { - let above_min = self.min.is_none_or(|m| tokens >= m); - let below_max = self.max.is_none_or(|m| tokens <= m); - above_min && below_max - } -} - -/// 延迟范围 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct LatencyRange { - /// 最小延迟(毫秒) - pub min_ms: Option, - /// 最大延迟(毫秒) - pub max_ms: Option, -} - -impl LatencyRange { - /// 检查延迟是否在范围内 - pub fn contains(&self, latency_ms: u64) -> bool { - let above_min = self.min_ms.is_none_or(|m| latency_ms >= m); - let below_max = self.max_ms.is_none_or(|m| latency_ms <= m); - above_min && below_max - } -} - -/// Flow 过滤器 -/// -/// 支持多维度过滤条件,用于查询 Flow。 -#[derive(Debug, Clone, Default, Serialize, Deserialize)] -pub struct FlowFilter { - /// 时间范围 - #[serde(skip_serializing_if = "Option::is_none")] - pub time_range: Option, - /// 提供商类型列表 - #[serde(skip_serializing_if = "Option::is_none")] - pub providers: Option>, - /// 模型名称列表(支持通配符 *) - #[serde(skip_serializing_if = "Option::is_none")] - pub models: Option>, - /// 状态列表 - #[serde(skip_serializing_if = "Option::is_none")] - pub states: Option>, - /// 是否有错误 - #[serde(skip_serializing_if = "Option::is_none")] - pub has_error: Option, - /// 是否有工具调用 - #[serde(skip_serializing_if = "Option::is_none")] - pub has_tool_calls: Option, - /// 是否有思维链 - #[serde(skip_serializing_if = "Option::is_none")] - pub has_thinking: Option, - /// 是否是流式响应 - #[serde(skip_serializing_if = "Option::is_none")] - pub is_streaming: Option, - /// 内容搜索(响应内容) - #[serde(skip_serializing_if = "Option::is_none")] - pub content_search: Option, - /// 请求搜索(请求内容) - #[serde(skip_serializing_if = "Option::is_none")] - pub request_search: Option, - /// Token 范围 - #[serde(skip_serializing_if = "Option::is_none")] - pub token_range: Option, - /// 延迟范围 - #[serde(skip_serializing_if = "Option::is_none")] - pub latency_range: Option, - /// 标签列表 - #[serde(skip_serializing_if = "Option::is_none")] - pub tags: Option>, - /// 仅收藏 - #[serde(default)] - pub starred_only: bool, - /// 凭证 ID - #[serde(skip_serializing_if = "Option::is_none")] - pub credential_id: Option, - /// Flow 类型 - #[serde(skip_serializing_if = "Option::is_none")] - pub flow_types: Option>, -} - -impl FlowFilter { - /// 创建空过滤器(匹配所有) - pub fn new() -> Self { - Self::default() - } - - /// 检查 Flow 是否匹配过滤条件 - pub fn matches(&self, flow: &LLMFlow) -> bool { - // 时间范围过滤 - if let Some(ref time_range) = self.time_range { - if !time_range.contains(&flow.timestamps.created) { - return false; - } - } - - // 提供商过滤 - if let Some(ref providers) = self.providers { - if !providers.contains(&flow.metadata.provider) { - return false; - } - } - - // 模型过滤(支持通配符) - if let Some(ref models) = self.models { - let model_matches = models - .iter() - .any(|pattern| Self::match_pattern(pattern, &flow.request.model)); - if !model_matches { - return false; - } - } - - // 状态过滤 - if let Some(ref states) = self.states { - if !states.contains(&flow.state) { - return false; - } - } - - // 错误过滤 - if let Some(has_error) = self.has_error { - let flow_has_error = flow.error.is_some(); - if has_error != flow_has_error { - return false; - } - } - - // 工具调用过滤 - if let Some(has_tool_calls) = self.has_tool_calls { - let flow_has_tool_calls = flow - .response - .as_ref() - .is_some_and(|r| !r.tool_calls.is_empty()); - if has_tool_calls != flow_has_tool_calls { - return false; - } - } - - // 思维链过滤 - if let Some(has_thinking) = self.has_thinking { - let flow_has_thinking = flow.response.as_ref().is_some_and(|r| r.thinking.is_some()); - if has_thinking != flow_has_thinking { - return false; - } - } - - // 流式响应过滤 - if let Some(is_streaming) = self.is_streaming { - let flow_is_streaming = flow.request.parameters.stream; - if is_streaming != flow_is_streaming { - return false; - } - } - - // 内容搜索(搜索响应内容、模型名称、提供商名称) - if let Some(ref search) = self.content_search { - let search_lower = search.to_lowercase(); - - // 搜索响应内容 - let content = flow - .response - .as_ref() - .map_or(String::new(), |r| r.content.clone()); - let content_matches = content.to_lowercase().contains(&search_lower); - - // 搜索模型名称 - let model_matches = flow.request.model.to_lowercase().contains(&search_lower); - - // 搜索提供商名称 - let provider_name = format!("{:?}", flow.metadata.provider).to_lowercase(); - let provider_matches = provider_name.contains(&search_lower); - - // 任一匹配即可 - if !content_matches && !model_matches && !provider_matches { - return false; - } - } - - // 请求搜索 - if let Some(ref search) = self.request_search { - let request_text = Self::get_request_text(flow); - if !request_text.to_lowercase().contains(&search.to_lowercase()) { - return false; - } - } - - // Token 范围过滤 - if let Some(ref token_range) = self.token_range { - let total_tokens = flow.response.as_ref().map_or(0, |r| r.usage.total_tokens); - if !token_range.contains(total_tokens) { - return false; - } - } - - // 延迟范围过滤 - if let Some(ref latency_range) = self.latency_range { - if !latency_range.contains(flow.timestamps.duration_ms) { - return false; - } - } - - // 标签过滤 - if let Some(ref tags) = self.tags { - let has_any_tag = tags.iter().any(|t| flow.annotations.tags.contains(t)); - if !has_any_tag { - return false; - } - } - - // 收藏过滤 - if self.starred_only && !flow.annotations.starred { - return false; - } - - // 凭证 ID 过滤 - if let Some(ref credential_id) = self.credential_id { - if flow.metadata.credential_id.as_ref() != Some(credential_id) { - return false; - } - } - - // Flow 类型过滤 - if let Some(ref flow_types) = self.flow_types { - if !flow_types.contains(&flow.flow_type) { - return false; - } - } - - true - } - - /// 模式匹配(支持 * 通配符) - fn match_pattern(pattern: &str, text: &str) -> bool { - if pattern == "*" { - return true; - } - - if pattern.contains('*') { - // 简单的通配符匹配 - let parts: Vec<&str> = pattern.split('*').collect(); - let mut pos = 0; - let text_lower = text.to_lowercase(); - - for (i, part) in parts.iter().enumerate() { - if part.is_empty() { - continue; - } - - let part_lower = part.to_lowercase(); - if let Some(found_pos) = text_lower[pos..].find(&part_lower) { - // 第一个部分必须从开头匹配 - if i == 0 && found_pos != 0 { - return false; - } - pos += found_pos + part.len(); - } else { - return false; - } - } - - // 最后一个部分必须匹配到结尾 - if !pattern.ends_with('*') && pos != text.len() { - return false; - } - - true - } else { - text.to_lowercase() == pattern.to_lowercase() - } - } - - /// 获取请求文本(用于搜索) - fn get_request_text(flow: &LLMFlow) -> String { - let mut text = String::new(); - - // 添加系统提示词 - if let Some(ref system) = flow.request.system_prompt { - text.push_str(system); - text.push('\n'); - } - - // 添加消息内容 - for msg in &flow.request.messages { - text.push_str(&msg.content.get_all_text()); - text.push('\n'); - } - - text - } -} - -// ============================================================================ -// 内存存储 -// ============================================================================ - -/// Flow 内存存储 -/// -/// 使用 LRU 策略管理内存中的 Flow 缓存。 -/// 线程安全,支持并发读写。 -pub struct FlowMemoryStore { - /// Flow 存储(ID -> Flow) - flows: HashMap>>, - /// 有序 ID 列表(用于 LRU 驱逐) - ordered_ids: VecDeque, - /// 最大缓存大小 - max_size: usize, -} - -impl FlowMemoryStore { - /// 创建新的内存存储 - /// - /// # 参数 - /// - `max_size`: 最大缓存 Flow 数量 - pub fn new(max_size: usize) -> Self { - Self { - flows: HashMap::with_capacity(max_size), - ordered_ids: VecDeque::with_capacity(max_size), - max_size, - } - } - - /// 获取当前缓存大小 - pub fn len(&self) -> usize { - self.flows.len() - } - - /// 检查缓存是否为空 - pub fn is_empty(&self) -> bool { - self.flows.is_empty() - } - - /// 获取最大缓存大小 - pub fn max_size(&self) -> usize { - self.max_size - } - - /// 添加 Flow 到缓存 - /// - /// 如果缓存已满,会驱逐最旧的 Flow。 - pub fn add(&mut self, flow: LLMFlow) { - let id = flow.id.clone(); - - eprintln!( - "[MEMORY_STORE] 添加 Flow: id={}, model={}, state={:?}", - id, flow.request.model, flow.state - ); - - // 如果已存在,先移除旧的 - if self.flows.contains_key(&id) { - self.ordered_ids.retain(|i| i != &id); - eprintln!("[MEMORY_STORE] 移除旧的 Flow: id={id}"); - } - - // 检查是否需要驱逐 - while self.flows.len() >= self.max_size { - eprintln!("[MEMORY_STORE] 缓存已满,驱逐最旧的 Flow"); - self.evict_oldest(); - } - - // 添加新 Flow - self.flows.insert(id.clone(), Arc::new(RwLock::new(flow))); - self.ordered_ids.push_back(id.clone()); - - eprintln!("[MEMORY_STORE] Flow 已添加,当前数量: {}", self.flows.len()); - } - - /// 获取 Flow - /// - /// 返回 Flow 的共享引用,可用于读取或更新。 - pub fn get(&self, id: &str) -> Option>> { - self.flows.get(id).cloned() - } - - /// 更新 Flow - /// - /// 使用提供的更新函数修改 Flow。 - /// - /// # 参数 - /// - `id`: Flow ID - /// - `updater`: 更新函数 - /// - /// # 返回 - /// - `true`: 更新成功 - /// - `false`: Flow 不存在 - pub fn update(&self, id: &str, updater: F) -> bool - where - F: FnOnce(&mut LLMFlow), - { - if let Some(flow_lock) = self.flows.get(id) { - if let Ok(mut flow) = flow_lock.write() { - updater(&mut flow); - return true; - } - } - false - } - - /// 获取最近的 Flow 列表 - /// - /// # 参数 - /// - `limit`: 最大返回数量 - /// - /// # 返回 - /// 按时间倒序排列的 Flow 列表 - pub fn get_recent(&self, limit: usize) -> Vec { - let mut flows: Vec = Vec::with_capacity(limit.min(self.flows.len())); - - // 从最新到最旧遍历 - for id in self.ordered_ids.iter().rev().take(limit) { - if let Some(flow_lock) = self.flows.get(id) { - if let Ok(flow) = flow_lock.read() { - flows.push(flow.clone()); - } - } - } - - flows - } - - /// 查询 Flow - /// - /// # 参数 - /// - `filter`: 过滤条件 - /// - /// # 返回 - /// 匹配过滤条件的 Flow 列表(按时间倒序) - pub fn query(&self, filter: &FlowFilter) -> Vec { - let mut results: Vec = Vec::new(); - - // 从最新到最旧遍历 - for id in self.ordered_ids.iter().rev() { - if let Some(flow_lock) = self.flows.get(id) { - if let Ok(flow) = flow_lock.read() { - if filter.matches(&flow) { - results.push(flow.clone()); - } - } - } - } - - results - } - - /// 删除 Flow - /// - /// # 返回 - /// - `true`: 删除成功 - /// - `false`: Flow 不存在 - pub fn remove(&mut self, id: &str) -> bool { - if self.flows.remove(id).is_some() { - self.ordered_ids.retain(|i| i != id); - true - } else { - false - } - } - - /// 清空所有 Flow - pub fn clear(&mut self) { - self.flows.clear(); - self.ordered_ids.clear(); - } - - /// 按时间清理 Flow - /// - /// 删除指定时间之前创建的所有 Flow。 - /// - /// # 参数 - /// - `before`: 截止时间,早于此时间的 Flow 将被删除 - /// - /// # 返回 - /// 删除的 Flow 数量 - pub fn cleanup_before(&mut self, before: chrono::DateTime) -> usize { - let mut to_remove = Vec::new(); - - for id in &self.ordered_ids { - if let Some(flow_lock) = self.flows.get(id) { - if let Ok(flow) = flow_lock.read() { - if flow.timestamps.created < before { - to_remove.push(id.clone()); - } - } - } - } - - let count = to_remove.len(); - for id in to_remove { - self.flows.remove(&id); - self.ordered_ids.retain(|i| i != &id); - } - - count - } - - /// 驱逐最旧的 Flow - fn evict_oldest(&mut self) { - if let Some(oldest_id) = self.ordered_ids.pop_front() { - self.flows.remove(&oldest_id); - } - } - - /// 获取所有 Flow ID - pub fn get_all_ids(&self) -> Vec { - self.ordered_ids.iter().cloned().collect() - } - - /// 检查 Flow 是否存在 - pub fn contains(&self, id: &str) -> bool { - self.flows.contains_key(id) - } -} - -// ============================================================================ -// 测试模块 -// ============================================================================ - -#[cfg(test)] -mod tests { - use super::*; - use crate::flow_monitor::models::{FlowMetadata, LLMRequest, RequestParameters}; - - /// 创建测试用的 Flow - fn create_test_flow(id: &str, model: &str, provider: ProviderType) -> LLMFlow { - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model: model.to_string(), - parameters: RequestParameters { - stream: false, - ..Default::default() - }, - ..Default::default() - }; - - let metadata = FlowMetadata { - provider, - ..Default::default() - }; - - LLMFlow::new(id.to_string(), FlowType::ChatCompletions, request, metadata) - } - - #[test] - fn test_memory_store_add_and_get() { - let mut store = FlowMemoryStore::new(10); - let flow = create_test_flow("test-1", "gpt-4", ProviderType::OpenAI); - - store.add(flow.clone()); - - assert_eq!(store.len(), 1); - assert!(store.contains("test-1")); - - let retrieved = store.get("test-1").unwrap(); - let retrieved_flow = retrieved.read().unwrap(); - assert_eq!(retrieved_flow.id, "test-1"); - assert_eq!(retrieved_flow.request.model, "gpt-4"); - } - - #[test] - fn test_memory_store_lru_eviction() { - let mut store = FlowMemoryStore::new(3); - - // 添加 3 个 Flow - store.add(create_test_flow("flow-1", "gpt-4", ProviderType::OpenAI)); - store.add(create_test_flow("flow-2", "gpt-4", ProviderType::OpenAI)); - store.add(create_test_flow("flow-3", "gpt-4", ProviderType::OpenAI)); - - assert_eq!(store.len(), 3); - - // 添加第 4 个,应该驱逐最旧的 - store.add(create_test_flow("flow-4", "gpt-4", ProviderType::OpenAI)); - - assert_eq!(store.len(), 3); - assert!(!store.contains("flow-1")); // 最旧的被驱逐 - assert!(store.contains("flow-2")); - assert!(store.contains("flow-3")); - assert!(store.contains("flow-4")); - } - - #[test] - fn test_memory_store_update() { - let mut store = FlowMemoryStore::new(10); - let flow = create_test_flow("test-1", "gpt-4", ProviderType::OpenAI); - - store.add(flow); - - // 更新 Flow - let updated = store.update("test-1", |f| { - f.state = FlowState::Completed; - }); - - assert!(updated); - - // 验证更新 - let retrieved = store.get("test-1").unwrap(); - let retrieved_flow = retrieved.read().unwrap(); - assert_eq!(retrieved_flow.state, FlowState::Completed); - } - - #[test] - fn test_memory_store_get_recent() { - let mut store = FlowMemoryStore::new(10); - - store.add(create_test_flow("flow-1", "gpt-4", ProviderType::OpenAI)); - store.add(create_test_flow("flow-2", "gpt-4", ProviderType::OpenAI)); - store.add(create_test_flow("flow-3", "gpt-4", ProviderType::OpenAI)); - - let recent = store.get_recent(2); - - assert_eq!(recent.len(), 2); - assert_eq!(recent[0].id, "flow-3"); // 最新的在前 - assert_eq!(recent[1].id, "flow-2"); - } - - #[test] - fn test_memory_store_remove() { - let mut store = FlowMemoryStore::new(10); - - store.add(create_test_flow("flow-1", "gpt-4", ProviderType::OpenAI)); - store.add(create_test_flow("flow-2", "gpt-4", ProviderType::OpenAI)); - - assert!(store.remove("flow-1")); - assert_eq!(store.len(), 1); - assert!(!store.contains("flow-1")); - assert!(store.contains("flow-2")); - - // 删除不存在的 - assert!(!store.remove("flow-999")); - } - - #[test] - fn test_memory_store_clear() { - let mut store = FlowMemoryStore::new(10); - - store.add(create_test_flow("flow-1", "gpt-4", ProviderType::OpenAI)); - store.add(create_test_flow("flow-2", "gpt-4", ProviderType::OpenAI)); - - store.clear(); - - assert!(store.is_empty()); - assert_eq!(store.len(), 0); - } - - #[test] - fn test_flow_filter_provider() { - let flow = create_test_flow("test-1", "gpt-4", ProviderType::OpenAI); - - let filter = FlowFilter { - providers: Some(vec![ProviderType::OpenAI]), - ..Default::default() - }; - assert!(filter.matches(&flow)); - - let filter = FlowFilter { - providers: Some(vec![ProviderType::Claude]), - ..Default::default() - }; - assert!(!filter.matches(&flow)); - } - - #[test] - fn test_flow_filter_model_wildcard() { - let flow = create_test_flow("test-1", "gpt-4-turbo", ProviderType::OpenAI); - - // 精确匹配 - let filter = FlowFilter { - models: Some(vec!["gpt-4-turbo".to_string()]), - ..Default::default() - }; - assert!(filter.matches(&flow)); - - // 通配符匹配 - let filter = FlowFilter { - models: Some(vec!["gpt-4*".to_string()]), - ..Default::default() - }; - assert!(filter.matches(&flow)); - - // 通配符不匹配 - let filter = FlowFilter { - models: Some(vec!["claude*".to_string()]), - ..Default::default() - }; - assert!(!filter.matches(&flow)); - } - - #[test] - fn test_flow_filter_state() { - let mut flow = create_test_flow("test-1", "gpt-4", ProviderType::OpenAI); - flow.state = FlowState::Completed; - - let filter = FlowFilter { - states: Some(vec![FlowState::Completed]), - ..Default::default() - }; - assert!(filter.matches(&flow)); - - let filter = FlowFilter { - states: Some(vec![FlowState::Pending]), - ..Default::default() - }; - assert!(!filter.matches(&flow)); - } - - #[test] - fn test_flow_filter_starred() { - let mut flow = create_test_flow("test-1", "gpt-4", ProviderType::OpenAI); - - let filter = FlowFilter { - starred_only: true, - ..Default::default() - }; - assert!(!filter.matches(&flow)); - - flow.annotations.starred = true; - assert!(filter.matches(&flow)); - } - - #[test] - fn test_memory_store_query() { - let mut store = FlowMemoryStore::new(10); - - store.add(create_test_flow("flow-1", "gpt-4", ProviderType::OpenAI)); - store.add(create_test_flow("flow-2", "claude-3", ProviderType::Claude)); - store.add(create_test_flow( - "flow-3", - "gpt-4-turbo", - ProviderType::OpenAI, - )); - - // 按提供商过滤 - let filter = FlowFilter { - providers: Some(vec![ProviderType::OpenAI]), - ..Default::default() - }; - let results = store.query(&filter); - assert_eq!(results.len(), 2); - - // 按模型通配符过滤 - let filter = FlowFilter { - models: Some(vec!["gpt-4*".to_string()]), - ..Default::default() - }; - let results = store.query(&filter); - assert_eq!(results.len(), 2); - } - - #[test] - fn test_time_range() { - let now = Utc::now(); - let past = now - chrono::Duration::hours(1); - let future = now + chrono::Duration::hours(1); - - let range = TimeRange::new(Some(past), Some(future)); - assert!(range.contains(&now)); - - let range = TimeRange::new(Some(future), None); - assert!(!range.contains(&now)); - - let range = TimeRange::new(None, Some(past)); - assert!(!range.contains(&now)); - } - - #[test] - fn test_token_range() { - let range = TokenRange { - min: Some(100), - max: Some(1000), - }; - - assert!(range.contains(500)); - assert!(range.contains(100)); - assert!(range.contains(1000)); - assert!(!range.contains(50)); - assert!(!range.contains(1500)); - } - - #[test] - fn test_latency_range() { - let range = LatencyRange { - min_ms: Some(100), - max_ms: Some(1000), - }; - - assert!(range.contains(500)); - assert!(range.contains(100)); - assert!(range.contains(1000)); - assert!(!range.contains(50)); - assert!(!range.contains(1500)); - } -} - -// ============================================================================ -// 属性测试模块 -// ============================================================================ - -#[cfg(test)] -mod property_tests { - #![allow(dead_code)] - use super::*; - use crate::flow_monitor::models::{FlowMetadata, LLMRequest, RequestParameters}; - use proptest::prelude::*; - - // ======================================================================== - // 生成器 - // ======================================================================== - - /// 生成随机的 ProviderType - fn arb_provider_type() -> impl Strategy { - prop_oneof![ - Just(ProviderType::Kiro), - Just(ProviderType::Gemini), - Just(ProviderType::OpenAI), - Just(ProviderType::Claude), - Just(ProviderType::Antigravity), - Just(ProviderType::Vertex), - Just(ProviderType::GeminiApiKey), - Just(ProviderType::Codex), - Just(ProviderType::ClaudeOAuth), - ] - } - - /// 生成随机的 FlowType - fn arb_flow_type() -> impl Strategy { - prop_oneof![ - Just(FlowType::ChatCompletions), - Just(FlowType::AnthropicMessages), - Just(FlowType::GeminiGenerateContent), - Just(FlowType::Embeddings), - "[a-z]{3,10}".prop_map(FlowType::Other), - ] - } - - /// 生成随机的 Flow ID - fn arb_flow_id() -> impl Strategy { - "[a-f0-9]{8}-[a-f0-9]{4}-[a-f0-9]{4}-[a-f0-9]{4}-[a-f0-9]{12}" - } - - /// 生成随机的模型名称 - fn arb_model_name() -> impl Strategy { - prop_oneof![ - Just("gpt-4".to_string()), - Just("gpt-4-turbo".to_string()), - Just("gpt-3.5-turbo".to_string()), - Just("claude-3-opus".to_string()), - Just("claude-3-sonnet".to_string()), - Just("gemini-pro".to_string()), - "[a-z]{3,10}-[0-9]{1,2}".prop_map(|s| s), - ] - } - - /// 生成随机的 LLMRequest - fn arb_llm_request() -> impl Strategy { - (arb_model_name(), any::()).prop_map(|(model, stream)| LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model, - parameters: RequestParameters { - stream, - ..Default::default() - }, - ..Default::default() - }) - } - - /// 生成随机的 FlowMetadata - fn arb_flow_metadata() -> impl Strategy { - arb_provider_type().prop_map(|provider| FlowMetadata { - provider, - ..Default::default() - }) - } - - /// 生成随机的 LLMFlow - fn arb_llm_flow() -> impl Strategy { - ( - arb_flow_id(), - arb_flow_type(), - arb_llm_request(), - arb_flow_metadata(), - ) - .prop_map(|(id, flow_type, request, metadata)| { - LLMFlow::new(id, flow_type, request, metadata) - }) - } - - /// 生成随机的缓存大小(1-100) - fn arb_cache_size() -> impl Strategy { - 1usize..=100usize - } - - /// 生成随机的 Flow 数量(用于测试) - fn arb_flow_count() -> impl Strategy { - 1usize..=200usize - } - - // ======================================================================== - // 属性测试 - // ======================================================================== - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: llm-flow-monitor, Property 4: 内存缓存大小不变量** - /// **Validates: Requirements 3.1, 3.2** - /// - /// *对于任意* 数量的 Flow 添加操作,内存缓存中的 Flow 数量应该永远不超过配置的最大值。 - #[test] - fn prop_memory_cache_size_invariant( - max_size in arb_cache_size(), - flow_count in arb_flow_count(), - ) { - let mut store = FlowMemoryStore::new(max_size); - - // 添加多个 Flow - for i in 0..flow_count { - let id = format!("flow-{i}"); - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model: "gpt-4".to_string(), - ..Default::default() - }; - let metadata = FlowMetadata::default(); - let flow = LLMFlow::new(id, FlowType::ChatCompletions, request, metadata); - - store.add(flow); - - // 验证不变量:缓存大小永远不超过 max_size - prop_assert!( - store.len() <= max_size, - "缓存大小 {} 超过了最大值 {}", - store.len(), - max_size - ); - } - - // 最终验证 - let expected_size = flow_count.min(max_size); - prop_assert_eq!( - store.len(), - expected_size, - "最终缓存大小应该是 min(flow_count, max_size)" - ); - } - - /// **Feature: llm-flow-monitor, Property 4b: LRU 驱逐正确性** - /// **Validates: Requirements 3.2** - /// - /// *对于任意* 缓存大小和 Flow 序列,当缓存满时应该驱逐最旧的 Flow。 - #[test] - fn prop_lru_eviction_correctness( - max_size in 2usize..=10usize, - ) { - let mut store = FlowMemoryStore::new(max_size); - - // 添加 max_size + 1 个 Flow - let total_flows = max_size + 1; - for i in 0..total_flows { - let id = format!("flow-{i}"); - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model: "gpt-4".to_string(), - ..Default::default() - }; - let metadata = FlowMetadata::default(); - let flow = LLMFlow::new(id, FlowType::ChatCompletions, request, metadata); - store.add(flow); - } - - // 验证最旧的 Flow 被驱逐 - prop_assert!( - !store.contains("flow-0"), - "最旧的 Flow (flow-0) 应该被驱逐" - ); - - // 验证最新的 Flow 仍然存在 - for i in 1..total_flows { - prop_assert!( - store.contains(&format!("flow-{i}")), - "Flow flow-{} 应该仍然存在", - i - ); - } - } - - /// **Feature: llm-flow-monitor, Property 4c: 存储 Round-Trip** - /// **Validates: Requirements 3.1** - /// - /// *对于任意* 有效的 LLMFlow,添加到缓存后再读取,读取的 Flow 应该与原始 Flow 等价。 - #[test] - fn prop_memory_store_roundtrip( - id in arb_flow_id(), - flow_type in arb_flow_type(), - request in arb_llm_request(), - metadata in arb_flow_metadata(), - ) { - let mut store = FlowMemoryStore::new(100); - - let original_flow = LLMFlow::new(id.clone(), flow_type, request, metadata); - - // 添加到缓存 - store.add(original_flow.clone()); - - // 读取 - let retrieved = store.get(&id).expect("Flow 应该存在"); - let retrieved_flow = retrieved.read().unwrap(); - - // 验证关键字段一致 - prop_assert_eq!(&retrieved_flow.id, &original_flow.id, "ID 应该一致"); - prop_assert_eq!(&retrieved_flow.state, &original_flow.state, "状态应该一致"); - prop_assert_eq!( - &retrieved_flow.request.model, - &original_flow.request.model, - "模型应该一致" - ); - prop_assert_eq!( - &retrieved_flow.metadata.provider, - &original_flow.metadata.provider, - "Provider 应该一致" - ); - } - - /// **Feature: llm-flow-monitor, Property 4d: 过滤正确性** - /// **Validates: Requirements 4.1-4.9** - /// - /// *对于任意* 过滤条件和 Flow 集合,查询返回的所有 Flow 都应该满足该过滤条件。 - #[test] - fn prop_filter_correctness( - provider in arb_provider_type(), - ) { - let mut store = FlowMemoryStore::new(100); - - // 添加不同 Provider 的 Flow - let providers = [ProviderType::OpenAI, - ProviderType::Claude, - ProviderType::Gemini, - ProviderType::Kiro]; - - for (i, p) in providers.iter().enumerate() { - let id = format!("flow-{i}"); - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model: "gpt-4".to_string(), - ..Default::default() - }; - let metadata = FlowMetadata { - provider: *p, - ..Default::default() - }; - let flow = LLMFlow::new(id, FlowType::ChatCompletions, request, metadata); - store.add(flow); - } - - // 按 Provider 过滤 - let filter = FlowFilter { - providers: Some(vec![provider]), - ..Default::default() - }; - - let results = store.query(&filter); - - // 验证所有结果都匹配过滤条件 - for flow in &results { - prop_assert_eq!( - flow.metadata.provider, - provider, - "查询结果的 Provider 应该匹配过滤条件" - ); - } - } - - /// **Feature: llm-flow-monitor, Property 4e: 模型通配符过滤正确性** - /// **Validates: Requirements 4.3** - /// - /// *对于任意* 模型通配符模式,查询返回的所有 Flow 的模型名称都应该匹配该模式。 - #[test] - fn prop_model_wildcard_filter_correctness( - prefix in "[a-z]{2,5}", - ) { - let mut store = FlowMemoryStore::new(100); - - // 添加不同模型的 Flow - // 使用数字前缀的模型名称,确保不会与随机生成的字母 prefix 冲突 - let models = [format!("{prefix}-model-1"), - format!("{prefix}-model-2"), - "123-non-matching-model".to_string(), - "456-another-non-matching".to_string()]; - - for (i, model) in models.iter().enumerate() { - let id = format!("flow-{i}"); - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model: model.clone(), - ..Default::default() - }; - let metadata = FlowMetadata::default(); - let flow = LLMFlow::new(id, FlowType::ChatCompletions, request, metadata); - store.add(flow); - } - - // 使用通配符过滤 - let pattern = format!("{prefix}*"); - let filter = FlowFilter { - models: Some(vec![pattern.clone()]), - ..Default::default() - }; - - let results = store.query(&filter); - - // 验证所有结果都匹配通配符模式 - for flow in &results { - prop_assert!( - flow.request.model.to_lowercase().starts_with(&prefix.to_lowercase()), - "模型 {} 应该以 {} 开头", - flow.request.model, - prefix - ); - } - - // 验证匹配数量正确(应该是 2 个以 prefix 开头的模型) - prop_assert_eq!(results.len(), 2, "应该有 2 个匹配的 Flow"); - } - - /// **Feature: llm-flow-monitor, Property 4f: 更新操作正确性** - /// **Validates: Requirements 3.1** - /// - /// *对于任意* Flow 和更新操作,更新后的 Flow 应该反映更新内容。 - #[test] - fn prop_update_correctness( - id in arb_flow_id(), - new_state in prop_oneof![ - Just(FlowState::Streaming), - Just(FlowState::Completed), - Just(FlowState::Failed), - ], - ) { - let mut store = FlowMemoryStore::new(100); - - let request = LLMRequest::default(); - let metadata = FlowMetadata::default(); - let flow = LLMFlow::new(id.clone(), FlowType::ChatCompletions, request, metadata); - - store.add(flow); - - // 更新状态 - let updated = store.update(&id, |f| { - f.state = new_state.clone(); - }); - - prop_assert!(updated, "更新应该成功"); - - // 验证更新生效 - let retrieved = store.get(&id).unwrap(); - let retrieved_flow = retrieved.read().unwrap(); - prop_assert_eq!( - &retrieved_flow.state, - &new_state, - "状态应该被更新" - ); - } - - /// **Feature: llm-flow-monitor, Property 4g: get_recent 顺序正确性** - /// **Validates: Requirements 3.1** - /// - /// *对于任意* Flow 序列,get_recent 返回的 Flow 应该按添加顺序倒序排列。 - #[test] - fn prop_get_recent_order( - count in 5usize..=20usize, - ) { - let mut store = FlowMemoryStore::new(100); - - // 添加多个 Flow - for i in 0..count { - let id = format!("flow-{i:03}"); - let request = LLMRequest::default(); - let metadata = FlowMetadata::default(); - let flow = LLMFlow::new(id, FlowType::ChatCompletions, request, metadata); - store.add(flow); - } - - // 获取最近的 Flow - let recent = store.get_recent(count); - - // 验证顺序(最新的在前) - for (i, flow) in recent.iter().enumerate() { - let expected_id = format!("flow-{:03}", count - 1 - i); - prop_assert_eq!( - &flow.id, - &expected_id, - "第 {} 个 Flow 应该是 {}", - i, - expected_id - ); - } - } - } -} diff --git a/src-tauri/src/flow_monitor/mod.rs b/src-tauri/src/flow_monitor/mod.rs deleted file mode 100644 index c99984857..000000000 --- a/src-tauri/src/flow_monitor/mod.rs +++ /dev/null @@ -1,148 +0,0 @@ -//! LLM Flow Monitor 模块 -//! -//! 该模块提供完整的 LLM API 流量监控功能,参考 mitmproxy 的 Flow 模型设计。 -//! 用于捕获、存储、分析和回放 AI Agent 与大模型之间的完整交互数据。 -//! -//! # 主要组件 -//! -//! - `models`: 核心数据模型,包括 LLMFlow、LLMRequest、LLMResponse 等 -//! - `stream_rebuilder`: SSE 流式响应重建器 -//! - `memory_store`: 内存存储,支持 LRU 驱逐策略 -//! - `file_store`: 文件存储,支持 JSONL 格式和 SQLite 索引 -//! - `query_service`: 查询服务,支持多维度过滤、排序、分页和全文搜索 -//! - `exporter`: 导出服务,支持 HAR、JSON、JSONL、Markdown、CSV 格式 -//! - `monitor`: 核心监控服务 -//! - `filter_parser`: 高级过滤表达式解析器,支持类似 mitmproxy 的语法 - -pub mod batch_ops; -pub mod bookmark; -pub mod code_exporter; -pub mod diff; -pub mod enhanced_stats; -pub mod exporter; -pub mod file_store; -pub mod filter_parser; -pub mod interceptor; -pub mod memory_store; -pub mod models; -pub mod monitor; -pub mod query_service; -pub mod quick_filter; -pub mod replayer; -pub mod session; -pub mod stream_rebuilder; - -// 重新导出核心类型 -pub use models::{ - ClientInfo, - ContentPart, - FlowAnnotations, - // 错误 - FlowError, - FlowErrorType, - // 元数据 - FlowMetadata, - FlowState, - FlowTimestamps, - FlowType, - // 核心 Flow 结构 - LLMFlow, - // 请求相关 - LLMRequest, - // 响应相关 - LLMResponse, - Message, - MessageContent, - MessageRole, - RequestParameters, - RoutingInfo, - StopReason, - StreamChunk, - StreamInfo, - ThinkingContent, - TokenUsage, - ToolCall, - ToolCallDelta, - ToolDefinition, - ToolResult, -}; - -// 重新导出流重建器 -pub use stream_rebuilder::{StreamFormat, StreamRebuilder, StreamRebuilderError}; - -// 重新导出内存存储 -pub use memory_store::{FlowFilter, FlowMemoryStore, LatencyRange, TimeRange, TokenRange}; - -// 重新导出文件存储 -pub use file_store::{ - CleanupResult, FileStoreError, FlowFileStore, FlowIndexRecord, FtsSearchResult, RotationConfig, -}; - -// 重新导出查询服务 -pub use query_service::{ - FlowQueryResult, FlowQueryService, FlowSearchResult, FlowSortBy, FlowStats, ModelStats, - ProviderStats, QueryWithExpressionError, StateStats, -}; - -// 重新导出导出服务 -pub use exporter::{ - default_redaction_rules, ExportFormat, ExportOptions, ExportResult, FlowExporter, HarArchive, - HarEntry, HarLlmExtension, HarLog, RedactionRule, Redactor, -}; - -// 重新导出监控服务 -pub use monitor::{ - FlowEvent, FlowMonitor, FlowMonitorConfig, FlowSummary, FlowUpdate, RequestRateTracker, - ThresholdCheckResult, ThresholdConfig, -}; - -// 重新导出过滤表达式解析器 -pub use filter_parser::{ - get_filter_help, Comparison, ComparisonOp, FilterExpr, FilterParseError, FilterParser, - FilterToken, FILTER_HELP, -}; - -// 重新导出拦截器 -pub use interceptor::{ - FlowInterceptor, InterceptAction, InterceptConfig, InterceptEvent, InterceptState, - InterceptType, InterceptedFlow, InterceptorError, ModifiedData, TimeoutAction, -}; - -// 重新导出重放器 -pub use replayer::{ - BatchReplayResult, FlowReplayer, ReplayConfig, ReplayResult, ReplayerError, RequestModification, -}; - -// 重新导出差异对比器 -pub use diff::{ - DiffConfig, DiffItem, DiffType, FlowDiff, FlowDiffResult, MessageDiffItem, TokenDiff, -}; - -// 重新导出会话管理器 -pub use session::{ - AutoSessionConfig, FlowSession, SessionError, SessionExportResult, SessionManager, -}; - -// 重新导出快速过滤器管理器 -pub use quick_filter::{ - QuickFilter, QuickFilterError, QuickFilterExport, QuickFilterManager, QuickFilterUpdate, - PRESET_FILTERS, -}; - -// 重新导出代码导出器 -pub use code_exporter::{CodeExporter, CodeFormat}; - -// 重新导出书签管理器 -pub use bookmark::{BookmarkError, BookmarkExport, BookmarkManager, FlowBookmark}; - -// 重新导出增强统计服务 -pub use enhanced_stats::{ - Distribution, EnhancedStats, EnhancedStatsService, ReportFormat, StatsTimeRange, - TimeSeriesPoint, TrendData, -}; - -// 重新导出批量操作服务 -pub use batch_ops::{BatchOperation, BatchOperations, BatchOpsError, BatchResult}; - -// 重新导出 ProviderType(从 lib.rs) -pub use crate::ProviderType; diff --git a/src-tauri/src/flow_monitor/models.rs b/src-tauri/src/flow_monitor/models.rs deleted file mode 100644 index 0d4888db4..000000000 --- a/src-tauri/src/flow_monitor/models.rs +++ /dev/null @@ -1,1274 +0,0 @@ -//! LLM Flow Monitor 核心数据模型 -//! -//! 定义 LLM 请求/响应流的完整数据结构,参考 mitmproxy 的 Flow 模型设计。 - -use chrono::{DateTime, Utc}; -use serde::{Deserialize, Serialize}; -use std::collections::HashMap; - -use crate::ProviderType; - -// ============================================================================ -// 核心 Flow 结构 -// ============================================================================ - -/// LLM 请求/响应流 -/// -/// 类似 mitmproxy 的 HTTPFlow,但专门针对 LLM API 优化。 -/// 包含完整的请求信息、响应信息、元数据和时间戳。 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct LLMFlow { - /// 唯一标识符 - pub id: String, - /// 流类型 - pub flow_type: FlowType, - /// 请求信息 - pub request: LLMRequest, - /// 响应信息(可能为空,如请求失败或正在进行中) - pub response: Option, - /// 错误信息(如果发生错误) - pub error: Option, - /// 元数据 - pub metadata: FlowMetadata, - /// 时间戳 - pub timestamps: FlowTimestamps, - /// 流状态 - pub state: FlowState, - /// 用户标记和注释 - pub annotations: FlowAnnotations, -} - -impl LLMFlow { - /// 创建新的 LLM Flow - pub fn new( - id: String, - flow_type: FlowType, - request: LLMRequest, - metadata: FlowMetadata, - ) -> Self { - let now = Utc::now(); - Self { - id, - flow_type, - request: request.clone(), - response: None, - error: None, - metadata, - timestamps: FlowTimestamps { - created: now, - request_start: request.timestamp, - request_end: None, - response_start: None, - response_end: None, - duration_ms: 0, - ttfb_ms: None, - }, - state: FlowState::Pending, - annotations: FlowAnnotations::default(), - } - } -} - -/// 流类型 -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] -pub enum FlowType { - /// OpenAI Chat Completions - #[default] - ChatCompletions, - /// Anthropic Messages - AnthropicMessages, - /// Gemini Generate Content - GeminiGenerateContent, - /// Embeddings - Embeddings, - /// 其他类型 - Other(String), -} - -/// 流状态 -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)] -pub enum FlowState { - /// 等待响应 - #[default] - Pending, - /// 正在流式传输 - Streaming, - /// 已完成 - Completed, - /// 失败 - Failed, - /// 已取消 - Cancelled, -} - -// ============================================================================ -// 请求数据结构 -// ============================================================================ - -/// LLM 请求 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct LLMRequest { - /// HTTP 方法 - pub method: String, - /// 请求路径 - pub path: String, - /// 请求头 - pub headers: HashMap, - /// 原始请求体(JSON) - pub body: serde_json::Value, - /// 解析后的消息列表 - pub messages: Vec, - /// 系统提示词(如果有) - pub system_prompt: Option, - /// 工具定义(如果有) - pub tools: Option>, - /// 请求的模型名称 - pub model: String, - /// 原始模型名称(别名解析前) - pub original_model: Option, - /// 请求参数 - pub parameters: RequestParameters, - /// 请求体大小(字节) - pub size_bytes: usize, - /// 请求开始时间戳 - pub timestamp: DateTime, -} - -impl Default for LLMRequest { - fn default() -> Self { - Self { - method: "POST".to_string(), - path: String::new(), - headers: HashMap::new(), - body: serde_json::Value::Null, - messages: Vec::new(), - system_prompt: None, - tools: None, - model: String::new(), - original_model: None, - parameters: RequestParameters::default(), - size_bytes: 0, - timestamp: Utc::now(), - } - } -} - -/// 消息结构 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct Message { - /// 消息角色 - pub role: MessageRole, - /// 消息内容 - pub content: MessageContent, - /// 工具调用(如果有) - pub tool_calls: Option>, - /// 工具结果(如果有) - pub tool_result: Option, - /// 消息名称(如果有) - pub name: Option, -} - -impl Default for Message { - fn default() -> Self { - Self { - role: MessageRole::User, - content: MessageContent::Text(String::new()), - tool_calls: None, - tool_result: None, - name: None, - } - } -} - -/// 消息角色 -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -#[serde(rename_all = "lowercase")] -#[derive(Default)] -pub enum MessageRole { - /// 系统消息 - System, - /// 用户消息 - #[default] - User, - /// 助手消息 - Assistant, - /// 工具消息 - Tool, - /// 函数消息(兼容旧版 OpenAI API) - Function, -} - -/// 消息内容(支持多模态) -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(untagged)] -pub enum MessageContent { - /// 纯文本内容 - Text(String), - /// 多模态内容(文本、图片等) - MultiModal(Vec), -} - -impl Default for MessageContent { - fn default() -> Self { - MessageContent::Text(String::new()) - } -} - -impl MessageContent { - /// 获取文本内容 - pub fn as_text(&self) -> Option<&str> { - match self { - MessageContent::Text(s) => Some(s), - MessageContent::MultiModal(_) => None, - } - } - - /// 获取所有文本内容(包括多模态中的文本部分) - pub fn get_all_text(&self) -> String { - match self { - MessageContent::Text(s) => s.clone(), - MessageContent::MultiModal(parts) => parts - .iter() - .filter_map(|p| { - if let ContentPart::Text { text } = p { - Some(text.as_str()) - } else { - None - } - }) - .collect::>() - .join("\n"), - } - } -} - -/// 内容部分(多模态消息的组成部分) -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(tag = "type", rename_all = "snake_case")] -pub enum ContentPart { - /// 文本部分 - Text { text: String }, - /// 图片部分 - ImageUrl { image_url: ImageUrl }, - /// 图片数据(base64) - Image { - #[serde(skip_serializing_if = "Option::is_none")] - media_type: Option, - #[serde(skip_serializing_if = "Option::is_none")] - data: Option, - #[serde(skip_serializing_if = "Option::is_none")] - url: Option, - }, -} - -/// 图片 URL -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ImageUrl { - pub url: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub detail: Option, -} - -/// 工具定义 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ToolDefinition { - /// 工具类型(通常为 "function") - #[serde(rename = "type")] - pub tool_type: String, - /// 函数定义 - pub function: FunctionDefinition, -} - -/// 函数定义 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct FunctionDefinition { - /// 函数名称 - pub name: String, - /// 函数描述 - #[serde(skip_serializing_if = "Option::is_none")] - pub description: Option, - /// 参数 schema - #[serde(skip_serializing_if = "Option::is_none")] - pub parameters: Option, -} - -/// 工具调用 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ToolCall { - /// 工具调用 ID - pub id: String, - /// 工具类型 - #[serde(rename = "type")] - pub tool_type: String, - /// 函数调用详情 - pub function: FunctionCall, -} - -/// 函数调用 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct FunctionCall { - /// 函数名称 - pub name: String, - /// 函数参数(JSON 字符串) - pub arguments: String, -} - -/// 工具结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ToolResult { - /// 工具调用 ID - pub tool_call_id: String, - /// 结果内容 - pub content: String, - /// 是否为错误结果 - #[serde(default)] - pub is_error: bool, -} - -/// 请求参数 -#[derive(Debug, Clone, Serialize, Deserialize, Default)] -pub struct RequestParameters { - /// 温度参数 - #[serde(skip_serializing_if = "Option::is_none")] - pub temperature: Option, - /// Top-p 参数 - #[serde(skip_serializing_if = "Option::is_none")] - pub top_p: Option, - /// 最大 Token 数 - #[serde(skip_serializing_if = "Option::is_none")] - pub max_tokens: Option, - /// 停止序列 - #[serde(skip_serializing_if = "Option::is_none")] - pub stop: Option>, - /// 是否流式响应 - #[serde(default)] - pub stream: bool, - /// 其他参数 - #[serde(flatten)] - pub extra: HashMap, -} - -// ============================================================================ -// 响应数据结构 -// ============================================================================ - -/// LLM 响应 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct LLMResponse { - /// HTTP 状态码 - pub status_code: u16, - /// 状态文本 - pub status_text: String, - /// 响应头 - pub headers: HashMap, - /// 原始响应体(完整 JSON,流式响应会被重建) - pub body: serde_json::Value, - /// 提取的文本内容 - pub content: String, - /// 思维链内容(如果有) - pub thinking: Option, - /// 工具调用(如果有) - pub tool_calls: Vec, - /// Token 使用统计 - pub usage: TokenUsage, - /// 停止原因 - pub stop_reason: Option, - /// 响应体大小(字节) - pub size_bytes: usize, - /// 响应开始时间戳 - pub timestamp_start: DateTime, - /// 响应结束时间戳 - pub timestamp_end: DateTime, - /// 流式响应信息(如果是流式) - pub stream_info: Option, -} - -impl Default for LLMResponse { - fn default() -> Self { - let now = Utc::now(); - Self { - status_code: 200, - status_text: "OK".to_string(), - headers: HashMap::new(), - body: serde_json::Value::Null, - content: String::new(), - thinking: None, - tool_calls: Vec::new(), - usage: TokenUsage::default(), - stop_reason: None, - size_bytes: 0, - timestamp_start: now, - timestamp_end: now, - stream_info: None, - } - } -} - -/// 思维链内容 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ThinkingContent { - /// 思维链文本 - pub text: String, - /// 思维链 Token 数 - #[serde(skip_serializing_if = "Option::is_none")] - pub tokens: Option, - /// 签名(用于验证) - #[serde(skip_serializing_if = "Option::is_none")] - pub signature: Option, -} - -/// Token 使用统计 -#[derive(Debug, Clone, Serialize, Deserialize, Default)] -pub struct TokenUsage { - /// 输入 Token 数 - pub input_tokens: u32, - /// 输出 Token 数 - pub output_tokens: u32, - /// 缓存读取 Token 数 - #[serde(skip_serializing_if = "Option::is_none")] - pub cache_read_tokens: Option, - /// 缓存写入 Token 数 - #[serde(skip_serializing_if = "Option::is_none")] - pub cache_write_tokens: Option, - /// 思维链 Token 数 - #[serde(skip_serializing_if = "Option::is_none")] - pub thinking_tokens: Option, - /// 总 Token 数 - pub total_tokens: u32, -} - -impl TokenUsage { - /// 计算总 Token 数 - pub fn calculate_total(&mut self) { - self.total_tokens = self.input_tokens + self.output_tokens; - } -} - -/// 停止原因 -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -#[serde(rename_all = "snake_case")] -pub enum StopReason { - /// 正常结束 - Stop, - /// 达到最大长度 - Length, - /// 工具调用 - ToolCalls, - /// 内容过滤 - ContentFilter, - /// 函数调用(兼容旧版) - FunctionCall, - /// 结束 Token - EndTurn, - /// 其他原因 - Other(String), -} - -/// 流式响应信息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct StreamInfo { - /// Chunk 数量 - pub chunk_count: u32, - /// 首个 Chunk 延迟(毫秒) - pub first_chunk_latency_ms: u64, - /// 平均 Chunk 间隔(毫秒) - pub avg_chunk_interval_ms: f64, - /// 原始 Chunks(可选,根据配置决定是否保存) - #[serde(skip_serializing_if = "Option::is_none")] - pub raw_chunks: Option>, -} - -/// 流式 Chunk -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct StreamChunk { - /// Chunk 索引 - pub index: u32, - /// 事件类型(SSE event) - pub event: Option, - /// 数据内容 - pub data: String, - /// 时间戳 - pub timestamp: DateTime, - /// 解析后的内容增量 - #[serde(skip_serializing_if = "Option::is_none")] - pub content_delta: Option, - /// 解析后的工具调用增量 - #[serde(skip_serializing_if = "Option::is_none")] - pub tool_call_delta: Option, - /// 解析后的思维链增量 - #[serde(skip_serializing_if = "Option::is_none")] - pub thinking_delta: Option, -} - -/// 工具调用增量 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ToolCallDelta { - /// 工具调用索引 - pub index: u32, - /// 工具调用 ID(首次出现时) - #[serde(skip_serializing_if = "Option::is_none")] - pub id: Option, - /// 函数名称(首次出现时) - #[serde(skip_serializing_if = "Option::is_none")] - pub function_name: Option, - /// 参数增量 - #[serde(skip_serializing_if = "Option::is_none")] - pub arguments_delta: Option, -} - -// ============================================================================ -// 元数据结构 -// ============================================================================ - -/// 流元数据 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct FlowMetadata { - /// 提供商类型 - pub provider: ProviderType, - /// 提供商 ID(实际的 provider ID,如 "deepseek", "moonshot" 等) - #[serde(skip_serializing_if = "Option::is_none")] - pub provider_id: Option, - /// 凭证 ID - #[serde(skip_serializing_if = "Option::is_none")] - pub credential_id: Option, - /// 凭证名称 - #[serde(skip_serializing_if = "Option::is_none")] - pub credential_name: Option, - /// 重试次数 - #[serde(default)] - pub retry_count: u32, - /// 客户端信息 - pub client_info: ClientInfo, - /// 路由信息 - pub routing_info: RoutingInfo, - /// 注入的参数 - #[serde(skip_serializing_if = "Option::is_none")] - pub injected_params: Option>, - /// 上下文使用百分比 - #[serde(skip_serializing_if = "Option::is_none")] - pub context_usage_percentage: Option, -} - -impl Default for FlowMetadata { - fn default() -> Self { - Self { - provider: ProviderType::Kiro, - provider_id: None, - credential_id: None, - credential_name: None, - retry_count: 0, - client_info: ClientInfo::default(), - routing_info: RoutingInfo::default(), - injected_params: None, - context_usage_percentage: None, - } - } -} - -/// 客户端信息 -#[derive(Debug, Clone, Serialize, Deserialize, Default)] -pub struct ClientInfo { - /// 客户端 IP - #[serde(skip_serializing_if = "Option::is_none")] - pub ip: Option, - /// User-Agent - #[serde(skip_serializing_if = "Option::is_none")] - pub user_agent: Option, - /// 请求 ID - #[serde(skip_serializing_if = "Option::is_none")] - pub request_id: Option, -} - -/// 路由信息 -#[derive(Debug, Clone, Serialize, Deserialize, Default)] -pub struct RoutingInfo { - /// 目标 URL - #[serde(skip_serializing_if = "Option::is_none")] - pub target_url: Option, - /// 使用的路由规则 - #[serde(skip_serializing_if = "Option::is_none")] - pub route_rule: Option, - /// 负载均衡策略 - #[serde(skip_serializing_if = "Option::is_none")] - pub load_balance_strategy: Option, -} - -/// 时间戳集合 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct FlowTimestamps { - /// 创建时间 - pub created: DateTime, - /// 请求开始时间 - pub request_start: DateTime, - /// 请求结束时间 - #[serde(skip_serializing_if = "Option::is_none")] - pub request_end: Option>, - /// 响应开始时间 - #[serde(skip_serializing_if = "Option::is_none")] - pub response_start: Option>, - /// 响应结束时间 - #[serde(skip_serializing_if = "Option::is_none")] - pub response_end: Option>, - /// 总耗时(毫秒) - pub duration_ms: u64, - /// 首字节时间(毫秒) - #[serde(skip_serializing_if = "Option::is_none")] - pub ttfb_ms: Option, -} - -impl Default for FlowTimestamps { - fn default() -> Self { - let now = Utc::now(); - Self { - created: now, - request_start: now, - request_end: None, - response_start: None, - response_end: None, - duration_ms: 0, - ttfb_ms: None, - } - } -} - -impl FlowTimestamps { - /// 计算耗时 - pub fn calculate_duration(&mut self) { - if let Some(end) = self.response_end { - self.duration_ms = (end - self.request_start).num_milliseconds().max(0) as u64; - } - } - - /// 计算 TTFB - pub fn calculate_ttfb(&mut self) { - if let Some(start) = self.response_start { - self.ttfb_ms = Some((start - self.request_start).num_milliseconds().max(0) as u64); - } - } -} - -/// 用户标注 -#[derive(Debug, Clone, Serialize, Deserialize, Default)] -pub struct FlowAnnotations { - /// 标记(如 ⭐、🔴、🟢) - #[serde(skip_serializing_if = "Option::is_none")] - pub marker: Option, - /// 评论 - #[serde(skip_serializing_if = "Option::is_none")] - pub comment: Option, - /// 标签 - #[serde(default)] - pub tags: Vec, - /// 是否收藏 - #[serde(default)] - pub starred: bool, -} - -// ============================================================================ -// 错误结构 -// ============================================================================ - -/// 流错误 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct FlowError { - /// 错误类型 - pub error_type: FlowErrorType, - /// 错误消息 - pub message: String, - /// HTTP 状态码(如果有) - #[serde(skip_serializing_if = "Option::is_none")] - pub status_code: Option, - /// 原始响应(如果有) - #[serde(skip_serializing_if = "Option::is_none")] - pub raw_response: Option, - /// 时间戳 - pub timestamp: DateTime, - /// 是否可重试 - pub retryable: bool, -} - -impl FlowError { - /// 创建新的错误 - pub fn new(error_type: FlowErrorType, message: impl Into) -> Self { - Self { - error_type, - message: message.into(), - status_code: None, - raw_response: None, - timestamp: Utc::now(), - retryable: false, - } - } - - /// 设置状态码 - pub fn with_status_code(mut self, code: u16) -> Self { - self.status_code = Some(code); - self - } - - /// 设置原始响应 - pub fn with_raw_response(mut self, response: impl Into) -> Self { - self.raw_response = Some(response.into()); - self - } - - /// 设置是否可重试 - pub fn with_retryable(mut self, retryable: bool) -> Self { - self.retryable = retryable; - self - } -} - -/// 错误类型 -#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] -#[serde(rename_all = "snake_case")] -#[derive(Default)] -pub enum FlowErrorType { - /// 网络错误 - Network, - /// 超时 - Timeout, - /// 认证错误 - Authentication, - /// 速率限制 - RateLimit, - /// 内容过滤 - ContentFilter, - /// 服务器错误 - ServerError, - /// 请求错误 - BadRequest, - /// 模型不可用 - ModelUnavailable, - /// Token 限制超出 - TokenLimitExceeded, - /// 请求被取消(用户拦截后取消) - Cancelled, - /// 其他错误 - #[default] - Other, -} - -impl FlowErrorType { - /// 根据 HTTP 状态码推断错误类型 - pub fn from_status_code(code: u16) -> Self { - match code { - 401 | 403 => FlowErrorType::Authentication, - 429 => FlowErrorType::RateLimit, - 400 => FlowErrorType::BadRequest, - 404 => FlowErrorType::ModelUnavailable, - 500..=599 => FlowErrorType::ServerError, - _ => FlowErrorType::Other, - } - } - - /// 判断是否可重试 - pub fn is_retryable(&self) -> bool { - matches!( - self, - FlowErrorType::Network - | FlowErrorType::Timeout - | FlowErrorType::RateLimit - | FlowErrorType::ServerError - ) - } -} - -// ============================================================================ -// 测试模块 -// ============================================================================ - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_flow_creation() { - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model: "gpt-4".to_string(), - ..Default::default() - }; - - let metadata = FlowMetadata { - provider: ProviderType::OpenAI, - ..Default::default() - }; - - let flow = LLMFlow::new( - "test-id".to_string(), - FlowType::ChatCompletions, - request, - metadata, - ); - - assert_eq!(flow.id, "test-id"); - assert_eq!(flow.state, FlowState::Pending); - assert_eq!(flow.flow_type, FlowType::ChatCompletions); - assert!(flow.response.is_none()); - assert!(flow.error.is_none()); - } - - #[test] - fn test_message_content_text() { - let content = MessageContent::Text("Hello, world!".to_string()); - assert_eq!(content.as_text(), Some("Hello, world!")); - assert_eq!(content.get_all_text(), "Hello, world!"); - } - - #[test] - fn test_message_content_multimodal() { - let content = MessageContent::MultiModal(vec![ - ContentPart::Text { - text: "First part".to_string(), - }, - ContentPart::Text { - text: "Second part".to_string(), - }, - ]); - assert!(content.as_text().is_none()); - assert_eq!(content.get_all_text(), "First part\nSecond part"); - } - - #[test] - fn test_token_usage_calculate_total() { - let mut usage = TokenUsage { - input_tokens: 100, - output_tokens: 50, - ..Default::default() - }; - usage.calculate_total(); - assert_eq!(usage.total_tokens, 150); - } - - #[test] - fn test_flow_error_type_from_status_code() { - assert_eq!( - FlowErrorType::from_status_code(401), - FlowErrorType::Authentication - ); - assert_eq!( - FlowErrorType::from_status_code(429), - FlowErrorType::RateLimit - ); - assert_eq!( - FlowErrorType::from_status_code(500), - FlowErrorType::ServerError - ); - assert_eq!(FlowErrorType::from_status_code(200), FlowErrorType::Other); - } - - #[test] - fn test_flow_error_type_is_retryable() { - assert!(FlowErrorType::Network.is_retryable()); - assert!(FlowErrorType::Timeout.is_retryable()); - assert!(FlowErrorType::RateLimit.is_retryable()); - assert!(FlowErrorType::ServerError.is_retryable()); - assert!(!FlowErrorType::Authentication.is_retryable()); - assert!(!FlowErrorType::BadRequest.is_retryable()); - } - - #[test] - fn test_flow_timestamps_calculate() { - let start = Utc::now(); - let response_start = start + chrono::Duration::milliseconds(100); - let end = start + chrono::Duration::milliseconds(500); - - let mut timestamps = FlowTimestamps { - created: start, - request_start: start, - request_end: Some(start + chrono::Duration::milliseconds(50)), - response_start: Some(response_start), - response_end: Some(end), - duration_ms: 0, - ttfb_ms: None, - }; - - timestamps.calculate_duration(); - timestamps.calculate_ttfb(); - - assert_eq!(timestamps.duration_ms, 500); - assert_eq!(timestamps.ttfb_ms, Some(100)); - } - - #[test] - fn test_flow_error_builder() { - let error = FlowError::new(FlowErrorType::RateLimit, "Too many requests") - .with_status_code(429) - .with_retryable(true); - - assert_eq!(error.error_type, FlowErrorType::RateLimit); - assert_eq!(error.message, "Too many requests"); - assert_eq!(error.status_code, Some(429)); - assert!(error.retryable); - } - - #[test] - fn test_serialization_roundtrip() { - let flow = LLMFlow::new( - "test-id".to_string(), - FlowType::ChatCompletions, - LLMRequest::default(), - FlowMetadata::default(), - ); - - let json = serde_json::to_string(&flow).unwrap(); - let deserialized: LLMFlow = serde_json::from_str(&json).unwrap(); - - assert_eq!(flow.id, deserialized.id); - assert_eq!(flow.state, deserialized.state); - } -} - -// ============================================================================ -// 属性测试模块 -// ============================================================================ - -#[cfg(test)] -mod property_tests { - use super::*; - use proptest::prelude::*; - - // ======================================================================== - // 生成器 - // ======================================================================== - - /// 生成随机的 ProviderType - fn arb_provider_type() -> impl Strategy { - prop_oneof![ - Just(ProviderType::Kiro), - Just(ProviderType::Gemini), - Just(ProviderType::OpenAI), - Just(ProviderType::Claude), - Just(ProviderType::Antigravity), - Just(ProviderType::Vertex), - Just(ProviderType::GeminiApiKey), - Just(ProviderType::Codex), - Just(ProviderType::ClaudeOAuth), - ] - } - - /// 生成随机的 FlowType - fn arb_flow_type() -> impl Strategy { - prop_oneof![ - Just(FlowType::ChatCompletions), - Just(FlowType::AnthropicMessages), - Just(FlowType::GeminiGenerateContent), - Just(FlowType::Embeddings), - "[a-z]{3,10}".prop_map(FlowType::Other), - ] - } - - /// 生成随机的 MessageRole - fn arb_message_role() -> impl Strategy { - prop_oneof![ - Just(MessageRole::System), - Just(MessageRole::User), - Just(MessageRole::Assistant), - Just(MessageRole::Tool), - Just(MessageRole::Function), - ] - } - - /// 生成随机的 MessageContent - fn arb_message_content() -> impl Strategy { - prop_oneof![ - ".*".prop_map(MessageContent::Text), - prop::collection::vec( - "[a-zA-Z0-9 ]{1,50}".prop_map(|text| ContentPart::Text { text }), - 1..5 - ) - .prop_map(MessageContent::MultiModal), - ] - } - - /// 生成随机的 Message - fn arb_message() -> impl Strategy { - (arb_message_role(), arb_message_content()).prop_map(|(role, content)| Message { - role, - content, - tool_calls: None, - tool_result: None, - name: None, - }) - } - - /// 生成随机的 RequestParameters - fn arb_request_parameters() -> impl Strategy { - ( - prop::option::of(0.0f32..2.0f32), - prop::option::of(0.0f32..1.0f32), - prop::option::of(1u32..4096u32), - any::(), - ) - .prop_map( - |(temperature, top_p, max_tokens, stream)| RequestParameters { - temperature, - top_p, - max_tokens, - stop: None, - stream, - extra: HashMap::new(), - }, - ) - } - - /// 生成随机的 LLMRequest - fn arb_llm_request() -> impl Strategy { - ( - "[a-z]{3,20}", // model - prop::collection::vec(arb_message(), 0..5), // messages - arb_request_parameters(), // parameters - prop::option::of("[a-zA-Z0-9 ]{10,100}"), // system_prompt - ) - .prop_map(|(model, messages, parameters, system_prompt)| LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - headers: HashMap::new(), - body: serde_json::Value::Null, - messages, - system_prompt, - tools: None, - model, - original_model: None, - parameters, - size_bytes: 0, - timestamp: Utc::now(), - }) - } - - /// 生成随机的 FlowMetadata - fn arb_flow_metadata() -> impl Strategy { - ( - arb_provider_type(), - prop::option::of("[a-f0-9]{8}"), - prop::option::of("[a-zA-Z0-9_]{3,20}"), - ) - .prop_map(|(provider, credential_id, credential_name)| FlowMetadata { - provider, - provider_id: None, - credential_id, - credential_name, - retry_count: 0, - client_info: ClientInfo::default(), - routing_info: RoutingInfo::default(), - injected_params: None, - context_usage_percentage: None, - }) - } - - /// 生成随机的 Flow ID - fn arb_flow_id() -> impl Strategy { - "[a-f0-9]{8}-[a-f0-9]{4}-[a-f0-9]{4}-[a-f0-9]{4}-[a-f0-9]{12}" - } - - // ======================================================================== - // 属性测试 - // ======================================================================== - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: llm-flow-monitor, Property 1: Flow 创建正确性** - /// **Validates: Requirements 1.1, 1.2** - /// - /// *对于任意* 有效的 API 请求,当 Flow_Monitor 创建新的 LLM_Flow 时, - /// 该 Flow 应该具有唯一的 ID、pending 状态,并且请求信息应该被正确提取和存储。 - #[test] - fn prop_flow_creation_correctness( - id in arb_flow_id(), - flow_type in arb_flow_type(), - request in arb_llm_request(), - metadata in arb_flow_metadata(), - ) { - // 创建 Flow - let flow = LLMFlow::new(id.clone(), flow_type.clone(), request.clone(), metadata.clone()); - - // 验证 ID 正确设置 - prop_assert_eq!(&flow.id, &id, "Flow ID 应该与输入 ID 相同"); - - // 验证初始状态为 Pending - prop_assert_eq!(flow.state, FlowState::Pending, "新创建的 Flow 状态应该是 Pending"); - - // 验证 FlowType 正确设置 - prop_assert_eq!(flow.flow_type, flow_type, "FlowType 应该正确设置"); - - // 验证请求信息正确存储 - prop_assert_eq!(flow.request.model, request.model, "模型名称应该正确存储"); - prop_assert_eq!(flow.request.method, request.method, "HTTP 方法应该正确存储"); - prop_assert_eq!(flow.request.path, request.path, "请求路径应该正确存储"); - prop_assert_eq!(flow.request.messages.len(), request.messages.len(), "消息列表长度应该一致"); - prop_assert_eq!(flow.request.system_prompt, request.system_prompt, "系统提示词应该正确存储"); - prop_assert_eq!(flow.request.parameters.stream, request.parameters.stream, "流式参数应该正确存储"); - - // 验证元数据正确存储 - prop_assert_eq!(flow.metadata.provider, metadata.provider, "Provider 类型应该正确存储"); - prop_assert_eq!(flow.metadata.credential_id, metadata.credential_id, "凭证 ID 应该正确存储"); - - // 验证响应和错误初始为空 - prop_assert!(flow.response.is_none(), "新创建的 Flow 响应应该为空"); - prop_assert!(flow.error.is_none(), "新创建的 Flow 错误应该为空"); - - // 验证时间戳已设置 - prop_assert!(flow.timestamps.created <= Utc::now(), "创建时间应该已设置"); - prop_assert!(flow.timestamps.request_start <= Utc::now(), "请求开始时间应该已设置"); - - // 验证标注初始为默认值 - prop_assert!(!flow.annotations.starred, "新创建的 Flow 不应该被收藏"); - prop_assert!(flow.annotations.tags.is_empty(), "新创建的 Flow 标签应该为空"); - } - - /// **Feature: llm-flow-monitor, Property 1b: Flow 序列化往返** - /// **Validates: Requirements 1.1, 1.2** - /// - /// *对于任意* 有效的 LLMFlow,序列化后再反序列化应该得到等价的对象。 - #[test] - fn prop_flow_serialization_roundtrip( - id in arb_flow_id(), - flow_type in arb_flow_type(), - request in arb_llm_request(), - metadata in arb_flow_metadata(), - ) { - let flow = LLMFlow::new(id, flow_type, request, metadata); - - // 序列化 - let json = serde_json::to_string(&flow).expect("序列化应该成功"); - - // 反序列化 - let deserialized: LLMFlow = serde_json::from_str(&json).expect("反序列化应该成功"); - - // 验证关键字段一致 - prop_assert_eq!(flow.id, deserialized.id, "ID 应该在往返后保持一致"); - prop_assert_eq!(flow.state, deserialized.state, "状态应该在往返后保持一致"); - prop_assert_eq!(flow.request.model, deserialized.request.model, "模型应该在往返后保持一致"); - prop_assert_eq!(flow.request.method, deserialized.request.method, "方法应该在往返后保持一致"); - prop_assert_eq!(flow.metadata.provider, deserialized.metadata.provider, "Provider 应该在往返后保持一致"); - } - - /// **Feature: llm-flow-monitor, Property 1c: 消息内容提取正确性** - /// **Validates: Requirements 1.2** - /// - /// *对于任意* 消息内容,get_all_text() 应该返回所有文本内容。 - #[test] - fn prop_message_content_text_extraction( - content in arb_message_content(), - ) { - let text = content.get_all_text(); - - match &content { - MessageContent::Text(s) => { - prop_assert_eq!(&text, s, "纯文本内容应该完整返回"); - } - MessageContent::MultiModal(parts) => { - // 验证所有文本部分都包含在结果中 - for part in parts { - if let ContentPart::Text { text: part_text } = part { - prop_assert!( - text.contains(part_text), - "多模态内容中的文本部分应该包含在结果中" - ); - } - } - } - } - } - - /// **Feature: llm-flow-monitor, Property 1d: 错误类型可重试判断** - /// **Validates: Requirements 1.8** - /// - /// *对于任意* 错误类型,is_retryable() 应该正确判断是否可重试。 - #[test] - fn prop_error_type_retryable_consistency( - status_code in 100u16..600u16, - ) { - let error_type = FlowErrorType::from_status_code(status_code); - let is_retryable = error_type.is_retryable(); - - // 验证可重试的错误类型 - match error_type { - FlowErrorType::Network - | FlowErrorType::Timeout - | FlowErrorType::RateLimit - | FlowErrorType::ServerError => { - prop_assert!(is_retryable, "{:?} 应该是可重试的", error_type); - } - FlowErrorType::Authentication - | FlowErrorType::BadRequest - | FlowErrorType::ContentFilter - | FlowErrorType::ModelUnavailable - | FlowErrorType::TokenLimitExceeded - | FlowErrorType::Cancelled - | FlowErrorType::Other => { - prop_assert!(!is_retryable, "{:?} 不应该是可重试的", error_type); - } - } - } - - /// **Feature: llm-flow-monitor, Property 1e: Token 使用量计算正确性** - /// **Validates: Requirements 1.9** - /// - /// *对于任意* Token 使用量,calculate_total() 应该正确计算总数。 - #[test] - fn prop_token_usage_total_calculation( - input_tokens in 0u32..100000u32, - output_tokens in 0u32..100000u32, - ) { - let mut usage = TokenUsage { - input_tokens, - output_tokens, - ..Default::default() - }; - - usage.calculate_total(); - - prop_assert_eq!( - usage.total_tokens, - input_tokens + output_tokens, - "总 Token 数应该等于输入 + 输出" - ); - } - - /// **Feature: llm-flow-monitor, Property 1f: 时间戳计算正确性** - /// **Validates: Requirements 1.9** - /// - /// *对于任意* 有效的时间戳序列,duration 和 ttfb 计算应该正确。 - #[test] - fn prop_timestamps_calculation( - ttfb_ms in 0i64..10000i64, - response_duration_ms in 0i64..100000i64, - ) { - let start = Utc::now(); - let response_start = start + chrono::Duration::milliseconds(ttfb_ms); - let end = response_start + chrono::Duration::milliseconds(response_duration_ms); - - let mut timestamps = FlowTimestamps { - created: start, - request_start: start, - request_end: Some(start + chrono::Duration::milliseconds(10)), - response_start: Some(response_start), - response_end: Some(end), - duration_ms: 0, - ttfb_ms: None, - }; - - timestamps.calculate_duration(); - timestamps.calculate_ttfb(); - - // 验证 TTFB 计算 - prop_assert_eq!( - timestamps.ttfb_ms, - Some(ttfb_ms as u64), - "TTFB 应该正确计算" - ); - - // 验证总耗时计算 - let expected_duration = ttfb_ms + response_duration_ms; - prop_assert_eq!( - timestamps.duration_ms, - expected_duration as u64, - "总耗时应该正确计算" - ); - } - } -} diff --git a/src-tauri/src/flow_monitor/monitor.rs b/src-tauri/src/flow_monitor/monitor.rs deleted file mode 100644 index ccbb2a108..000000000 --- a/src-tauri/src/flow_monitor/monitor.rs +++ /dev/null @@ -1,2787 +0,0 @@ -//! Flow 核心监控服务 -//! -//! 该模块实现 LLM Flow 的核心监控功能,包括: -//! - Flow 生命周期管理(创建、更新、完成、失败) -//! - 流式响应处理 -//! - 实时事件发送 -//! - 标注管理 -//! - 阈值检测(延迟、Token 使用量) -//! - 请求速率计算 - -#![allow(dead_code)] - -use chrono::{DateTime, Duration, Utc}; -use serde::{Deserialize, Serialize}; -use std::collections::{HashMap, VecDeque}; -use std::sync::Arc; -use tokio::sync::{broadcast, RwLock}; -use uuid::Uuid; - -use super::file_store::FlowFileStore; -use super::memory_store::FlowMemoryStore; -use super::models::{ - FlowAnnotations, FlowError, FlowMetadata, FlowState, FlowType, LLMFlow, LLMRequest, - LLMResponse, TokenUsage, -}; -use super::stream_rebuilder::{StreamFormat, StreamRebuilder}; - -// ============================================================================ -// 配置结构 -// ============================================================================ - -/// Flow 监控配置 -/// -/// 控制 Flow Monitor 的行为,包括启用/禁用、缓存大小、持久化等。 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct FlowMonitorConfig { - /// 是否启用监控 - #[serde(default = "default_enabled")] - pub enabled: bool, - /// 最大内存 Flow 数量 - #[serde(default = "default_max_memory_flows")] - pub max_memory_flows: usize, - /// 是否持久化到文件 - #[serde(default = "default_persist_to_file")] - pub persist_to_file: bool, - /// 保留天数 - #[serde(default = "default_retention_days")] - pub retention_days: u32, - /// 是否保存原始流式 chunks - #[serde(default)] - pub save_stream_chunks: bool, - /// 最大请求体大小(字节) - #[serde(default = "default_max_request_body_size")] - pub max_request_body_size: usize, - /// 最大响应体大小(字节) - #[serde(default = "default_max_response_body_size")] - pub max_response_body_size: usize, - /// 是否保存图片内容 - #[serde(default)] - pub save_image_content: bool, - /// 缩略图大小 - #[serde(default = "default_thumbnail_size")] - pub thumbnail_size: (u32, u32), - /// 采样率(0.0-1.0,1.0 表示全部采样) - #[serde(default = "default_sampling_rate")] - pub sampling_rate: f32, - /// 排除的模型列表(支持通配符) - #[serde(default)] - pub excluded_models: Vec, - /// 排除的路径列表(支持通配符) - #[serde(default)] - pub excluded_paths: Vec, -} - -fn default_enabled() -> bool { - true -} - -fn default_max_memory_flows() -> usize { - 1000 -} - -fn default_persist_to_file() -> bool { - true -} - -fn default_retention_days() -> u32 { - 7 -} - -fn default_max_request_body_size() -> usize { - 10 * 1024 * 1024 // 10MB -} - -fn default_max_response_body_size() -> usize { - 10 * 1024 * 1024 // 10MB -} - -fn default_thumbnail_size() -> (u32, u32) { - (128, 128) -} - -fn default_sampling_rate() -> f32 { - 1.0 -} - -impl Default for FlowMonitorConfig { - fn default() -> Self { - Self { - enabled: default_enabled(), - max_memory_flows: default_max_memory_flows(), - persist_to_file: default_persist_to_file(), - retention_days: default_retention_days(), - save_stream_chunks: false, - max_request_body_size: default_max_request_body_size(), - max_response_body_size: default_max_response_body_size(), - save_image_content: false, - thumbnail_size: default_thumbnail_size(), - sampling_rate: default_sampling_rate(), - excluded_models: Vec::new(), - excluded_paths: Vec::new(), - } - } -} - -impl FlowMonitorConfig { - /// 检查是否应该监控该请求 - pub fn should_monitor(&self, model: &str, path: &str) -> bool { - if !self.enabled { - return false; - } - - // 检查采样率 - if self.sampling_rate < 1.0 { - let random: f32 = rand::random(); - if random > self.sampling_rate { - return false; - } - } - - // 检查排除的模型 - for pattern in &self.excluded_models { - if Self::match_pattern(pattern, model) { - return false; - } - } - - // 检查排除的路径 - for pattern in &self.excluded_paths { - if Self::match_pattern(pattern, path) { - return false; - } - } - - true - } - - /// 模式匹配(支持 * 通配符) - fn match_pattern(pattern: &str, text: &str) -> bool { - if pattern == "*" { - return true; - } - - if pattern.contains('*') { - let parts: Vec<&str> = pattern.split('*').collect(); - let mut pos = 0; - let text_lower = text.to_lowercase(); - - for (i, part) in parts.iter().enumerate() { - if part.is_empty() { - continue; - } - - let part_lower = part.to_lowercase(); - if let Some(found_pos) = text_lower[pos..].find(&part_lower) { - if i == 0 && found_pos != 0 { - return false; - } - pos += found_pos + part.len(); - } else { - return false; - } - } - - if !pattern.ends_with('*') && pos != text.len() { - return false; - } - - true - } else { - text.to_lowercase() == pattern.to_lowercase() - } - } -} - -// ============================================================================ -// 阈值配置 -// ============================================================================ - -/// 阈值配置 -/// -/// 用于配置延迟和 Token 使用量的警告阈值。 -/// -/// **Validates: Requirements 10.3, 10.4** -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ThresholdConfig { - /// 是否启用阈值检测 - #[serde(default = "default_threshold_enabled")] - pub enabled: bool, - /// 延迟阈值(毫秒) - #[serde(default = "default_latency_threshold")] - pub latency_threshold_ms: u64, - /// Token 使用量阈值 - #[serde(default = "default_token_threshold")] - pub token_threshold: u32, - /// 输入 Token 阈值(可选) - #[serde(skip_serializing_if = "Option::is_none")] - pub input_token_threshold: Option, - /// 输出 Token 阈值(可选) - #[serde(skip_serializing_if = "Option::is_none")] - pub output_token_threshold: Option, -} - -fn default_threshold_enabled() -> bool { - true -} - -fn default_latency_threshold() -> u64 { - 5000 // 5 秒 -} - -fn default_token_threshold() -> u32 { - 10000 -} - -impl Default for ThresholdConfig { - fn default() -> Self { - Self { - enabled: default_threshold_enabled(), - latency_threshold_ms: default_latency_threshold(), - token_threshold: default_token_threshold(), - input_token_threshold: None, - output_token_threshold: None, - } - } -} - -/// 阈值检测结果 -/// -/// 表示 Flow 是否超过了配置的阈值。 -#[derive(Debug, Clone, Serialize, Deserialize, Default)] -pub struct ThresholdCheckResult { - /// 是否超过延迟阈值 - pub latency_exceeded: bool, - /// 是否超过 Token 阈值 - pub token_exceeded: bool, - /// 是否超过输入 Token 阈值 - pub input_token_exceeded: bool, - /// 是否超过输出 Token 阈值 - pub output_token_exceeded: bool, - /// 实际延迟(毫秒) - pub actual_latency_ms: u64, - /// 实际 Token 使用量 - pub actual_tokens: u32, - /// 实际输入 Token - pub actual_input_tokens: u32, - /// 实际输出 Token - pub actual_output_tokens: u32, -} - -impl ThresholdCheckResult { - /// 检查是否有任何阈值被超过 - pub fn any_exceeded(&self) -> bool { - self.latency_exceeded - || self.token_exceeded - || self.input_token_exceeded - || self.output_token_exceeded - } -} - -// ============================================================================ -// 通知配置 -// ============================================================================ - -/// 通知类型 -/// -/// **Validates: Requirements 10.1, 10.2** -#[derive(Debug, Clone, Serialize, Deserialize)] -pub enum NotificationType { - /// 新 Flow 通知 - NewFlow, - /// 错误 Flow 通知 - ErrorFlow, - /// 延迟阈值警告 - LatencyWarning, - /// Token 阈值警告 - TokenWarning, -} - -/// 通知配置 -/// -/// 用于配置各种通知的启用状态和行为。 -/// -/// **Validates: Requirements 10.1, 10.2** -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct NotificationConfig { - /// 是否启用通知 - #[serde(default = "default_notification_enabled")] - pub enabled: bool, - /// 新 Flow 通知配置 - #[serde(default)] - pub new_flow: NotificationSettings, - /// 错误 Flow 通知配置 - #[serde(default = "default_error_notification")] - pub error_flow: NotificationSettings, - /// 延迟警告通知配置 - #[serde(default = "default_latency_warning")] - pub latency_warning: NotificationSettings, - /// Token 警告通知配置 - #[serde(default = "default_token_warning")] - pub token_warning: NotificationSettings, -} - -/// 通知设置 -#[derive(Debug, Clone, Serialize, Deserialize, Default)] -pub struct NotificationSettings { - /// 是否启用 - pub enabled: bool, - /// 是否显示桌面通知 - pub desktop: bool, - /// 是否播放声音 - pub sound: bool, - /// 声音文件路径(可选) - pub sound_file: Option, -} - -fn default_notification_enabled() -> bool { - true -} - -fn default_error_notification() -> NotificationSettings { - NotificationSettings { - enabled: true, - desktop: true, - sound: true, - sound_file: None, - } -} - -fn default_latency_warning() -> NotificationSettings { - NotificationSettings { - enabled: true, - desktop: false, - sound: false, - sound_file: None, - } -} - -fn default_token_warning() -> NotificationSettings { - NotificationSettings { - enabled: true, - desktop: false, - sound: false, - sound_file: None, - } -} - -impl Default for NotificationConfig { - fn default() -> Self { - Self { - enabled: default_notification_enabled(), - new_flow: NotificationSettings::default(), - error_flow: default_error_notification(), - latency_warning: default_latency_warning(), - token_warning: default_token_warning(), - } - } -} - -/// 通知事件 -/// -/// 表示需要发送的通知。 -/// -/// **Validates: Requirements 10.1, 10.2, 10.3, 10.4** -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct NotificationEvent { - /// 通知类型 - pub notification_type: NotificationType, - /// 通知标题 - pub title: String, - /// 通知内容 - pub message: String, - /// 关联的 Flow ID - pub flow_id: String, - /// 通知时间 - pub timestamp: DateTime, - /// 是否需要桌面通知 - pub desktop: bool, - /// 是否需要声音 - pub sound: bool, - /// 声音文件路径 - pub sound_file: Option, -} - -impl NotificationEvent { - /// 创建新 Flow 通知 - pub fn new_flow(flow_id: String, model: String, settings: &NotificationSettings) -> Self { - Self { - notification_type: NotificationType::NewFlow, - title: "新的 LLM 请求".to_string(), - message: format!("模型: {model}"), - flow_id, - timestamp: Utc::now(), - desktop: settings.desktop, - sound: settings.sound, - sound_file: settings.sound_file.clone(), - } - } - - /// 创建错误 Flow 通知 - pub fn error_flow( - flow_id: String, - model: String, - error: String, - settings: &NotificationSettings, - ) -> Self { - Self { - notification_type: NotificationType::ErrorFlow, - title: "LLM 请求失败".to_string(), - message: format!("模型: {model}, 错误: {error}"), - flow_id, - timestamp: Utc::now(), - desktop: settings.desktop, - sound: settings.sound, - sound_file: settings.sound_file.clone(), - } - } - - /// 创建延迟警告通知 - pub fn latency_warning( - flow_id: String, - model: String, - actual_ms: u64, - threshold_ms: u64, - settings: &NotificationSettings, - ) -> Self { - Self { - notification_type: NotificationType::LatencyWarning, - title: "延迟警告".to_string(), - message: format!("模型: {model}, 延迟: {actual_ms}ms (阈值: {threshold_ms}ms)"), - flow_id, - timestamp: Utc::now(), - desktop: settings.desktop, - sound: settings.sound, - sound_file: settings.sound_file.clone(), - } - } - - /// 创建 Token 警告通知 - pub fn token_warning( - flow_id: String, - model: String, - actual_tokens: u32, - threshold_tokens: u32, - settings: &NotificationSettings, - ) -> Self { - Self { - notification_type: NotificationType::TokenWarning, - title: "Token 使用警告".to_string(), - message: format!("模型: {model}, Token: {actual_tokens} (阈值: {threshold_tokens})"), - flow_id, - timestamp: Utc::now(), - desktop: settings.desktop, - sound: settings.sound, - sound_file: settings.sound_file.clone(), - } - } -} - -// ============================================================================ -// 请求速率追踪器 -// ============================================================================ - -/// 请求速率追踪器 -/// -/// 用于计算指定时间窗口内的请求速率。 -/// -/// **Validates: Requirements 10.7** -#[derive(Debug)] -pub struct RequestRateTracker { - /// 请求时间戳队列 - timestamps: VecDeque>, - /// 时间窗口(秒) - window_seconds: i64, -} - -impl RequestRateTracker { - /// 创建新的请求速率追踪器 - /// - /// # Arguments - /// * `window_seconds` - 时间窗口(秒) - pub fn new(window_seconds: i64) -> Self { - Self { - timestamps: VecDeque::new(), - window_seconds, - } - } - - /// 记录一个新请求 - pub fn record_request(&mut self) { - self.record_request_at(Utc::now()); - } - - /// 在指定时间记录一个新请求 - pub fn record_request_at(&mut self, timestamp: DateTime) { - self.timestamps.push_back(timestamp); - self.cleanup_old_entries(timestamp); - } - - /// 清理过期的条目 - fn cleanup_old_entries(&mut self, now: DateTime) { - let cutoff = now - Duration::seconds(self.window_seconds); - while let Some(front) = self.timestamps.front() { - if *front < cutoff { - self.timestamps.pop_front(); - } else { - break; - } - } - } - - /// 获取当前请求速率(每秒) - pub fn get_rate(&self) -> f64 { - self.get_rate_at(Utc::now()) - } - - /// 获取指定时间点的请求速率(每秒) - pub fn get_rate_at(&self, now: DateTime) -> f64 { - let cutoff = now - Duration::seconds(self.window_seconds); - let count = self.timestamps.iter().filter(|&&ts| ts >= cutoff).count(); - - if self.window_seconds > 0 { - count as f64 / self.window_seconds as f64 - } else { - 0.0 - } - } - - /// 获取时间窗口内的请求数量 - pub fn get_count(&self) -> usize { - self.get_count_at(Utc::now()) - } - - /// 获取指定时间点的时间窗口内的请求数量 - pub fn get_count_at(&self, now: DateTime) -> usize { - let cutoff = now - Duration::seconds(self.window_seconds); - self.timestamps.iter().filter(|&&ts| ts >= cutoff).count() - } - - /// 获取时间窗口(秒) - pub fn window_seconds(&self) -> i64 { - self.window_seconds - } - - /// 设置时间窗口(秒) - pub fn set_window_seconds(&mut self, window_seconds: i64) { - self.window_seconds = window_seconds; - self.cleanup_old_entries(Utc::now()); - } - - /// 清空所有记录 - pub fn clear(&mut self) { - self.timestamps.clear(); - } -} - -impl Default for RequestRateTracker { - fn default() -> Self { - Self::new(60) // 默认 60 秒窗口 - } -} - -impl Clone for RequestRateTracker { - fn clone(&self) -> Self { - Self { - timestamps: self.timestamps.clone(), - window_seconds: self.window_seconds, - } - } -} - -// ============================================================================ -// 事件类型 -// ============================================================================ - -/// Flow 摘要信息 -/// -/// 用于事件通知,包含 Flow 的关键信息。 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct FlowSummary { - /// Flow ID - pub id: String, - /// 流类型 - pub flow_type: FlowType, - /// 模型名称 - pub model: String, - /// 提供商 - pub provider: String, - /// 状态 - pub state: FlowState, - /// 创建时间 - pub created_at: DateTime, - /// 耗时(毫秒) - pub duration_ms: Option, - /// Token 使用量 - pub usage: Option, - /// 是否有错误 - pub has_error: bool, - /// 是否有工具调用 - pub has_tool_calls: bool, - /// 是否有思维链 - pub has_thinking: bool, -} - -impl From<&LLMFlow> for FlowSummary { - fn from(flow: &LLMFlow) -> Self { - Self { - id: flow.id.clone(), - flow_type: flow.flow_type.clone(), - model: flow.request.model.clone(), - provider: format!("{:?}", flow.metadata.provider), - state: flow.state.clone(), - created_at: flow.timestamps.created, - duration_ms: if flow.timestamps.duration_ms > 0 { - Some(flow.timestamps.duration_ms) - } else { - None - }, - usage: flow.response.as_ref().map(|r| r.usage.clone()), - has_error: flow.error.is_some(), - has_tool_calls: flow - .response - .as_ref() - .is_some_and(|r| !r.tool_calls.is_empty()), - has_thinking: flow.response.as_ref().is_some_and(|r| r.thinking.is_some()), - } - } -} - -/// Flow 更新信息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct FlowUpdate { - /// 新状态 - pub state: Option, - /// 内容增量 - pub content_delta: Option, - /// 当前内容长度 - pub content_length: Option, - /// 当前 chunk 数量 - pub chunk_count: Option, -} - -/// 实时 Flow 事件 -#[derive(Debug, Clone, Serialize)] -#[serde(tag = "type")] -pub enum FlowEvent { - /// Flow 开始 - FlowStarted { flow: FlowSummary }, - /// Flow 更新 - FlowUpdated { id: String, update: FlowUpdate }, - /// Flow 完成 - FlowCompleted { id: String, summary: FlowSummary }, - /// Flow 失败 - FlowFailed { id: String, error: FlowError }, - /// 阈值警告 - /// - /// **Validates: Requirements 10.3, 10.4** - ThresholdWarning { - id: String, - result: ThresholdCheckResult, - }, - /// 通知事件 - /// - /// **Validates: Requirements 10.1, 10.2, 10.3, 10.4** - Notification { notification: NotificationEvent }, - /// 请求速率更新 - /// - /// **Validates: Requirements 10.7** - RequestRateUpdate { rate: f64, count: usize }, -} - -// ============================================================================ -// 活跃 Flow 状态 -// ============================================================================ - -/// 活跃 Flow 状态 -/// -/// 用于跟踪正在进行中的 Flow,包括流式响应重建器。 -struct ActiveFlow { - /// Flow 数据 - flow: LLMFlow, - /// 流式响应重建器(如果是流式响应) - stream_rebuilder: Option, - /// 请求开始时间 - request_start: DateTime, -} - -// ============================================================================ -// 核心监控服务 -// ============================================================================ - -/// Flow 监控服务 -/// -/// 负责捕获和管理 LLM Flow 的核心服务。 -pub struct FlowMonitor { - /// 配置 - config: RwLock, - /// 内存存储 - memory_store: Arc>, - /// 文件存储(可选) - file_store: Option>, - /// 活跃 Flow(正在进行中的请求) - active_flows: RwLock>, - /// 事件发送器 - event_sender: broadcast::Sender, - /// 阈值配置 - threshold_config: RwLock, - /// 请求速率追踪器 - rate_tracker: RwLock, - /// 通知配置 - notification_config: RwLock, -} - -impl FlowMonitor { - /// 创建新的 Flow 监控服务 - /// - /// # 参数 - /// - `config`: 监控配置 - /// - `file_store`: 文件存储(可选) - pub fn new(config: FlowMonitorConfig, file_store: Option>) -> Self { - let memory_store = Arc::new(RwLock::new(FlowMemoryStore::new(config.max_memory_flows))); - let (event_sender, _) = broadcast::channel(1000); - - Self { - config: RwLock::new(config), - memory_store, - file_store, - active_flows: RwLock::new(HashMap::new()), - event_sender, - threshold_config: RwLock::new(ThresholdConfig::default()), - rate_tracker: RwLock::new(RequestRateTracker::default()), - notification_config: RwLock::new(NotificationConfig::default()), - } - } - - /// 创建带通知配置的 Flow 监控服务 - /// - /// # 参数 - /// - `config`: 监控配置 - /// - `file_store`: 文件存储(可选) - /// - `threshold_config`: 阈值配置 - /// - `notification_config`: 通知配置 - pub fn with_notification_config( - config: FlowMonitorConfig, - file_store: Option>, - threshold_config: ThresholdConfig, - notification_config: NotificationConfig, - ) -> Self { - let memory_store = Arc::new(RwLock::new(FlowMemoryStore::new(config.max_memory_flows))); - let (event_sender, _) = broadcast::channel(1000); - - Self { - config: RwLock::new(config), - memory_store, - file_store, - active_flows: RwLock::new(HashMap::new()), - event_sender, - threshold_config: RwLock::new(threshold_config), - rate_tracker: RwLock::new(RequestRateTracker::default()), - notification_config: RwLock::new(notification_config), - } - } - - /// 创建带完整配置的 Flow 监控服务 - /// - /// # 参数 - /// - `config`: 监控配置 - /// - `file_store`: 文件存储(可选) - /// - `threshold_config`: 阈值配置 - /// - `notification_config`: 通知配置 - pub fn with_full_config( - config: FlowMonitorConfig, - file_store: Option>, - threshold_config: ThresholdConfig, - notification_config: NotificationConfig, - ) -> Self { - let memory_store = Arc::new(RwLock::new(FlowMemoryStore::new(config.max_memory_flows))); - let (event_sender, _) = broadcast::channel(1000); - - Self { - config: RwLock::new(config), - memory_store, - file_store, - active_flows: RwLock::new(HashMap::new()), - event_sender, - threshold_config: RwLock::new(threshold_config), - rate_tracker: RwLock::new(RequestRateTracker::default()), - notification_config: RwLock::new(notification_config), - } - } - - /// 获取内存存储的引用 - pub fn memory_store(&self) -> Arc> { - self.memory_store.clone() - } - - /// 获取文件存储的引用 - pub fn file_store(&self) -> Option> { - self.file_store.clone() - } - - /// 获取当前配置 - pub async fn config(&self) -> FlowMonitorConfig { - self.config.read().await.clone() - } - - /// 更新配置 - pub async fn update_config(&self, config: FlowMonitorConfig) { - let mut current = self.config.write().await; - - // 如果缓存大小改变,需要调整内存存储 - if current.max_memory_flows != config.max_memory_flows { - // 创建新的内存存储(旧数据会丢失) - // 实际应用中可能需要更复杂的迁移逻辑 - let mut store = self.memory_store.write().await; - *store = FlowMemoryStore::new(config.max_memory_flows); - } - - *current = config; - } - - /// 获取阈值配置 - /// - /// **Validates: Requirements 10.3, 10.4** - pub async fn threshold_config(&self) -> ThresholdConfig { - self.threshold_config.read().await.clone() - } - - /// 更新阈值配置 - /// - /// **Validates: Requirements 10.3, 10.4** - pub async fn update_threshold_config(&self, config: ThresholdConfig) { - let mut current = self.threshold_config.write().await; - *current = config; - } - - /// 获取当前请求速率(每秒) - /// - /// **Validates: Requirements 10.7** - pub async fn get_request_rate(&self) -> f64 { - self.rate_tracker.read().await.get_rate() - } - - /// 获取时间窗口内的请求数量 - /// - /// **Validates: Requirements 10.7** - pub async fn get_request_count(&self) -> usize { - self.rate_tracker.read().await.get_count() - } - - /// 设置请求速率追踪器的时间窗口 - /// - /// **Validates: Requirements 10.7** - pub async fn set_rate_window(&self, window_seconds: i64) { - self.rate_tracker - .write() - .await - .set_window_seconds(window_seconds); - } - - /// 获取通知配置 - /// - /// **Validates: Requirements 10.1, 10.2** - pub async fn notification_config(&self) -> NotificationConfig { - self.notification_config.read().await.clone() - } - - /// 更新通知配置 - /// - /// **Validates: Requirements 10.1, 10.2** - pub async fn update_notification_config(&self, config: NotificationConfig) { - let mut current = self.notification_config.write().await; - *current = config; - } - - /// 触发通知 - /// - /// **Validates: Requirements 10.1, 10.2, 10.3, 10.4** - /// - /// # Arguments - /// * `notification` - 通知事件 - async fn trigger_notification(&self, notification: NotificationEvent) { - let config = self.notification_config.read().await; - - if !config.enabled { - return; - } - - // 发送通知事件 - let _ = self.event_sender.send(FlowEvent::Notification { - notification: notification.clone(), - }); - } - - /// 检查并触发新 Flow 通知 - /// - /// **Validates: Requirements 10.1** - async fn check_new_flow_notification(&self, flow: &LLMFlow) { - let config = self.notification_config.read().await; - - if config.new_flow.enabled { - let notification = NotificationEvent::new_flow( - flow.id.clone(), - flow.request.model.clone(), - &config.new_flow, - ); - drop(config); - self.trigger_notification(notification).await; - } - } - - /// 检查并触发错误 Flow 通知 - /// - /// **Validates: Requirements 10.2** - async fn check_error_flow_notification(&self, flow: &LLMFlow, error: &FlowError) { - let config = self.notification_config.read().await; - - if config.error_flow.enabled { - let notification = NotificationEvent::error_flow( - flow.id.clone(), - flow.request.model.clone(), - error.message.clone(), - &config.error_flow, - ); - drop(config); - self.trigger_notification(notification).await; - } - } - - /// 检查并触发阈值警告通知 - /// - /// **Validates: Requirements 10.3, 10.4** - async fn check_threshold_notifications(&self, flow: &LLMFlow, result: &ThresholdCheckResult) { - let config = self.notification_config.read().await; - let threshold_config = self.threshold_config.read().await; - - // 延迟警告通知 - if result.latency_exceeded && config.latency_warning.enabled { - let notification = NotificationEvent::latency_warning( - flow.id.clone(), - flow.request.model.clone(), - result.actual_latency_ms, - threshold_config.latency_threshold_ms, - &config.latency_warning, - ); - drop(config); - drop(threshold_config); - self.trigger_notification(notification).await; - return; - } - - // Token 警告通知 - if result.token_exceeded && config.token_warning.enabled { - let notification = NotificationEvent::token_warning( - flow.id.clone(), - flow.request.model.clone(), - result.actual_tokens, - threshold_config.token_threshold, - &config.token_warning, - ); - drop(config); - drop(threshold_config); - self.trigger_notification(notification).await; - } - } - - /// 发送请求速率更新事件 - /// - /// **Validates: Requirements 10.7** - async fn send_rate_update(&self) { - let tracker = self.rate_tracker.read().await; - let rate = tracker.get_rate(); - let count = tracker.get_count(); - drop(tracker); - - let _ = self - .event_sender - .send(FlowEvent::RequestRateUpdate { rate, count }); - } - - /// 订阅实时事件 - pub fn subscribe(&self) -> broadcast::Receiver { - self.event_sender.subscribe() - } - - /// 开始捕获一个新的 Flow - /// - /// # 参数 - /// - `request`: LLM 请求 - /// - `metadata`: Flow 元数据 - /// - /// # 返回 - /// - `Some(flow_id)`: 成功创建 Flow,返回 Flow ID - /// - `None`: 根据配置跳过监控 - pub async fn start_flow(&self, request: LLMRequest, metadata: FlowMetadata) -> Option { - let config = self.config.read().await; - - // 检查是否应该监控 - if !config.should_monitor(&request.model, &request.path) { - eprintln!( - "[FLOW_MONITOR] 跳过监控: model={}, path={}", - request.model, request.path - ); - return None; - } - - // 记录请求到速率追踪器 - { - let mut tracker = self.rate_tracker.write().await; - tracker.record_request(); - } - - // 生成唯一 ID - let flow_id = Uuid::new_v4().to_string(); - - eprintln!( - "[FLOW_MONITOR] 创建新 Flow: id={}, model={}, provider={:?}", - flow_id, request.model, metadata.provider - ); - - // 确定 Flow 类型 - let flow_type = Self::determine_flow_type(&request.path); - - // 创建 Flow - let flow = LLMFlow::new(flow_id.clone(), flow_type, request.clone(), metadata); - - // 创建活跃 Flow 状态 - let active_flow = ActiveFlow { - flow: flow.clone(), - stream_rebuilder: None, - request_start: Utc::now(), - }; - - // 添加到活跃 Flow - { - let mut active = self.active_flows.write().await; - active.insert(flow_id.clone(), active_flow); - eprintln!("[FLOW_MONITOR] 活跃 Flow 数量: {}", active.len()); - } - - // 发送事件 - let summary = FlowSummary::from(&flow); - let _ = self - .event_sender - .send(FlowEvent::FlowStarted { flow: summary }); - - // 检查新 Flow 通知 - self.check_new_flow_notification(&flow).await; - - // 发送请求速率更新 - self.send_rate_update().await; - - Some(flow_id) - } - - /// 根据路径确定 Flow 类型 - fn determine_flow_type(path: &str) -> FlowType { - let path_lower = path.to_lowercase(); - - if path_lower.contains("/chat/completions") { - FlowType::ChatCompletions - } else if path_lower.contains("/messages") { - FlowType::AnthropicMessages - } else if path_lower.contains(":generatecontent") || path_lower.contains("/generate") { - FlowType::GeminiGenerateContent - } else if path_lower.contains("/embeddings") { - FlowType::Embeddings - } else { - FlowType::Other(path.to_string()) - } - } - - /// 设置 Flow 为流式模式 - /// - /// # 参数 - /// - `flow_id`: Flow ID - /// - `format`: 流式响应格式 - pub async fn set_streaming(&self, flow_id: &str, format: StreamFormat) { - let config = self.config.read().await; - let save_chunks = config.save_stream_chunks; - drop(config); - - let mut active = self.active_flows.write().await; - if let Some(active_flow) = active.get_mut(flow_id) { - active_flow.flow.state = FlowState::Streaming; - active_flow.stream_rebuilder = - Some(StreamRebuilder::new(format).with_save_raw_chunks(save_chunks)); - - // 发送更新事件 - let _ = self.event_sender.send(FlowEvent::FlowUpdated { - id: flow_id.to_string(), - update: FlowUpdate { - state: Some(FlowState::Streaming), - content_delta: None, - content_length: None, - chunk_count: None, - }, - }); - } - } - - /// 处理流式 chunk - /// - /// # 参数 - /// - `flow_id`: Flow ID - /// - `event`: SSE 事件类型(可选) - /// - `data`: SSE 数据内容 - pub async fn process_chunk(&self, flow_id: &str, event: Option<&str>, data: &str) { - let mut active = self.active_flows.write().await; - if let Some(active_flow) = active.get_mut(flow_id) { - if let Some(ref mut rebuilder) = active_flow.stream_rebuilder { - // 处理 chunk - if let Err(e) = rebuilder.process_event(event, data) { - tracing::warn!("处理流式 chunk 失败: {}", e); - } - - // 发送更新事件(可选,根据需要调整频率) - // 这里简化处理,每个 chunk 都发送事件 - // 实际应用中可能需要节流 - } - } - } - - /// 完成 Flow - /// - /// # 参数 - /// - `flow_id`: Flow ID - /// - `response`: LLM 响应(如果是非流式响应) - pub async fn complete_flow(&self, flow_id: &str, response: Option) { - eprintln!( - "[FLOW_MONITOR] 准备完成 Flow: id={}, has_response={}", - flow_id, - response.is_some() - ); - - let mut active = self.active_flows.write().await; - - if let Some(mut active_flow) = active.remove(flow_id) { - let now = Utc::now(); - - // 如果有流式重建器,使用重建的响应 - let final_response = if let Some(rebuilder) = active_flow.stream_rebuilder.take() { - Some(rebuilder.finish()) - } else { - response - }; - - // 更新 Flow - active_flow.flow.response = final_response.clone(); - active_flow.flow.state = FlowState::Completed; - active_flow.flow.timestamps.response_end = Some(now); - active_flow.flow.timestamps.calculate_duration(); - active_flow.flow.timestamps.calculate_ttfb(); - - eprintln!( - "[FLOW_MONITOR] Flow 状态更新: id={}, state={:?}, duration_ms={}", - flow_id, active_flow.flow.state, active_flow.flow.timestamps.duration_ms - ); - - // 检查阈值 - let threshold_result = self.check_threshold(&active_flow.flow).await; - - // 保存到内存存储 - { - let mut store = self.memory_store.write().await; - store.add(active_flow.flow.clone()); - eprintln!( - "[FLOW_MONITOR] 已保存到内存存储: id={}, 内存中 Flow 数量={}", - flow_id, - store.len() - ); - } - - // 保存到文件存储 - if let Some(ref file_store) = self.file_store { - if let Err(e) = file_store.write(&active_flow.flow) { - tracing::error!("保存 Flow 到文件失败: {}", e); - eprintln!("[FLOW_MONITOR] 保存到文件失败: id={flow_id}, error={e}"); - } else { - eprintln!("[FLOW_MONITOR] 已保存到文件存储: id={flow_id}"); - } - } else { - eprintln!("[FLOW_MONITOR] 文件存储未启用"); - } - - // 发送完成事件 - let summary = FlowSummary::from(&active_flow.flow); - let _ = self.event_sender.send(FlowEvent::FlowCompleted { - id: flow_id.to_string(), - summary, - }); - - // 如果超过阈值,发送警告事件 - if threshold_result.any_exceeded() { - let _ = self.event_sender.send(FlowEvent::ThresholdWarning { - id: flow_id.to_string(), - result: threshold_result.clone(), - }); - - // 检查并触发阈值通知 - self.check_threshold_notifications(&active_flow.flow, &threshold_result) - .await; - } - - eprintln!("[FLOW_MONITOR] Flow 完成处理完毕: id={flow_id}"); - } else { - eprintln!("[FLOW_MONITOR] 警告: 未找到活跃 Flow: id={flow_id}"); - } - } - - /// 标记 Flow 失败 - /// - /// # 参数 - /// - `flow_id`: Flow ID - /// - `error`: 错误信息 - pub async fn fail_flow(&self, flow_id: &str, error: FlowError) { - let mut active = self.active_flows.write().await; - - if let Some(mut active_flow) = active.remove(flow_id) { - let now = Utc::now(); - - // 更新 Flow - active_flow.flow.error = Some(error.clone()); - active_flow.flow.state = FlowState::Failed; - active_flow.flow.timestamps.response_end = Some(now); - active_flow.flow.timestamps.calculate_duration(); - - // 保存到内存存储 - { - let mut store = self.memory_store.write().await; - store.add(active_flow.flow.clone()); - } - - // 保存到文件存储 - if let Some(ref file_store) = self.file_store { - if let Err(e) = file_store.write(&active_flow.flow) { - tracing::error!("保存 Flow 到文件失败: {}", e); - } - } - - // 发送失败事件 - let _ = self.event_sender.send(FlowEvent::FlowFailed { - id: flow_id.to_string(), - error: error.clone(), - }); - - // 检查错误 Flow 通知 - self.check_error_flow_notification(&active_flow.flow, &error) - .await; - } - } - - /// 取消 Flow - /// - /// # 参数 - /// - `flow_id`: Flow ID - pub async fn cancel_flow(&self, flow_id: &str) { - let mut active = self.active_flows.write().await; - - if let Some(mut active_flow) = active.remove(flow_id) { - let now = Utc::now(); - - // 更新 Flow - active_flow.flow.state = FlowState::Cancelled; - active_flow.flow.timestamps.response_end = Some(now); - active_flow.flow.timestamps.calculate_duration(); - - // 保存到内存存储 - { - let mut store = self.memory_store.write().await; - store.add(active_flow.flow.clone()); - } - - // 保存到文件存储 - if let Some(ref file_store) = self.file_store { - if let Err(e) = file_store.write(&active_flow.flow) { - tracing::error!("保存 Flow 到文件失败: {}", e); - } - } - } - } - - /// 更新 Flow 标注 - /// - /// # 参数 - /// - `flow_id`: Flow ID - /// - `annotations`: 新的标注信息 - /// - /// # 返回 - /// - `true`: 更新成功 - /// - `false`: Flow 不存在 - pub async fn update_annotations(&self, flow_id: &str, annotations: FlowAnnotations) -> bool { - // 先尝试更新内存中的 Flow - let updated = { - let store = self.memory_store.read().await; - store.update(flow_id, |flow| { - flow.annotations = annotations.clone(); - }) - }; - - // 如果内存中存在,同时更新文件存储的索引 - if updated { - if let Some(ref file_store) = self.file_store { - if let Err(e) = file_store.update_annotations(flow_id, &annotations) { - tracing::error!("更新文件存储标注失败: {}", e); - } - } - } - - updated - } - - /// 收藏/取消收藏 Flow - pub async fn toggle_starred(&self, flow_id: &str) -> bool { - let store = self.memory_store.read().await; - store.update(flow_id, |flow| { - flow.annotations.starred = !flow.annotations.starred; - }) - } - - /// 添加评论 - pub async fn add_comment(&self, flow_id: &str, comment: String) -> bool { - let store = self.memory_store.read().await; - store.update(flow_id, |flow| { - flow.annotations.comment = Some(comment); - }) - } - - /// 添加标签 - pub async fn add_tag(&self, flow_id: &str, tag: String) -> bool { - let store = self.memory_store.read().await; - store.update(flow_id, |flow| { - if !flow.annotations.tags.contains(&tag) { - flow.annotations.tags.push(tag); - } - }) - } - - /// 移除标签 - pub async fn remove_tag(&self, flow_id: &str, tag: &str) -> bool { - let store = self.memory_store.read().await; - store.update(flow_id, |flow| { - flow.annotations.tags.retain(|t| t != tag); - }) - } - - /// 设置标记 - pub async fn set_marker(&self, flow_id: &str, marker: Option) -> bool { - let store = self.memory_store.read().await; - store.update(flow_id, |flow| { - flow.annotations.marker = marker; - }) - } - - /// 获取活跃 Flow 数量 - pub async fn active_flow_count(&self) -> usize { - self.active_flows.read().await.len() - } - - /// 获取内存中的 Flow 数量 - pub async fn memory_flow_count(&self) -> usize { - self.memory_store.read().await.len() - } - - /// 检查监控是否启用 - pub async fn is_enabled(&self) -> bool { - self.config.read().await.enabled - } - - /// 启用监控 - pub async fn enable(&self) { - self.config.write().await.enabled = true; - } - - /// 禁用监控 - pub async fn disable(&self) { - self.config.write().await.enabled = false; - } - - /// 检查 Flow 是否超过阈值 - /// - /// **Validates: Requirements 10.3, 10.4** - /// - /// # Arguments - /// * `flow` - 要检查的 Flow - /// - /// # Returns - /// 阈值检测结果 - pub async fn check_threshold(&self, flow: &LLMFlow) -> ThresholdCheckResult { - let config = self.threshold_config.read().await; - Self::check_threshold_with_config(flow, &config) - } - - /// 使用指定配置检查 Flow 是否超过阈值 - /// - /// **Validates: Requirements 10.3, 10.4** - /// - /// # Arguments - /// * `flow` - 要检查的 Flow - /// * `config` - 阈值配置 - /// - /// # Returns - /// 阈值检测结果 - pub fn check_threshold_with_config( - flow: &LLMFlow, - config: &ThresholdConfig, - ) -> ThresholdCheckResult { - if !config.enabled { - return ThresholdCheckResult::default(); - } - - let actual_latency_ms = flow.timestamps.duration_ms; - let (actual_input_tokens, actual_output_tokens, actual_tokens) = - if let Some(ref response) = flow.response { - ( - response.usage.input_tokens, - response.usage.output_tokens, - response.usage.total_tokens, - ) - } else { - (0, 0, 0) - }; - - let latency_exceeded = actual_latency_ms > config.latency_threshold_ms; - let token_exceeded = actual_tokens > config.token_threshold; - let input_token_exceeded = config - .input_token_threshold - .is_some_and(|threshold| actual_input_tokens > threshold); - let output_token_exceeded = config - .output_token_threshold - .is_some_and(|threshold| actual_output_tokens > threshold); - - ThresholdCheckResult { - latency_exceeded, - token_exceeded, - input_token_exceeded, - output_token_exceeded, - actual_latency_ms, - actual_tokens, - actual_input_tokens, - actual_output_tokens, - } - } - - /// 计算指定时间窗口内的请求速率 - /// - /// **Validates: Requirements 10.7** - /// - /// # Arguments - /// * `timestamps` - 请求时间戳列表 - /// * `window_seconds` - 时间窗口(秒) - /// - /// # Returns - /// 请求速率(每秒) - pub fn calculate_request_rate(timestamps: &[DateTime], window_seconds: i64) -> f64 { - if timestamps.is_empty() || window_seconds <= 0 { - return 0.0; - } - - let now = Utc::now(); - let cutoff = now - Duration::seconds(window_seconds); - let count = timestamps.iter().filter(|&&ts| ts >= cutoff).count(); - - count as f64 / window_seconds as f64 - } - - /// 计算指定时间点的请求速率 - /// - /// **Validates: Requirements 10.7** - /// - /// # Arguments - /// * `timestamps` - 请求时间戳列表 - /// * `window_seconds` - 时间窗口(秒) - /// * `at_time` - 计算时间点 - /// - /// # Returns - /// 请求速率(每秒) - pub fn calculate_request_rate_at( - timestamps: &[DateTime], - window_seconds: i64, - at_time: DateTime, - ) -> f64 { - if timestamps.is_empty() || window_seconds <= 0 { - return 0.0; - } - - let cutoff = at_time - Duration::seconds(window_seconds); - let count = timestamps - .iter() - .filter(|&&ts| ts >= cutoff && ts <= at_time) - .count(); - - count as f64 / window_seconds as f64 - } -} - -// ============================================================================ -// 测试模块 -// ============================================================================ - -#[cfg(test)] -mod tests { - use super::*; - use crate::flow_monitor::models::{ - FlowMetadata, LLMRequest, Message, MessageContent, MessageRole, RequestParameters, - }; - use crate::ProviderType; - - /// 创建测试用的 LLMRequest - fn create_test_request(model: &str, path: &str) -> LLMRequest { - LLMRequest { - method: "POST".to_string(), - path: path.to_string(), - headers: HashMap::new(), - body: serde_json::Value::Null, - messages: vec![Message { - role: MessageRole::User, - content: MessageContent::Text("Hello".to_string()), - tool_calls: None, - tool_result: None, - name: None, - }], - system_prompt: None, - tools: None, - model: model.to_string(), - original_model: None, - parameters: RequestParameters::default(), - size_bytes: 0, - timestamp: Utc::now(), - } - } - - /// 创建测试用的 FlowMetadata - fn create_test_metadata(provider: ProviderType) -> FlowMetadata { - FlowMetadata { - provider, - credential_id: Some("test-cred".to_string()), - credential_name: Some("Test Credential".to_string()), - ..Default::default() - } - } - - #[tokio::test] - async fn test_flow_monitor_creation() { - let config = FlowMonitorConfig::default(); - let monitor = FlowMonitor::new(config, None); - - assert!(monitor.is_enabled().await); - assert_eq!(monitor.active_flow_count().await, 0); - assert_eq!(monitor.memory_flow_count().await, 0); - } - - #[tokio::test] - async fn test_start_flow() { - let config = FlowMonitorConfig::default(); - let monitor = FlowMonitor::new(config, None); - - let request = create_test_request("gpt-4", "/v1/chat/completions"); - let metadata = create_test_metadata(ProviderType::OpenAI); - - let flow_id = monitor.start_flow(request, metadata).await; - - assert!(flow_id.is_some()); - assert_eq!(monitor.active_flow_count().await, 1); - } - - #[tokio::test] - async fn test_complete_flow() { - let config = FlowMonitorConfig::default(); - let monitor = FlowMonitor::new(config, None); - - let request = create_test_request("gpt-4", "/v1/chat/completions"); - let metadata = create_test_metadata(ProviderType::OpenAI); - - let flow_id = monitor.start_flow(request, metadata).await.unwrap(); - - // 完成 Flow - monitor.complete_flow(&flow_id, None).await; - - assert_eq!(monitor.active_flow_count().await, 0); - assert_eq!(monitor.memory_flow_count().await, 1); - } - - #[tokio::test] - async fn test_fail_flow() { - let config = FlowMonitorConfig::default(); - let monitor = FlowMonitor::new(config, None); - - let request = create_test_request("gpt-4", "/v1/chat/completions"); - let metadata = create_test_metadata(ProviderType::OpenAI); - - let flow_id = monitor.start_flow(request, metadata).await.unwrap(); - - // 失败 Flow - let error = FlowError::new( - crate::flow_monitor::models::FlowErrorType::Network, - "Connection failed", - ); - monitor.fail_flow(&flow_id, error).await; - - assert_eq!(monitor.active_flow_count().await, 0); - assert_eq!(monitor.memory_flow_count().await, 1); - } - - #[tokio::test] - async fn test_config_should_monitor() { - let config = FlowMonitorConfig { - enabled: true, - sampling_rate: 1.0, - excluded_models: vec!["test-*".to_string()], - excluded_paths: vec!["/health".to_string()], - ..Default::default() - }; - - // 正常请求应该被监控 - assert!(config.should_monitor("gpt-4", "/v1/chat/completions")); - - // 排除的模型不应该被监控 - assert!(!config.should_monitor("test-model", "/v1/chat/completions")); - - // 排除的路径不应该被监控 - assert!(!config.should_monitor("gpt-4", "/health")); - } - - #[tokio::test] - async fn test_disabled_monitor() { - let config = FlowMonitorConfig { - enabled: false, - ..Default::default() - }; - let monitor = FlowMonitor::new(config, None); - - let request = create_test_request("gpt-4", "/v1/chat/completions"); - let metadata = create_test_metadata(ProviderType::OpenAI); - - // 禁用时不应该创建 Flow - let flow_id = monitor.start_flow(request, metadata).await; - assert!(flow_id.is_none()); - } - - #[tokio::test] - async fn test_event_subscription() { - let config = FlowMonitorConfig::default(); - let monitor = FlowMonitor::new(config, None); - - let mut receiver = monitor.subscribe(); - - let request = create_test_request("gpt-4", "/v1/chat/completions"); - let metadata = create_test_metadata(ProviderType::OpenAI); - - let flow_id = monitor.start_flow(request, metadata).await.unwrap(); - - // 应该收到 FlowStarted 事件 - let event = receiver.try_recv(); - assert!(event.is_ok()); - if let FlowEvent::FlowStarted { flow } = event.unwrap() { - assert_eq!(flow.id, flow_id); - assert_eq!(flow.model, "gpt-4"); - } else { - panic!("Expected FlowStarted event"); - } - } - - #[tokio::test] - async fn test_flow_type_detection() { - assert_eq!( - FlowMonitor::determine_flow_type("/v1/chat/completions"), - FlowType::ChatCompletions - ); - assert_eq!( - FlowMonitor::determine_flow_type("/v1/messages"), - FlowType::AnthropicMessages - ); - assert_eq!( - FlowMonitor::determine_flow_type("/v1/models/gemini-pro:generatecontent"), - FlowType::GeminiGenerateContent - ); - assert_eq!( - FlowMonitor::determine_flow_type("/v1/embeddings"), - FlowType::Embeddings - ); - } - - #[tokio::test] - async fn test_annotations_update() { - let config = FlowMonitorConfig::default(); - let monitor = FlowMonitor::new(config, None); - - let request = create_test_request("gpt-4", "/v1/chat/completions"); - let metadata = create_test_metadata(ProviderType::OpenAI); - - let flow_id = monitor.start_flow(request, metadata).await.unwrap(); - monitor.complete_flow(&flow_id, None).await; - - // 测试收藏 - assert!(monitor.toggle_starred(&flow_id).await); - - // 测试添加评论 - assert!( - monitor - .add_comment(&flow_id, "Test comment".to_string()) - .await - ); - - // 测试添加标签 - assert!(monitor.add_tag(&flow_id, "important".to_string()).await); - - // 测试设置标记 - assert!(monitor.set_marker(&flow_id, Some("⭐".to_string())).await); - } -} - -// ============================================================================ -// 属性测试模块 -// ============================================================================ - -#[cfg(test)] -mod property_tests { - use super::*; - use crate::flow_monitor::models::{ - FlowErrorType, FlowMetadata, LLMRequest, Message, MessageContent, MessageRole, - RequestParameters, - }; - use crate::ProviderType; - use proptest::prelude::*; - use tokio::runtime::Runtime; - - // ======================================================================== - // 生成器 - // ======================================================================== - - /// 生成随机的 ProviderType - fn arb_provider_type() -> impl Strategy { - prop_oneof![ - Just(ProviderType::Kiro), - Just(ProviderType::Gemini), - Just(ProviderType::OpenAI), - Just(ProviderType::Claude), - Just(ProviderType::Antigravity), - ] - } - - /// 生成随机的模型名称 - fn arb_model_name() -> impl Strategy { - prop_oneof![ - Just("gpt-4".to_string()), - Just("gpt-4-turbo".to_string()), - Just("gpt-3.5-turbo".to_string()), - Just("claude-3-opus".to_string()), - Just("claude-3-sonnet".to_string()), - Just("gemini-pro".to_string()), - ] - } - - /// 生成随机的路径 - fn arb_path() -> impl Strategy { - prop_oneof![ - Just("/v1/chat/completions".to_string()), - Just("/v1/messages".to_string()), - Just("/v1/embeddings".to_string()), - ] - } - - /// 生成随机的 LLMRequest - fn arb_llm_request() -> impl Strategy { - (arb_model_name(), arb_path()).prop_map(|(model, path)| LLMRequest { - method: "POST".to_string(), - path, - headers: HashMap::new(), - body: serde_json::Value::Null, - messages: vec![Message { - role: MessageRole::User, - content: MessageContent::Text("Test message".to_string()), - tool_calls: None, - tool_result: None, - name: None, - }], - system_prompt: None, - tools: None, - model, - original_model: None, - parameters: RequestParameters::default(), - size_bytes: 0, - timestamp: Utc::now(), - }) - } - - /// 生成随机的 FlowMetadata - fn arb_flow_metadata() -> impl Strategy { - arb_provider_type().prop_map(|provider| FlowMetadata { - provider, - credential_id: Some("test-cred".to_string()), - credential_name: Some("Test Credential".to_string()), - ..Default::default() - }) - } - - /// 生成随机的 FlowErrorType - fn arb_flow_error_type() -> impl Strategy { - prop_oneof![ - Just(FlowErrorType::Network), - Just(FlowErrorType::Timeout), - Just(FlowErrorType::Authentication), - Just(FlowErrorType::RateLimit), - Just(FlowErrorType::ServerError), - Just(FlowErrorType::BadRequest), - ] - } - - // ======================================================================== - // 属性测试 - // ======================================================================== - - proptest! { - #![proptest_config(ProptestConfig::with_cases(20))] - - /// **Feature: llm-flow-monitor, Property 9: 事件发送正确性** - /// **Validates: Requirements 6.1, 6.2, 6.3, 6.4** - /// - /// *对于任意* Flow 生命周期操作(开始、更新、完成、失败), - /// 应该发出对应的事件,且事件内容应该正确反映 Flow 状态。 - #[test] - fn prop_event_emission_correctness( - request in arb_llm_request(), - metadata in arb_flow_metadata(), - ) { - let rt = Runtime::new().unwrap(); - rt.block_on(async { - let config = FlowMonitorConfig::default(); - - // 创建禁用通知的配置 - let notification_config = NotificationConfig { - enabled: false, - new_flow: NotificationSettings::default(), - error_flow: NotificationSettings::default(), - latency_warning: NotificationSettings::default(), - token_warning: NotificationSettings::default(), - }; - - let monitor = FlowMonitor::with_notification_config( - config, - None, - ThresholdConfig::default(), - notification_config - ); - - let mut receiver = monitor.subscribe(); - - // 开始 Flow - let flow_id = monitor.start_flow(request.clone(), metadata.clone()).await; - prop_assert!(flow_id.is_some(), "Flow 应该被创建"); - let flow_id = flow_id.unwrap(); - - // 验证 FlowStarted 事件 - let event = receiver.try_recv(); - prop_assert!(event.is_ok(), "应该收到 FlowStarted 事件"); - if let FlowEvent::FlowStarted { flow } = event.unwrap() { - prop_assert_eq!(flow.id, flow_id.clone(), "事件中的 Flow ID 应该正确"); - prop_assert_eq!(flow.model, request.model, "事件中的模型应该正确"); - prop_assert_eq!( - flow.state, - FlowState::Pending, - "新 Flow 状态应该是 Pending" - ); - } else { - prop_assert!(false, "应该是 FlowStarted 事件"); - } - - // 可能有 RequestRateUpdate 事件,消费它 - let _ = receiver.try_recv(); - - // 完成 Flow - monitor.complete_flow(&flow_id, None).await; - - // 验证 FlowCompleted 事件(可能需要跳过其他事件) - let mut found_completed = false; - for _ in 0..3 { // 最多尝试 3 次 - let event = receiver.try_recv(); - if event.is_ok() { - if let FlowEvent::FlowCompleted { id, summary } = event.unwrap() { - prop_assert_eq!(id, flow_id.clone(), "事件中的 Flow ID 应该正确"); - prop_assert_eq!( - summary.state, - FlowState::Completed, - "完成后状态应该是 Completed" - ); - found_completed = true; - break; - } - // 如果不是 FlowCompleted 事件,继续尝试下一个 - } else { - break; - } - } - prop_assert!(found_completed, "应该收到 FlowCompleted 事件"); - - Ok(()) - })?; - } - - /// **Feature: llm-flow-monitor, Property 9b: 失败事件发送正确性** - /// **Validates: Requirements 6.4** - /// - /// *对于任意* Flow 失败操作,应该发出 FlowFailed 事件, - /// 且事件内容应该包含正确的错误信息。 - #[test] - fn prop_failure_event_correctness( - request in arb_llm_request(), - metadata in arb_flow_metadata(), - error_type in arb_flow_error_type(), - ) { - let rt = Runtime::new().unwrap(); - rt.block_on(async { - let config = FlowMonitorConfig::default(); - - // 创建禁用通知的配置 - let notification_config = NotificationConfig { - enabled: false, - new_flow: NotificationSettings::default(), - error_flow: NotificationSettings::default(), - latency_warning: NotificationSettings::default(), - token_warning: NotificationSettings::default(), - }; - - let monitor = FlowMonitor::with_notification_config( - config, - None, - ThresholdConfig::default(), - notification_config - ); - - let mut receiver = monitor.subscribe(); - - // 开始 Flow - let flow_id = monitor.start_flow(request, metadata).await.unwrap(); - - // 消费 FlowStarted 事件 - let _ = receiver.try_recv(); - // 可能有 RequestRateUpdate 事件,消费它 - let _ = receiver.try_recv(); - - // 失败 Flow - let error = FlowError::new(error_type.clone(), "Test error message"); - monitor.fail_flow(&flow_id, error.clone()).await; - - // 验证 FlowFailed 事件(可能需要跳过其他事件) - let mut found_failed = false; - for _ in 0..3 { // 最多尝试 3 次 - let event = receiver.try_recv(); - if event.is_ok() { - if let FlowEvent::FlowFailed { id, error: evt_error } = event.unwrap() { - prop_assert_eq!(id, flow_id, "事件中的 Flow ID 应该正确"); - prop_assert_eq!( - evt_error.error_type, - error_type, - "事件中的错误类型应该正确" - ); - prop_assert_eq!( - evt_error.message, - "Test error message", - "事件中的错误消息应该正确" - ); - found_failed = true; - break; - } - // 如果不是 FlowFailed 事件,继续尝试下一个 - } else { - break; - } - } - prop_assert!(found_failed, "应该收到 FlowFailed 事件"); - - Ok(()) - })?; - } - - /// **Feature: llm-flow-monitor, Property 10: 标注 Round-Trip** - /// **Validates: Requirements 7.1, 7.2, 7.3, 7.4** - /// - /// *对于任意* Flow 和标注操作(收藏、评论、标签、标记), - /// 更新后再读取,标注信息应该与设置的值一致。 - #[test] - fn prop_annotation_roundtrip( - request in arb_llm_request(), - metadata in arb_flow_metadata(), - starred in any::(), - comment in prop::option::of("[a-zA-Z0-9 ]{1,50}"), - marker in prop::option::of(prop_oneof![ - Just("⭐".to_string()), - Just("🔴".to_string()), - Just("🟢".to_string()), - ]), - tags in prop::collection::vec("[a-z]{3,10}", 0..3), - ) { - let rt = Runtime::new().unwrap(); - rt.block_on(async { - let config = FlowMonitorConfig::default(); - let monitor = FlowMonitor::new(config, None); - - // 创建并完成 Flow - let flow_id = monitor.start_flow(request, metadata).await.unwrap(); - monitor.complete_flow(&flow_id, None).await; - - // 设置标注 - let annotations = FlowAnnotations { - starred, - comment: comment.clone(), - marker: marker.clone(), - tags: tags.clone(), - }; - - let updated = monitor.update_annotations(&flow_id, annotations.clone()).await; - prop_assert!(updated, "标注更新应该成功"); - - // 读取并验证 - let store = monitor.memory_store.read().await; - let flow_lock = store.get(&flow_id); - prop_assert!(flow_lock.is_some(), "Flow 应该存在"); - - let binding = flow_lock.unwrap(); - let flow = binding.read().unwrap(); - prop_assert_eq!(flow.annotations.starred, starred, "收藏状态应该一致"); - prop_assert_eq!(flow.annotations.comment.clone(), comment, "评论应该一致"); - prop_assert_eq!(flow.annotations.marker.clone(), marker, "标记应该一致"); - prop_assert_eq!(flow.annotations.tags.clone(), tags, "标签应该一致"); - - Ok(()) - })?; - } - - /// **Feature: llm-flow-monitor, Property 12: 配置生效属性** - /// **Validates: Requirements 11.1, 11.2, 11.7, 11.8** - /// - /// *对于任意* 监控配置(启用/禁用、缓存大小、采样率、排除规则), - /// Flow_Monitor 的行为应该符合配置。 - #[test] - fn prop_config_effectiveness( - enabled in any::(), - max_memory_flows in 10usize..100usize, - excluded_model in prop::option::of("[a-z]{3,10}"), - ) { - let rt = Runtime::new().unwrap(); - rt.block_on(async { - // 构建配置 - let excluded_models = excluded_model - .clone() - .map(|m| vec![format!("{}*", m)]) - .unwrap_or_default(); - - let config = FlowMonitorConfig { - enabled, - max_memory_flows, - sampling_rate: 1.0, // 确保采样率为 100% - excluded_models: excluded_models.clone(), - ..Default::default() - }; - - let monitor = FlowMonitor::new(config, None); - - // 验证启用/禁用配置 - prop_assert_eq!( - monitor.is_enabled().await, - enabled, - "监控启用状态应该与配置一致" - ); - - // 测试排除模型配置 - if let Some(ref excluded) = excluded_model { - let excluded_model_name = format!("{excluded}-test"); - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model: excluded_model_name, - ..Default::default() - }; - let metadata = FlowMetadata::default(); - - let flow_id = monitor.start_flow(request, metadata).await; - - if enabled { - // 启用时,排除的模型不应该被监控 - prop_assert!( - flow_id.is_none(), - "排除的模型不应该被监控" - ); - } else { - // 禁用时,任何模型都不应该被监控 - prop_assert!( - flow_id.is_none(), - "禁用时不应该监控任何模型" - ); - } - } - - // 测试非排除模型 - if enabled { - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model: "gpt-4".to_string(), - ..Default::default() - }; - let metadata = FlowMetadata::default(); - - let flow_id = monitor.start_flow(request, metadata).await; - prop_assert!( - flow_id.is_some(), - "启用时,非排除的模型应该被监控" - ); - } - - Ok(()) - })?; - } - - /// **Feature: llm-flow-monitor, Property 12b: 缓存大小配置生效** - /// **Validates: Requirements 11.2** - /// - /// *对于任意* 缓存大小配置,内存存储的最大大小应该与配置一致。 - #[test] - fn prop_cache_size_config( - max_memory_flows in 10usize..100usize, - ) { - let rt = Runtime::new().unwrap(); - rt.block_on(async { - let config = FlowMonitorConfig { - enabled: true, - max_memory_flows, - sampling_rate: 1.0, - ..Default::default() - }; - - let monitor = FlowMonitor::new(config, None); - - // 验证内存存储的最大大小 - let store = monitor.memory_store.read().await; - prop_assert_eq!( - store.max_size(), - max_memory_flows, - "内存存储的最大大小应该与配置一致" - ); - - Ok(()) - })?; - } - - /// **Feature: flow-monitor-enhancement, Property 18: 阈值检测正确性** - /// **Validates: Requirements 10.3, 10.4** - /// - /// *对于任意* 阈值配置和 Flow,阈值检测应该正确判断是否超过阈值。 - #[test] - fn prop_threshold_detection_correctness( - latency_threshold_ms in 100u64..10000u64, - token_threshold in 100u32..50000u32, - actual_latency_ms in 0u64..20000u64, - actual_input_tokens in 0u32..30000u32, - actual_output_tokens in 0u32..30000u32, - input_token_threshold in prop::option::of(100u32..50000u32), - output_token_threshold in prop::option::of(100u32..50000u32), - ) { - use crate::flow_monitor::models::{LLMResponse, TokenUsage}; - - // 创建阈值配置 - let config = ThresholdConfig { - enabled: true, - latency_threshold_ms, - token_threshold, - input_token_threshold, - output_token_threshold, - }; - - // 创建测试 Flow - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model: "gpt-4".to_string(), - ..Default::default() - }; - let metadata = FlowMetadata::default(); - let mut flow = LLMFlow::new( - "test-flow".to_string(), - FlowType::ChatCompletions, - request, - metadata, - ); - - // 设置延迟 - flow.timestamps.duration_ms = actual_latency_ms; - - // 设置 Token 使用量 - let actual_total_tokens = actual_input_tokens + actual_output_tokens; - flow.response = Some(LLMResponse { - usage: TokenUsage { - input_tokens: actual_input_tokens, - output_tokens: actual_output_tokens, - total_tokens: actual_total_tokens, - ..Default::default() - }, - ..Default::default() - }); - - // 执行阈值检测 - let result = FlowMonitor::check_threshold_with_config(&flow, &config); - - // 验证延迟阈值检测 - let expected_latency_exceeded = actual_latency_ms > latency_threshold_ms; - prop_assert_eq!( - result.latency_exceeded, - expected_latency_exceeded, - "延迟阈值检测应该正确: 实际延迟 {} ms, 阈值 {} ms", - actual_latency_ms, - latency_threshold_ms - ); - - // 验证 Token 阈值检测 - let expected_token_exceeded = actual_total_tokens > token_threshold; - prop_assert_eq!( - result.token_exceeded, - expected_token_exceeded, - "Token 阈值检测应该正确: 实际 Token {}, 阈值 {}", - actual_total_tokens, - token_threshold - ); - - // 验证输入 Token 阈值检测 - let expected_input_exceeded = input_token_threshold - .is_some_and(|threshold| actual_input_tokens > threshold); - prop_assert_eq!( - result.input_token_exceeded, - expected_input_exceeded, - "输入 Token 阈值检测应该正确" - ); - - // 验证输出 Token 阈值检测 - let expected_output_exceeded = output_token_threshold - .is_some_and(|threshold| actual_output_tokens > threshold); - prop_assert_eq!( - result.output_token_exceeded, - expected_output_exceeded, - "输出 Token 阈值检测应该正确" - ); - - // 验证实际值记录正确 - prop_assert_eq!( - result.actual_latency_ms, - actual_latency_ms, - "实际延迟应该正确记录" - ); - prop_assert_eq!( - result.actual_tokens, - actual_total_tokens, - "实际 Token 数应该正确记录" - ); - prop_assert_eq!( - result.actual_input_tokens, - actual_input_tokens, - "实际输入 Token 数应该正确记录" - ); - prop_assert_eq!( - result.actual_output_tokens, - actual_output_tokens, - "实际输出 Token 数应该正确记录" - ); - - // 验证 any_exceeded 方法 - let expected_any_exceeded = expected_latency_exceeded - || expected_token_exceeded - || expected_input_exceeded - || expected_output_exceeded; - prop_assert_eq!( - result.any_exceeded(), - expected_any_exceeded, - "any_exceeded 应该正确反映是否有任何阈值被超过" - ); - } - - /// **Feature: flow-monitor-enhancement, Property 18b: 禁用阈值检测** - /// **Validates: Requirements 10.3, 10.4** - /// - /// *对于任意* Flow,当阈值检测禁用时,所有检测结果应该为 false。 - #[test] - fn prop_threshold_detection_disabled( - actual_latency_ms in 0u64..20000u64, - actual_input_tokens in 0u32..30000u32, - actual_output_tokens in 0u32..30000u32, - ) { - use crate::flow_monitor::models::{LLMResponse, TokenUsage}; - - // 创建禁用的阈值配置 - let config = ThresholdConfig { - enabled: false, - latency_threshold_ms: 100, // 很低的阈值 - token_threshold: 100, // 很低的阈值 - input_token_threshold: Some(100), - output_token_threshold: Some(100), - }; - - // 创建测试 Flow - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model: "gpt-4".to_string(), - ..Default::default() - }; - let metadata = FlowMetadata::default(); - let mut flow = LLMFlow::new( - "test-flow".to_string(), - FlowType::ChatCompletions, - request, - metadata, - ); - - // 设置延迟和 Token(超过阈值) - flow.timestamps.duration_ms = actual_latency_ms; - flow.response = Some(LLMResponse { - usage: TokenUsage { - input_tokens: actual_input_tokens, - output_tokens: actual_output_tokens, - total_tokens: actual_input_tokens + actual_output_tokens, - ..Default::default() - }, - ..Default::default() - }); - - // 执行阈值检测 - let result = FlowMonitor::check_threshold_with_config(&flow, &config); - - // 验证所有检测结果都为 false - prop_assert!( - !result.latency_exceeded, - "禁用时延迟阈值检测应该为 false" - ); - prop_assert!( - !result.token_exceeded, - "禁用时 Token 阈值检测应该为 false" - ); - prop_assert!( - !result.input_token_exceeded, - "禁用时输入 Token 阈值检测应该为 false" - ); - prop_assert!( - !result.output_token_exceeded, - "禁用时输出 Token 阈值检测应该为 false" - ); - prop_assert!( - !result.any_exceeded(), - "禁用时 any_exceeded 应该为 false" - ); - } - - /// **Feature: flow-monitor-enhancement, Property 20: 通知触发正确性** - /// **Validates: Requirements 10.1, 10.2, 10.3, 10.4** - /// - /// *对于任意* 通知配置和 Flow 事件,当通知启用时应该触发相应的通知事件。 - #[test] - fn prop_notification_trigger_correctness( - request in arb_llm_request(), - metadata in arb_flow_metadata(), - new_flow_enabled in any::(), - error_flow_enabled in any::(), - latency_warning_enabled in any::(), - token_warning_enabled in any::(), - ) { - let rt = Runtime::new().unwrap(); - rt.block_on(async { - // 创建通知配置 - let notification_config = NotificationConfig { - enabled: true, - new_flow: NotificationSettings { - enabled: new_flow_enabled, - desktop: true, - sound: false, - sound_file: None, - }, - error_flow: NotificationSettings { - enabled: error_flow_enabled, - desktop: true, - sound: false, - sound_file: None, - }, - latency_warning: NotificationSettings { - enabled: latency_warning_enabled, - desktop: false, - sound: false, - sound_file: None, - }, - token_warning: NotificationSettings { - enabled: token_warning_enabled, - desktop: false, - sound: false, - sound_file: None, - }, - }; - - // 创建阈值配置(低阈值,容易触发) - let threshold_config = ThresholdConfig { - enabled: true, - latency_threshold_ms: 100, - token_threshold: 100, - input_token_threshold: None, - output_token_threshold: None, - }; - - let config = FlowMonitorConfig::default(); - let monitor = FlowMonitor::with_full_config( - config, - None, - threshold_config, - notification_config, - ); - - let mut receiver = monitor.subscribe(); - - // 开始 Flow - let flow_id = monitor.start_flow(request.clone(), metadata.clone()).await; - prop_assert!(flow_id.is_some(), "Flow 应该被创建"); - let flow_id = flow_id.unwrap(); - - // 消费 FlowStarted 事件 - let _ = receiver.try_recv(); - - // 检查新 Flow 通知 - if new_flow_enabled { - // 应该有 RequestRateUpdate 事件 - let event = receiver.try_recv(); - if event.is_ok() { - let event_value = event.unwrap(); - if let FlowEvent::RequestRateUpdate { .. } = event_value { - // 这是速率更新事件,继续检查通知事件 - let notification_event = receiver.try_recv(); - if notification_event.is_ok() { - if let FlowEvent::Notification { notification } = notification_event.unwrap() { - prop_assert_eq!( - notification.flow_id, - flow_id.clone(), - "通知中的 Flow ID 应该正确" - ); - prop_assert!( - matches!(notification.notification_type, NotificationType::NewFlow), - "应该是新 Flow 通知" - ); - } - } - } else if let FlowEvent::Notification { notification } = event_value { - prop_assert_eq!( - notification.flow_id, - flow_id.clone(), - "通知中的 Flow ID 应该正确" - ); - prop_assert!( - matches!(notification.notification_type, NotificationType::NewFlow), - "应该是新 Flow 通知" - ); - } - } - } - - // 测试错误通知 - if error_flow_enabled { - let error = FlowError::new(FlowErrorType::Network, "Test error"); - monitor.fail_flow(&flow_id, error).await; - - // 消费 FlowFailed 事件 - let _ = receiver.try_recv(); - - // 检查错误通知 - let event = receiver.try_recv(); - if event.is_ok() { - if let FlowEvent::Notification { notification } = event.unwrap() { - prop_assert_eq!( - notification.flow_id, - flow_id.clone(), - "错误通知中的 Flow ID 应该正确" - ); - prop_assert!( - matches!(notification.notification_type, NotificationType::ErrorFlow), - "应该是错误 Flow 通知" - ); - } - } - } - - Ok(()) - })?; - } - - /// **Feature: flow-monitor-enhancement, Property 20b: 禁用通知不触发** - /// **Validates: Requirements 10.1, 10.2** - /// - /// *对于任意* Flow 事件,当通知禁用时不应该触发通知事件。 - #[test] - fn prop_disabled_notifications_not_triggered( - request in arb_llm_request(), - metadata in arb_flow_metadata(), - ) { - let rt = Runtime::new().unwrap(); - rt.block_on(async { - // 创建禁用的通知配置 - let notification_config = NotificationConfig { - enabled: false, // 全局禁用 - new_flow: NotificationSettings { - enabled: true, // 即使启用也不应该触发 - desktop: true, - sound: false, - sound_file: None, - }, - error_flow: NotificationSettings { - enabled: true, // 即使启用也不应该触发 - desktop: true, - sound: false, - sound_file: None, - }, - ..Default::default() - }; - - let config = FlowMonitorConfig::default(); - let monitor = FlowMonitor::with_full_config( - config, - None, - ThresholdConfig::default(), - notification_config, - ); - - let mut receiver = monitor.subscribe(); - - // 开始 Flow - let flow_id = monitor.start_flow(request, metadata).await; - prop_assert!(flow_id.is_some(), "Flow 应该被创建"); - let flow_id = flow_id.unwrap(); - - // 消费 FlowStarted 事件 - let _ = receiver.try_recv(); - - // 可能有 RequestRateUpdate 事件,消费它 - let event = receiver.try_recv(); - if event.is_ok() { - let event_value = event.unwrap(); - if let FlowEvent::RequestRateUpdate { .. } = event_value { - // 这是速率更新事件,检查是否还有其他事件 - let next_event = receiver.try_recv(); - prop_assert!( - next_event.is_err() || !matches!(next_event.unwrap(), FlowEvent::Notification { .. }), - "禁用通知时不应该有通知事件" - ); - } else { - prop_assert!( - !matches!(event_value, FlowEvent::Notification { .. }), - "禁用通知时不应该有通知事件" - ); - } - } - - // 测试错误情况 - let error = FlowError::new(FlowErrorType::Network, "Test error"); - monitor.fail_flow(&flow_id, error).await; - - // 消费 FlowFailed 事件 - let _ = receiver.try_recv(); - - // 检查不应该有通知事件 - let event = receiver.try_recv(); - if event.is_ok() { - let event = event.unwrap(); - prop_assert!( - !matches!(event, FlowEvent::Notification { .. }), - "禁用通知时不应该有错误通知事件" - ); - } - - Ok(()) - })?; - } - } -} - -// ============================================================================ -// 请求速率追踪器属性测试 -// ============================================================================ - -#[cfg(test)] -mod rate_tracker_property_tests { - use super::*; - use proptest::prelude::*; - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: flow-monitor-enhancement, Property 19: 请求速率计算正确性** - /// **Validates: Requirements 10.7** - /// - /// *对于任意* 时间窗口内的请求集合,请求速率计算应该正确反映该窗口内的请求数量。 - #[test] - fn prop_request_rate_calculation_correctness( - window_seconds in 10i64..120i64, - request_count in 0usize..100usize, - ) { - let mut tracker = RequestRateTracker::new(window_seconds); - let now = Utc::now(); - - // 在时间窗口内添加请求 - for i in 0..request_count { - // 在窗口内均匀分布请求 - let offset_seconds = if request_count > 1 { - (i as i64 * (window_seconds - 1)) / (request_count as i64 - 1).max(1) - } else { - 0 - }; - let timestamp = now - Duration::seconds(window_seconds - 1 - offset_seconds); - tracker.record_request_at(timestamp); - } - - // 计算速率 - let rate = tracker.get_rate_at(now); - let count = tracker.get_count_at(now); - - // 验证请求数量 - prop_assert_eq!( - count, - request_count, - "时间窗口内的请求数量应该正确" - ); - - // 验证速率计算 - let expected_rate = request_count as f64 / window_seconds as f64; - prop_assert!( - (rate - expected_rate).abs() < 0.0001, - "请求速率计算应该正确: 期望 {}, 实际 {}", - expected_rate, - rate - ); - } - - /// **Feature: flow-monitor-enhancement, Property 19b: 过期请求清理** - /// **Validates: Requirements 10.7** - /// - /// *对于任意* 请求集合,超出时间窗口的请求应该被正确排除。 - #[test] - fn prop_expired_requests_excluded( - window_seconds in 10i64..60i64, - in_window_count in 0usize..50usize, - out_window_count in 0usize..50usize, - ) { - let mut tracker = RequestRateTracker::new(window_seconds); - let now = Utc::now(); - - // 添加窗口内的请求 - for i in 0..in_window_count { - let offset = (i as i64 * (window_seconds - 1)) / (in_window_count as i64).max(1); - let timestamp = now - Duration::seconds(offset); - tracker.record_request_at(timestamp); - } - - // 添加窗口外的请求(过期的) - for i in 0..out_window_count { - let offset = window_seconds + 1 + i as i64; - let timestamp = now - Duration::seconds(offset); - tracker.record_request_at(timestamp); - } - - // 验证只计算窗口内的请求 - let count = tracker.get_count_at(now); - prop_assert_eq!( - count, - in_window_count, - "只应该计算时间窗口内的请求: 期望 {}, 实际 {}", - in_window_count, - count - ); - - // 验证速率只基于窗口内的请求 - let rate = tracker.get_rate_at(now); - let expected_rate = in_window_count as f64 / window_seconds as f64; - prop_assert!( - (rate - expected_rate).abs() < 0.0001, - "请求速率应该只基于窗口内的请求" - ); - } - - /// **Feature: flow-monitor-enhancement, Property 19c: 空窗口处理** - /// **Validates: Requirements 10.7** - /// - /// *对于任意* 空的请求集合,请求速率应该为 0。 - #[test] - fn prop_empty_window_rate_zero( - window_seconds in 1i64..120i64, - ) { - let tracker = RequestRateTracker::new(window_seconds); - - // 验证空窗口的速率为 0 - let rate = tracker.get_rate(); - prop_assert_eq!( - rate, - 0.0, - "空窗口的请求速率应该为 0" - ); - - // 验证空窗口的请求数量为 0 - let count = tracker.get_count(); - prop_assert_eq!( - count, - 0, - "空窗口的请求数量应该为 0" - ); - } - - /// **Feature: flow-monitor-enhancement, Property 19d: 窗口大小变更** - /// **Validates: Requirements 10.7** - /// - /// *对于任意* 窗口大小变更,请求计数应该正确反映新窗口内的请求。 - #[test] - fn prop_window_size_change( - initial_window in 30i64..60i64, - new_window in 10i64..30i64, - request_count in 10usize..50usize, - ) { - let mut tracker = RequestRateTracker::new(initial_window); - let now = Utc::now(); - - // 在初始窗口内均匀添加请求 - // 请求时间从 now 到 now - (initial_window - 1) 秒 - for i in 0..request_count { - let offset = (i as i64 * (initial_window - 1)) / (request_count as i64).max(1); - let timestamp = now - Duration::seconds(offset); - tracker.record_request_at(timestamp); - } - - // 验证初始窗口内的请求数量 - let initial_count = tracker.get_count_at(now); - prop_assert_eq!( - initial_count, - request_count, - "初始窗口内的请求数量应该正确" - ); - - // 更改窗口大小 - tracker.set_window_seconds(new_window); - - // 计算新窗口内应该有多少请求 - // cutoff = now - new_window,所以 timestamp >= cutoff 意味着 offset <= new_window - let expected_new_count = (0..request_count) - .filter(|&i| { - let offset = (i as i64 * (initial_window - 1)) / (request_count as i64).max(1); - // 请求在新窗口内的条件是 offset < new_window(严格小于) - // 因为 cutoff = now - new_window,timestamp = now - offset - // timestamp >= cutoff 等价于 now - offset >= now - new_window - // 即 offset <= new_window - offset < new_window - }) - .count(); - - // 验证新窗口内的请求数量 - let new_count = tracker.get_count_at(now); - - // 由于整数除法舍入和边界条件,允许 ±2 的误差 - // 边界情况:当 offset 恰好等于 new_window 时,由于整数除法的舍入 - // 可能导致多个请求落在边界附近 - let diff = (new_count as i64 - expected_new_count as i64).abs(); - prop_assert!( - diff <= 2, - "新窗口内的请求数量应该接近预期: 期望 {}, 实际 {}, 差异 {}", - expected_new_count, - new_count, - diff - ); - } - } -} diff --git a/src-tauri/src/flow_monitor/query_service.rs b/src-tauri/src/flow_monitor/query_service.rs deleted file mode 100644 index f630591de..000000000 --- a/src-tauri/src/flow_monitor/query_service.rs +++ /dev/null @@ -1,1337 +0,0 @@ -//! Flow 查询服务 -//! -//! 该模块实现 LLM Flow 的查询服务,支持多维度过滤、排序、分页和全文搜索。 -//! 查询时先检查内存缓存,再检查文件存储。 - -use chrono::{DateTime, Utc}; -use serde::{Deserialize, Serialize}; -use std::cmp::Ordering; -use std::sync::Arc; -use thiserror::Error; -use tokio::sync::RwLock; - -use super::file_store::{FileStoreError, FlowFileStore}; -use super::filter_parser::{FilterParseError, FilterParser}; -use super::memory_store::{FlowFilter, FlowMemoryStore}; -use super::models::{FlowState, LLMFlow}; - -// ============================================================================ -// 错误类型 -// ============================================================================ - -/// 使用过滤表达式查询时的错误 -#[derive(Debug, Error)] -pub enum QueryWithExpressionError { - /// 过滤表达式解析错误 - #[error("过滤表达式解析错误: {0}")] - ParseError(#[from] FilterParseError), - /// 文件存储错误 - #[error("文件存储错误: {0}")] - FileStoreError(#[from] FileStoreError), -} - -// ============================================================================ -// 排序选项 -// ============================================================================ - -/// 排序字段 -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] -#[serde(rename_all = "snake_case")] -#[derive(Default)] -pub enum FlowSortBy { - /// 按创建时间排序 - #[default] - CreatedAt, - /// 按耗时排序 - Duration, - /// 按总 Token 数排序 - TotalTokens, - /// 按响应内容长度排序 - ContentLength, - /// 按模型名称排序 - Model, -} - -// ============================================================================ -// 查询结果 -// ============================================================================ - -/// 查询结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct FlowQueryResult { - /// 匹配的 Flow 列表 - pub flows: Vec, - /// 总数(不含分页) - pub total: usize, - /// 当前页码 - pub page: usize, - /// 每页大小 - pub page_size: usize, - /// 总页数 - pub total_pages: usize, - /// 是否有下一页 - pub has_next: bool, - /// 是否有上一页 - pub has_prev: bool, -} - -impl FlowQueryResult { - /// 创建空结果 - pub fn empty(page: usize, page_size: usize) -> Self { - Self { - flows: Vec::new(), - total: 0, - page, - page_size, - total_pages: 0, - has_next: false, - has_prev: false, - } - } -} - -/// 搜索结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct FlowSearchResult { - /// Flow ID - pub id: String, - /// 创建时间 - pub created_at: DateTime, - /// 模型名称 - pub model: String, - /// 提供商 - pub provider: String, - /// 匹配的内容片段 - pub snippet: String, - /// 匹配分数 - pub score: f64, -} - -// ============================================================================ -// 统计信息 -// ============================================================================ - -/// Flow 统计信息 -#[derive(Debug, Clone, Default, Serialize, Deserialize)] -pub struct FlowStats { - /// 总请求数 - pub total_requests: usize, - /// 成功请求数 - pub successful_requests: usize, - /// 失败请求数 - pub failed_requests: usize, - /// 成功率 - pub success_rate: f64, - /// 平均延迟(毫秒) - pub avg_latency_ms: f64, - /// 最小延迟(毫秒) - pub min_latency_ms: u64, - /// 最大延迟(毫秒) - pub max_latency_ms: u64, - /// 总输入 Token 数 - pub total_input_tokens: u64, - /// 总输出 Token 数 - pub total_output_tokens: u64, - /// 平均输入 Token 数 - pub avg_input_tokens: f64, - /// 平均输出 Token 数 - pub avg_output_tokens: f64, - /// 按提供商统计 - pub by_provider: Vec, - /// 按模型统计 - pub by_model: Vec, - /// 按状态统计 - pub by_state: Vec, -} - -/// 按提供商统计 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ProviderStats { - pub provider: String, - pub count: usize, - pub success_rate: f64, - pub avg_latency_ms: f64, -} - -/// 按模型统计 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ModelStats { - pub model: String, - pub count: usize, - pub success_rate: f64, - pub avg_latency_ms: f64, -} - -/// 按状态统计 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct StateStats { - pub state: String, - pub count: usize, -} - -// ============================================================================ -// 查询服务 -// ============================================================================ - -/// Flow 查询服务 -/// -/// 提供统一的查询接口,先查内存缓存,再查文件存储。 -pub struct FlowQueryService { - /// 内存存储 - memory_store: Arc>, - /// 文件存储 - file_store: Arc, -} - -impl FlowQueryService { - /// 创建新的查询服务 - pub fn new(memory_store: Arc>, file_store: Arc) -> Self { - Self { - memory_store, - file_store, - } - } - - /// 查询 Flow - /// - /// # 参数 - /// - `filter`: 过滤条件 - /// - `sort_by`: 排序字段 - /// - `sort_desc`: 是否降序 - /// - `page`: 页码(从 1 开始) - /// - `page_size`: 每页大小 - pub async fn query( - &self, - filter: FlowFilter, - sort_by: FlowSortBy, - sort_desc: bool, - page: usize, - page_size: usize, - ) -> Result { - eprintln!( - "[QUERY_SERVICE] 开始查询: filter={filter:?}, sort_by={sort_by:?}, page={page}, page_size={page_size}" - ); - - // 先从内存获取 - let memory_flows = { - let store = self.memory_store.read().await; - let flows = store.query(&filter); - eprintln!("[QUERY_SERVICE] 内存中查询到 {} 条记录", flows.len()); - flows - }; - - // 再从文件获取(如果需要更多数据) - // 这里简化处理:如果内存数据足够,就不查文件 - // 实际应用中可能需要更复杂的合并逻辑 - let mut all_flows = memory_flows; - - // 如果内存数据不足,从文件补充 - let memory_count = all_flows.len(); - let needed = page * page_size; - - if memory_count < needed { - eprintln!( - "[QUERY_SERVICE] 内存数据不足,从文件补充: memory_count={memory_count}, needed={needed}" - ); - - // 从文件存储获取更多数据 - let file_flows = self.file_store.query(&filter, needed * 2, 0)?; - eprintln!("[QUERY_SERVICE] 文件中查询到 {} 条记录", file_flows.len()); - - // 合并并去重(以 ID 为准) - let memory_ids: std::collections::HashSet<_> = - all_flows.iter().map(|f| f.id.clone()).collect(); - - for flow in file_flows { - if !memory_ids.contains(&flow.id) { - all_flows.push(flow); - } - } - } - - eprintln!("[QUERY_SERVICE] 合并后总记录数: {}", all_flows.len()); - - // 排序 - Self::sort_flows(&mut all_flows, sort_by, sort_desc); - - // 计算分页 - let total = all_flows.len(); - let total_pages = if page_size > 0 { - total.div_ceil(page_size) - } else { - 0 - }; - - // 应用分页 - let page = page.max(1); - let start = (page - 1) * page_size; - let end = (start + page_size).min(total); - - let flows = if start < total { - all_flows[start..end].to_vec() - } else { - Vec::new() - }; - - Ok(FlowQueryResult { - flows, - total, - page, - page_size, - total_pages, - has_next: page < total_pages, - has_prev: page > 1, - }) - } - - /// 使用过滤表达式查询 Flow - /// - /// 支持类似 mitmproxy 的过滤表达式语法,如: - /// - `~m claude` - 模型名称包含 "claude" - /// - `~p kiro & ~m claude` - 提供商为 kiro 且模型包含 claude - /// - `~e | ~latency >5s` - 有错误或延迟超过 5 秒 - /// - /// # 参数 - /// - `filter_expr`: 过滤表达式字符串 - /// - `sort_by`: 排序字段 - /// - `sort_desc`: 是否降序 - /// - `page`: 页码(从 1 开始) - /// - `page_size`: 每页大小 - /// - /// # 返回 - /// - `Ok(FlowQueryResult)` - 查询结果 - /// - `Err(QueryWithExpressionError)` - 解析或查询错误 - pub async fn query_with_expression( - &self, - filter_expr: &str, - sort_by: FlowSortBy, - sort_desc: bool, - page: usize, - page_size: usize, - ) -> Result { - // 解析过滤表达式 - let expr = FilterParser::parse(filter_expr)?; - - // 编译为过滤函数 - let filter_fn = FilterParser::compile(&expr); - - // 从内存获取所有 Flow 并应用过滤 - let memory_flows = { - let store = self.memory_store.read().await; - let all_flows = store.query(&FlowFilter::default()); - all_flows - .into_iter() - .filter(|f| filter_fn(f)) - .collect::>() - }; - - let mut all_flows = memory_flows; - - // 如果内存数据不足,从文件补充 - let memory_count = all_flows.len(); - let needed = page * page_size; - - if memory_count < needed { - // 从文件存储获取更多数据 - let file_flows = self - .file_store - .query(&FlowFilter::default(), needed * 2, 0)?; - - // 合并并去重(以 ID 为准),同时应用过滤 - let memory_ids: std::collections::HashSet<_> = - all_flows.iter().map(|f| f.id.clone()).collect(); - - for flow in file_flows { - if !memory_ids.contains(&flow.id) && filter_fn(&flow) { - all_flows.push(flow); - } - } - } - - // 排序 - Self::sort_flows(&mut all_flows, sort_by, sort_desc); - - // 计算分页 - let total = all_flows.len(); - let total_pages = if page_size > 0 { - total.div_ceil(page_size) - } else { - 0 - }; - - // 应用分页 - let page = page.max(1); - let start = (page - 1) * page_size; - let end = (start + page_size).min(total); - - let flows = if start < total { - all_flows[start..end].to_vec() - } else { - Vec::new() - }; - - Ok(FlowQueryResult { - flows, - total, - page, - page_size, - total_pages, - has_next: page < total_pages, - has_prev: page > 1, - }) - } - - /// 排序 Flow 列表 - fn sort_flows(flows: &mut [LLMFlow], sort_by: FlowSortBy, desc: bool) { - flows.sort_by(|a, b| { - let cmp = match sort_by { - FlowSortBy::CreatedAt => a.timestamps.created.cmp(&b.timestamps.created), - FlowSortBy::Duration => a.timestamps.duration_ms.cmp(&b.timestamps.duration_ms), - FlowSortBy::TotalTokens => { - let a_tokens = a.response.as_ref().map_or(0, |r| r.usage.total_tokens); - let b_tokens = b.response.as_ref().map_or(0, |r| r.usage.total_tokens); - a_tokens.cmp(&b_tokens) - } - FlowSortBy::ContentLength => { - let a_len = a.response.as_ref().map_or(0, |r| r.content.len()); - let b_len = b.response.as_ref().map_or(0, |r| r.content.len()); - a_len.cmp(&b_len) - } - FlowSortBy::Model => a.request.model.cmp(&b.request.model), - }; - - if desc { - cmp.reverse() - } else { - cmp - } - }); - } - - /// 全文搜索 - /// - /// 使用 SQLite FTS5 进行全文搜索 - /// - /// # 参数 - /// - `query`: 搜索关键词 - /// - `limit`: 最大返回数量 - pub async fn search( - &self, - query: &str, - limit: usize, - ) -> Result, FileStoreError> { - // 先在内存中搜索 - let memory_results = self.search_in_memory(query, limit).await; - - // 如果内存结果不足,在文件中搜索 - if memory_results.len() < limit { - let file_results = self.search_in_file(query, limit - memory_results.len())?; - - // 合并结果 - let mut all_results = memory_results; - let existing_ids: std::collections::HashSet<_> = - all_results.iter().map(|r| r.id.clone()).collect(); - - for result in file_results { - if !existing_ids.contains(&result.id) { - all_results.push(result); - } - } - - // 按分数排序 - all_results.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap_or(Ordering::Equal)); - - Ok(all_results) - } else { - Ok(memory_results) - } - } - - /// 在内存中搜索 - async fn search_in_memory(&self, query: &str, limit: usize) -> Vec { - let store = self.memory_store.read().await; - let query_lower = query.to_lowercase(); - - let mut results = Vec::new(); - - // 获取所有 Flow 并手动搜索 - let all_flows = store.get_recent(10000); // 获取足够多的 Flow - - for flow in all_flows { - // 检查是否匹配搜索条件 - let mut matches = false; - let mut match_text = String::new(); - - // 搜索 Flow ID - if flow.id.to_lowercase().contains(&query_lower) { - matches = true; - match_text = flow.id.clone(); - } - - // 搜索模型名称 - if !matches && flow.request.model.to_lowercase().contains(&query_lower) { - matches = true; - match_text = flow.request.model.clone(); - } - - // 搜索响应内容 - if !matches { - if let Some(ref response) = flow.response { - if response.content.to_lowercase().contains(&query_lower) { - matches = true; - match_text = response.content.clone(); - } - } - } - - // 搜索请求消息 - if !matches { - for message in &flow.request.messages { - let message_text = message.content.get_all_text(); - if message_text.to_lowercase().contains(&query_lower) { - matches = true; - match_text = message_text; - break; - } - } - } - - if matches { - let snippet = Self::extract_snippet(&match_text, &query_lower, 100); - let score = Self::calculate_score(&match_text, &query_lower); - - results.push(FlowSearchResult { - id: flow.id, - created_at: flow.timestamps.created, - model: flow.request.model, - provider: format!("{:?}", flow.metadata.provider), - snippet, - score, - }); - - if results.len() >= limit { - break; - } - } - } - - // 按分数排序 - results.sort_by(|a, b| { - b.score - .partial_cmp(&a.score) - .unwrap_or(std::cmp::Ordering::Equal) - }); - - results - } - - /// 在文件中搜索(使用 SQLite FTS5) - fn search_in_file( - &self, - query: &str, - limit: usize, - ) -> Result, FileStoreError> { - let fts_results = self.file_store.search(query, limit)?; - - let results: Vec = fts_results - .into_iter() - .filter_map(|r| { - // 解析创建时间 - let created_at = chrono::DateTime::parse_from_rfc3339(&r.created_at) - .ok()? - .with_timezone(&Utc); - - Some(FlowSearchResult { - id: r.id, - created_at, - model: r.model, - provider: r.provider, - snippet: r.snippet, - score: 1.0, // FTS5 已经按 rank 排序 - }) - }) - .collect(); - - Ok(results) - } - - /// 提取匹配片段 - fn extract_snippet(content: &str, query: &str, max_len: usize) -> String { - let content_lower = content.to_lowercase(); - - if let Some(pos) = content_lower.find(query) { - let start = pos.saturating_sub(max_len / 2); - let end = (pos + query.len() + max_len / 2).min(content.len()); - - let mut snippet = String::new(); - if start > 0 { - snippet.push_str("..."); - } - snippet.push_str(&content[start..end]); - if end < content.len() { - snippet.push_str("..."); - } - snippet - } else { - content.chars().take(max_len).collect() - } - } - - /// 计算匹配分数 - fn calculate_score(content: &str, query: &str) -> f64 { - let content_lower = content.to_lowercase(); - let count = content_lower.matches(query).count(); - - // 简单的 TF 分数 - if content.is_empty() { - 0.0 - } else { - (count as f64) / (content.len() as f64) * 1000.0 - } - } - - /// 获取统计信息 - /// - /// # 参数 - /// - `filter`: 过滤条件(可选) - pub async fn get_stats(&self, filter: &FlowFilter) -> FlowStats { - // 从内存获取 Flow - let flows = { - let store = self.memory_store.read().await; - store.query(filter) - }; - - Self::calculate_stats(&flows) - } - - /// 计算统计信息 - fn calculate_stats(flows: &[LLMFlow]) -> FlowStats { - if flows.is_empty() { - return FlowStats::default(); - } - - let total = flows.len(); - let mut successful = 0; - let mut failed = 0; - let mut total_latency: u64 = 0; - let mut min_latency = u64::MAX; - let mut max_latency = 0u64; - let mut total_input_tokens: u64 = 0; - let mut total_output_tokens: u64 = 0; - - // 按提供商和模型分组 - let mut provider_map: std::collections::HashMap = - std::collections::HashMap::new(); - let mut model_map: std::collections::HashMap = - std::collections::HashMap::new(); - let mut state_map: std::collections::HashMap = - std::collections::HashMap::new(); - - for flow in flows { - // 状态统计 - let state_str = format!("{:?}", flow.state); - *state_map.entry(state_str).or_insert(0) += 1; - - // 成功/失败统计 - match flow.state { - FlowState::Completed => successful += 1, - FlowState::Failed => failed += 1, - _ => {} - } - - // 延迟统计 - let latency = flow.timestamps.duration_ms; - total_latency += latency; - min_latency = min_latency.min(latency); - max_latency = max_latency.max(latency); - - // Token 统计 - if let Some(ref response) = flow.response { - total_input_tokens += response.usage.input_tokens as u64; - total_output_tokens += response.usage.output_tokens as u64; - } - - // 按提供商分组 - let provider_str = format!("{:?}", flow.metadata.provider); - let provider_entry = provider_map.entry(provider_str).or_insert((0, 0, 0)); - provider_entry.0 += 1; - if flow.state == FlowState::Completed { - provider_entry.1 += 1; - } - provider_entry.2 += latency; - - // 按模型分组 - let model_entry = model_map - .entry(flow.request.model.clone()) - .or_insert((0, 0, 0)); - model_entry.0 += 1; - if flow.state == FlowState::Completed { - model_entry.1 += 1; - } - model_entry.2 += latency; - } - - // 构建统计结果 - let by_provider: Vec = provider_map - .into_iter() - .map(|(provider, (count, success, latency))| ProviderStats { - provider, - count, - success_rate: if count > 0 { - success as f64 / count as f64 - } else { - 0.0 - }, - avg_latency_ms: if count > 0 { - latency as f64 / count as f64 - } else { - 0.0 - }, - }) - .collect(); - - let by_model: Vec = model_map - .into_iter() - .map(|(model, (count, success, latency))| ModelStats { - model, - count, - success_rate: if count > 0 { - success as f64 / count as f64 - } else { - 0.0 - }, - avg_latency_ms: if count > 0 { - latency as f64 / count as f64 - } else { - 0.0 - }, - }) - .collect(); - - let by_state: Vec = state_map - .into_iter() - .map(|(state, count)| StateStats { state, count }) - .collect(); - - FlowStats { - total_requests: total, - successful_requests: successful, - failed_requests: failed, - success_rate: if total > 0 { - successful as f64 / total as f64 - } else { - 0.0 - }, - avg_latency_ms: if total > 0 { - total_latency as f64 / total as f64 - } else { - 0.0 - }, - min_latency_ms: if min_latency == u64::MAX { - 0 - } else { - min_latency - }, - max_latency_ms: max_latency, - total_input_tokens, - total_output_tokens, - avg_input_tokens: if total > 0 { - total_input_tokens as f64 / total as f64 - } else { - 0.0 - }, - avg_output_tokens: if total > 0 { - total_output_tokens as f64 / total as f64 - } else { - 0.0 - }, - by_provider, - by_model, - by_state, - } - } - - /// 根据 ID 获取单个 Flow - pub async fn get_flow(&self, id: &str) -> Result, FileStoreError> { - // 先从内存查找 - { - let store = self.memory_store.read().await; - if let Some(flow_lock) = store.get(id) { - if let Ok(flow) = flow_lock.read() { - return Ok(Some(flow.clone())); - } - } - } - - // 从文件查找 - self.file_store.get(id) - } - - /// 获取最近的 Flow - pub async fn get_recent(&self, limit: usize) -> Vec { - let store = self.memory_store.read().await; - store.get_recent(limit) - } -} - -// ============================================================================ -// 测试模块 -// ============================================================================ - -#[cfg(test)] -mod tests { - use super::*; - use crate::flow_monitor::models::{ - FlowMetadata, FlowType, LLMRequest, LLMResponse, RequestParameters, TokenUsage, - }; - use crate::ProviderType; - - /// 创建测试用的 Flow - fn create_test_flow( - id: &str, - model: &str, - provider: ProviderType, - state: FlowState, - ) -> LLMFlow { - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model: model.to_string(), - parameters: RequestParameters { - stream: false, - ..Default::default() - }, - ..Default::default() - }; - - let metadata = FlowMetadata { - provider, - ..Default::default() - }; - - let mut flow = LLMFlow::new(id.to_string(), FlowType::ChatCompletions, request, metadata); - flow.state = state; - flow - } - - #[test] - fn test_flow_sort_by_created_at() { - let mut flows = vec![ - create_test_flow( - "flow-1", - "gpt-4", - ProviderType::OpenAI, - FlowState::Completed, - ), - create_test_flow( - "flow-2", - "gpt-4", - ProviderType::OpenAI, - FlowState::Completed, - ), - create_test_flow( - "flow-3", - "gpt-4", - ProviderType::OpenAI, - FlowState::Completed, - ), - ]; - - // 设置不同的创建时间 - flows[0].timestamps.created = Utc::now() - chrono::Duration::hours(2); - flows[1].timestamps.created = Utc::now() - chrono::Duration::hours(1); - flows[2].timestamps.created = Utc::now(); - - // 升序排序 - FlowQueryService::sort_flows(&mut flows, FlowSortBy::CreatedAt, false); - assert_eq!(flows[0].id, "flow-1"); - assert_eq!(flows[1].id, "flow-2"); - assert_eq!(flows[2].id, "flow-3"); - - // 降序排序 - FlowQueryService::sort_flows(&mut flows, FlowSortBy::CreatedAt, true); - assert_eq!(flows[0].id, "flow-3"); - assert_eq!(flows[1].id, "flow-2"); - assert_eq!(flows[2].id, "flow-1"); - } - - #[test] - fn test_flow_sort_by_duration() { - let mut flows = vec![ - create_test_flow( - "flow-1", - "gpt-4", - ProviderType::OpenAI, - FlowState::Completed, - ), - create_test_flow( - "flow-2", - "gpt-4", - ProviderType::OpenAI, - FlowState::Completed, - ), - create_test_flow( - "flow-3", - "gpt-4", - ProviderType::OpenAI, - FlowState::Completed, - ), - ]; - - flows[0].timestamps.duration_ms = 100; - flows[1].timestamps.duration_ms = 300; - flows[2].timestamps.duration_ms = 200; - - // 升序排序 - FlowQueryService::sort_flows(&mut flows, FlowSortBy::Duration, false); - assert_eq!(flows[0].timestamps.duration_ms, 100); - assert_eq!(flows[1].timestamps.duration_ms, 200); - assert_eq!(flows[2].timestamps.duration_ms, 300); - } - - #[test] - fn test_flow_sort_by_model() { - let mut flows = vec![ - create_test_flow( - "flow-1", - "gpt-4", - ProviderType::OpenAI, - FlowState::Completed, - ), - create_test_flow( - "flow-2", - "claude-3", - ProviderType::Claude, - FlowState::Completed, - ), - create_test_flow( - "flow-3", - "gemini-pro", - ProviderType::Gemini, - FlowState::Completed, - ), - ]; - - // 升序排序 - FlowQueryService::sort_flows(&mut flows, FlowSortBy::Model, false); - assert_eq!(flows[0].request.model, "claude-3"); - assert_eq!(flows[1].request.model, "gemini-pro"); - assert_eq!(flows[2].request.model, "gpt-4"); - } - - #[test] - fn test_calculate_stats() { - let mut flows = vec![ - create_test_flow( - "flow-1", - "gpt-4", - ProviderType::OpenAI, - FlowState::Completed, - ), - create_test_flow("flow-2", "gpt-4", ProviderType::OpenAI, FlowState::Failed), - create_test_flow( - "flow-3", - "claude-3", - ProviderType::Claude, - FlowState::Completed, - ), - ]; - - // 设置延迟 - flows[0].timestamps.duration_ms = 100; - flows[1].timestamps.duration_ms = 200; - flows[2].timestamps.duration_ms = 150; - - // 设置响应 - flows[0].response = Some(LLMResponse { - usage: TokenUsage { - input_tokens: 100, - output_tokens: 50, - total_tokens: 150, - ..Default::default() - }, - ..Default::default() - }); - flows[2].response = Some(LLMResponse { - usage: TokenUsage { - input_tokens: 200, - output_tokens: 100, - total_tokens: 300, - ..Default::default() - }, - ..Default::default() - }); - - let stats = FlowQueryService::calculate_stats(&flows); - - assert_eq!(stats.total_requests, 3); - assert_eq!(stats.successful_requests, 2); - assert_eq!(stats.failed_requests, 1); - assert!((stats.success_rate - 2.0 / 3.0).abs() < 0.001); - assert_eq!(stats.min_latency_ms, 100); - assert_eq!(stats.max_latency_ms, 200); - assert_eq!(stats.total_input_tokens, 300); - assert_eq!(stats.total_output_tokens, 150); - } - - #[test] - fn test_extract_snippet() { - let content = "This is a test content with some keywords for searching."; - - let snippet = FlowQueryService::extract_snippet(content, "keywords", 20); - assert!(snippet.contains("keywords")); - - let snippet = FlowQueryService::extract_snippet(content, "notfound", 20); - assert_eq!(snippet, "This is a test conte"); - } - - #[test] - fn test_calculate_score() { - let content = "hello world hello"; - let score = FlowQueryService::calculate_score(content, "hello"); - assert!(score > 0.0); - - let score_empty = FlowQueryService::calculate_score("", "hello"); - assert_eq!(score_empty, 0.0); - } - - #[test] - fn test_flow_query_result_empty() { - let result = FlowQueryResult::empty(1, 10); - assert!(result.flows.is_empty()); - assert_eq!(result.total, 0); - assert_eq!(result.page, 1); - assert_eq!(result.page_size, 10); - assert!(!result.has_next); - assert!(!result.has_prev); - } -} - -// ============================================================================ -// 属性测试模块 -// ============================================================================ - -#[cfg(test)] -mod property_tests { - #![allow(dead_code)] - use super::*; - use crate::flow_monitor::models::{ - FlowMetadata, FlowType, LLMRequest, LLMResponse, RequestParameters, TokenUsage, - }; - use crate::ProviderType; - use proptest::prelude::*; - - // ======================================================================== - // 生成器 - // ======================================================================== - - /// 生成随机的 ProviderType - fn arb_provider_type() -> impl Strategy { - prop_oneof![ - Just(ProviderType::Kiro), - Just(ProviderType::Gemini), - Just(ProviderType::OpenAI), - Just(ProviderType::Claude), - Just(ProviderType::Antigravity), - ] - } - - /// 生成随机的 FlowState - fn arb_flow_state() -> impl Strategy { - prop_oneof![ - Just(FlowState::Pending), - Just(FlowState::Streaming), - Just(FlowState::Completed), - Just(FlowState::Failed), - Just(FlowState::Cancelled), - ] - } - - /// 生成随机的 FlowSortBy - fn arb_sort_by() -> impl Strategy { - prop_oneof![ - Just(FlowSortBy::CreatedAt), - Just(FlowSortBy::Duration), - Just(FlowSortBy::TotalTokens), - Just(FlowSortBy::ContentLength), - Just(FlowSortBy::Model), - ] - } - - /// 生成随机的模型名称 - fn arb_model_name() -> impl Strategy { - prop_oneof![ - Just("gpt-4".to_string()), - Just("gpt-4-turbo".to_string()), - Just("gpt-3.5-turbo".to_string()), - Just("claude-3-opus".to_string()), - Just("claude-3-sonnet".to_string()), - Just("gemini-pro".to_string()), - ] - } - - /// 生成随机的 LLMFlow - fn arb_llm_flow() -> impl Strategy { - ( - "[a-f0-9]{8}", - arb_model_name(), - arb_provider_type(), - arb_flow_state(), - 0u64..10000u64, - 0u32..1000u32, - 0u32..500u32, - ) - .prop_map( - |(id, model, provider, state, duration, input_tokens, output_tokens)| { - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model, - parameters: RequestParameters { - stream: false, - ..Default::default() - }, - ..Default::default() - }; - - let metadata = FlowMetadata { - provider, - ..Default::default() - }; - - let mut flow = LLMFlow::new(id, FlowType::ChatCompletions, request, metadata); - flow.state = state; - flow.timestamps.duration_ms = duration; - - if flow.state == FlowState::Completed { - flow.response = Some(LLMResponse { - usage: TokenUsage { - input_tokens, - output_tokens, - total_tokens: input_tokens + output_tokens, - ..Default::default() - }, - ..Default::default() - }); - } - - flow - }, - ) - } - - // ======================================================================== - // 属性测试 - // ======================================================================== - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: llm-flow-monitor, Property 5: 过滤正确性** - /// **Validates: Requirements 4.1-4.9** - /// - /// *对于任意* 过滤条件和 Flow 集合,查询返回的所有 Flow 都应该满足该过滤条件。 - #[test] - fn prop_filter_correctness( - provider in arb_provider_type(), - ) { - // 创建不同 Provider 的 Flow - let providers = [ProviderType::OpenAI, - ProviderType::Claude, - ProviderType::Gemini, - ProviderType::Kiro]; - - let mut flows: Vec = Vec::new(); - for (i, p) in providers.iter().enumerate() { - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model: "gpt-4".to_string(), - ..Default::default() - }; - let metadata = FlowMetadata { - provider: *p, - ..Default::default() - }; - let flow = LLMFlow::new(format!("flow-{i}"), FlowType::ChatCompletions, request, metadata); - flows.push(flow); - } - - // 按 Provider 过滤 - let filter = FlowFilter { - providers: Some(vec![provider]), - ..Default::default() - }; - - let filtered: Vec<&LLMFlow> = flows.iter().filter(|f| filter.matches(f)).collect(); - - // 验证所有结果都匹配过滤条件 - for flow in &filtered { - prop_assert_eq!( - flow.metadata.provider, - provider, - "查询结果的 Provider 应该匹配过滤条件" - ); - } - } - - /// **Feature: llm-flow-monitor, Property 6: 排序正确性** - /// **Validates: Requirements 4.10** - /// - /// *对于任意* 排序选项和 Flow 集合,查询返回的 Flow 列表应该按指定字段正确排序。 - #[test] - fn prop_sort_correctness( - sort_by in arb_sort_by(), - desc in any::(), - ) { - // 创建多个 Flow - let mut flows: Vec = Vec::new(); - for i in 0..10 { - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model: format!("model-{}", i % 3), - ..Default::default() - }; - let metadata = FlowMetadata::default(); - let mut flow = LLMFlow::new(format!("flow-{i}"), FlowType::ChatCompletions, request, metadata); - flow.timestamps.duration_ms = (i * 100) as u64; - flow.timestamps.created = Utc::now() - chrono::Duration::minutes(i as i64); - - if i % 2 == 0 { - flow.response = Some(LLMResponse { - content: "x".repeat(i * 10), - usage: TokenUsage { - input_tokens: (i * 10) as u32, - output_tokens: (i * 5) as u32, - total_tokens: (i * 15) as u32, - ..Default::default() - }, - ..Default::default() - }); - } - - flows.push(flow); - } - - // 排序 - FlowQueryService::sort_flows(&mut flows, sort_by, desc); - - // 验证排序正确性 - for i in 1..flows.len() { - let cmp = match sort_by { - FlowSortBy::CreatedAt => flows[i-1].timestamps.created.cmp(&flows[i].timestamps.created), - FlowSortBy::Duration => flows[i-1].timestamps.duration_ms.cmp(&flows[i].timestamps.duration_ms), - FlowSortBy::TotalTokens => { - let a = flows[i-1].response.as_ref().map_or(0, |r| r.usage.total_tokens); - let b = flows[i].response.as_ref().map_or(0, |r| r.usage.total_tokens); - a.cmp(&b) - } - FlowSortBy::ContentLength => { - let a = flows[i-1].response.as_ref().map_or(0, |r| r.content.len()); - let b = flows[i].response.as_ref().map_or(0, |r| r.content.len()); - a.cmp(&b) - } - FlowSortBy::Model => flows[i-1].request.model.cmp(&flows[i].request.model), - }; - - let expected = if desc { - cmp != std::cmp::Ordering::Less - } else { - cmp != std::cmp::Ordering::Greater - }; - - prop_assert!( - expected, - "排序不正确: {:?} vs {:?} (sort_by={:?}, desc={})", - flows[i-1].id, - flows[i].id, - sort_by, - desc - ); - } - } - - /// **Feature: llm-flow-monitor, Property 7: 分页正确性** - /// **Validates: Requirements 4.11** - /// - /// *对于任意* 分页参数(page, page_size)和 Flow 集合, - /// 返回的结果应该是正确的分页切片,且总数应该正确。 - #[test] - fn prop_pagination_correctness( - total_count in 1usize..=100usize, - page_size in 1usize..=20usize, - page in 1usize..=10usize, - ) { - // 创建 Flow 列表 - let mut all_flows: Vec = Vec::new(); - for i in 0..total_count { - let request = LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - model: "gpt-4".to_string(), - ..Default::default() - }; - let metadata = FlowMetadata::default(); - let flow = LLMFlow::new(format!("flow-{i:04}"), FlowType::ChatCompletions, request, metadata); - all_flows.push(flow); - } - - // 计算分页 - let total = all_flows.len(); - let total_pages = if page_size > 0 { - total.div_ceil(page_size) - } else { - 0 - }; - - let start = (page - 1) * page_size; - let end = (start + page_size).min(total); - - let page_flows = if start < total { - all_flows[start..end].to_vec() - } else { - Vec::new() - }; - - // 验证分页结果 - let expected_count = if start < total { - (end - start).min(page_size) - } else { - 0 - }; - - prop_assert_eq!( - page_flows.len(), - expected_count, - "分页结果数量不正确" - ); - - // 验证 has_next 和 has_prev - let has_next = page < total_pages; - let has_prev = page > 1; - - prop_assert_eq!( - has_next, - page < total_pages, - "has_next 不正确" - ); - - prop_assert_eq!( - has_prev, - page > 1, - "has_prev 不正确" - ); - - // 验证分页内容正确 - for (i, flow) in page_flows.iter().enumerate() { - let expected_id = format!("flow-{:04}", start + i); - prop_assert_eq!( - &flow.id, - &expected_id, - "分页内容不正确" - ); - } - } - } -} diff --git a/src-tauri/src/flow_monitor/quick_filter.rs b/src-tauri/src/flow_monitor/quick_filter.rs deleted file mode 100644 index 9f337a96d..000000000 --- a/src-tauri/src/flow_monitor/quick_filter.rs +++ /dev/null @@ -1,1320 +0,0 @@ -//! 快速过滤器管理器 -//! -//! 该模块实现快速过滤器功能,支持保存和使用常用的过滤条件, -//! 便于快速筛选 Flow。 -//! -//! **Validates: Requirements 6.1-6.7** - -use chrono::{DateTime, Utc}; -use rusqlite::{params, Connection, OptionalExtension}; -use serde::{Deserialize, Serialize}; -use std::path::PathBuf; -use std::sync::Mutex; -use thiserror::Error; -use uuid::Uuid; - -use super::filter_parser::FilterParser; - -// ============================================================================ -// 错误类型 -// ============================================================================ - -/// 快速过滤器错误 -#[derive(Debug, Error)] -pub enum QuickFilterError { - #[error("SQLite 错误: {0}")] - Sqlite(#[from] rusqlite::Error), - - #[error("快速过滤器不存在: {0}")] - FilterNotFound(String), - - #[error("无效的过滤表达式: {0}")] - InvalidFilterExpr(String), - - #[error("JSON 序列化错误: {0}")] - Json(#[from] serde_json::Error), - - #[error("IO 错误: {0}")] - Io(#[from] std::io::Error), - - #[error("无法删除预设过滤器")] - CannotDeletePreset, - - #[error("过滤器名称已存在: {0}")] - DuplicateName(String), -} - -pub type Result = std::result::Result; - -// ============================================================================ -// 数据结构 -// ============================================================================ - -/// 快速过滤器 -/// -/// **Validates: Requirements 6.1, 6.3** -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] -pub struct QuickFilter { - /// 唯一标识符 - pub id: String, - /// 过滤器名称 - pub name: String, - /// 过滤器描述 - #[serde(skip_serializing_if = "Option::is_none")] - pub description: Option, - /// 过滤表达式 - pub filter_expr: String, - /// 分组名称 - #[serde(skip_serializing_if = "Option::is_none")] - pub group: Option, - /// 排序顺序 - pub order: i32, - /// 是否为预设过滤器 - pub is_preset: bool, - /// 创建时间 - pub created_at: DateTime, -} - -impl QuickFilter { - /// 创建新的快速过滤器 - pub fn new( - name: impl Into, - filter_expr: impl Into, - description: Option, - group: Option, - ) -> Self { - Self { - id: Uuid::new_v4().to_string(), - name: name.into(), - description, - filter_expr: filter_expr.into(), - group, - order: 0, - is_preset: false, - created_at: Utc::now(), - } - } - - /// 创建预设过滤器 - pub fn preset( - name: impl Into, - filter_expr: impl Into, - description: impl Into, - order: i32, - ) -> Self { - Self { - id: Uuid::new_v4().to_string(), - name: name.into(), - description: Some(description.into()), - filter_expr: filter_expr.into(), - group: Some("预设".to_string()), - order, - is_preset: true, - created_at: Utc::now(), - } - } -} - -/// 快速过滤器更新 -#[derive(Debug, Clone, Serialize, Deserialize, Default)] -pub struct QuickFilterUpdate { - /// 新名称 - #[serde(skip_serializing_if = "Option::is_none")] - pub name: Option, - /// 新描述 - #[serde(skip_serializing_if = "Option::is_none")] - pub description: Option>, - /// 新过滤表达式 - #[serde(skip_serializing_if = "Option::is_none")] - pub filter_expr: Option, - /// 新分组 - #[serde(skip_serializing_if = "Option::is_none")] - pub group: Option>, - /// 新排序顺序 - #[serde(skip_serializing_if = "Option::is_none")] - pub order: Option, -} - -/// 快速过滤器导出数据 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct QuickFilterExport { - /// 版本号 - pub version: String, - /// 导出时间 - pub exported_at: DateTime, - /// 过滤器列表 - pub filters: Vec, -} - -impl QuickFilterExport { - pub fn new(filters: Vec) -> Self { - Self { - version: "1.0".to_string(), - exported_at: Utc::now(), - filters, - } - } -} - -// ============================================================================ -// 预设过滤器 -// ============================================================================ - -/// 预设快速过滤器 -/// -/// **Validates: Requirements 6.6** -pub const PRESET_FILTERS: &[(&str, &str, &str)] = &[ - ("最近失败", "~e", "显示所有失败的请求"), - ("高延迟", "~latency >5s", "延迟超过 5 秒的请求"), - ("大 Token", "~tokens >10000", "Token 数超过 10000 的请求"), - ("有工具调用", "~t", "包含工具调用的请求"), - ("有思维链", "~k", "包含思维链的请求"), - ("已收藏", "~starred", "已收藏的请求"), -]; - -// ============================================================================ -// 快速过滤器管理器 -// ============================================================================ - -/// 快速过滤器管理器 -/// -/// **Validates: Requirements 6.1-6.7** -pub struct QuickFilterManager { - /// SQLite 连接 - db: Mutex, -} - -impl QuickFilterManager { - /// 创建新的快速过滤器管理器 - /// - /// # Arguments - /// * `db_path` - SQLite 数据库路径 - pub fn new(db_path: PathBuf) -> Result { - // 确保目录存在 - if let Some(parent) = db_path.parent() { - std::fs::create_dir_all(parent)?; - } - - let conn = Connection::open(&db_path)?; - Self::init_database(&conn)?; - - let manager = Self { - db: Mutex::new(conn), - }; - - // 初始化预设过滤器 - manager.init_presets()?; - - Ok(manager) - } - - /// 从现有连接创建快速过滤器管理器(用于测试) - pub fn from_connection(conn: Connection) -> Result { - Self::init_database(&conn)?; - - let manager = Self { - db: Mutex::new(conn), - }; - - Ok(manager) - } - - /// 初始化数据库表 - fn init_database(conn: &Connection) -> Result<()> { - conn.execute_batch( - r#" - -- 快速过滤器表 - CREATE TABLE IF NOT EXISTS quick_filters ( - id TEXT PRIMARY KEY, - name TEXT NOT NULL, - description TEXT, - filter_expr TEXT NOT NULL, - group_name TEXT, - sort_order INTEGER DEFAULT 0, - is_preset INTEGER DEFAULT 0, - created_at TEXT NOT NULL - ); - - CREATE INDEX IF NOT EXISTS idx_quick_filters_name ON quick_filters(name); - CREATE INDEX IF NOT EXISTS idx_quick_filters_group ON quick_filters(group_name); - CREATE INDEX IF NOT EXISTS idx_quick_filters_order ON quick_filters(sort_order); - "#, - )?; - - Ok(()) - } - - /// 初始化预设过滤器 - /// - /// **Validates: Requirements 6.6** - pub fn init_presets(&self) -> Result<()> { - let conn = self.db.lock().unwrap(); - - // 检查是否已有预设过滤器 - let count: i64 = conn.query_row( - "SELECT COUNT(*) FROM quick_filters WHERE is_preset = 1", - [], - |row| row.get(0), - )?; - - if count > 0 { - return Ok(()); - } - - // 插入预设过滤器 - for (i, (name, expr, desc)) in PRESET_FILTERS.iter().enumerate() { - let filter = QuickFilter::preset(*name, *expr, *desc, i as i32); - conn.execute( - r#" - INSERT INTO quick_filters (id, name, description, filter_expr, group_name, sort_order, is_preset, created_at) - VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8) - "#, - params![ - filter.id, - filter.name, - filter.description, - filter.filter_expr, - filter.group, - filter.order, - filter.is_preset as i32, - filter.created_at.to_rfc3339(), - ], - )?; - } - - Ok(()) - } - - /// 保存快速过滤器 - /// - /// **Validates: Requirements 6.1** - /// - /// # Arguments - /// * `name` - 过滤器名称 - /// * `filter_expr` - 过滤表达式 - /// * `description` - 描述(可选) - /// * `group` - 分组(可选) - /// - /// # Returns - /// 新创建的快速过滤器 - pub fn save( - &self, - name: impl Into, - filter_expr: impl Into, - description: Option<&str>, - group: Option<&str>, - ) -> Result { - let name = name.into(); - let filter_expr = filter_expr.into(); - - // 验证过滤表达式 - FilterParser::validate(&filter_expr) - .map_err(|e| QuickFilterError::InvalidFilterExpr(e.to_string()))?; - - let filter = QuickFilter::new( - name, - filter_expr, - description.map(String::from), - group.map(String::from), - ); - - let conn = self.db.lock().unwrap(); - - conn.execute( - r#" - INSERT INTO quick_filters (id, name, description, filter_expr, group_name, sort_order, is_preset, created_at) - VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8) - "#, - params![ - filter.id, - filter.name, - filter.description, - filter.filter_expr, - filter.group, - filter.order, - filter.is_preset as i32, - filter.created_at.to_rfc3339(), - ], - )?; - - Ok(filter) - } - - /// 获取快速过滤器 - /// - /// # Arguments - /// * `id` - 过滤器 ID - /// - /// # Returns - /// 快速过滤器(如果存在) - pub fn get(&self, id: &str) -> Result> { - let conn = self.db.lock().unwrap(); - - let filter: Option<(String, String, Option, String, Option, i32, i32, String)> = conn - .query_row( - r#" - SELECT id, name, description, filter_expr, group_name, sort_order, is_preset, created_at - FROM quick_filters - WHERE id = ?1 - "#, - params![id], - |row| { - Ok(( - row.get(0)?, - row.get(1)?, - row.get(2)?, - row.get(3)?, - row.get(4)?, - row.get(5)?, - row.get(6)?, - row.get(7)?, - )) - }, - ) - .optional()?; - - match filter { - Some((id, name, description, filter_expr, group, order, is_preset, created_at)) => { - Ok(Some(QuickFilter { - id, - name, - description, - filter_expr, - group, - order, - is_preset: is_preset != 0, - created_at: DateTime::parse_from_rfc3339(&created_at) - .map(|dt| dt.with_timezone(&Utc)) - .unwrap_or_else(|_| Utc::now()), - })) - } - None => Ok(None), - } - } - - /// 更新快速过滤器 - /// - /// **Validates: Requirements 6.4** - /// - /// # Arguments - /// * `id` - 过滤器 ID - /// * `updates` - 更新内容 - /// - /// # Returns - /// 更新后的快速过滤器 - pub fn update(&self, id: &str, updates: QuickFilterUpdate) -> Result { - // 验证新的过滤表达式(如果有) - if let Some(ref expr) = updates.filter_expr { - FilterParser::validate(expr) - .map_err(|e| QuickFilterError::InvalidFilterExpr(e.to_string()))?; - } - - let conn = self.db.lock().unwrap(); - - // 检查过滤器是否存在 - let exists: bool = conn - .query_row( - "SELECT 1 FROM quick_filters WHERE id = ?1", - params![id], - |_| Ok(true), - ) - .optional()? - .unwrap_or(false); - - if !exists { - return Err(QuickFilterError::FilterNotFound(id.to_string())); - } - - // 更新各字段 - if let Some(ref name) = updates.name { - conn.execute( - "UPDATE quick_filters SET name = ?1 WHERE id = ?2", - params![name, id], - )?; - } - - if let Some(ref description) = updates.description { - conn.execute( - "UPDATE quick_filters SET description = ?1 WHERE id = ?2", - params![description, id], - )?; - } - - if let Some(ref filter_expr) = updates.filter_expr { - conn.execute( - "UPDATE quick_filters SET filter_expr = ?1 WHERE id = ?2", - params![filter_expr, id], - )?; - } - - if let Some(ref group) = updates.group { - conn.execute( - "UPDATE quick_filters SET group_name = ?1 WHERE id = ?2", - params![group, id], - )?; - } - - if let Some(order) = updates.order { - conn.execute( - "UPDATE quick_filters SET sort_order = ?1 WHERE id = ?2", - params![order, id], - )?; - } - - drop(conn); - - // 返回更新后的过滤器 - self.get(id)? - .ok_or_else(|| QuickFilterError::FilterNotFound(id.to_string())) - } - - /// 删除快速过滤器 - /// - /// **Validates: Requirements 6.4** - /// - /// # Arguments - /// * `id` - 过滤器 ID - pub fn delete(&self, id: &str) -> Result<()> { - let conn = self.db.lock().unwrap(); - - // 检查是否为预设过滤器 - let is_preset: Option = conn - .query_row( - "SELECT is_preset FROM quick_filters WHERE id = ?1", - params![id], - |row| row.get(0), - ) - .optional()?; - - match is_preset { - Some(1) => return Err(QuickFilterError::CannotDeletePreset), - None => return Err(QuickFilterError::FilterNotFound(id.to_string())), - _ => {} - } - - conn.execute("DELETE FROM quick_filters WHERE id = ?1", params![id])?; - - Ok(()) - } - - /// 列出所有快速过滤器 - /// - /// **Validates: Requirements 6.2, 6.5** - /// - /// # Returns - /// 快速过滤器列表(按分组和排序顺序排列) - pub fn list(&self) -> Result> { - let conn = self.db.lock().unwrap(); - - let mut stmt = conn.prepare( - r#" - SELECT id, name, description, filter_expr, group_name, sort_order, is_preset, created_at - FROM quick_filters - ORDER BY group_name ASC, sort_order ASC, created_at ASC - "#, - )?; - - let filters = stmt - .query_map([], |row| { - Ok(QuickFilter { - id: row.get(0)?, - name: row.get(1)?, - description: row.get(2)?, - filter_expr: row.get(3)?, - group: row.get(4)?, - order: row.get(5)?, - is_preset: row.get::<_, i32>(6)? != 0, - created_at: DateTime::parse_from_rfc3339(&row.get::<_, String>(7)?) - .map(|dt| dt.with_timezone(&Utc)) - .unwrap_or_else(|_| Utc::now()), - }) - })? - .filter_map(|r| r.ok()) - .collect(); - - Ok(filters) - } - - /// 按分组列出快速过滤器 - /// - /// **Validates: Requirements 6.5** - /// - /// # Arguments - /// * `group` - 分组名称(None 表示无分组的过滤器) - /// - /// # Returns - /// 快速过滤器列表 - pub fn list_by_group(&self, group: Option<&str>) -> Result> { - let conn = self.db.lock().unwrap(); - - let mut stmt = if group.is_some() { - conn.prepare( - r#" - SELECT id, name, description, filter_expr, group_name, sort_order, is_preset, created_at - FROM quick_filters - WHERE group_name = ?1 - ORDER BY sort_order ASC, created_at ASC - "#, - )? - } else { - conn.prepare( - r#" - SELECT id, name, description, filter_expr, group_name, sort_order, is_preset, created_at - FROM quick_filters - WHERE group_name IS NULL - ORDER BY sort_order ASC, created_at ASC - "#, - )? - }; - - let filters = if let Some(g) = group { - stmt.query_map(params![g], |row| { - Ok(QuickFilter { - id: row.get(0)?, - name: row.get(1)?, - description: row.get(2)?, - filter_expr: row.get(3)?, - group: row.get(4)?, - order: row.get(5)?, - is_preset: row.get::<_, i32>(6)? != 0, - created_at: DateTime::parse_from_rfc3339(&row.get::<_, String>(7)?) - .map(|dt| dt.with_timezone(&Utc)) - .unwrap_or_else(|_| Utc::now()), - }) - })? - .filter_map(|r| r.ok()) - .collect() - } else { - stmt.query_map([], |row| { - Ok(QuickFilter { - id: row.get(0)?, - name: row.get(1)?, - description: row.get(2)?, - filter_expr: row.get(3)?, - group: row.get(4)?, - order: row.get(5)?, - is_preset: row.get::<_, i32>(6)? != 0, - created_at: DateTime::parse_from_rfc3339(&row.get::<_, String>(7)?) - .map(|dt| dt.with_timezone(&Utc)) - .unwrap_or_else(|_| Utc::now()), - }) - })? - .filter_map(|r| r.ok()) - .collect() - }; - - Ok(filters) - } - - /// 获取所有分组名称 - /// - /// **Validates: Requirements 6.5** - /// - /// # Returns - /// 分组名称列表 - pub fn list_groups(&self) -> Result> { - let conn = self.db.lock().unwrap(); - - let mut stmt = conn.prepare( - r#" - SELECT DISTINCT group_name - FROM quick_filters - WHERE group_name IS NOT NULL - ORDER BY group_name ASC - "#, - )?; - - let groups: Vec = stmt - .query_map([], |row| row.get(0))? - .filter_map(|r| r.ok()) - .collect(); - - Ok(groups) - } - - /// 导出快速过滤器 - /// - /// **Validates: Requirements 6.7** - /// - /// # Arguments - /// * `include_presets` - 是否包含预设过滤器 - /// - /// # Returns - /// JSON 格式的导出数据 - pub fn export(&self, include_presets: bool) -> Result { - let filters = if include_presets { - self.list()? - } else { - self.list()?.into_iter().filter(|f| !f.is_preset).collect() - }; - - let export_data = QuickFilterExport::new(filters); - let json = serde_json::to_string_pretty(&export_data)?; - - Ok(json) - } - - /// 导入快速过滤器 - /// - /// **Validates: Requirements 6.7** - /// - /// # Arguments - /// * `data` - JSON 格式的导入数据 - /// * `overwrite` - 是否覆盖同名过滤器 - /// - /// # Returns - /// 导入的快速过滤器列表 - pub fn import(&self, data: &str, overwrite: bool) -> Result> { - let export_data: QuickFilterExport = serde_json::from_str(data)?; - - let mut imported = Vec::new(); - let conn = self.db.lock().unwrap(); - - for mut filter in export_data.filters { - // 跳过预设过滤器 - if filter.is_preset { - continue; - } - - // 验证过滤表达式 - if FilterParser::validate(&filter.filter_expr).is_err() { - continue; - } - - // 检查是否存在同名过滤器 - let existing_id: Option = conn - .query_row( - "SELECT id FROM quick_filters WHERE name = ?1 AND is_preset = 0", - params![filter.name], - |row| row.get(0), - ) - .optional()?; - - if let Some(existing) = existing_id { - if overwrite { - // 更新现有过滤器 - conn.execute( - r#" - UPDATE quick_filters - SET description = ?1, filter_expr = ?2, group_name = ?3, sort_order = ?4 - WHERE id = ?5 - "#, - params![ - filter.description, - filter.filter_expr, - filter.group, - filter.order, - existing, - ], - )?; - filter.id = existing; - } else { - // 跳过已存在的过滤器 - continue; - } - } else { - // 生成新 ID - filter.id = Uuid::new_v4().to_string(); - filter.created_at = Utc::now(); - - conn.execute( - r#" - INSERT INTO quick_filters (id, name, description, filter_expr, group_name, sort_order, is_preset, created_at) - VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8) - "#, - params![ - filter.id, - filter.name, - filter.description, - filter.filter_expr, - filter.group, - filter.order, - filter.is_preset as i32, - filter.created_at.to_rfc3339(), - ], - )?; - } - - imported.push(filter); - } - - Ok(imported) - } - - /// 获取过滤器数量 - pub fn count(&self) -> Result { - let conn = self.db.lock().unwrap(); - let count: i64 = - conn.query_row("SELECT COUNT(*) FROM quick_filters", [], |row| row.get(0))?; - Ok(count as usize) - } - - /// 获取非预设过滤器数量 - pub fn count_custom(&self) -> Result { - let conn = self.db.lock().unwrap(); - let count: i64 = conn.query_row( - "SELECT COUNT(*) FROM quick_filters WHERE is_preset = 0", - [], - |row| row.get(0), - )?; - Ok(count as usize) - } - - /// 按名称查找过滤器 - pub fn find_by_name(&self, name: &str) -> Result> { - let conn = self.db.lock().unwrap(); - - let filter: Option<(String, String, Option, String, Option, i32, i32, String)> = conn - .query_row( - r#" - SELECT id, name, description, filter_expr, group_name, sort_order, is_preset, created_at - FROM quick_filters - WHERE name = ?1 - "#, - params![name], - |row| { - Ok(( - row.get(0)?, - row.get(1)?, - row.get(2)?, - row.get(3)?, - row.get(4)?, - row.get(5)?, - row.get(6)?, - row.get(7)?, - )) - }, - ) - .optional()?; - - match filter { - Some((id, name, description, filter_expr, group, order, is_preset, created_at)) => { - Ok(Some(QuickFilter { - id, - name, - description, - filter_expr, - group, - order, - is_preset: is_preset != 0, - created_at: DateTime::parse_from_rfc3339(&created_at) - .map(|dt| dt.with_timezone(&Utc)) - .unwrap_or_else(|_| Utc::now()), - })) - } - None => Ok(None), - } - } - - /// 清除所有非预设过滤器(用于测试) - #[cfg(test)] - pub fn clear_custom(&self) -> Result<()> { - let conn = self.db.lock().unwrap(); - conn.execute("DELETE FROM quick_filters WHERE is_preset = 0", [])?; - Ok(()) - } -} - -// ============================================================================ -// 测试模块 -// ============================================================================ - -#[cfg(test)] -mod tests { - use super::*; - - fn create_test_manager() -> QuickFilterManager { - let conn = Connection::open_in_memory().unwrap(); - QuickFilterManager::from_connection(conn).unwrap() - } - - #[test] - fn test_save_quick_filter() { - let manager = create_test_manager(); - - let filter = manager - .save( - "Test Filter", - "~e", - Some("A test filter"), - Some("Test Group"), - ) - .unwrap(); - - assert!(!filter.id.is_empty()); - assert_eq!(filter.name, "Test Filter"); - assert_eq!(filter.filter_expr, "~e"); - assert_eq!(filter.description, Some("A test filter".to_string())); - assert_eq!(filter.group, Some("Test Group".to_string())); - assert!(!filter.is_preset); - } - - #[test] - fn test_get_quick_filter() { - let manager = create_test_manager(); - - let created = manager.save("Test", "~e", None, None).unwrap(); - let retrieved = manager.get(&created.id).unwrap(); - - assert!(retrieved.is_some()); - let retrieved = retrieved.unwrap(); - assert_eq!(retrieved.id, created.id); - assert_eq!(retrieved.name, "Test"); - assert_eq!(retrieved.filter_expr, "~e"); - } - - #[test] - fn test_list_quick_filters() { - let manager = create_test_manager(); - - manager.save("Filter 1", "~e", None, None).unwrap(); - manager.save("Filter 2", "~t", None, None).unwrap(); - - let filters = manager.list().unwrap(); - // 包含预设过滤器 - assert!(filters.len() >= 2); - } - - #[test] - fn test_update_quick_filter() { - let manager = create_test_manager(); - - let filter = manager.save("Original", "~e", None, None).unwrap(); - - let updates = QuickFilterUpdate { - name: Some("Updated".to_string()), - description: Some(Some("New description".to_string())), - filter_expr: Some("~t".to_string()), - group: Some(Some("New Group".to_string())), - order: Some(10), - }; - - let updated = manager.update(&filter.id, updates).unwrap(); - - assert_eq!(updated.name, "Updated"); - assert_eq!(updated.description, Some("New description".to_string())); - assert_eq!(updated.filter_expr, "~t"); - assert_eq!(updated.group, Some("New Group".to_string())); - assert_eq!(updated.order, 10); - } - - #[test] - fn test_delete_quick_filter() { - let manager = create_test_manager(); - - let filter = manager.save("Test", "~e", None, None).unwrap(); - manager.delete(&filter.id).unwrap(); - - let retrieved = manager.get(&filter.id).unwrap(); - assert!(retrieved.is_none()); - } - - #[test] - fn test_cannot_delete_preset() { - let manager = create_test_manager(); - manager.init_presets().unwrap(); - - let presets: Vec<_> = manager - .list() - .unwrap() - .into_iter() - .filter(|f| f.is_preset) - .collect(); - assert!(!presets.is_empty()); - - let result = manager.delete(&presets[0].id); - assert!(matches!(result, Err(QuickFilterError::CannotDeletePreset))); - } - - #[test] - fn test_invalid_filter_expr() { - let manager = create_test_manager(); - - let result = manager.save("Invalid", "invalid expression", None, None); - assert!(matches!( - result, - Err(QuickFilterError::InvalidFilterExpr(_)) - )); - } - - #[test] - fn test_filter_not_found() { - let manager = create_test_manager(); - - let result = manager.update("non-existent", QuickFilterUpdate::default()); - assert!(matches!(result, Err(QuickFilterError::FilterNotFound(_)))); - } - - #[test] - fn test_list_by_group() { - let manager = create_test_manager(); - - manager - .save("Filter 1", "~e", None, Some("Group A")) - .unwrap(); - manager - .save("Filter 2", "~t", None, Some("Group A")) - .unwrap(); - manager - .save("Filter 3", "~k", None, Some("Group B")) - .unwrap(); - - let group_a = manager.list_by_group(Some("Group A")).unwrap(); - assert_eq!(group_a.len(), 2); - - let group_b = manager.list_by_group(Some("Group B")).unwrap(); - assert_eq!(group_b.len(), 1); - } - - #[test] - fn test_list_groups() { - let manager = create_test_manager(); - manager.init_presets().unwrap(); - - manager - .save("Filter 1", "~e", None, Some("Custom")) - .unwrap(); - - let groups = manager.list_groups().unwrap(); - assert!(groups.contains(&"预设".to_string())); - assert!(groups.contains(&"Custom".to_string())); - } - - #[test] - fn test_export_import() { - let manager = create_test_manager(); - - manager - .save("Export Test 1", "~e", Some("Desc 1"), Some("Group")) - .unwrap(); - manager - .save("Export Test 2", "~t", Some("Desc 2"), None) - .unwrap(); - - // 导出(不包含预设) - let exported = manager.export(false).unwrap(); - - // 创建新管理器并导入 - let manager2 = create_test_manager(); - let imported = manager2.import(&exported, false).unwrap(); - - assert_eq!(imported.len(), 2); - - // 验证导入的过滤器 - let filter1 = manager2.find_by_name("Export Test 1").unwrap().unwrap(); - assert_eq!(filter1.filter_expr, "~e"); - assert_eq!(filter1.description, Some("Desc 1".to_string())); - } - - #[test] - fn test_import_overwrite() { - let manager = create_test_manager(); - - manager - .save("Test Filter", "~e", Some("Original"), None) - .unwrap(); - - // 创建导出数据 - let export_data = QuickFilterExport::new(vec![QuickFilter::new( - "Test Filter", - "~t", - Some("Updated".to_string()), - None, - )]); - let json = serde_json::to_string(&export_data).unwrap(); - - // 导入并覆盖 - manager.import(&json, true).unwrap(); - - let filter = manager.find_by_name("Test Filter").unwrap().unwrap(); - assert_eq!(filter.filter_expr, "~t"); - assert_eq!(filter.description, Some("Updated".to_string())); - } - - #[test] - fn test_import_no_overwrite() { - let manager = create_test_manager(); - - manager - .save("Test Filter", "~e", Some("Original"), None) - .unwrap(); - - // 创建导出数据 - let export_data = QuickFilterExport::new(vec![QuickFilter::new( - "Test Filter", - "~t", - Some("Updated".to_string()), - None, - )]); - let json = serde_json::to_string(&export_data).unwrap(); - - // 导入但不覆盖 - let imported = manager.import(&json, false).unwrap(); - assert!(imported.is_empty()); - - let filter = manager.find_by_name("Test Filter").unwrap().unwrap(); - assert_eq!(filter.filter_expr, "~e"); - assert_eq!(filter.description, Some("Original".to_string())); - } - - #[test] - fn test_preset_filters_initialized() { - let manager = create_test_manager(); - manager.init_presets().unwrap(); - - let presets: Vec<_> = manager - .list() - .unwrap() - .into_iter() - .filter(|f| f.is_preset) - .collect(); - assert_eq!(presets.len(), PRESET_FILTERS.len()); - - // 验证预设过滤器内容 - for (name, expr, _) in PRESET_FILTERS { - let filter = manager.find_by_name(name).unwrap(); - assert!(filter.is_some(), "Preset filter '{name}' should exist"); - let filter = filter.unwrap(); - assert_eq!(filter.filter_expr, *expr); - assert!(filter.is_preset); - } - } - - #[test] - fn test_count() { - let manager = create_test_manager(); - manager.init_presets().unwrap(); - - let initial_count = manager.count().unwrap(); - assert_eq!(initial_count, PRESET_FILTERS.len()); - - manager.save("Custom 1", "~e", None, None).unwrap(); - manager.save("Custom 2", "~t", None, None).unwrap(); - - assert_eq!(manager.count().unwrap(), initial_count + 2); - assert_eq!(manager.count_custom().unwrap(), 2); - } -} - -// ============================================================================ -// 属性测试模块 -// ============================================================================ - -#[cfg(test)] -mod property_tests { - use super::*; - use proptest::prelude::*; - - // ======================================================================== - // 生成器 - // ======================================================================== - - /// 生成随机的过滤器名称 - fn arb_filter_name() -> impl Strategy { - "[a-zA-Z0-9 _-]{1,50}".prop_filter("Name should not be empty", |s| !s.trim().is_empty()) - } - - /// 生成随机的过滤器描述 - fn arb_filter_description() -> impl Strategy> { - prop::option::of("[a-zA-Z0-9 _-]{0,200}") - } - - /// 生成随机的分组名称 - fn arb_group_name() -> impl Strategy> { - prop::option::of("[a-zA-Z0-9 _-]{1,30}") - } - - /// 生成有效的过滤表达式 - fn arb_valid_filter_expr() -> impl Strategy { - prop_oneof![ - Just("~e".to_string()), - Just("~t".to_string()), - Just("~k".to_string()), - Just("~starred".to_string()), - Just("~latency >5s".to_string()), - Just("~tokens >1000".to_string()), - Just("~s completed".to_string()), - Just("~s failed".to_string()), - Just("~e | ~t".to_string()), - Just("~e & ~t".to_string()), - Just("!~e".to_string()), - "[a-zA-Z0-9_-]{1,20}".prop_map(|s| format!("~m {s}")), - "[a-zA-Z0-9_-]{1,20}".prop_map(|s| format!("~p {s}")), - "[a-zA-Z0-9_-]{1,20}".prop_map(|s| format!("~tag {s}")), - ] - } - - /// 生成随机的快速过滤器 - fn arb_quick_filter() -> impl Strategy, Option)> - { - ( - arb_filter_name(), - arb_valid_filter_expr(), - arb_filter_description(), - arb_group_name(), - ) - } - - /// 生成多个快速过滤器 - fn arb_quick_filters( - max_len: usize, - ) -> impl Strategy, Option)>> { - prop::collection::vec(arb_quick_filter(), 1..max_len) - } - - // ======================================================================== - // 属性测试 - // ======================================================================== - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: flow-monitor-enhancement, Property 11: 快速过滤器 Round-Trip** - /// **Validates: Requirements 6.1, 6.2** - /// - /// *对于任意* 快速过滤器,保存后再加载应该得到等价的过滤器。 - #[test] - fn prop_quick_filter_roundtrip( - (name, filter_expr, description, group) in arb_quick_filter() - ) { - let manager = create_test_manager(); - - // 保存过滤器 - let saved = manager.save(&name, &filter_expr, description.as_deref(), group.as_deref()).unwrap(); - - // 加载过滤器 - let loaded = manager.get(&saved.id).unwrap().unwrap(); - - // 验证等价性 - prop_assert_eq!(saved.id, loaded.id); - prop_assert_eq!(saved.name, loaded.name); - prop_assert_eq!(saved.filter_expr, loaded.filter_expr); - prop_assert_eq!(saved.description, loaded.description); - prop_assert_eq!(saved.group, loaded.group); - prop_assert_eq!(saved.is_preset, loaded.is_preset); - } - - /// **Feature: flow-monitor-enhancement, Property 12: 快速过滤器导入导出 Round-Trip** - /// **Validates: Requirements 6.7** - /// - /// *对于任意* 快速过滤器集合,导出后再导入应该得到等价的集合。 - #[test] - fn prop_quick_filter_export_import_roundtrip( - filters in arb_quick_filters(10) - ) { - let manager1 = create_test_manager(); - - // 保存所有过滤器 - let mut saved_filters = Vec::new(); - for (name, filter_expr, description, group) in &filters { - // 使用唯一名称避免冲突 - let unique_name = format!("{}_{}", name, saved_filters.len()); - let filter = manager1.save(&unique_name, filter_expr, description.as_deref(), group.as_deref()).unwrap(); - saved_filters.push(filter); - } - - // 导出(不包含预设) - let exported = manager1.export(false).unwrap(); - - // 创建新管理器并导入 - let manager2 = create_test_manager(); - let imported = manager2.import(&exported, false).unwrap(); - - // 验证导入数量 - prop_assert_eq!(imported.len(), saved_filters.len()); - - // 验证每个过滤器的内容 - for saved in &saved_filters { - let found = manager2.find_by_name(&saved.name).unwrap(); - prop_assert!(found.is_some(), "Filter '{}' should be imported", saved.name); - - let found = found.unwrap(); - prop_assert_eq!(&saved.name, &found.name); - prop_assert_eq!(&saved.filter_expr, &found.filter_expr); - prop_assert_eq!(&saved.description, &found.description); - prop_assert_eq!(&saved.group, &found.group); - } - } - - /// 过滤器更新后应该保持一致性 - #[test] - fn prop_filter_update_consistency( - (name, filter_expr, description, group) in arb_quick_filter(), - (new_name, new_filter_expr, new_description, new_group) in arb_quick_filter() - ) { - let manager = create_test_manager(); - - // 保存原始过滤器 - let original = manager.save(&name, &filter_expr, description.as_deref(), group.as_deref()).unwrap(); - - // 更新过滤器 - let updates = QuickFilterUpdate { - name: Some(new_name.clone()), - filter_expr: Some(new_filter_expr.clone()), - description: Some(new_description.clone()), - group: Some(new_group.clone()), - order: None, - }; - - let updated = manager.update(&original.id, updates).unwrap(); - - // 验证更新后的值 - prop_assert_eq!(updated.id, original.id); - prop_assert_eq!(updated.name, new_name); - prop_assert_eq!(updated.filter_expr, new_filter_expr); - prop_assert_eq!(updated.description, new_description); - prop_assert_eq!(updated.group, new_group); - } - - /// 删除后过滤器应该不存在 - #[test] - fn prop_filter_delete( - (name, filter_expr, description, group) in arb_quick_filter() - ) { - let manager = create_test_manager(); - - // 保存过滤器 - let filter = manager.save(&name, &filter_expr, description.as_deref(), group.as_deref()).unwrap(); - - // 删除过滤器 - manager.delete(&filter.id).unwrap(); - - // 验证不存在 - let found = manager.get(&filter.id).unwrap(); - prop_assert!(found.is_none()); - } - - /// 列表应该包含所有保存的过滤器 - #[test] - fn prop_list_contains_all( - filters in arb_quick_filters(5) - ) { - let manager = create_test_manager(); - - // 保存所有过滤器 - let mut saved_ids = Vec::new(); - for (i, (name, filter_expr, description, group)) in filters.iter().enumerate() { - let unique_name = format!("{name}_{i}"); - let filter = manager.save(&unique_name, filter_expr, description.as_deref(), group.as_deref()).unwrap(); - saved_ids.push(filter.id); - } - - // 获取列表 - let list = manager.list().unwrap(); - - // 验证所有保存的过滤器都在列表中 - for id in &saved_ids { - prop_assert!( - list.iter().any(|f| &f.id == id), - "Filter with id '{}' should be in list", - id - ); - } - } - } - - fn create_test_manager() -> QuickFilterManager { - let conn = Connection::open_in_memory().unwrap(); - QuickFilterManager::from_connection(conn).unwrap() - } -} diff --git a/src-tauri/src/flow_monitor/replayer.rs b/src-tauri/src/flow_monitor/replayer.rs deleted file mode 100644 index ffe5142c5..000000000 --- a/src-tauri/src/flow_monitor/replayer.rs +++ /dev/null @@ -1,1017 +0,0 @@ -//! Flow 重放器 -//! -//! 该模块实现 LLM Flow 的重放功能,允许用户重新发送历史请求。 -//! -//! # 功能 -//! -//! - 重放单个 Flow -//! - 批量重放多个 Flow -//! - 支持修改请求参数后重放 -//! - 支持选择不同的凭证 -//! - 重放的 Flow 会被标记为 "replay" - -use chrono::{DateTime, Utc}; -use reqwest::Client; -use serde::{Deserialize, Serialize}; -use std::collections::HashMap; -use std::sync::Arc; -use std::time::Duration; -use tokio::time::sleep; -use uuid::Uuid; - -use super::models::{ - FlowAnnotations, FlowMetadata, FlowState, FlowTimestamps, LLMFlow, LLMRequest, LLMResponse, - Message, RequestParameters, TokenUsage, -}; -use super::monitor::FlowMonitor; -use crate::database::DbConnection; -use crate::ProviderPoolService; -use crate::ProviderType; - -// ============================================================================ -// 配置结构 -// ============================================================================ - -/// 重放配置 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ReplayConfig { - /// 使用的凭证 ID(可选,为空时使用原始凭证或自动选择) - #[serde(skip_serializing_if = "Option::is_none")] - pub credential_id: Option, - /// 请求修改(可选) - #[serde(skip_serializing_if = "Option::is_none")] - pub modify_request: Option, - /// 重放间隔(毫秒),用于批量重放时避免触发速率限制 - #[serde(default = "default_interval_ms")] - pub interval_ms: u64, -} - -fn default_interval_ms() -> u64 { - 1000 // 默认 1 秒间隔 -} - -impl Default for ReplayConfig { - fn default() -> Self { - Self { - credential_id: None, - modify_request: None, - interval_ms: default_interval_ms(), - } - } -} - -/// 请求修改 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct RequestModification { - /// 修改模型名称 - #[serde(skip_serializing_if = "Option::is_none")] - pub model: Option, - /// 修改消息列表 - #[serde(skip_serializing_if = "Option::is_none")] - pub messages: Option>, - /// 修改请求参数 - #[serde(skip_serializing_if = "Option::is_none")] - pub parameters: Option, - /// 修改系统提示词 - #[serde(skip_serializing_if = "Option::is_none")] - pub system_prompt: Option, -} - -// ============================================================================ -// 重放结果 -// ============================================================================ - -/// 重放结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ReplayResult { - /// 原始 Flow ID - pub original_flow_id: String, - /// 重放生成的新 Flow ID - pub replay_flow_id: String, - /// 是否成功 - pub success: bool, - /// 错误信息(如果失败) - #[serde(skip_serializing_if = "Option::is_none")] - pub error: Option, - /// 重放开始时间 - pub started_at: DateTime, - /// 重放结束时间 - pub completed_at: DateTime, - /// 耗时(毫秒) - pub duration_ms: u64, -} - -impl ReplayResult { - /// 创建成功的重放结果 - pub fn success( - original_flow_id: String, - replay_flow_id: String, - started_at: DateTime, - completed_at: DateTime, - ) -> Self { - let duration_ms = (completed_at - started_at).num_milliseconds().max(0) as u64; - Self { - original_flow_id, - replay_flow_id, - success: true, - error: None, - started_at, - completed_at, - duration_ms, - } - } - - /// 创建失败的重放结果 - pub fn failure( - original_flow_id: String, - error: String, - started_at: DateTime, - completed_at: DateTime, - ) -> Self { - let duration_ms = (completed_at - started_at).num_milliseconds().max(0) as u64; - Self { - original_flow_id, - replay_flow_id: String::new(), - success: false, - error: Some(error), - started_at, - completed_at, - duration_ms, - } - } -} - -// ============================================================================ -// 批量重放结果 -// ============================================================================ - -/// 批量重放结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct BatchReplayResult { - /// 总数 - pub total: usize, - /// 成功数 - pub success_count: usize, - /// 失败数 - pub failure_count: usize, - /// 各个 Flow 的重放结果 - pub results: Vec, - /// 批量重放开始时间 - pub started_at: DateTime, - /// 批量重放结束时间 - pub completed_at: DateTime, - /// 总耗时(毫秒) - pub total_duration_ms: u64, -} - -// ============================================================================ -// 重放器错误 -// ============================================================================ - -/// 重放器错误 -#[derive(Debug, Clone, thiserror::Error, Serialize, Deserialize)] -pub enum ReplayerError { - /// Flow 不存在 - #[error("Flow '{0}' 不存在")] - FlowNotFound(String), - /// 凭证不可用 - #[error("凭证 '{0}' 不可用")] - CredentialUnavailable(String), - /// 请求失败 - #[error("请求失败: {0}")] - RequestFailed(String), - /// 内部错误 - #[error("内部错误: {0}")] - Internal(String), -} - -// ============================================================================ -// Flow 重放器 -// ============================================================================ - -/// Flow 重放器 -/// -/// 负责重放历史 LLM Flow 的核心服务。 -pub struct FlowReplayer { - /// HTTP 客户端 - client: Client, - /// Flow 监控服务 - flow_monitor: Arc, - /// 凭证池服务 - provider_pool: Arc, - /// 数据库连接 - db: DbConnection, -} - -impl FlowReplayer { - /// 创建新的重放器 - pub fn new( - flow_monitor: Arc, - provider_pool: Arc, - db: DbConnection, - ) -> Self { - let client = Client::builder() - .timeout(Duration::from_secs(120)) - .build() - .unwrap_or_default(); - - Self { - client, - flow_monitor, - provider_pool, - db, - } - } - - /// 重放单个 Flow - /// - /// **Validates: Requirements 3.1, 3.3, 3.4** - /// - /// # Arguments - /// * `flow_id` - 要重放的 Flow ID - /// * `config` - 重放配置 - /// - /// # Returns - /// * `Ok(ReplayResult)` - 重放结果 - /// * `Err(ReplayerError)` - 重放失败 - pub async fn replay( - &self, - flow_id: &str, - config: ReplayConfig, - ) -> Result { - let started_at = Utc::now(); - - // 获取原始 Flow - let original_flow = self.get_flow(flow_id).await?; - - // 应用请求修改 - let request = self.apply_modifications(&original_flow.request, &config.modify_request); - - // 确定使用的凭证 - let credential_id = self.resolve_credential(&original_flow, &config).await?; - - // 创建重放 Flow - let replay_flow_id = self - .create_replay_flow(&original_flow, &request, &credential_id) - .await; - - // 执行重放请求 - match self - .execute_replay(&request, &original_flow.metadata, &credential_id) - .await - { - Ok(response) => { - // 更新重放 Flow 的响应 - self.complete_replay_flow(&replay_flow_id, Some(response)) - .await; - let completed_at = Utc::now(); - Ok(ReplayResult::success( - flow_id.to_string(), - replay_flow_id, - started_at, - completed_at, - )) - } - Err(e) => { - // 标记重放 Flow 失败 - self.fail_replay_flow(&replay_flow_id, &e.to_string()).await; - let completed_at = Utc::now(); - Ok(ReplayResult::failure( - flow_id.to_string(), - e.to_string(), - started_at, - completed_at, - )) - } - } - } - - /// 批量重放多个 Flow - /// - /// **Validates: Requirements 3.6, 3.7** - /// - /// # Arguments - /// * `flow_ids` - 要重放的 Flow ID 列表 - /// * `config` - 重放配置 - /// - /// # Returns - /// * `BatchReplayResult` - 批量重放结果 - pub async fn replay_batch( - &self, - flow_ids: &[String], - config: ReplayConfig, - ) -> BatchReplayResult { - let started_at = Utc::now(); - let mut results = Vec::with_capacity(flow_ids.len()); - let mut success_count = 0; - let mut failure_count = 0; - - for (i, flow_id) in flow_ids.iter().enumerate() { - // 执行重放 - let result = match self.replay(flow_id, config.clone()).await { - Ok(r) => r, - Err(e) => { - ReplayResult::failure(flow_id.clone(), e.to_string(), Utc::now(), Utc::now()) - } - }; - - if result.success { - success_count += 1; - } else { - failure_count += 1; - } - - results.push(result); - - // 如果不是最后一个,等待间隔时间 - if i < flow_ids.len() - 1 && config.interval_ms > 0 { - sleep(Duration::from_millis(config.interval_ms)).await; - } - } - - let completed_at = Utc::now(); - let total_duration_ms = (completed_at - started_at).num_milliseconds().max(0) as u64; - - BatchReplayResult { - total: flow_ids.len(), - success_count, - failure_count, - results, - started_at, - completed_at, - total_duration_ms, - } - } - - /// 获取 Flow - async fn get_flow(&self, flow_id: &str) -> Result { - // 先从内存存储获取 - let store = self.flow_monitor.memory_store(); - let store_guard = store.read().await; - - if let Some(flow_lock) = store_guard.get(flow_id) { - let flow = flow_lock.read().unwrap().clone(); - return Ok(flow); - } - drop(store_guard); - - // 再从文件存储获取 - if let Some(file_store) = self.flow_monitor.file_store() { - if let Ok(Some(flow)) = file_store.get(flow_id) { - return Ok(flow); - } - } - - Err(ReplayerError::FlowNotFound(flow_id.to_string())) - } - - /// 应用请求修改 - fn apply_modifications( - &self, - original: &LLMRequest, - modification: &Option, - ) -> LLMRequest { - let mut request = original.clone(); - - if let Some(mod_config) = modification { - // 修改模型 - if let Some(ref model) = mod_config.model { - request.model = model.clone(); - } - - // 修改消息 - if let Some(ref messages) = mod_config.messages { - request.messages = messages.clone(); - } - - // 修改参数 - if let Some(ref params) = mod_config.parameters { - request.parameters = params.clone(); - } - - // 修改系统提示词 - if let Some(ref system_prompt) = mod_config.system_prompt { - request.system_prompt = Some(system_prompt.clone()); - } - } - - // 更新时间戳 - request.timestamp = Utc::now(); - - request - } - - /// 解析凭证 - async fn resolve_credential( - &self, - original_flow: &LLMFlow, - config: &ReplayConfig, - ) -> Result, ReplayerError> { - // 如果配置中指定了凭证,使用指定的凭证 - if let Some(ref cred_id) = config.credential_id { - return Ok(Some(cred_id.clone())); - } - - // 否则使用原始 Flow 的凭证 - Ok(original_flow.metadata.credential_id.clone()) - } - - /// 创建重放 Flow - /// - /// **Validates: Requirements 3.2** - async fn create_replay_flow( - &self, - original_flow: &LLMFlow, - request: &LLMRequest, - credential_id: &Option, - ) -> String { - let replay_flow_id = Uuid::new_v4().to_string(); - let now = Utc::now(); - - // 创建重放 Flow 的元数据 - let mut metadata = original_flow.metadata.clone(); - metadata.credential_id = credential_id.clone(); - - // 创建重放 Flow - let replay_flow = LLMFlow { - id: replay_flow_id.clone(), - flow_type: original_flow.flow_type.clone(), - request: request.clone(), - response: None, - error: None, - metadata, - timestamps: FlowTimestamps { - created: now, - request_start: now, - request_end: None, - response_start: None, - response_end: None, - duration_ms: 0, - ttfb_ms: None, - }, - state: FlowState::Pending, - annotations: FlowAnnotations { - marker: Some("🔄".to_string()), // 重放标记 - comment: Some(format!("重放自 Flow: {}", original_flow.id)), - tags: vec!["replay".to_string()], - starred: false, - }, - }; - - // 保存到内存存储 - { - let store = self.flow_monitor.memory_store(); - let mut store_guard = store.write().await; - store_guard.add(replay_flow.clone()); - } - - // 保存到文件存储 - if let Some(file_store) = self.flow_monitor.file_store() { - if let Err(e) = file_store.write(&replay_flow) { - tracing::error!("保存重放 Flow 到文件失败: {}", e); - } - } - - replay_flow_id - } - - /// 执行重放请求 - async fn execute_replay( - &self, - request: &LLMRequest, - metadata: &FlowMetadata, - credential_id: &Option, - ) -> Result { - // 构建请求 URL - let base_url = self.get_base_url(&metadata.provider); - let url = format!("{}{}", base_url, request.path); - - // 获取认证信息 - let auth_header = self - .get_auth_header(&metadata.provider, credential_id) - .await?; - - // 构建请求 - let mut req_builder = self.client.post(&url); - - // 添加认证头 - if let Some(auth) = auth_header { - req_builder = req_builder.header("Authorization", auth); - } - - // 添加其他头 - req_builder = req_builder - .header("Content-Type", "application/json") - .header("Accept", "application/json"); - - // 添加请求体 - req_builder = req_builder.json(&request.body); - - // 发送请求 - let start_time = Utc::now(); - let response = req_builder - .send() - .await - .map_err(|e| ReplayerError::RequestFailed(e.to_string()))?; - - let end_time = Utc::now(); - let status_code = response.status().as_u16(); - let status_text = response.status().to_string(); - - // 获取响应头 - let mut headers = HashMap::new(); - for (key, value) in response.headers() { - if let Ok(v) = value.to_str() { - headers.insert(key.to_string(), v.to_string()); - } - } - - // 获取响应体 - let body_bytes = response - .bytes() - .await - .map_err(|e| ReplayerError::RequestFailed(e.to_string()))?; - let size_bytes = body_bytes.len(); - - // 解析响应体 - let body: serde_json::Value = serde_json::from_slice(&body_bytes).unwrap_or_else(|_| { - serde_json::Value::String(String::from_utf8_lossy(&body_bytes).to_string()) - }); - - // 提取内容 - let content = self.extract_content(&body, &metadata.provider); - - // 提取 token 使用量 - let usage = self.extract_usage(&body, &metadata.provider); - - Ok(LLMResponse { - status_code, - status_text, - headers, - body, - content, - thinking: None, - tool_calls: Vec::new(), - usage, - stop_reason: None, - size_bytes, - timestamp_start: start_time, - timestamp_end: end_time, - stream_info: None, - }) - } - - /// 获取基础 URL - fn get_base_url(&self, provider: &ProviderType) -> String { - match provider { - ProviderType::OpenAI => "https://api.openai.com".to_string(), - ProviderType::Claude => "https://api.anthropic.com".to_string(), - ProviderType::Gemini | ProviderType::GeminiApiKey => { - "https://generativelanguage.googleapis.com".to_string() - } - ProviderType::Kiro => "https://codewhisperer.us-east-1.amazonaws.com".to_string(), - _ => "https://api.openai.com".to_string(), // 默认使用 OpenAI 兼容 API - } - } - - /// 获取认证头 - async fn get_auth_header( - &self, - provider: &ProviderType, - credential_id: &Option, - ) -> Result, ReplayerError> { - // 如果没有指定凭证,尝试从凭证池选择 - let _cred_id = if let Some(id) = credential_id { - id.clone() - } else { - // 尝试从凭证池选择 - let provider_type_str = format!("{provider:?}"); - if let Ok(Some(cred)) = - self.provider_pool - .select_credential(&self.db, &provider_type_str, None) - { - cred.uuid - } else { - return Ok(None); - } - }; - - // TODO: 根据凭证 ID 获取实际的认证信息 - // 这里需要根据具体的凭证类型来获取 token - // 目前返回 None,实际实现需要从凭证池获取 token - Ok(None) - } - - /// 提取响应内容 - fn extract_content(&self, body: &serde_json::Value, provider: &ProviderType) -> String { - match provider { - ProviderType::OpenAI | ProviderType::Kiro => { - // OpenAI 格式 - body["choices"][0]["message"]["content"] - .as_str() - .unwrap_or("") - .to_string() - } - ProviderType::Claude | ProviderType::ClaudeOAuth => { - // Claude 格式 - body["content"][0]["text"] - .as_str() - .unwrap_or("") - .to_string() - } - ProviderType::Gemini | ProviderType::GeminiApiKey => { - // Gemini 格式 - body["candidates"][0]["content"]["parts"][0]["text"] - .as_str() - .unwrap_or("") - .to_string() - } - _ => { - // 尝试通用格式 - body["choices"][0]["message"]["content"] - .as_str() - .or_else(|| body["content"][0]["text"].as_str()) - .unwrap_or("") - .to_string() - } - } - } - - /// 提取 token 使用量 - fn extract_usage(&self, body: &serde_json::Value, provider: &ProviderType) -> TokenUsage { - let usage = &body["usage"]; - - match provider { - ProviderType::OpenAI | ProviderType::Kiro => TokenUsage { - input_tokens: usage["prompt_tokens"].as_u64().unwrap_or(0) as u32, - output_tokens: usage["completion_tokens"].as_u64().unwrap_or(0) as u32, - total_tokens: usage["total_tokens"].as_u64().unwrap_or(0) as u32, - ..Default::default() - }, - ProviderType::Claude | ProviderType::ClaudeOAuth => TokenUsage { - input_tokens: usage["input_tokens"].as_u64().unwrap_or(0) as u32, - output_tokens: usage["output_tokens"].as_u64().unwrap_or(0) as u32, - total_tokens: (usage["input_tokens"].as_u64().unwrap_or(0) - + usage["output_tokens"].as_u64().unwrap_or(0)) - as u32, - ..Default::default() - }, - _ => TokenUsage::default(), - } - } - - /// 完成重放 Flow - async fn complete_replay_flow(&self, flow_id: &str, response: Option) { - let now = Utc::now(); - - // 更新内存存储中的 Flow - let store = self.flow_monitor.memory_store(); - let store_guard = store.read().await; - - if let Some(flow_lock) = store_guard.get(flow_id) { - let mut flow = flow_lock.write().unwrap(); - flow.response = response; - flow.state = FlowState::Completed; - flow.timestamps.response_end = Some(now); - flow.timestamps.calculate_duration(); - } - } - - /// 标记重放 Flow 失败 - async fn fail_replay_flow(&self, flow_id: &str, error: &str) { - let now = Utc::now(); - - // 更新内存存储中的 Flow - let store = self.flow_monitor.memory_store(); - let store_guard = store.read().await; - - if let Some(flow_lock) = store_guard.get(flow_id) { - let mut flow = flow_lock.write().unwrap(); - flow.state = FlowState::Failed; - flow.error = Some(super::models::FlowError::new( - super::models::FlowErrorType::Other, - error, - )); - flow.timestamps.response_end = Some(now); - flow.timestamps.calculate_duration(); - } - } - - /// 检查 Flow 是否为重放 Flow - /// - /// **Validates: Requirements 3.2** - pub fn is_replay_flow(flow: &LLMFlow) -> bool { - flow.annotations.tags.contains(&"replay".to_string()) - } - - /// 获取原始 Flow ID(从重放 Flow 的注释中提取) - pub fn get_original_flow_id(flow: &LLMFlow) -> Option { - if let Some(ref comment) = flow.annotations.comment { - if comment.starts_with("重放自 Flow: ") { - return Some(comment.replace("重放自 Flow: ", "")); - } - } - None - } -} - -// ============================================================================ -// 单元测试 -// ============================================================================ - -#[cfg(test)] -mod tests { - use super::super::models::FlowType; - use super::*; - - #[test] - fn test_replay_config_default() { - let config = ReplayConfig::default(); - assert!(config.credential_id.is_none()); - assert!(config.modify_request.is_none()); - assert_eq!(config.interval_ms, 1000); - } - - #[test] - fn test_replay_result_success() { - let started_at = Utc::now(); - let completed_at = started_at + chrono::Duration::milliseconds(500); - - let result = ReplayResult::success( - "original-id".to_string(), - "replay-id".to_string(), - started_at, - completed_at, - ); - - assert!(result.success); - assert_eq!(result.original_flow_id, "original-id"); - assert_eq!(result.replay_flow_id, "replay-id"); - assert!(result.error.is_none()); - assert_eq!(result.duration_ms, 500); - } - - #[test] - fn test_replay_result_failure() { - let started_at = Utc::now(); - let completed_at = started_at + chrono::Duration::milliseconds(100); - - let result = ReplayResult::failure( - "original-id".to_string(), - "Connection failed".to_string(), - started_at, - completed_at, - ); - - assert!(!result.success); - assert_eq!(result.original_flow_id, "original-id"); - assert!(result.replay_flow_id.is_empty()); - assert_eq!(result.error, Some("Connection failed".to_string())); - } - - #[test] - fn test_is_replay_flow() { - let mut flow = LLMFlow::new( - "test-id".to_string(), - FlowType::ChatCompletions, - LLMRequest::default(), - FlowMetadata::default(), - ); - - // 没有 replay 标签 - assert!(!FlowReplayer::is_replay_flow(&flow)); - - // 添加 replay 标签 - flow.annotations.tags.push("replay".to_string()); - assert!(FlowReplayer::is_replay_flow(&flow)); - } - - #[test] - fn test_get_original_flow_id() { - let mut flow = LLMFlow::new( - "replay-id".to_string(), - FlowType::ChatCompletions, - LLMRequest::default(), - FlowMetadata::default(), - ); - - // 没有注释 - assert!(FlowReplayer::get_original_flow_id(&flow).is_none()); - - // 添加重放注释 - flow.annotations.comment = Some("重放自 Flow: original-id".to_string()); - assert_eq!( - FlowReplayer::get_original_flow_id(&flow), - Some("original-id".to_string()) - ); - } - - #[test] - fn test_request_modification_serialization() { - let modification = RequestModification { - model: Some("gpt-4-turbo".to_string()), - messages: None, - parameters: None, - system_prompt: Some("You are a helpful assistant.".to_string()), - }; - - let json = serde_json::to_string(&modification).unwrap(); - let deserialized: RequestModification = serde_json::from_str(&json).unwrap(); - - assert_eq!(deserialized.model, Some("gpt-4-turbo".to_string())); - assert_eq!( - deserialized.system_prompt, - Some("You are a helpful assistant.".to_string()) - ); - } -} - -// ============================================================================ -// 属性测试 -// ============================================================================ - -#[cfg(test)] -mod property_tests { - use super::super::models::FlowType; - use super::*; - use proptest::prelude::*; - - // ======================================================================== - // 生成器 - // ======================================================================== - - /// 生成随机的 Flow ID - fn arb_flow_id() -> impl Strategy { - "[a-f0-9]{8}-[a-f0-9]{4}-[a-f0-9]{4}-[a-f0-9]{4}-[a-f0-9]{12}" - } - - /// 生成随机的模型名称 - fn arb_model_name() -> impl Strategy { - prop_oneof![ - Just("gpt-4".to_string()), - Just("gpt-4-turbo".to_string()), - Just("gpt-3.5-turbo".to_string()), - Just("claude-3-opus".to_string()), - Just("claude-3-sonnet".to_string()), - Just("gemini-pro".to_string()), - ] - } - - /// 生成随机的 LLMRequest - fn arb_llm_request() -> impl Strategy { - arb_model_name().prop_map(|model| LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - headers: std::collections::HashMap::new(), - body: serde_json::Value::Null, - messages: Vec::new(), - system_prompt: None, - tools: None, - model, - original_model: None, - parameters: RequestParameters::default(), - size_bytes: 0, - timestamp: Utc::now(), - }) - } - - /// 生成随机的 FlowMetadata - fn arb_flow_metadata() -> impl Strategy { - prop_oneof![ - Just(crate::ProviderType::OpenAI), - Just(crate::ProviderType::Claude), - Just(crate::ProviderType::Gemini), - Just(crate::ProviderType::Kiro), - ] - .prop_map(|provider| FlowMetadata { - provider, - credential_id: Some("test-cred".to_string()), - credential_name: Some("Test Credential".to_string()), - ..Default::default() - }) - } - - /// 生成随机的 LLMFlow - fn arb_llm_flow() -> impl Strategy { - (arb_flow_id(), arb_llm_request(), arb_flow_metadata()).prop_map( - |(id, request, metadata)| { - LLMFlow::new(id, FlowType::ChatCompletions, request, metadata) - }, - ) - } - - // ======================================================================== - // 属性测试 - // ======================================================================== - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: flow-monitor-enhancement, Property 5: 重放 Flow 标记正确性** - /// **Validates: Requirements 3.2** - /// - /// *对于任意* 重放操作,新创建的 Flow 应该被正确标记为 "replay", - /// 并且包含原始 Flow 的引用。 - #[test] - fn prop_replay_flow_marking_correctness( - original_flow in arb_llm_flow(), - ) { - // 保存原始 Flow ID 的副本 - let original_flow_id = original_flow.id.clone(); - - // 创建一个模拟的重放 Flow(模拟 create_replay_flow 的行为) - let replay_flow_id = uuid::Uuid::new_v4().to_string(); - let now = Utc::now(); - - // 创建重放 Flow 的元数据 - let metadata = original_flow.metadata.clone(); - - // 创建重放 Flow(模拟 create_replay_flow 的逻辑) - let replay_flow = LLMFlow { - id: replay_flow_id.clone(), - flow_type: original_flow.flow_type.clone(), - request: original_flow.request.clone(), - response: None, - error: None, - metadata, - timestamps: FlowTimestamps { - created: now, - request_start: now, - request_end: None, - response_start: None, - response_end: None, - duration_ms: 0, - ttfb_ms: None, - }, - state: FlowState::Pending, - annotations: FlowAnnotations { - marker: Some("🔄".to_string()), // 重放标记 - comment: Some(format!("重放自 Flow: {original_flow_id}")), - tags: vec!["replay".to_string()], - starred: false, - }, - }; - - // 验证 1: 重放 Flow 应该有 "replay" 标签 - prop_assert!( - FlowReplayer::is_replay_flow(&replay_flow), - "重放 Flow 应该被标记为 replay" - ); - - // 验证 2: 重放 Flow 应该包含原始 Flow ID 的引用 - let extracted_original_id = FlowReplayer::get_original_flow_id(&replay_flow); - prop_assert!( - extracted_original_id.is_some(), - "重放 Flow 应该包含原始 Flow ID 的引用" - ); - prop_assert_eq!( - extracted_original_id.unwrap(), - original_flow_id.clone(), - "提取的原始 Flow ID 应该与实际原始 Flow ID 一致" - ); - - // 验证 3: 重放 Flow 应该有重放标记 emoji - prop_assert_eq!( - replay_flow.annotations.marker, - Some("🔄".to_string()), - "重放 Flow 应该有重放标记 emoji" - ); - - // 验证 4: 重放 Flow 的 ID 应该与原始 Flow 不同 - prop_assert_ne!( - replay_flow.id, - original_flow_id, - "重放 Flow 的 ID 应该与原始 Flow 不同" - ); - - // 验证 5: 原始 Flow 不应该被标记为 replay(除非它本身就是重放) - if !original_flow.annotations.tags.contains(&"replay".to_string()) { - prop_assert!( - !FlowReplayer::is_replay_flow(&original_flow), - "原始 Flow 不应该被标记为 replay" - ); - } - } - - /// **Feature: flow-monitor-enhancement, Property 5b: 非重放 Flow 标记正确性** - /// **Validates: Requirements 3.2** - /// - /// *对于任意* 普通 Flow(非重放),is_replay_flow 应该返回 false。 - #[test] - fn prop_non_replay_flow_not_marked( - flow in arb_llm_flow(), - ) { - // 普通 Flow 不应该被标记为 replay - prop_assert!( - !FlowReplayer::is_replay_flow(&flow), - "普通 Flow 不应该被标记为 replay" - ); - - // 普通 Flow 不应该有原始 Flow ID - prop_assert!( - FlowReplayer::get_original_flow_id(&flow).is_none(), - "普通 Flow 不应该有原始 Flow ID" - ); - } - } -} diff --git a/src-tauri/src/flow_monitor/session.rs b/src-tauri/src/flow_monitor/session.rs deleted file mode 100644 index d11058db2..000000000 --- a/src-tauri/src/flow_monitor/session.rs +++ /dev/null @@ -1,1191 +0,0 @@ -//! 会话管理器 -//! -//! 该模块实现 Flow 会话管理功能,支持将相关的 Flow 组织成会话, -//! 便于管理和分析交互历史。 -//! -//! **Validates: Requirements 5.1-5.7** - -use chrono::{DateTime, Utc}; -use rusqlite::{params, Connection, OptionalExtension}; -use serde::{Deserialize, Serialize}; -use std::collections::HashMap; -use std::path::PathBuf; -use std::sync::Mutex; -use thiserror::Error; -use uuid::Uuid; - -use super::exporter::{ExportFormat, ExportOptions, FlowExporter}; -use super::models::LLMFlow; - -// ============================================================================ -// 错误类型 -// ============================================================================ - -/// 会话管理错误 -#[derive(Debug, Error)] -pub enum SessionError { - #[error("SQLite 错误: {0}")] - Sqlite(#[from] rusqlite::Error), - - #[error("会话不存在: {0}")] - SessionNotFound(String), - - #[error("Flow 不存在: {0}")] - FlowNotFound(String), - - #[error("JSON 序列化错误: {0}")] - Json(#[from] serde_json::Error), - - #[error("IO 错误: {0}")] - Io(#[from] std::io::Error), -} - -pub type Result = std::result::Result; - -// ============================================================================ -// 数据结构 -// ============================================================================ - -/// Flow 会话 -/// -/// **Validates: Requirements 5.1, 5.5** -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct FlowSession { - /// 唯一标识符 - pub id: String, - /// 会话名称 - pub name: String, - /// 会话描述 - #[serde(skip_serializing_if = "Option::is_none")] - pub description: Option, - /// 关联的 Flow ID 列表 - pub flow_ids: Vec, - /// 创建时间 - pub created_at: DateTime, - /// 更新时间 - pub updated_at: DateTime, - /// 是否已归档 - pub archived: bool, -} - -impl FlowSession { - /// 创建新会话 - pub fn new(name: impl Into, description: Option) -> Self { - let now = Utc::now(); - Self { - id: Uuid::new_v4().to_string(), - name: name.into(), - description, - flow_ids: Vec::new(), - created_at: now, - updated_at: now, - archived: false, - } - } -} - -/// 自动会话检测配置 -/// -/// **Validates: Requirements 5.4** -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct AutoSessionConfig { - /// 是否启用自动会话检测 - pub enabled: bool, - /// 时间窗口(毫秒)- 在此时间内的请求会被归入同一会话 - pub time_window_ms: u64, - /// 是否按客户端分组 - pub group_by_client: bool, -} - -impl Default for AutoSessionConfig { - fn default() -> Self { - Self { - enabled: false, - time_window_ms: 30_000, // 30 秒 - group_by_client: true, - } - } -} - -/// 会话导出结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct SessionExportResult { - /// 会话信息 - pub session: FlowSession, - /// 导出的数据 - pub data: String, - /// 导出格式 - pub format: ExportFormat, - /// 导出的 Flow 数量 - pub flow_count: usize, -} - -// ============================================================================ -// 会话管理器 -// ============================================================================ - -/// 会话管理器 -/// -/// **Validates: Requirements 5.1-5.7** -pub struct SessionManager { - /// SQLite 连接 - db: Mutex, - /// 自动会话检测配置 - auto_config: Mutex, - /// 最近活跃会话缓存(用于自动检测) - /// key: client_id 或 "default", value: (session_id, last_activity_time) - active_sessions: Mutex)>>, -} - -impl SessionManager { - /// 创建新的会话管理器 - /// - /// # Arguments - /// * `db_path` - SQLite 数据库路径 - pub fn new(db_path: PathBuf) -> Result { - // 确保目录存在 - if let Some(parent) = db_path.parent() { - std::fs::create_dir_all(parent)?; - } - - let conn = Connection::open(&db_path)?; - Self::init_database(&conn)?; - - Ok(Self { - db: Mutex::new(conn), - auto_config: Mutex::new(AutoSessionConfig::default()), - active_sessions: Mutex::new(HashMap::new()), - }) - } - - /// 从现有连接创建会话管理器(用于测试) - pub fn from_connection(conn: Connection) -> Result { - Self::init_database(&conn)?; - - Ok(Self { - db: Mutex::new(conn), - auto_config: Mutex::new(AutoSessionConfig::default()), - active_sessions: Mutex::new(HashMap::new()), - }) - } - - /// 初始化数据库表 - fn init_database(conn: &Connection) -> Result<()> { - conn.execute_batch( - r#" - -- 会话表 - CREATE TABLE IF NOT EXISTS flow_sessions ( - id TEXT PRIMARY KEY, - name TEXT NOT NULL, - description TEXT, - created_at TEXT NOT NULL, - updated_at TEXT NOT NULL, - archived INTEGER DEFAULT 0 - ); - - -- 会话-Flow 关联表 - CREATE TABLE IF NOT EXISTS session_flows ( - session_id TEXT NOT NULL, - flow_id TEXT NOT NULL, - added_at TEXT NOT NULL, - PRIMARY KEY (session_id, flow_id), - FOREIGN KEY (session_id) REFERENCES flow_sessions(id) ON DELETE CASCADE - ); - - CREATE INDEX IF NOT EXISTS idx_session_flows_session ON session_flows(session_id); - CREATE INDEX IF NOT EXISTS idx_session_flows_flow ON session_flows(flow_id); - CREATE INDEX IF NOT EXISTS idx_sessions_archived ON flow_sessions(archived); - CREATE INDEX IF NOT EXISTS idx_sessions_created ON flow_sessions(created_at); - "#, - )?; - - Ok(()) - } - - /// 创建新会话 - /// - /// **Validates: Requirements 5.1** - /// - /// # Arguments - /// * `name` - 会话名称 - /// * `description` - 会话描述(可选) - /// - /// # Returns - /// 新创建的会话 - pub fn create_session( - &self, - name: impl Into, - description: Option<&str>, - ) -> Result { - let session = FlowSession::new(name, description.map(String::from)); - - let conn = self.db.lock().unwrap(); - conn.execute( - r#" - INSERT INTO flow_sessions (id, name, description, created_at, updated_at, archived) - VALUES (?1, ?2, ?3, ?4, ?5, ?6) - "#, - params![ - session.id, - session.name, - session.description, - session.created_at.to_rfc3339(), - session.updated_at.to_rfc3339(), - session.archived as i32, - ], - )?; - - Ok(session) - } - - /// 获取会话 - /// - /// # Arguments - /// * `session_id` - 会话 ID - /// - /// # Returns - /// 会话信息(如果存在) - pub fn get_session(&self, session_id: &str) -> Result> { - let conn = self.db.lock().unwrap(); - - let session: Option<(String, String, Option, String, String, i32)> = conn - .query_row( - r#" - SELECT id, name, description, created_at, updated_at, archived - FROM flow_sessions - WHERE id = ?1 - "#, - params![session_id], - |row| { - Ok(( - row.get(0)?, - row.get(1)?, - row.get(2)?, - row.get(3)?, - row.get(4)?, - row.get(5)?, - )) - }, - ) - .optional()?; - - match session { - Some((id, name, description, created_at, updated_at, archived)) => { - // 获取关联的 Flow ID - let flow_ids = self.get_session_flow_ids_internal(&conn, &id)?; - - Ok(Some(FlowSession { - id, - name, - description, - flow_ids, - created_at: DateTime::parse_from_rfc3339(&created_at) - .map(|dt| dt.with_timezone(&Utc)) - .unwrap_or_else(|_| Utc::now()), - updated_at: DateTime::parse_from_rfc3339(&updated_at) - .map(|dt| dt.with_timezone(&Utc)) - .unwrap_or_else(|_| Utc::now()), - archived: archived != 0, - })) - } - None => Ok(None), - } - } - - /// 获取会话关联的 Flow ID(内部方法) - fn get_session_flow_ids_internal( - &self, - conn: &Connection, - session_id: &str, - ) -> Result> { - let mut stmt = conn.prepare( - r#" - SELECT flow_id FROM session_flows - WHERE session_id = ?1 - ORDER BY added_at ASC - "#, - )?; - - let flow_ids: Vec = stmt - .query_map(params![session_id], |row| row.get(0))? - .filter_map(|r| r.ok()) - .collect(); - - Ok(flow_ids) - } - - /// 列出所有会话 - /// - /// # Arguments - /// * `include_archived` - 是否包含已归档的会话 - /// - /// # Returns - /// 会话列表 - pub fn list_sessions(&self, include_archived: bool) -> Result> { - let conn = self.db.lock().unwrap(); - - let sql = if include_archived { - "SELECT id, name, description, created_at, updated_at, archived FROM flow_sessions ORDER BY updated_at DESC" - } else { - "SELECT id, name, description, created_at, updated_at, archived FROM flow_sessions WHERE archived = 0 ORDER BY updated_at DESC" - }; - - let mut stmt = conn.prepare(sql)?; - let rows = stmt.query_map([], |row| { - Ok(( - row.get::<_, String>(0)?, - row.get::<_, String>(1)?, - row.get::<_, Option>(2)?, - row.get::<_, String>(3)?, - row.get::<_, String>(4)?, - row.get::<_, i32>(5)?, - )) - })?; - - let mut sessions = Vec::new(); - for row in rows { - let (id, name, description, created_at, updated_at, archived) = row?; - let flow_ids = self.get_session_flow_ids_internal(&conn, &id)?; - - sessions.push(FlowSession { - id, - name, - description, - flow_ids, - created_at: DateTime::parse_from_rfc3339(&created_at) - .map(|dt| dt.with_timezone(&Utc)) - .unwrap_or_else(|_| Utc::now()), - updated_at: DateTime::parse_from_rfc3339(&updated_at) - .map(|dt| dt.with_timezone(&Utc)) - .unwrap_or_else(|_| Utc::now()), - archived: archived != 0, - }); - } - - Ok(sessions) - } - - /// 添加 Flow 到会话 - /// - /// **Validates: Requirements 5.2** - /// - /// # Arguments - /// * `session_id` - 会话 ID - /// * `flow_id` - Flow ID - pub fn add_flow(&self, session_id: &str, flow_id: &str) -> Result<()> { - let conn = self.db.lock().unwrap(); - - // 检查会话是否存在 - let exists: bool = conn - .query_row( - "SELECT 1 FROM flow_sessions WHERE id = ?1", - params![session_id], - |_| Ok(true), - ) - .optional()? - .unwrap_or(false); - - if !exists { - return Err(SessionError::SessionNotFound(session_id.to_string())); - } - - // 添加关联(忽略重复) - conn.execute( - r#" - INSERT OR IGNORE INTO session_flows (session_id, flow_id, added_at) - VALUES (?1, ?2, ?3) - "#, - params![session_id, flow_id, Utc::now().to_rfc3339()], - )?; - - // 更新会话的更新时间 - conn.execute( - "UPDATE flow_sessions SET updated_at = ?1 WHERE id = ?2", - params![Utc::now().to_rfc3339(), session_id], - )?; - - Ok(()) - } - - /// 从会话移除 Flow - /// - /// **Validates: Requirements 5.2** - /// - /// # Arguments - /// * `session_id` - 会话 ID - /// * `flow_id` - Flow ID - pub fn remove_flow(&self, session_id: &str, flow_id: &str) -> Result<()> { - let conn = self.db.lock().unwrap(); - - conn.execute( - "DELETE FROM session_flows WHERE session_id = ?1 AND flow_id = ?2", - params![session_id, flow_id], - )?; - - // 更新会话的更新时间 - conn.execute( - "UPDATE flow_sessions SET updated_at = ?1 WHERE id = ?2", - params![Utc::now().to_rfc3339(), session_id], - )?; - - Ok(()) - } - - /// 更新会话信息 - /// - /// **Validates: Requirements 5.5** - /// - /// # Arguments - /// * `session_id` - 会话 ID - /// * `name` - 新名称(可选) - /// * `description` - 新描述(可选) - pub fn update_session( - &self, - session_id: &str, - name: Option<&str>, - description: Option>, - ) -> Result<()> { - let conn = self.db.lock().unwrap(); - - // 检查会话是否存在 - let exists: bool = conn - .query_row( - "SELECT 1 FROM flow_sessions WHERE id = ?1", - params![session_id], - |_| Ok(true), - ) - .optional()? - .unwrap_or(false); - - if !exists { - return Err(SessionError::SessionNotFound(session_id.to_string())); - } - - if let Some(new_name) = name { - conn.execute( - "UPDATE flow_sessions SET name = ?1, updated_at = ?2 WHERE id = ?3", - params![new_name, Utc::now().to_rfc3339(), session_id], - )?; - } - - if let Some(new_desc) = description { - conn.execute( - "UPDATE flow_sessions SET description = ?1, updated_at = ?2 WHERE id = ?3", - params![new_desc, Utc::now().to_rfc3339(), session_id], - )?; - } - - Ok(()) - } - - /// 归档会话 - /// - /// **Validates: Requirements 5.7** - /// - /// # Arguments - /// * `session_id` - 会话 ID - pub fn archive_session(&self, session_id: &str) -> Result<()> { - let conn = self.db.lock().unwrap(); - - let rows_affected = conn.execute( - "UPDATE flow_sessions SET archived = 1, updated_at = ?1 WHERE id = ?2", - params![Utc::now().to_rfc3339(), session_id], - )?; - - if rows_affected == 0 { - return Err(SessionError::SessionNotFound(session_id.to_string())); - } - - Ok(()) - } - - /// 取消归档会话 - /// - /// # Arguments - /// * `session_id` - 会话 ID - pub fn unarchive_session(&self, session_id: &str) -> Result<()> { - let conn = self.db.lock().unwrap(); - - let rows_affected = conn.execute( - "UPDATE flow_sessions SET archived = 0, updated_at = ?1 WHERE id = ?2", - params![Utc::now().to_rfc3339(), session_id], - )?; - - if rows_affected == 0 { - return Err(SessionError::SessionNotFound(session_id.to_string())); - } - - Ok(()) - } - - /// 删除会话 - /// - /// **Validates: Requirements 5.7** - /// - /// # Arguments - /// * `session_id` - 会话 ID - pub fn delete_session(&self, session_id: &str) -> Result<()> { - let conn = self.db.lock().unwrap(); - - // 删除关联 - conn.execute( - "DELETE FROM session_flows WHERE session_id = ?1", - params![session_id], - )?; - - // 删除会话 - let rows_affected = conn.execute( - "DELETE FROM flow_sessions WHERE id = ?1", - params![session_id], - )?; - - if rows_affected == 0 { - return Err(SessionError::SessionNotFound(session_id.to_string())); - } - - Ok(()) - } - - /// 获取会话中的 Flow ID 列表 - /// - /// # Arguments - /// * `session_id` - 会话 ID - /// - /// # Returns - /// Flow ID 列表 - pub fn get_session_flow_ids(&self, session_id: &str) -> Result> { - let conn = self.db.lock().unwrap(); - self.get_session_flow_ids_internal(&conn, session_id) - } - - /// 检查 Flow 是否在会话中 - /// - /// # Arguments - /// * `session_id` - 会话 ID - /// * `flow_id` - Flow ID - /// - /// # Returns - /// 是否在会话中 - pub fn is_flow_in_session(&self, session_id: &str, flow_id: &str) -> Result { - let conn = self.db.lock().unwrap(); - - let exists: bool = conn - .query_row( - "SELECT 1 FROM session_flows WHERE session_id = ?1 AND flow_id = ?2", - params![session_id, flow_id], - |_| Ok(true), - ) - .optional()? - .unwrap_or(false); - - Ok(exists) - } - - /// 获取 Flow 所属的会话列表 - /// - /// # Arguments - /// * `flow_id` - Flow ID - /// - /// # Returns - /// 会话 ID 列表 - pub fn get_sessions_for_flow(&self, flow_id: &str) -> Result> { - let conn = self.db.lock().unwrap(); - - let mut stmt = conn.prepare("SELECT session_id FROM session_flows WHERE flow_id = ?1")?; - - let session_ids: Vec = stmt - .query_map(params![flow_id], |row| row.get(0))? - .filter_map(|r| r.ok()) - .collect(); - - Ok(session_ids) - } - - /// 获取会话中的 Flow 数量 - /// - /// # Arguments - /// * `session_id` - 会话 ID - /// - /// # Returns - /// Flow 数量 - pub fn get_session_flow_count(&self, session_id: &str) -> Result { - let conn = self.db.lock().unwrap(); - - let count: i64 = conn.query_row( - "SELECT COUNT(*) FROM session_flows WHERE session_id = ?1", - params![session_id], - |row| row.get(0), - )?; - - Ok(count as usize) - } - - // ======================================================================== - // 自动会话检测 - // ======================================================================== - - /// 获取自动会话检测配置 - pub fn get_auto_config(&self) -> AutoSessionConfig { - self.auto_config.lock().unwrap().clone() - } - - /// 设置自动会话检测配置 - pub fn set_auto_config(&self, config: AutoSessionConfig) { - *self.auto_config.lock().unwrap() = config; - } - - /// 自动检测会话 - /// - /// **Validates: Requirements 5.4** - /// - /// 根据配置自动检测 Flow 应该归属的会话。 - /// 如果在时间窗口内有活跃会话,则返回该会话 ID; - /// 否则返回 None,表示应该创建新会话或不归入任何会话。 - /// - /// # Arguments - /// * `flow` - LLM Flow - /// - /// # Returns - /// 会话 ID(如果检测到应该归入某个会话) - pub fn detect_session(&self, flow: &LLMFlow) -> Option { - let config = self.auto_config.lock().unwrap().clone(); - - if !config.enabled { - return None; - } - - let now = Utc::now(); - let time_window = chrono::Duration::milliseconds(config.time_window_ms as i64); - - // 确定客户端标识 - let client_key = if config.group_by_client { - flow.metadata - .client_info - .ip - .clone() - .or_else(|| flow.metadata.client_info.request_id.clone()) - .unwrap_or_else(|| "default".to_string()) - } else { - "default".to_string() - }; - - let mut active_sessions = self.active_sessions.lock().unwrap(); - - // 检查是否有活跃会话 - if let Some((session_id, last_activity)) = active_sessions.get(&client_key) { - if now - *last_activity < time_window { - // 更新最后活动时间 - let session_id = session_id.clone(); - active_sessions.insert(client_key, (session_id.clone(), now)); - return Some(session_id); - } - } - - // 没有活跃会话 - None - } - - /// 注册活跃会话(用于自动检测) - /// - /// # Arguments - /// * `session_id` - 会话 ID - /// * `client_key` - 客户端标识(可选,默认为 "default") - pub fn register_active_session(&self, session_id: &str, client_key: Option<&str>) { - let key = client_key.unwrap_or("default").to_string(); - let mut active_sessions = self.active_sessions.lock().unwrap(); - active_sessions.insert(key, (session_id.to_string(), Utc::now())); - } - - /// 清除活跃会话缓存 - pub fn clear_active_sessions(&self) { - let mut active_sessions = self.active_sessions.lock().unwrap(); - active_sessions.clear(); - } - - // ======================================================================== - // 会话导出 - // ======================================================================== - - /// 导出会话 - /// - /// **Validates: Requirements 5.6** - /// - /// # Arguments - /// * `session_id` - 会话 ID - /// * `flows` - 会话中的 Flow 列表 - /// * `format` - 导出格式 - /// - /// # Returns - /// 导出结果 - pub fn export_session( - &self, - session_id: &str, - flows: &[LLMFlow], - format: ExportFormat, - ) -> Result { - // 获取会话信息 - let session = self - .get_session(session_id)? - .ok_or_else(|| SessionError::SessionNotFound(session_id.to_string()))?; - - // 创建导出器 - let options = ExportOptions { - format, - filter: None, - include_raw: true, - include_stream_chunks: false, - redact_sensitive: false, - redaction_rules: Vec::new(), - compress: false, - }; - let exporter = FlowExporter::new(options); - - // 导出数据 - let data = match format { - ExportFormat::HAR => { - let har = exporter.export_har(flows); - serde_json::to_string_pretty(&har)? - } - ExportFormat::JSON => { - // 包含会话信息的完整导出 - let export_data = serde_json::json!({ - "session": session, - "flows": flows, - }); - serde_json::to_string_pretty(&export_data)? - } - ExportFormat::JSONL => exporter.export_jsonl(flows), - ExportFormat::Markdown => { - let mut md = format!( - "# 会话: {}\n\n**ID**: {}\n**创建时间**: {}\n**Flow 数量**: {}\n\n", - session.name, - session.id, - session.created_at.format("%Y-%m-%d %H:%M:%S UTC"), - flows.len() - ); - if let Some(ref desc) = session.description { - md.push_str(&format!("**描述**: {desc}\n\n")); - } - md.push_str("---\n\n"); - md.push_str(&exporter.export_markdown_multiple(flows)); - md - } - ExportFormat::CSV => exporter.export_csv(flows), - }; - - Ok(SessionExportResult { - session, - data, - format, - flow_count: flows.len(), - }) - } - - /// 获取会话数量 - pub fn session_count(&self) -> Result { - let conn = self.db.lock().unwrap(); - let count: i64 = - conn.query_row("SELECT COUNT(*) FROM flow_sessions", [], |row| row.get(0))?; - Ok(count as usize) - } - - /// 获取所有会话 ID - pub fn get_all_session_ids(&self) -> Result> { - let conn = self.db.lock().unwrap(); - let mut stmt = conn.prepare("SELECT id FROM flow_sessions")?; - let ids: Vec = stmt - .query_map([], |row| row.get(0))? - .filter_map(|r| r.ok()) - .collect(); - Ok(ids) - } -} - -// ============================================================================ -// 测试模块 -// ============================================================================ - -#[cfg(test)] -mod tests { - use super::*; - - fn create_test_manager() -> SessionManager { - let conn = Connection::open_in_memory().unwrap(); - SessionManager::from_connection(conn).unwrap() - } - - #[test] - fn test_create_session() { - let manager = create_test_manager(); - - let session = manager - .create_session("Test Session", Some("A test session")) - .unwrap(); - - assert!(!session.id.is_empty()); - assert_eq!(session.name, "Test Session"); - assert_eq!(session.description, Some("A test session".to_string())); - assert!(session.flow_ids.is_empty()); - assert!(!session.archived); - } - - #[test] - fn test_get_session() { - let manager = create_test_manager(); - - let created = manager.create_session("Test", None).unwrap(); - let retrieved = manager.get_session(&created.id).unwrap(); - - assert!(retrieved.is_some()); - let retrieved = retrieved.unwrap(); - assert_eq!(retrieved.id, created.id); - assert_eq!(retrieved.name, "Test"); - } - - #[test] - fn test_list_sessions() { - let manager = create_test_manager(); - - manager.create_session("Session 1", None).unwrap(); - manager.create_session("Session 2", None).unwrap(); - - let sessions = manager.list_sessions(false).unwrap(); - assert_eq!(sessions.len(), 2); - } - - #[test] - fn test_add_and_remove_flow() { - let manager = create_test_manager(); - - let session = manager.create_session("Test", None).unwrap(); - - // 添加 Flow - manager.add_flow(&session.id, "flow-1").unwrap(); - manager.add_flow(&session.id, "flow-2").unwrap(); - - let flow_ids = manager.get_session_flow_ids(&session.id).unwrap(); - assert_eq!(flow_ids.len(), 2); - assert!(flow_ids.contains(&"flow-1".to_string())); - assert!(flow_ids.contains(&"flow-2".to_string())); - - // 移除 Flow - manager.remove_flow(&session.id, "flow-1").unwrap(); - - let flow_ids = manager.get_session_flow_ids(&session.id).unwrap(); - assert_eq!(flow_ids.len(), 1); - assert!(!flow_ids.contains(&"flow-1".to_string())); - } - - #[test] - fn test_archive_session() { - let manager = create_test_manager(); - - let session = manager.create_session("Test", None).unwrap(); - - // 归档 - manager.archive_session(&session.id).unwrap(); - - let retrieved = manager.get_session(&session.id).unwrap().unwrap(); - assert!(retrieved.archived); - - // 列表不包含已归档 - let sessions = manager.list_sessions(false).unwrap(); - assert!(sessions.is_empty()); - - // 列表包含已归档 - let sessions = manager.list_sessions(true).unwrap(); - assert_eq!(sessions.len(), 1); - } - - #[test] - fn test_delete_session() { - let manager = create_test_manager(); - - let session = manager.create_session("Test", None).unwrap(); - manager.add_flow(&session.id, "flow-1").unwrap(); - - manager.delete_session(&session.id).unwrap(); - - let retrieved = manager.get_session(&session.id).unwrap(); - assert!(retrieved.is_none()); - } - - #[test] - fn test_session_not_found() { - let manager = create_test_manager(); - - let result = manager.add_flow("non-existent", "flow-1"); - assert!(result.is_err()); - assert!(matches!( - result.unwrap_err(), - SessionError::SessionNotFound(_) - )); - } - - #[test] - fn test_is_flow_in_session() { - let manager = create_test_manager(); - - let session = manager.create_session("Test", None).unwrap(); - manager.add_flow(&session.id, "flow-1").unwrap(); - - assert!(manager.is_flow_in_session(&session.id, "flow-1").unwrap()); - assert!(!manager.is_flow_in_session(&session.id, "flow-2").unwrap()); - } - - #[test] - fn test_get_sessions_for_flow() { - let manager = create_test_manager(); - - let session1 = manager.create_session("Session 1", None).unwrap(); - let session2 = manager.create_session("Session 2", None).unwrap(); - - manager.add_flow(&session1.id, "flow-1").unwrap(); - manager.add_flow(&session2.id, "flow-1").unwrap(); - - let sessions = manager.get_sessions_for_flow("flow-1").unwrap(); - assert_eq!(sessions.len(), 2); - } - - #[test] - fn test_update_session() { - let manager = create_test_manager(); - - let session = manager.create_session("Original", None).unwrap(); - - manager - .update_session(&session.id, Some("Updated"), Some(Some("New description"))) - .unwrap(); - - let retrieved = manager.get_session(&session.id).unwrap().unwrap(); - assert_eq!(retrieved.name, "Updated"); - assert_eq!(retrieved.description, Some("New description".to_string())); - } - - #[test] - fn test_session_id_uniqueness() { - let manager = create_test_manager(); - - let mut ids = std::collections::HashSet::new(); - for i in 0..100 { - let session = manager - .create_session(format!("Session {i}"), None) - .unwrap(); - assert!(ids.insert(session.id), "Session ID should be unique"); - } - } -} - -// ============================================================================ -// 属性测试模块 -// ============================================================================ - -#[cfg(test)] -mod property_tests { - use super::*; - use proptest::prelude::*; - - // ======================================================================== - // 生成器 - // ======================================================================== - - /// 生成随机的会话名称 - fn arb_session_name() -> impl Strategy { - "[a-zA-Z0-9 _-]{1,50}".prop_filter("Name should not be empty", |s| !s.trim().is_empty()) - } - - /// 生成随机的会话描述 - fn arb_session_description() -> impl Strategy> { - prop::option::of("[a-zA-Z0-9 _-]{0,200}") - } - - /// 生成随机的 Flow ID - fn arb_flow_id() -> impl Strategy { - "[a-f0-9]{8}-[a-f0-9]{4}-[a-f0-9]{4}-[a-f0-9]{4}-[a-f0-9]{12}" - } - - /// 生成随机的 Flow ID 列表 - fn arb_flow_ids(max_len: usize) -> impl Strategy> { - prop::collection::vec(arb_flow_id(), 0..max_len) - } - - // ======================================================================== - // 属性测试 - // ======================================================================== - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: flow-monitor-enhancement, Property 8: 会话 ID 唯一性** - /// **Validates: Requirements 5.1** - /// - /// *对于任意* 数量的会话创建操作,每个会话的 ID 应该是唯一的。 - #[test] - fn prop_session_id_uniqueness( - names in prop::collection::vec(arb_session_name(), 1..50) - ) { - let manager = create_test_manager(); - let mut ids = std::collections::HashSet::new(); - - for name in names { - let session = manager.create_session(&name, None).unwrap(); - prop_assert!( - ids.insert(session.id.clone()), - "Session ID '{}' should be unique", - session.id - ); - } - } - - /// **Feature: flow-monitor-enhancement, Property 9: 会话 Flow 关联正确性** - /// **Validates: Requirements 5.2** - /// - /// *对于任意* 会话和 Flow 添加操作,添加后查询该会话应该包含所有添加的 Flow。 - #[test] - fn prop_session_flow_association( - name in arb_session_name(), - flow_ids in arb_flow_ids(20) - ) { - let manager = create_test_manager(); - let session = manager.create_session(&name, None).unwrap(); - - // 添加所有 Flow - for flow_id in &flow_ids { - manager.add_flow(&session.id, flow_id).unwrap(); - } - - // 验证所有 Flow 都在会话中 - let retrieved_ids = manager.get_session_flow_ids(&session.id).unwrap(); - - // 去重后的 flow_ids(因为可能有重复) - let unique_flow_ids: std::collections::HashSet<_> = flow_ids.iter().collect(); - - prop_assert_eq!( - retrieved_ids.len(), - unique_flow_ids.len(), - "Session should contain all unique added flows" - ); - - for flow_id in &flow_ids { - prop_assert!( - retrieved_ids.contains(flow_id), - "Flow '{}' should be in session", - flow_id - ); - } - } - - /// **Feature: flow-monitor-enhancement, Property 10: 会话导出完整性** - /// **Validates: Requirements 5.6** - /// - /// *对于任意* 会话,导出应该包含该会话中的所有 Flow。 - /// 注意:这里我们测试导出的元数据正确性,因为实际 Flow 数据需要从外部获取。 - #[test] - fn prop_session_export_completeness( - name in arb_session_name(), - description in arb_session_description(), - flow_ids in arb_flow_ids(10) - ) { - let manager = create_test_manager(); - let session = manager.create_session(&name, description.as_deref()).unwrap(); - - // 添加 Flow - for flow_id in &flow_ids { - manager.add_flow(&session.id, flow_id).unwrap(); - } - - // 获取会话信息 - let retrieved = manager.get_session(&session.id).unwrap().unwrap(); - - // 验证会话信息完整 - prop_assert_eq!(retrieved.name.clone(), name.clone()); - prop_assert_eq!(retrieved.description, description); - - // 验证 Flow 数量 - let unique_count = flow_ids.iter().collect::>().len(); - prop_assert_eq!( - retrieved.flow_ids.len(), - unique_count, - "Session should contain all unique flows" - ); - - // 导出会话(使用空 Flow 列表测试导出功能) - let result = manager.export_session(&session.id, &[], ExportFormat::JSON).unwrap(); - prop_assert_eq!(result.session.id, session.id); - prop_assert_eq!(result.session.name, name); - prop_assert_eq!(result.flow_count, 0); - } - - /// 会话创建和检索的 Round-Trip 测试 - #[test] - fn prop_session_roundtrip( - name in arb_session_name(), - description in arb_session_description() - ) { - let manager = create_test_manager(); - - let created = manager.create_session(&name, description.as_deref()).unwrap(); - let retrieved = manager.get_session(&created.id).unwrap().unwrap(); - - prop_assert_eq!(created.id, retrieved.id); - prop_assert_eq!(created.name, retrieved.name); - prop_assert_eq!(created.description, retrieved.description); - prop_assert_eq!(created.archived, retrieved.archived); - } - - /// Flow 添加和移除的正确性测试 - #[test] - fn prop_flow_add_remove( - name in arb_session_name(), - flow_ids in arb_flow_ids(10) - ) { - let manager = create_test_manager(); - let session = manager.create_session(&name, None).unwrap(); - - // 添加所有 Flow - for flow_id in &flow_ids { - manager.add_flow(&session.id, flow_id).unwrap(); - } - - // 移除所有 Flow - for flow_id in &flow_ids { - manager.remove_flow(&session.id, flow_id).unwrap(); - } - - // 验证会话为空 - let retrieved_ids = manager.get_session_flow_ids(&session.id).unwrap(); - prop_assert!( - retrieved_ids.is_empty(), - "Session should be empty after removing all flows" - ); - } - - /// 归档和取消归档的正确性测试 - #[test] - fn prop_archive_unarchive( - name in arb_session_name() - ) { - let manager = create_test_manager(); - let session = manager.create_session(&name, None).unwrap(); - - // 初始状态:未归档 - let retrieved = manager.get_session(&session.id).unwrap().unwrap(); - prop_assert!(!retrieved.archived); - - // 归档 - manager.archive_session(&session.id).unwrap(); - let retrieved = manager.get_session(&session.id).unwrap().unwrap(); - prop_assert!(retrieved.archived); - - // 取消归档 - manager.unarchive_session(&session.id).unwrap(); - let retrieved = manager.get_session(&session.id).unwrap().unwrap(); - prop_assert!(!retrieved.archived); - } - } - - fn create_test_manager() -> SessionManager { - let conn = Connection::open_in_memory().unwrap(); - SessionManager::from_connection(conn).unwrap() - } -} diff --git a/src-tauri/src/flow_monitor/stream_rebuilder.rs b/src-tauri/src/flow_monitor/stream_rebuilder.rs deleted file mode 100644 index ab9648840..000000000 --- a/src-tauri/src/flow_monitor/stream_rebuilder.rs +++ /dev/null @@ -1,1703 +0,0 @@ -//! SSE 流式响应重建器 -//! -//! 该模块负责将分散的 SSE chunks 合并为完整的 LLM 响应。 -//! 支持 OpenAI、Anthropic、Gemini 等多种流式响应格式。 - -use chrono::{DateTime, Utc}; -use serde::{Deserialize, Serialize}; -use std::collections::HashMap; -use thiserror::Error; - -use super::models::{ - LLMResponse, StopReason, StreamChunk, StreamInfo, ThinkingContent, TokenUsage, ToolCall, - ToolCallDelta, -}; - -// ============================================================================ -// 错误类型 -// ============================================================================ - -/// 流重建错误 -#[derive(Debug, Error)] -pub enum StreamRebuilderError { - /// JSON 解析错误 - #[error("JSON 解析错误: {0}")] - JsonParseError(#[from] serde_json::Error), - - /// 无效的事件格式 - #[error("无效的事件格式: {0}")] - InvalidEventFormat(String), - - /// 未知的流格式 - #[error("未知的流格式")] - UnknownFormat, -} - -// ============================================================================ -// 流格式枚举 -// ============================================================================ - -/// 流式响应格式 -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] -pub enum StreamFormat { - /// OpenAI 格式 (data: {...}) - OpenAI, - /// Anthropic 格式 (event: xxx, data: {...}) - Anthropic, - /// Gemini 格式 - Gemini, - /// 未知格式 - #[default] - Unknown, -} - -// ============================================================================ -// 工具调用构建器 -// ============================================================================ - -/// 工具调用构建器,用于累积流式工具调用数据 -#[derive(Debug, Clone, Default)] -struct ToolCallBuilder { - /// 工具调用 ID - id: Option, - /// 工具类型 - tool_type: String, - /// 函数名称 - function_name: Option, - /// 函数参数(累积的 JSON 字符串) - arguments: String, -} - -impl ToolCallBuilder { - fn build(self) -> Option { - let id = self.id?; - let name = self.function_name?; - - Some(ToolCall { - id, - tool_type: self.tool_type, - function: super::models::FunctionCall { - name, - arguments: self.arguments, - }, - }) - } -} - -// ============================================================================ -// 流重建器 -// ============================================================================ - -/// SSE 流重建器 -/// -/// 将分散的 SSE chunks 合并为完整的 LLM 响应。 -/// 支持多种流式响应格式的解析和重建。 -#[derive(Debug)] -pub struct StreamRebuilder { - /// 累积的 chunks - chunks: Vec, - /// 内容缓冲区 - content_buffer: String, - /// 工具调用构建器(按索引) - tool_calls_buffer: HashMap, - /// 思维链缓冲区 - thinking_buffer: Option, - /// 首个 chunk 时间 - first_chunk_time: Option>, - /// 最后一个 chunk 时间 - last_chunk_time: Option>, - /// 流格式 - format: StreamFormat, - /// chunk 计数器 - chunk_index: u32, - /// 停止原因 - stop_reason: Option, - /// Token 使用量 - usage: TokenUsage, - /// 响应 ID - response_id: Option, - /// 模型名称 - model: Option, - /// 是否保存原始 chunks - save_raw_chunks: bool, - /// 当前内容块索引(Anthropic 格式) - current_content_block_index: Option, - /// 当前内容块类型(Anthropic 格式) - current_content_block_type: Option, -} - -impl StreamRebuilder { - /// 创建新的流重建器 - pub fn new(format: StreamFormat) -> Self { - Self { - chunks: Vec::new(), - content_buffer: String::new(), - tool_calls_buffer: HashMap::new(), - thinking_buffer: None, - first_chunk_time: None, - last_chunk_time: None, - format, - chunk_index: 0, - stop_reason: None, - usage: TokenUsage::default(), - response_id: None, - model: None, - save_raw_chunks: false, - current_content_block_index: None, - current_content_block_type: None, - } - } - - /// 设置是否保存原始 chunks - pub fn with_save_raw_chunks(mut self, save: bool) -> Self { - self.save_raw_chunks = save; - self - } - - /// 处理 SSE 事件 - /// - /// # 参数 - /// - `event`: SSE 事件类型(可选,如 "message", "content_block_delta" 等) - /// - `data`: SSE 数据内容 - /// - /// # 返回 - /// - `Ok(())`: 处理成功 - /// - `Err(StreamRebuilderError)`: 处理失败 - pub fn process_event( - &mut self, - event: Option<&str>, - data: &str, - ) -> Result<(), StreamRebuilderError> { - let now = Utc::now(); - - // 记录时间 - if self.first_chunk_time.is_none() { - self.first_chunk_time = Some(now); - } - self.last_chunk_time = Some(now); - - // 创建 chunk 记录 - let mut chunk = StreamChunk { - index: self.chunk_index, - event: event.map(|s| s.to_string()), - data: data.to_string(), - timestamp: now, - content_delta: None, - tool_call_delta: None, - thinking_delta: None, - }; - - // 根据格式处理 - let result = match self.format { - StreamFormat::OpenAI => self.process_openai_chunk(data, &mut chunk), - StreamFormat::Anthropic => self.process_anthropic_chunk(event, data, &mut chunk), - StreamFormat::Gemini => self.process_gemini_chunk(data, &mut chunk), - StreamFormat::Unknown => { - // 尝试自动检测格式 - if let Some(evt) = event { - if evt.starts_with("message_") || evt.starts_with("content_block") { - self.format = StreamFormat::Anthropic; - self.process_anthropic_chunk(Some(evt), data, &mut chunk) - } else { - // 尝试 OpenAI 格式 - self.format = StreamFormat::OpenAI; - self.process_openai_chunk(data, &mut chunk) - } - } else { - // 尝试 OpenAI 格式 - self.format = StreamFormat::OpenAI; - self.process_openai_chunk(data, &mut chunk) - } - } - }; - - // 保存 chunk - if self.save_raw_chunks { - self.chunks.push(chunk); - } - - self.chunk_index += 1; - result - } - - /// 处理 OpenAI 格式的 chunk - /// - /// OpenAI 流式响应格式: - /// ```text - /// data: {"id":"chatcmpl-xxx","object":"chat.completion.chunk","created":1234567890, - /// "model":"gpt-4","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]} - /// data: [DONE] - /// ``` - fn process_openai_chunk( - &mut self, - data: &str, - chunk: &mut StreamChunk, - ) -> Result<(), StreamRebuilderError> { - let data = data.trim(); - - // 处理 [DONE] 终止信号 - if data == "[DONE]" { - return Ok(()); - } - - // 解析 JSON - let json: serde_json::Value = serde_json::from_str(data)?; - - // 提取响应 ID 和模型 - if self.response_id.is_none() { - self.response_id = json - .get("id") - .and_then(|v| v.as_str()) - .map(|s| s.to_string()); - } - if self.model.is_none() { - self.model = json - .get("model") - .and_then(|v| v.as_str()) - .map(|s| s.to_string()); - } - - // 处理 choices - if let Some(choices) = json.get("choices").and_then(|v| v.as_array()) { - for choice in choices { - // 处理 delta - if let Some(delta) = choice.get("delta") { - // 处理内容增量 - if let Some(content) = delta.get("content").and_then(|v| v.as_str()) { - self.content_buffer.push_str(content); - chunk.content_delta = Some(content.to_string()); - } - - // 处理工具调用增量 - if let Some(tool_calls) = delta.get("tool_calls").and_then(|v| v.as_array()) { - for tc in tool_calls { - self.process_openai_tool_call_delta(tc, chunk)?; - } - } - } - - // 处理 finish_reason - if let Some(finish_reason) = choice.get("finish_reason").and_then(|v| v.as_str()) { - self.stop_reason = Some(Self::parse_openai_stop_reason(finish_reason)); - } - } - } - - // 处理 usage(某些 API 在最后一个 chunk 中包含 usage) - if let Some(usage) = json.get("usage") { - self.parse_openai_usage(usage); - } - - Ok(()) - } - - /// 处理 OpenAI 工具调用增量 - fn process_openai_tool_call_delta( - &mut self, - tc: &serde_json::Value, - chunk: &mut StreamChunk, - ) -> Result<(), StreamRebuilderError> { - let index = tc.get("index").and_then(|v| v.as_u64()).unwrap_or(0) as u32; - - let builder = self.tool_calls_buffer.entry(index).or_default(); - - // 提取 ID - if let Some(id) = tc.get("id").and_then(|v| v.as_str()) { - builder.id = Some(id.to_string()); - } - - // 提取函数信息 - if let Some(function) = tc.get("function") { - if let Some(name) = function.get("name").and_then(|v| v.as_str()) { - builder.function_name = Some(name.to_string()); - } - if let Some(args) = function.get("arguments").and_then(|v| v.as_str()) { - builder.arguments.push_str(args); - - // 记录增量 - chunk.tool_call_delta = Some(ToolCallDelta { - index, - id: builder.id.clone(), - function_name: builder.function_name.clone(), - arguments_delta: Some(args.to_string()), - }); - } - } - - Ok(()) - } - - /// 解析 OpenAI 停止原因 - fn parse_openai_stop_reason(reason: &str) -> StopReason { - match reason { - "stop" => StopReason::Stop, - "length" => StopReason::Length, - "tool_calls" => StopReason::ToolCalls, - "content_filter" => StopReason::ContentFilter, - "function_call" => StopReason::FunctionCall, - other => StopReason::Other(other.to_string()), - } - } - - /// 解析 OpenAI usage - fn parse_openai_usage(&mut self, usage: &serde_json::Value) { - if let Some(prompt_tokens) = usage.get("prompt_tokens").and_then(|v| v.as_u64()) { - self.usage.input_tokens = prompt_tokens as u32; - } - if let Some(completion_tokens) = usage.get("completion_tokens").and_then(|v| v.as_u64()) { - self.usage.output_tokens = completion_tokens as u32; - } - if let Some(total_tokens) = usage.get("total_tokens").and_then(|v| v.as_u64()) { - self.usage.total_tokens = total_tokens as u32; - } - } - - /// 处理 Anthropic 格式的 chunk - /// - /// Anthropic 流式响应格式: - /// ```text - /// event: message_start - /// data: {"type":"message_start","message":{"id":"msg_xxx","type":"message","role":"assistant","model":"claude-3"}} - /// - /// event: content_block_start - /// data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}} - /// - /// event: content_block_delta - /// data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hello"}} - /// - /// event: content_block_stop - /// data: {"type":"content_block_stop","index":0} - /// - /// event: message_delta - /// data: {"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":10}} - /// - /// event: message_stop - /// data: {"type":"message_stop"} - /// ``` - fn process_anthropic_chunk( - &mut self, - event: Option<&str>, - data: &str, - chunk: &mut StreamChunk, - ) -> Result<(), StreamRebuilderError> { - let data = data.trim(); - - // 空数据跳过 - if data.is_empty() { - return Ok(()); - } - - // 解析 JSON - let json: serde_json::Value = serde_json::from_str(data)?; - - // 根据事件类型处理 - let event_type = event.or_else(|| json.get("type").and_then(|v| v.as_str())); - - match event_type { - Some("message_start") => { - self.process_anthropic_message_start(&json)?; - } - Some("content_block_start") => { - self.process_anthropic_content_block_start(&json)?; - } - Some("content_block_delta") => { - self.process_anthropic_content_block_delta(&json, chunk)?; - } - Some("content_block_stop") => { - self.process_anthropic_content_block_stop(&json)?; - } - Some("message_delta") => { - self.process_anthropic_message_delta(&json)?; - } - Some("message_stop") => { - // 消息结束,无需特殊处理 - } - Some("ping") => { - // 心跳,忽略 - } - Some("error") => { - // 错误事件 - if let Some(error) = json.get("error") { - let msg = error - .get("message") - .and_then(|v| v.as_str()) - .unwrap_or("Unknown error"); - return Err(StreamRebuilderError::InvalidEventFormat(msg.to_string())); - } - } - _ => { - // 未知事件类型,忽略 - } - } - - Ok(()) - } - - /// 处理 Anthropic message_start 事件 - fn process_anthropic_message_start( - &mut self, - json: &serde_json::Value, - ) -> Result<(), StreamRebuilderError> { - if let Some(message) = json.get("message") { - self.response_id = message - .get("id") - .and_then(|v| v.as_str()) - .map(|s| s.to_string()); - self.model = message - .get("model") - .and_then(|v| v.as_str()) - .map(|s| s.to_string()); - - // 处理 usage(input_tokens) - if let Some(usage) = message.get("usage") { - if let Some(input_tokens) = usage.get("input_tokens").and_then(|v| v.as_u64()) { - self.usage.input_tokens = input_tokens as u32; - } - } - } - Ok(()) - } - - /// 处理 Anthropic content_block_start 事件 - fn process_anthropic_content_block_start( - &mut self, - json: &serde_json::Value, - ) -> Result<(), StreamRebuilderError> { - let index = json.get("index").and_then(|v| v.as_u64()).unwrap_or(0) as u32; - self.current_content_block_index = Some(index); - - if let Some(content_block) = json.get("content_block") { - let block_type = content_block - .get("type") - .and_then(|v| v.as_str()) - .unwrap_or("text"); - self.current_content_block_type = Some(block_type.to_string()); - - match block_type { - "tool_use" => { - // 工具调用开始 - let builder = self.tool_calls_buffer.entry(index).or_default(); - builder.id = content_block - .get("id") - .and_then(|v| v.as_str()) - .map(|s| s.to_string()); - builder.function_name = content_block - .get("name") - .and_then(|v| v.as_str()) - .map(|s| s.to_string()); - } - "thinking" => { - // 思维链开始 - if self.thinking_buffer.is_none() { - self.thinking_buffer = Some(String::new()); - } - } - _ => {} - } - } - - Ok(()) - } - - /// 处理 Anthropic content_block_delta 事件 - fn process_anthropic_content_block_delta( - &mut self, - json: &serde_json::Value, - chunk: &mut StreamChunk, - ) -> Result<(), StreamRebuilderError> { - let index = json.get("index").and_then(|v| v.as_u64()).unwrap_or(0) as u32; - - if let Some(delta) = json.get("delta") { - let delta_type = delta.get("type").and_then(|v| v.as_str()).unwrap_or(""); - - match delta_type { - "text_delta" => { - // 文本增量 - if let Some(text) = delta.get("text").and_then(|v| v.as_str()) { - self.content_buffer.push_str(text); - chunk.content_delta = Some(text.to_string()); - } - } - "thinking_delta" => { - // 思维链增量 - if let Some(thinking) = delta.get("thinking").and_then(|v| v.as_str()) { - if let Some(ref mut buffer) = self.thinking_buffer { - buffer.push_str(thinking); - } else { - self.thinking_buffer = Some(thinking.to_string()); - } - chunk.thinking_delta = Some(thinking.to_string()); - } - } - "input_json_delta" => { - // 工具调用参数增量 - if let Some(partial_json) = delta.get("partial_json").and_then(|v| v.as_str()) { - if let Some(builder) = self.tool_calls_buffer.get_mut(&index) { - builder.arguments.push_str(partial_json); - - chunk.tool_call_delta = Some(ToolCallDelta { - index, - id: builder.id.clone(), - function_name: builder.function_name.clone(), - arguments_delta: Some(partial_json.to_string()), - }); - } - } - } - "signature_delta" => { - // 签名增量(用于思维链验证) - // 暂时忽略 - } - _ => {} - } - } - - Ok(()) - } - - /// 处理 Anthropic content_block_stop 事件 - fn process_anthropic_content_block_stop( - &mut self, - _json: &serde_json::Value, - ) -> Result<(), StreamRebuilderError> { - self.current_content_block_index = None; - self.current_content_block_type = None; - Ok(()) - } - - /// 处理 Anthropic message_delta 事件 - fn process_anthropic_message_delta( - &mut self, - json: &serde_json::Value, - ) -> Result<(), StreamRebuilderError> { - // 处理停止原因 - if let Some(delta) = json.get("delta") { - if let Some(stop_reason) = delta.get("stop_reason").and_then(|v| v.as_str()) { - self.stop_reason = Some(Self::parse_anthropic_stop_reason(stop_reason)); - } - } - - // 处理 usage - if let Some(usage) = json.get("usage") { - if let Some(output_tokens) = usage.get("output_tokens").and_then(|v| v.as_u64()) { - self.usage.output_tokens = output_tokens as u32; - } - } - - Ok(()) - } - - /// 解析 Anthropic 停止原因 - fn parse_anthropic_stop_reason(reason: &str) -> StopReason { - match reason { - "end_turn" => StopReason::EndTurn, - "stop_sequence" => StopReason::Stop, - "max_tokens" => StopReason::Length, - "tool_use" => StopReason::ToolCalls, - other => StopReason::Other(other.to_string()), - } - } - - /// 处理 Gemini 格式的 chunk - /// - /// Gemini 流式响应格式: - /// ```text - /// data: {"candidates":[{"content":{"parts":[{"text":"Hello"}],"role":"model"}, - /// "finishReason":"STOP","index":0}],"usageMetadata":{"promptTokenCount":10,"candidatesTokenCount":5}} - /// ``` - fn process_gemini_chunk( - &mut self, - data: &str, - chunk: &mut StreamChunk, - ) -> Result<(), StreamRebuilderError> { - let data = data.trim(); - - // 空数据跳过 - if data.is_empty() { - return Ok(()); - } - - // 解析 JSON - let json: serde_json::Value = serde_json::from_str(data)?; - - // 处理 candidates - if let Some(candidates) = json.get("candidates").and_then(|v| v.as_array()) { - for candidate in candidates { - // 处理内容 - if let Some(content) = candidate.get("content") { - if let Some(parts) = content.get("parts").and_then(|v| v.as_array()) { - for part in parts { - if let Some(text) = part.get("text").and_then(|v| v.as_str()) { - self.content_buffer.push_str(text); - chunk.content_delta = Some(text.to_string()); - } - - // 处理函数调用 - if let Some(function_call) = part.get("functionCall") { - self.process_gemini_function_call(function_call, chunk)?; - } - } - } - } - - // 处理 finishReason - if let Some(finish_reason) = candidate.get("finishReason").and_then(|v| v.as_str()) - { - self.stop_reason = Some(Self::parse_gemini_stop_reason(finish_reason)); - } - } - } - - // 处理 usageMetadata - if let Some(usage) = json.get("usageMetadata") { - self.parse_gemini_usage(usage); - } - - Ok(()) - } - - /// 处理 Gemini 函数调用 - fn process_gemini_function_call( - &mut self, - function_call: &serde_json::Value, - chunk: &mut StreamChunk, - ) -> Result<(), StreamRebuilderError> { - let index = self.tool_calls_buffer.len() as u32; - let builder = self.tool_calls_buffer.entry(index).or_default(); - - // Gemini 的函数调用通常是完整的,不是增量的 - if let Some(name) = function_call.get("name").and_then(|v| v.as_str()) { - builder.function_name = Some(name.to_string()); - builder.id = Some(format!("call_{}", uuid::Uuid::new_v4())); - } - - if let Some(args) = function_call.get("args") { - let args_str = serde_json::to_string(args)?; - builder.arguments = args_str.clone(); - - chunk.tool_call_delta = Some(ToolCallDelta { - index, - id: builder.id.clone(), - function_name: builder.function_name.clone(), - arguments_delta: Some(args_str), - }); - } - - Ok(()) - } - - /// 解析 Gemini 停止原因 - fn parse_gemini_stop_reason(reason: &str) -> StopReason { - match reason { - "STOP" => StopReason::Stop, - "MAX_TOKENS" => StopReason::Length, - "SAFETY" => StopReason::ContentFilter, - "RECITATION" => StopReason::ContentFilter, - "FUNCTION_CALL" => StopReason::ToolCalls, - other => StopReason::Other(other.to_string()), - } - } - - /// 解析 Gemini usage - fn parse_gemini_usage(&mut self, usage: &serde_json::Value) { - if let Some(prompt_tokens) = usage.get("promptTokenCount").and_then(|v| v.as_u64()) { - self.usage.input_tokens = prompt_tokens as u32; - } - if let Some(candidates_tokens) = usage.get("candidatesTokenCount").and_then(|v| v.as_u64()) - { - self.usage.output_tokens = candidates_tokens as u32; - } - if let Some(total_tokens) = usage.get("totalTokenCount").and_then(|v| v.as_u64()) { - self.usage.total_tokens = total_tokens as u32; - } - } - - /// 完成流重建,返回完整的 LLM 响应 - /// - /// 合并累积的内容、工具调用、思维链,计算流式统计信息。 - pub fn finish(self) -> LLMResponse { - let now = Utc::now(); - - // 计算流式统计信息 - let stream_info = self.calculate_stream_info(); - - // 构建思维链内容 - let thinking = self.thinking_buffer.clone().map(|text| ThinkingContent { - text, - tokens: self.usage.thinking_tokens, - signature: None, - }); - - // 构建工具调用列表 - let mut tool_calls: Vec = self - .tool_calls_buffer - .iter() - .filter_map(|(_, builder)| builder.clone().build()) - .collect(); - - // 按索引排序(如果有多个工具调用) - tool_calls.sort_by_key(|tc| tc.id.clone()); - - // 构建响应体 JSON - let body = self.build_response_body(&tool_calls, &thinking); - - // 计算 Token 总数 - let mut usage = self.usage.clone(); - usage.calculate_total(); - - // 确定时间戳 - let timestamp_start = self.first_chunk_time.unwrap_or(now); - let timestamp_end = self.last_chunk_time.unwrap_or(now); - - LLMResponse { - status_code: 200, - status_text: "OK".to_string(), - headers: HashMap::new(), - body, - content: self.content_buffer, - thinking, - tool_calls, - usage, - stop_reason: self.stop_reason, - size_bytes: 0, // 将在外部计算 - timestamp_start, - timestamp_end, - stream_info: Some(stream_info), - } - } - - /// 计算流式统计信息 - fn calculate_stream_info(&self) -> StreamInfo { - let chunk_count = self.chunk_index; - - // 计算首个 chunk 延迟 - let first_chunk_latency_ms = 0; // 需要外部提供请求开始时间 - - // 计算平均 chunk 间隔 - let avg_chunk_interval_ms = if chunk_count > 1 { - if let (Some(first), Some(last)) = (self.first_chunk_time, self.last_chunk_time) { - let total_ms = (last - first).num_milliseconds() as f64; - total_ms / (chunk_count - 1) as f64 - } else { - 0.0 - } - } else { - 0.0 - }; - - StreamInfo { - chunk_count, - first_chunk_latency_ms, - avg_chunk_interval_ms, - raw_chunks: if self.save_raw_chunks { - Some(self.chunks.clone()) - } else { - None - }, - } - } - - /// 构建响应体 JSON - fn build_response_body( - &self, - tool_calls: &[ToolCall], - thinking: &Option, - ) -> serde_json::Value { - match self.format { - StreamFormat::OpenAI => self.build_openai_response_body(tool_calls), - StreamFormat::Anthropic => self.build_anthropic_response_body(tool_calls, thinking), - StreamFormat::Gemini => self.build_gemini_response_body(tool_calls), - StreamFormat::Unknown => serde_json::json!({ - "content": self.content_buffer, - "tool_calls": tool_calls, - }), - } - } - - /// 构建 OpenAI 格式响应体 - fn build_openai_response_body(&self, tool_calls: &[ToolCall]) -> serde_json::Value { - let mut message = serde_json::json!({ - "role": "assistant", - "content": if self.content_buffer.is_empty() { serde_json::Value::Null } else { serde_json::json!(self.content_buffer) }, - }); - - if !tool_calls.is_empty() { - let tc_json: Vec = tool_calls - .iter() - .map(|tc| { - serde_json::json!({ - "id": tc.id, - "type": tc.tool_type, - "function": { - "name": tc.function.name, - "arguments": tc.function.arguments, - } - }) - }) - .collect(); - message["tool_calls"] = serde_json::json!(tc_json); - } - - serde_json::json!({ - "id": self.response_id.clone().unwrap_or_default(), - "object": "chat.completion", - "model": self.model.clone().unwrap_or_default(), - "choices": [{ - "index": 0, - "message": message, - "finish_reason": self.stop_reason.as_ref().map(|r| format!("{r:?}").to_lowercase()), - }], - "usage": { - "prompt_tokens": self.usage.input_tokens, - "completion_tokens": self.usage.output_tokens, - "total_tokens": self.usage.total_tokens, - } - }) - } - - /// 构建 Anthropic 格式响应体 - fn build_anthropic_response_body( - &self, - tool_calls: &[ToolCall], - thinking: &Option, - ) -> serde_json::Value { - let mut content: Vec = Vec::new(); - - // 添加思维链内容 - if let Some(ref thinking_content) = thinking { - content.push(serde_json::json!({ - "type": "thinking", - "thinking": thinking_content.text, - })); - } - - // 添加文本内容 - if !self.content_buffer.is_empty() { - content.push(serde_json::json!({ - "type": "text", - "text": self.content_buffer, - })); - } - - // 添加工具调用 - for tc in tool_calls { - let input: serde_json::Value = - serde_json::from_str(&tc.function.arguments).unwrap_or(serde_json::json!({})); - content.push(serde_json::json!({ - "type": "tool_use", - "id": tc.id, - "name": tc.function.name, - "input": input, - })); - } - - serde_json::json!({ - "id": self.response_id.clone().unwrap_or_default(), - "type": "message", - "role": "assistant", - "model": self.model.clone().unwrap_or_default(), - "content": content, - "stop_reason": self.stop_reason.as_ref().map(|r| match r { - StopReason::EndTurn => "end_turn", - StopReason::Stop => "stop_sequence", - StopReason::Length => "max_tokens", - StopReason::ToolCalls => "tool_use", - _ => "end_turn", - }), - "usage": { - "input_tokens": self.usage.input_tokens, - "output_tokens": self.usage.output_tokens, - } - }) - } - - /// 构建 Gemini 格式响应体 - fn build_gemini_response_body(&self, tool_calls: &[ToolCall]) -> serde_json::Value { - let mut parts: Vec = Vec::new(); - - // 添加文本内容 - if !self.content_buffer.is_empty() { - parts.push(serde_json::json!({ - "text": self.content_buffer, - })); - } - - // 添加函数调用 - for tc in tool_calls { - let args: serde_json::Value = - serde_json::from_str(&tc.function.arguments).unwrap_or(serde_json::json!({})); - parts.push(serde_json::json!({ - "functionCall": { - "name": tc.function.name, - "args": args, - } - })); - } - - serde_json::json!({ - "candidates": [{ - "content": { - "parts": parts, - "role": "model", - }, - "finishReason": self.stop_reason.as_ref().map(|r| match r { - StopReason::Stop => "STOP", - StopReason::Length => "MAX_TOKENS", - StopReason::ContentFilter => "SAFETY", - StopReason::ToolCalls => "FUNCTION_CALL", - _ => "STOP", - }), - }], - "usageMetadata": { - "promptTokenCount": self.usage.input_tokens, - "candidatesTokenCount": self.usage.output_tokens, - "totalTokenCount": self.usage.total_tokens, - } - }) - } - - /// 获取当前格式 - pub fn format(&self) -> StreamFormat { - self.format - } - - /// 获取当前内容 - pub fn content(&self) -> &str { - &self.content_buffer - } - - /// 获取 chunk 数量 - pub fn chunk_count(&self) -> u32 { - self.chunk_index - } -} - -// ============================================================================ -// 单元测试 -// ============================================================================ - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_openai_simple_stream() { - let mut rebuilder = StreamRebuilder::new(StreamFormat::OpenAI); - - // 模拟 OpenAI 流式响应 - let chunks = vec![ - r#"{"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4","choices":[{"index":0,"delta":{"role":"assistant","content":""},"finish_reason":null}]}"#, - r#"{"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4","choices":[{"index":0,"delta":{"content":"Hello"},"finish_reason":null}]}"#, - r#"{"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4","choices":[{"index":0,"delta":{"content":" world"},"finish_reason":null}]}"#, - r#"{"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4","choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}"#, - "[DONE]", - ]; - - for chunk in chunks { - rebuilder.process_event(None, chunk).unwrap(); - } - - let response = rebuilder.finish(); - assert_eq!(response.content, "Hello world"); - assert_eq!(response.stop_reason, Some(StopReason::Stop)); - } - - #[test] - fn test_openai_tool_calls_stream() { - let mut rebuilder = StreamRebuilder::new(StreamFormat::OpenAI); - - let chunks = vec![ - r#"{"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4","choices":[{"index":0,"delta":{"role":"assistant","content":null,"tool_calls":[{"index":0,"id":"call_abc123","type":"function","function":{"name":"get_weather","arguments":""}}]},"finish_reason":null}]}"#, - r#"{"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"arguments":"{\"lo"}}]},"finish_reason":null}]}"#, - r#"{"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"arguments":"cation\":"}}]},"finish_reason":null}]}"#, - r#"{"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"arguments":"\"NYC\"}"}}]},"finish_reason":null}]}"#, - r#"{"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4","choices":[{"index":0,"delta":{},"finish_reason":"tool_calls"}]}"#, - "[DONE]", - ]; - - for chunk in chunks { - rebuilder.process_event(None, chunk).unwrap(); - } - - let response = rebuilder.finish(); - assert_eq!(response.tool_calls.len(), 1); - assert_eq!(response.tool_calls[0].function.name, "get_weather"); - assert_eq!( - response.tool_calls[0].function.arguments, - r#"{"location":"NYC"}"# - ); - assert_eq!(response.stop_reason, Some(StopReason::ToolCalls)); - } - - #[test] - fn test_anthropic_simple_stream() { - let mut rebuilder = StreamRebuilder::new(StreamFormat::Anthropic); - - let events = vec![ - ( - "message_start", - r#"{"type":"message_start","message":{"id":"msg_123","type":"message","role":"assistant","model":"claude-3-opus-20240229","usage":{"input_tokens":10}}}"#, - ), - ( - "content_block_start", - r#"{"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}"#, - ), - ( - "content_block_delta", - r#"{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hello"}}"#, - ), - ( - "content_block_delta", - r#"{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":" world"}}"#, - ), - ( - "content_block_stop", - r#"{"type":"content_block_stop","index":0}"#, - ), - ( - "message_delta", - r#"{"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":5}}"#, - ), - ("message_stop", r#"{"type":"message_stop"}"#), - ]; - - for (event, data) in events { - rebuilder.process_event(Some(event), data).unwrap(); - } - - let response = rebuilder.finish(); - assert_eq!(response.content, "Hello world"); - assert_eq!(response.stop_reason, Some(StopReason::EndTurn)); - assert_eq!(response.usage.input_tokens, 10); - assert_eq!(response.usage.output_tokens, 5); - } - - #[test] - fn test_anthropic_tool_use_stream() { - let mut rebuilder = StreamRebuilder::new(StreamFormat::Anthropic); - - let events = vec![ - ( - "message_start", - r#"{"type":"message_start","message":{"id":"msg_123","type":"message","role":"assistant","model":"claude-3","usage":{"input_tokens":10}}}"#, - ), - ( - "content_block_start", - r#"{"type":"content_block_start","index":0,"content_block":{"type":"tool_use","id":"toolu_123","name":"get_weather"}}"#, - ), - ( - "content_block_delta", - r#"{"type":"content_block_delta","index":0,"delta":{"type":"input_json_delta","partial_json":"{\"loc"}}"#, - ), - ( - "content_block_delta", - r#"{"type":"content_block_delta","index":0,"delta":{"type":"input_json_delta","partial_json":"ation\":\"NYC\"}"}}"#, - ), - ( - "content_block_stop", - r#"{"type":"content_block_stop","index":0}"#, - ), - ( - "message_delta", - r#"{"type":"message_delta","delta":{"stop_reason":"tool_use"},"usage":{"output_tokens":20}}"#, - ), - ("message_stop", r#"{"type":"message_stop"}"#), - ]; - - for (event, data) in events { - rebuilder.process_event(Some(event), data).unwrap(); - } - - let response = rebuilder.finish(); - assert_eq!(response.tool_calls.len(), 1); - assert_eq!(response.tool_calls[0].id, "toolu_123"); - assert_eq!(response.tool_calls[0].function.name, "get_weather"); - assert_eq!( - response.tool_calls[0].function.arguments, - r#"{"location":"NYC"}"# - ); - assert_eq!(response.stop_reason, Some(StopReason::ToolCalls)); - } - - #[test] - fn test_anthropic_thinking_stream() { - let mut rebuilder = StreamRebuilder::new(StreamFormat::Anthropic); - - let events = vec![ - ( - "message_start", - r#"{"type":"message_start","message":{"id":"msg_123","type":"message","role":"assistant","model":"claude-3","usage":{"input_tokens":10}}}"#, - ), - ( - "content_block_start", - r#"{"type":"content_block_start","index":0,"content_block":{"type":"thinking","thinking":""}}"#, - ), - ( - "content_block_delta", - r#"{"type":"content_block_delta","index":0,"delta":{"type":"thinking_delta","thinking":"Let me think..."}}"#, - ), - ( - "content_block_stop", - r#"{"type":"content_block_stop","index":0}"#, - ), - ( - "content_block_start", - r#"{"type":"content_block_start","index":1,"content_block":{"type":"text","text":""}}"#, - ), - ( - "content_block_delta", - r#"{"type":"content_block_delta","index":1,"delta":{"type":"text_delta","text":"The answer is 42."}}"#, - ), - ( - "content_block_stop", - r#"{"type":"content_block_stop","index":1}"#, - ), - ( - "message_delta", - r#"{"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":15}}"#, - ), - ("message_stop", r#"{"type":"message_stop"}"#), - ]; - - for (event, data) in events { - rebuilder.process_event(Some(event), data).unwrap(); - } - - let response = rebuilder.finish(); - assert_eq!(response.content, "The answer is 42."); - assert!(response.thinking.is_some()); - assert_eq!(response.thinking.unwrap().text, "Let me think..."); - } - - #[test] - fn test_gemini_simple_stream() { - let mut rebuilder = StreamRebuilder::new(StreamFormat::Gemini); - - let chunks = vec![ - r#"{"candidates":[{"content":{"parts":[{"text":"Hello"}],"role":"model"},"index":0}]}"#, - r#"{"candidates":[{"content":{"parts":[{"text":" world"}],"role":"model"},"index":0}]}"#, - r#"{"candidates":[{"content":{"parts":[{"text":"!"}],"role":"model"},"finishReason":"STOP","index":0}],"usageMetadata":{"promptTokenCount":10,"candidatesTokenCount":5,"totalTokenCount":15}}"#, - ]; - - for chunk in chunks { - rebuilder.process_event(None, chunk).unwrap(); - } - - let response = rebuilder.finish(); - assert_eq!(response.content, "Hello world!"); - assert_eq!(response.stop_reason, Some(StopReason::Stop)); - assert_eq!(response.usage.input_tokens, 10); - assert_eq!(response.usage.output_tokens, 5); - } - - #[test] - fn test_done_signal() { - let mut rebuilder = StreamRebuilder::new(StreamFormat::OpenAI); - - // [DONE] 信号应该被正确处理 - rebuilder.process_event(None, "[DONE]").unwrap(); - - let response = rebuilder.finish(); - assert!(response.content.is_empty()); - } - - #[test] - fn test_auto_detect_format() { - // 测试自动检测 Anthropic 格式 - let mut rebuilder = StreamRebuilder::new(StreamFormat::Unknown); - rebuilder - .process_event( - Some("message_start"), - r#"{"type":"message_start","message":{"id":"msg_123","type":"message","role":"assistant","model":"claude-3","usage":{"input_tokens":10}}}"#, - ) - .unwrap(); - assert_eq!(rebuilder.format(), StreamFormat::Anthropic); - - // 测试自动检测 OpenAI 格式 - let mut rebuilder = StreamRebuilder::new(StreamFormat::Unknown); - rebuilder - .process_event( - None, - r#"{"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4","choices":[{"index":0,"delta":{"content":"Hi"},"finish_reason":null}]}"#, - ) - .unwrap(); - assert_eq!(rebuilder.format(), StreamFormat::OpenAI); - } - - #[test] - fn test_stream_info_calculation() { - let mut rebuilder = StreamRebuilder::new(StreamFormat::OpenAI).with_save_raw_chunks(true); - - let chunks = vec![ - r#"{"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4","choices":[{"index":0,"delta":{"content":"A"},"finish_reason":null}]}"#, - r#"{"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4","choices":[{"index":0,"delta":{"content":"B"},"finish_reason":null}]}"#, - r#"{"id":"chatcmpl-123","object":"chat.completion.chunk","created":1234567890,"model":"gpt-4","choices":[{"index":0,"delta":{"content":"C"},"finish_reason":null}]}"#, - "[DONE]", - ]; - - for chunk in chunks { - rebuilder.process_event(None, chunk).unwrap(); - } - - let response = rebuilder.finish(); - assert!(response.stream_info.is_some()); - let stream_info = response.stream_info.unwrap(); - assert_eq!(stream_info.chunk_count, 4); - assert!(stream_info.raw_chunks.is_some()); - assert_eq!(stream_info.raw_chunks.unwrap().len(), 4); - } -} - -// ============================================================================ -// 属性测试 -// ============================================================================ - -#[cfg(test)] -mod property_tests { - use super::*; - use proptest::prelude::*; - - // ======================================================================== - // 生成器 - // ======================================================================== - - /// 生成随机的文本内容(用于模拟 LLM 响应) - fn arb_content() -> impl Strategy { - prop::collection::vec("[a-zA-Z0-9 .,!?\\n]{1,20}", 1..10).prop_map(|parts| parts.join("")) - } - - /// 生成随机的工具调用 - fn arb_tool_call() -> impl Strategy { - ( - "[a-z_]{3,15}", // function name - "[a-f0-9]{8}", // call id suffix - prop::option::of("[a-zA-Z0-9_]{1,20}"), // argument value - ) - .prop_map(|(name, id_suffix, arg_value)| { - let id = format!("call_{id_suffix}"); - let args = match arg_value { - Some(val) => format!(r#"{{"value":"{val}"}}"#), - None => "{}".to_string(), - }; - (id, name, args) - }) - } - - /// 生成 OpenAI 格式的流式 chunks - fn generate_openai_chunks( - content: &str, - tool_calls: &[(String, String, String)], - ) -> Vec { - let mut chunks = Vec::new(); - let model = "gpt-4"; - let id = "chatcmpl-test123"; - - // 初始 chunk - chunks.push(format!( - r#"{{"id":"{id}","object":"chat.completion.chunk","created":1234567890,"model":"{model}","choices":[{{"index":0,"delta":{{"role":"assistant","content":""}},"finish_reason":null}}]}}"# - )); - - // 内容 chunks(每个字符一个 chunk) - for ch in content.chars() { - let escaped = match ch { - '"' => "\\\"".to_string(), - '\\' => "\\\\".to_string(), - '\n' => "\\n".to_string(), - _ => ch.to_string(), - }; - chunks.push(format!( - r#"{{"id":"{id}","object":"chat.completion.chunk","created":1234567890,"model":"{model}","choices":[{{"index":0,"delta":{{"content":"{escaped}"}},"finish_reason":null}}]}}"# - )); - } - - // 工具调用 chunks - for (idx, (call_id, name, args)) in tool_calls.iter().enumerate() { - // 工具调用开始 - chunks.push(format!( - r#"{{"id":"{id}","object":"chat.completion.chunk","created":1234567890,"model":"{model}","choices":[{{"index":0,"delta":{{"tool_calls":[{{"index":{idx},"id":"{call_id}","type":"function","function":{{"name":"{name}","arguments":""}}}}]}},"finish_reason":null}}]}}"# - )); - - // 工具调用参数(一次性发送,避免分块导致的转义问题) - let args_escaped = args.replace('\\', "\\\\").replace('"', "\\\""); - chunks.push(format!( - r#"{{"id":"{id}","object":"chat.completion.chunk","created":1234567890,"model":"{model}","choices":[{{"index":0,"delta":{{"tool_calls":[{{"index":{idx},"function":{{"arguments":"{args_escaped}"}}}}]}},"finish_reason":null}}]}}"# - )); - } - - // 结束 chunk - let finish_reason = if tool_calls.is_empty() { - "stop" - } else { - "tool_calls" - }; - chunks.push(format!( - r#"{{"id":"{id}","object":"chat.completion.chunk","created":1234567890,"model":"{model}","choices":[{{"index":0,"delta":{{}},"finish_reason":"{finish_reason}"}}]}}"# - )); - - // [DONE] 信号 - chunks.push("[DONE]".to_string()); - - chunks - } - - /// 生成 Anthropic 格式的流式 chunks - fn generate_anthropic_chunks( - content: &str, - tool_calls: &[(String, String, String)], - ) -> Vec<(String, String)> { - let mut events = Vec::new(); - let model = "claude-3-opus-20240229"; - let id = "msg_test123"; - - // message_start - events.push(( - "message_start".to_string(), - format!( - r#"{{"type":"message_start","message":{{"id":"{id}","type":"message","role":"assistant","model":"{model}","usage":{{"input_tokens":10}}}}}}"# - ), - )); - - // 文本内容块 - if !content.is_empty() { - events.push(( - "content_block_start".to_string(), - r#"{"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}"#.to_string(), - )); - - // 内容 delta(每个字符一个) - for ch in content.chars() { - let escaped = match ch { - '"' => "\\\"".to_string(), - '\\' => "\\\\".to_string(), - '\n' => "\\n".to_string(), - _ => ch.to_string(), - }; - events.push(( - "content_block_delta".to_string(), - format!( - r#"{{"type":"content_block_delta","index":0,"delta":{{"type":"text_delta","text":"{escaped}"}}}}"# - ), - )); - } - - events.push(( - "content_block_stop".to_string(), - r#"{"type":"content_block_stop","index":0}"#.to_string(), - )); - } - - // 工具调用块 - for (idx, (call_id, name, args)) in tool_calls.iter().enumerate() { - let block_idx = if content.is_empty() { idx } else { idx + 1 }; - - events.push(( - "content_block_start".to_string(), - format!( - r#"{{"type":"content_block_start","index":{block_idx},"content_block":{{"type":"tool_use","id":"{call_id}","name":"{name}"}}}}"# - ), - )); - - // 参数 delta(一次性发送,避免分块导致的转义问题) - let args_escaped = args.replace('\\', "\\\\").replace('"', "\\\""); - events.push(( - "content_block_delta".to_string(), - format!( - r#"{{"type":"content_block_delta","index":{block_idx},"delta":{{"type":"input_json_delta","partial_json":"{args_escaped}"}}}}"# - ), - )); - - events.push(( - "content_block_stop".to_string(), - format!(r#"{{"type":"content_block_stop","index":{block_idx}}}"#), - )); - } - - // message_delta - let stop_reason = if tool_calls.is_empty() { - "end_turn" - } else { - "tool_use" - }; - events.push(( - "message_delta".to_string(), - format!( - r#"{{"type":"message_delta","delta":{{"stop_reason":"{stop_reason}"}},"usage":{{"output_tokens":20}}}}"# - ), - )); - - // message_stop - events.push(( - "message_stop".to_string(), - r#"{"type":"message_stop"}"#.to_string(), - )); - - events - } - - /// 生成 Gemini 格式的流式 chunks - fn generate_gemini_chunks(content: &str) -> Vec { - let mut chunks = Vec::new(); - - // 内容 chunks(每 5 个字符一个 chunk) - for chunk_str in content.as_bytes().chunks(5) { - let chunk_content = String::from_utf8_lossy(chunk_str); - let escaped = chunk_content - .replace('\\', "\\\\") - .replace('"', "\\\"") - .replace('\n', "\\n"); - chunks.push(format!( - r#"{{"candidates":[{{"content":{{"parts":[{{"text":"{escaped}"}}],"role":"model"}},"index":0}}]}}"# - )); - } - - // 最后一个 chunk 包含 finishReason 和 usage - if chunks.is_empty() { - chunks.push( - r#"{"candidates":[{"content":{"parts":[],"role":"model"},"finishReason":"STOP","index":0}],"usageMetadata":{"promptTokenCount":10,"candidatesTokenCount":5,"totalTokenCount":15}}"#.to_string() - ); - } else { - // 修改最后一个 chunk 添加 finishReason - let last = chunks.pop().unwrap(); - let modified = last.replace( - r#""index":0}]}"#, - r#""finishReason":"STOP","index":0}],"usageMetadata":{"promptTokenCount":10,"candidatesTokenCount":5,"totalTokenCount":15}}"# - ); - chunks.push(modified); - } - - chunks - } - - // ======================================================================== - // 属性测试 - // ======================================================================== - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: llm-flow-monitor, Property 2: 流式响应重建 Round-Trip** - /// **Validates: Requirements 1.4, 1.5, 2.1, 2.2, 2.3** - /// - /// *对于任意* 有效的 LLM 响应内容,将其拆分为 SSE chunks 后通过 Stream_Rebuilder 重建, - /// 重建后的内容应该与原始内容等价(包括文本内容、工具调用和思维链)。 - #[test] - fn prop_openai_stream_roundtrip( - content in arb_content(), - ) { - // 生成 OpenAI 格式的 chunks - let chunks = generate_openai_chunks(&content, &[]); - - // 重建 - let mut rebuilder = StreamRebuilder::new(StreamFormat::OpenAI); - for chunk in chunks { - rebuilder.process_event(None, &chunk).unwrap(); - } - let response = rebuilder.finish(); - - // 验证内容一致 - prop_assert_eq!( - response.content, - content, - "OpenAI 流式重建后的内容应该与原始内容一致" - ); - - // 验证停止原因 - prop_assert_eq!( - response.stop_reason, - Some(StopReason::Stop), - "无工具调用时停止原因应该是 Stop" - ); - } - - /// **Feature: llm-flow-monitor, Property 2b: OpenAI 工具调用流式重建** - /// **Validates: Requirements 1.4, 1.5, 2.1** - #[test] - fn prop_openai_tool_calls_roundtrip( - tool_call in arb_tool_call(), - ) { - let (call_id, name, args) = tool_call; - let tool_calls = vec![(call_id.clone(), name.clone(), args.clone())]; - - // 生成 OpenAI 格式的 chunks - let chunks = generate_openai_chunks("", &tool_calls); - - // 重建 - let mut rebuilder = StreamRebuilder::new(StreamFormat::OpenAI); - for chunk in chunks { - rebuilder.process_event(None, &chunk).unwrap(); - } - let response = rebuilder.finish(); - - // 验证工具调用 - prop_assert_eq!( - response.tool_calls.len(), - 1, - "应该有一个工具调用" - ); - prop_assert_eq!( - &response.tool_calls[0].id, - &call_id, - "工具调用 ID 应该一致" - ); - prop_assert_eq!( - &response.tool_calls[0].function.name, - &name, - "函数名称应该一致" - ); - prop_assert_eq!( - &response.tool_calls[0].function.arguments, - &args, - "函数参数应该一致" - ); - prop_assert_eq!( - response.stop_reason, - Some(StopReason::ToolCalls), - "有工具调用时停止原因应该是 ToolCalls" - ); - } - - /// **Feature: llm-flow-monitor, Property 2c: Anthropic 流式重建** - /// **Validates: Requirements 1.4, 1.5, 2.2** - #[test] - fn prop_anthropic_stream_roundtrip( - content in arb_content(), - ) { - // 生成 Anthropic 格式的 events - let events = generate_anthropic_chunks(&content, &[]); - - // 重建 - let mut rebuilder = StreamRebuilder::new(StreamFormat::Anthropic); - for (event, data) in events { - rebuilder.process_event(Some(&event), &data).unwrap(); - } - let response = rebuilder.finish(); - - // 验证内容一致 - prop_assert_eq!( - response.content, - content, - "Anthropic 流式重建后的内容应该与原始内容一致" - ); - - // 验证停止原因 - prop_assert_eq!( - response.stop_reason, - Some(StopReason::EndTurn), - "无工具调用时停止原因应该是 EndTurn" - ); - } - - /// **Feature: llm-flow-monitor, Property 2d: Anthropic 工具调用流式重建** - /// **Validates: Requirements 1.4, 1.5, 2.2** - #[test] - fn prop_anthropic_tool_calls_roundtrip( - tool_call in arb_tool_call(), - ) { - let (call_id, name, args) = tool_call; - let tool_calls = vec![(call_id.clone(), name.clone(), args.clone())]; - - // 生成 Anthropic 格式的 events - let events = generate_anthropic_chunks("", &tool_calls); - - // 重建 - let mut rebuilder = StreamRebuilder::new(StreamFormat::Anthropic); - for (event, data) in events { - rebuilder.process_event(Some(&event), &data).unwrap(); - } - let response = rebuilder.finish(); - - // 验证工具调用 - prop_assert_eq!( - response.tool_calls.len(), - 1, - "应该有一个工具调用" - ); - prop_assert_eq!( - &response.tool_calls[0].id, - &call_id, - "工具调用 ID 应该一致" - ); - prop_assert_eq!( - &response.tool_calls[0].function.name, - &name, - "函数名称应该一致" - ); - prop_assert_eq!( - &response.tool_calls[0].function.arguments, - &args, - "函数参数应该一致" - ); - prop_assert_eq!( - response.stop_reason, - Some(StopReason::ToolCalls), - "有工具调用时停止原因应该是 ToolCalls" - ); - } - - /// **Feature: llm-flow-monitor, Property 2e: Gemini 流式重建** - /// **Validates: Requirements 1.4, 1.5, 2.3** - #[test] - fn prop_gemini_stream_roundtrip( - content in arb_content(), - ) { - // 生成 Gemini 格式的 chunks - let chunks = generate_gemini_chunks(&content); - - // 重建 - let mut rebuilder = StreamRebuilder::new(StreamFormat::Gemini); - for chunk in chunks { - rebuilder.process_event(None, &chunk).unwrap(); - } - let response = rebuilder.finish(); - - // 验证内容一致 - prop_assert_eq!( - response.content, - content, - "Gemini 流式重建后的内容应该与原始内容一致" - ); - - // 验证停止原因 - prop_assert_eq!( - response.stop_reason, - Some(StopReason::Stop), - "停止原因应该是 Stop" - ); - } - - /// **Feature: llm-flow-monitor, Property 2f: 流式统计信息正确性** - /// **Validates: Requirements 1.5** - #[test] - fn prop_stream_info_correctness( - content in arb_content(), - ) { - // 生成 OpenAI 格式的 chunks - let chunks = generate_openai_chunks(&content, &[]); - let expected_chunk_count = chunks.len() as u32; - - // 重建(保存原始 chunks) - let mut rebuilder = StreamRebuilder::new(StreamFormat::OpenAI).with_save_raw_chunks(true); - for chunk in &chunks { - rebuilder.process_event(None, chunk).unwrap(); - } - let response = rebuilder.finish(); - - // 验证流式统计信息 - prop_assert!(response.stream_info.is_some(), "应该有流式统计信息"); - let stream_info = response.stream_info.unwrap(); - - prop_assert_eq!( - stream_info.chunk_count, - expected_chunk_count, - "chunk 数量应该正确" - ); - - prop_assert!(stream_info.raw_chunks.is_some(), "应该保存原始 chunks"); - prop_assert_eq!( - stream_info.raw_chunks.unwrap().len(), - expected_chunk_count as usize, - "保存的 chunks 数量应该正确" - ); - } - } -} diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 20ebb5730..2ebbf37e5 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -2,12 +2,13 @@ //! //! 这是一个 Tauri 应用,提供 AI API 的代理和管理功能。 //! -//! ## Workspace 结构(方案 A - 最小化拆分) +//! ## Workspace 结构(渐进式拆分) //! -//! 采用最小化拆分策略,只迁移真正独立的模块: -//! - ✅ proxycast-core crate(models, data, logger) +//! - ✅ proxycast-core crate(models, data, logger, errors, backends, config, connect, +//! middleware, orchestrator, plugin, session 部分, session_files) //! - ✅ proxycast-infra crate(proxy, resilience, injection, telemetry) -//! - 主 crate 保留所有业务逻辑模块(包括 plugin,因依赖 Tauri) +//! - ✅ proxycast-providers crate(providers, converter, streaming, translator, stream, session 部分) +//! - 主 crate 保留 Tauri 相关业务逻辑 // 抑制 objc crate 宏内部的 unexpected_cfgs 警告 // 该警告来自 cocoa/objc 依赖的 msg_send! 宏,是已知的 issue @@ -26,26 +27,31 @@ pub use proxycast_infra::{ TimeoutController, TokenSource, TokenStatsSummary, TokenTracker, TokenUsageRecord, }; +// 从 providers crate 重新导出(保持 crate::xxx 路径兼容) +pub use proxycast_providers::converter; +pub use proxycast_providers::providers; +pub use proxycast_providers::stream; +pub use proxycast_providers::streaming; +pub use proxycast_providers::translator; + +// 从 core crate 重新导出(保持 crate::xxx 路径兼容) +pub use proxycast_core::backends; +pub use proxycast_core::connect; +pub use proxycast_core::orchestrator; +pub use proxycast_core::session_files; + // 核心模块 pub mod agent; pub mod app; -pub mod backends; -pub mod browser_interceptor; -pub mod connect; pub mod content; pub mod credential; pub mod database; -pub mod flow_monitor; pub mod memory; -pub mod orchestrator; pub mod plugin; pub mod screenshot; pub mod services; pub mod session; -pub mod session_files; -pub mod stream; pub mod terminal; -pub mod translator; pub mod tray; pub mod voice; pub mod workspace; @@ -59,22 +65,21 @@ pub mod mcp; // 内部模块 mod commands; mod config; -mod converter; mod data; #[cfg(debug_assertions)] mod dev_bridge; -mod errors; mod logger; mod models; -mod providers; mod server_utils; +// 从 core crate 重新导出 errors +pub use proxycast_core::errors; + // 服务器相关模块 mod middleware; mod processor; mod router; mod server; -mod streaming; mod websocket; // 重新导出核心类型以保持向后兼容 diff --git a/src-tauri/src/middleware/mod.rs b/src-tauri/src/middleware/mod.rs index 2dad8eb31..643b4d354 100644 --- a/src-tauri/src/middleware/mod.rs +++ b/src-tauri/src/middleware/mod.rs @@ -1,10 +1,5 @@ //! Middleware 模块 //! -//! 提供 HTTP 请求处理的中间件组件 +//! 从 proxycast-core 重新导出 -pub mod management_auth; - -#[cfg(test)] -mod tests; - -pub use management_auth::ManagementAuthLayer; +pub use proxycast_core::middleware::*; diff --git a/src-tauri/src/models/anthropic.rs b/src-tauri/src/models/anthropic.rs deleted file mode 100644 index 87926833f..000000000 --- a/src-tauri/src/models/anthropic.rs +++ /dev/null @@ -1,142 +0,0 @@ -//! Anthropic/Claude API 数据模型 -use serde::{Deserialize, Serialize}; - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(tag = "type")] -pub enum AnthropicContentBlock { - #[serde(rename = "text")] - Text { text: String }, - #[serde(rename = "tool_use")] - ToolUse { - id: String, - name: String, - input: serde_json::Value, - }, - #[serde(rename = "tool_result")] - ToolResult { - tool_use_id: String, - content: serde_json::Value, - }, - #[serde(rename = "image")] - Image { source: ImageSource }, - /// Extended Thinking 块 - #[serde(rename = "thinking")] - Thinking { - thinking: String, - /// 签名字段,用于验证思维内容的完整性 - signature: String, - }, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ImageSource { - #[serde(rename = "type")] - pub source_type: String, - pub media_type: String, - pub data: String, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct AnthropicMessage { - pub role: String, - pub content: serde_json::Value, // Can be string or array -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct AnthropicTool { - pub name: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub description: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub input_schema: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct AnthropicMessagesRequest { - pub model: String, - pub messages: Vec, - #[serde(skip_serializing_if = "Option::is_none")] - pub max_tokens: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub system: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub temperature: Option, - #[serde(default)] - pub stream: bool, - #[serde(skip_serializing_if = "Option::is_none")] - pub tools: Option>, - #[serde(skip_serializing_if = "Option::is_none")] - pub tool_choice: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct AnthropicUsage { - pub input_tokens: u32, - pub output_tokens: u32, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[allow(dead_code)] -pub struct AnthropicMessagesResponse { - pub id: String, - #[serde(rename = "type")] - pub response_type: String, - pub role: String, - pub content: Vec, - pub model: String, - pub stop_reason: Option, - pub usage: AnthropicUsage, -} - -// Streaming events -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(tag = "type")] -pub enum AnthropicStreamEvent { - #[serde(rename = "message_start")] - MessageStart { message: AnthropicMessageStart }, - #[serde(rename = "content_block_start")] - ContentBlockStart { - index: u32, - content_block: AnthropicContentBlock, - }, - #[serde(rename = "content_block_delta")] - ContentBlockDelta { index: u32, delta: AnthropicDelta }, - #[serde(rename = "content_block_stop")] - ContentBlockStop { index: u32 }, - #[serde(rename = "message_delta")] - MessageDelta { - delta: AnthropicMessageDelta, - usage: AnthropicUsage, - }, - #[serde(rename = "message_stop")] - MessageStop, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct AnthropicMessageStart { - pub id: String, - #[serde(rename = "type")] - pub msg_type: String, - pub role: String, - pub model: String, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(tag = "type")] -pub enum AnthropicDelta { - #[serde(rename = "text_delta")] - TextDelta { text: String }, - #[serde(rename = "input_json_delta")] - InputJsonDelta { partial_json: String }, - /// Extended Thinking delta - #[serde(rename = "thinking_delta")] - ThinkingDelta { thinking: String }, - /// Signature delta for thinking blocks - #[serde(rename = "signature_delta")] - SignatureDelta { signature: String }, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct AnthropicMessageDelta { - pub stop_reason: Option, -} diff --git a/src-tauri/src/models/app_type.rs b/src-tauri/src/models/app_type.rs deleted file mode 100644 index 4ca757550..000000000 --- a/src-tauri/src/models/app_type.rs +++ /dev/null @@ -1,41 +0,0 @@ -use serde::{Deserialize, Serialize}; - -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] -#[serde(rename_all = "lowercase")] -pub enum AppType { - ProxyCast, - Claude, - Codex, - Gemini, -} - -impl AppType { - pub fn as_str(&self) -> &'static str { - match self { - AppType::ProxyCast => "proxycast", - AppType::Claude => "claude", - AppType::Codex => "codex", - AppType::Gemini => "gemini", - } - } -} - -impl std::str::FromStr for AppType { - type Err = String; - - fn from_str(s: &str) -> Result { - match s.to_lowercase().as_str() { - "proxycast" => Ok(AppType::ProxyCast), - "claude" => Ok(AppType::Claude), - "codex" => Ok(AppType::Codex), - "gemini" => Ok(AppType::Gemini), - _ => Err(format!("Invalid app type: {s}")), - } - } -} - -impl std::fmt::Display for AppType { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!(f, "{}", self.as_str()) - } -} diff --git a/src-tauri/src/models/codewhisperer.rs b/src-tauri/src/models/codewhisperer.rs deleted file mode 100644 index 98197e0ab..000000000 --- a/src-tauri/src/models/codewhisperer.rs +++ /dev/null @@ -1,173 +0,0 @@ -//! CodeWhisperer/Kiro API 数据模型 -//! -//! 支持标准工具和特殊工具类型(如 web_search)。 -//! -//! # 更新日志 -//! -//! - 2025-12-27: 添加 CWWebSearchTool 支持,修复 Issue #49 -use serde::{Deserialize, Serialize}; - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct CodeWhispererRequest { - pub conversation_state: ConversationState, - #[serde(skip_serializing_if = "Option::is_none")] - pub profile_arn: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct ConversationState { - pub chat_trigger_type: String, - pub conversation_id: String, - pub current_message: CurrentMessage, - #[serde(skip_serializing_if = "Option::is_none")] - pub history: Option>, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct CurrentMessage { - pub user_input_message: UserInputMessage, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct UserInputMessage { - pub content: String, - pub model_id: String, - pub origin: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub images: Option>, - #[serde(skip_serializing_if = "Option::is_none")] - pub user_input_message_context: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct UserInputMessageContext { - #[serde(skip_serializing_if = "Option::is_none")] - pub tools: Option>, - #[serde(skip_serializing_if = "Option::is_none")] - pub tool_results: Option>, -} - -/// CodeWhisperer 工具项 -/// -/// 支持两种类型: -/// - 标准工具(带 tool_specification) -/// - 联网搜索工具(仅 type 字段) -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(untagged)] -pub enum CWToolItem { - /// 标准工具定义 - Standard(CWTool), - /// 联网搜索工具 - WebSearch(CWWebSearchTool), -} - -/// 标准工具定义 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct CWTool { - pub tool_specification: ToolSpecification, -} - -/// 联网搜索工具 -/// -/// Codex/Kiro API 支持的特殊工具类型,用于联网搜索。 -/// 格式:`{"type": "web_search"}` -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CWWebSearchTool { - #[serde(rename = "type")] - pub tool_type: String, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct ToolSpecification { - pub name: String, - pub description: String, - pub input_schema: InputSchema, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct InputSchema { - pub json: serde_json::Value, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct CWToolResult { - pub content: Vec, - pub status: String, - pub tool_use_id: String, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CWTextContent { - pub text: String, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CWImage { - pub format: String, - pub source: CWImageSource, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CWImageSource { - pub bytes: String, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(untagged)] -pub enum HistoryItem { - User(UserHistoryItem), - Assistant(AssistantHistoryItem), -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct UserHistoryItem { - pub user_input_message: UserInputMessage, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct AssistantHistoryItem { - pub assistant_response_message: AssistantResponseMessage, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct AssistantResponseMessage { - pub content: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub tool_uses: Option>, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct CWToolUse { - pub input: serde_json::Value, - pub name: String, - pub tool_use_id: String, -} - -// Response types -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct CWStreamEvent { - #[serde(skip_serializing_if = "Option::is_none")] - pub assistant_response_event: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "camelCase")] -pub struct AssistantResponseEvent { - #[serde(skip_serializing_if = "Option::is_none")] - pub content: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub tool_use: Option, -} diff --git a/src-tauri/src/models/kiro_fingerprint.rs b/src-tauri/src/models/kiro_fingerprint.rs deleted file mode 100644 index 384eb743c..000000000 --- a/src-tauri/src/models/kiro_fingerprint.rs +++ /dev/null @@ -1,199 +0,0 @@ -//! Kiro 凭证指纹绑定模型 -//! -//! 为每个 Kiro 凭证存储独立的 Machine ID,实现多账号指纹隔离。 - -#![allow(dead_code)] - -use chrono::{DateTime, Utc}; -use serde::{Deserialize, Serialize}; -use std::collections::HashMap; -use std::fs; -use std::path::PathBuf; - -/// Kiro 凭证指纹绑定 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct KiroFingerprintBinding { - /// 凭证 UUID - pub credential_uuid: String, - /// 绑定的 Machine ID - pub machine_id: String, - /// 创建时间 - pub created_at: DateTime, - /// 最后切换时间 - pub last_switched_at: Option>, -} - -/// 指纹绑定存储 -#[derive(Debug, Clone, Serialize, Deserialize, Default)] -pub struct KiroFingerprintStore { - /// 凭证 UUID -> 指纹绑定 - pub bindings: HashMap, -} - -impl KiroFingerprintStore { - /// 获取存储文件路径 - pub fn get_storage_path() -> Result { - let app_data_dir = dirs::data_dir() - .ok_or_else(|| "无法获取应用数据目录".to_string())? - .join("proxycast"); - - // 确保目录存在 - if !app_data_dir.exists() { - fs::create_dir_all(&app_data_dir).map_err(|e| format!("创建应用数据目录失败: {e}"))?; - } - - Ok(app_data_dir.join("kiro_fingerprints.json")) - } - - /// 从文件加载 - pub fn load() -> Result { - let path = Self::get_storage_path()?; - - if !path.exists() { - return Ok(Self::default()); - } - - let content = - fs::read_to_string(&path).map_err(|e| format!("读取指纹存储文件失败: {e}"))?; - - serde_json::from_str(&content).map_err(|e| format!("解析指纹存储文件失败: {e}")) - } - - /// 保存到文件 - pub fn save(&self) -> Result<(), String> { - let path = Self::get_storage_path()?; - let content = - serde_json::to_string_pretty(self).map_err(|e| format!("序列化指纹存储失败: {e}"))?; - - fs::write(&path, content).map_err(|e| format!("写入指纹存储文件失败: {e}")) - } - - /// 获取凭证的指纹绑定 - pub fn get_binding(&self, credential_uuid: &str) -> Option<&KiroFingerprintBinding> { - self.bindings.get(credential_uuid) - } - - /// 获取或创建凭证的指纹绑定 - /// - /// 如果凭证没有绑定指纹,会基于凭证信息生成一个新的 Machine ID - pub fn get_or_create_binding( - &mut self, - credential_uuid: &str, - profile_arn: Option<&str>, - client_id: Option<&str>, - ) -> Result<&KiroFingerprintBinding, String> { - if !self.bindings.contains_key(credential_uuid) { - // 生成基于凭证的 Machine ID - let machine_id = generate_stable_machine_id(credential_uuid, profile_arn, client_id); - - let binding = KiroFingerprintBinding { - credential_uuid: credential_uuid.to_string(), - machine_id, - created_at: Utc::now(), - last_switched_at: None, - }; - - self.bindings.insert(credential_uuid.to_string(), binding); - self.save()?; - } - - Ok(self.bindings.get(credential_uuid).unwrap()) - } - - /// 更新最后切换时间 - pub fn update_last_switched(&mut self, credential_uuid: &str) -> Result<(), String> { - if let Some(binding) = self.bindings.get_mut(credential_uuid) { - binding.last_switched_at = Some(Utc::now()); - self.save()?; - } - Ok(()) - } - - /// 删除凭证的指纹绑定 - pub fn remove_binding(&mut self, credential_uuid: &str) -> Result<(), String> { - self.bindings.remove(credential_uuid); - self.save() - } -} - -/// 生成稳定的 Machine ID -/// -/// 基于凭证信息生成一个稳定的 UUID 格式 Machine ID。 -/// 同一凭证每次生成的 Machine ID 相同,确保账号身份一致。 -fn generate_stable_machine_id( - credential_uuid: &str, - profile_arn: Option<&str>, - client_id: Option<&str>, -) -> String { - use sha2::{Digest, Sha256}; - - // 使用凭证相关信息作为种子 - let seed = format!( - "kiro_fingerprint:{}:{}:{}", - credential_uuid, - profile_arn.unwrap_or(""), - client_id.unwrap_or("") - ); - - let mut hasher = Sha256::new(); - hasher.update(seed.as_bytes()); - let result = hasher.finalize(); - - // 将哈希结果转换为 UUID 格式 - let hex = format!("{result:x}"); - format!( - "{}-{}-{}-{}-{}", - &hex[0..8], - &hex[8..12], - &hex[12..16], - &hex[16..20], - &hex[20..32] - ) -} - -/// 切换到本地的结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct SwitchToLocalResult { - /// 是否成功 - pub success: bool, - /// 结果消息 - pub message: String, - /// 是否需要用户操作(如需管理员权限) - pub requires_action: bool, - /// 切换的 Machine ID - pub machine_id: Option, - /// 是否需要重启 Kiro IDE - pub requires_kiro_restart: bool, -} - -impl SwitchToLocalResult { - pub fn success(message: impl Into, machine_id: String) -> Self { - Self { - success: true, - message: message.into(), - requires_action: false, - machine_id: Some(machine_id), - requires_kiro_restart: true, - } - } - - pub fn error(message: impl Into) -> Self { - Self { - success: false, - message: message.into(), - requires_action: false, - machine_id: None, - requires_kiro_restart: false, - } - } - - pub fn requires_admin(message: impl Into) -> Self { - Self { - success: false, - message: message.into(), - requires_action: true, - machine_id: None, - requires_kiro_restart: false, - } - } -} diff --git a/src-tauri/src/models/machine_id.rs b/src-tauri/src/models/machine_id.rs deleted file mode 100644 index 0bc295d26..000000000 --- a/src-tauri/src/models/machine_id.rs +++ /dev/null @@ -1,209 +0,0 @@ -use serde::{Deserialize, Serialize}; - -/// 机器码信息结构 - v0.20.0 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct MachineIdInfo { - /// 当前机器码 - pub current_id: String, - /// 原始机器码(如果有备份) - pub original_id: Option, - /// 操作系统平台 - pub platform: String, - /// 是否可以修改 - pub can_modify: bool, - /// 是否需要管理员权限 - pub requires_admin: bool, - /// 是否存在备份 - pub backup_exists: bool, - /// 机器码格式类型 - pub format_type: MachineIdFormat, -} - -/// 机器码操作结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct MachineIdResult { - /// 操作是否成功 - pub success: bool, - /// 结果消息 - pub message: String, - /// 是否需要重启 - pub requires_restart: bool, - /// 是否需要管理员权限 - pub requires_admin: bool, - /// 新的机器码(如果操作成功) - pub new_machine_id: Option, -} - -/// 管理员权限状态 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct AdminStatus { - /// 是否具有管理员权限 - pub is_admin: bool, - /// 操作系统平台 - pub platform: String, - /// 权限提升方法说明 - pub elevation_method: Option, - /// 权限检查是否成功 - pub check_success: bool, -} - -/// 机器码格式类型 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(rename_all = "lowercase")] -pub enum MachineIdFormat { - /// UUID 格式 (xxxxxxxx-xxxx-xxxx-xxxx-xxxxxxxxxxxx) - Uuid, - /// 32位十六进制格式 (xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx) - #[serde(rename = "hex32")] - Hex32, - /// 其他格式 - #[serde(rename = "unknown")] - Unknown, -} - -/// 机器码备份信息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct MachineIdBackup { - /// 备份的机器码 - pub machine_id: String, - /// 备份时间戳 - pub timestamp: i64, - /// 操作系统平台 - pub platform: String, - /// 机器码格式 - pub format: MachineIdFormat, - /// 备份描述 - pub description: Option, -} - -/// 机器码历史记录 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct MachineIdHistory { - /// 记录ID - pub id: String, - /// 机器码 - pub machine_id: String, - /// 操作时间戳 - pub timestamp: String, - /// 操作系统平台 - pub platform: String, - /// 备份路径(可选) - pub backup_path: Option, -} - -/// 机器码操作类型 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[allow(dead_code)] -pub enum MachineIdOperation { - /// 获取当前机器码 - Get, - /// 设置新机器码 - Set, - /// 生成随机机器码 - Generate, - /// 备份机器码 - Backup, - /// 恢复机器码 - Restore, - /// 重置为原始机器码 - Reset, -} - -/// 机器码验证结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct MachineIdValidation { - /// 是否有效 - pub is_valid: bool, - /// 检测到的格式 - pub detected_format: MachineIdFormat, - /// 验证错误信息 - pub error_message: Option, - /// 格式化后的机器码(如果有效) - pub formatted_id: Option, -} - -impl MachineIdFormat { - /// 从字符串检测机器码格式 - pub fn detect(machine_id: &str) -> Self { - let cleaned = machine_id.replace("-", "").replace(" ", "").to_lowercase(); - - // 检查UUID格式:8-4-4-4-12个十六进制字符 - if machine_id.contains("-") && machine_id.len() == 36 { - let parts: Vec<&str> = machine_id.split('-').collect(); - if parts.len() == 5 - && parts[0].len() == 8 - && parts[1].len() == 4 - && parts[2].len() == 4 - && parts[3].len() == 4 - && parts[4].len() == 12 - && cleaned.chars().all(|c| c.is_ascii_hexdigit()) - { - return MachineIdFormat::Uuid; - } - } - - // 检查32位十六进制格式 - if cleaned.len() == 32 && cleaned.chars().all(|c| c.is_ascii_hexdigit()) { - return MachineIdFormat::Hex32; - } - - MachineIdFormat::Unknown - } - - /// 格式化机器码为标准格式 - pub fn format_machine_id(&self, machine_id: &str) -> Result { - let cleaned = machine_id.replace("-", "").replace(" ", "").to_lowercase(); - - match self { - MachineIdFormat::Uuid => { - if cleaned.len() != 32 { - return Err("UUID format requires 32 hex characters".to_string()); - } - if !cleaned.chars().all(|c| c.is_ascii_hexdigit()) { - return Err("UUID format requires hex characters only".to_string()); - } - Ok(format!( - "{}-{}-{}-{}-{}", - &cleaned[0..8], - &cleaned[8..12], - &cleaned[12..16], - &cleaned[16..20], - &cleaned[20..32] - )) - } - MachineIdFormat::Hex32 => { - if cleaned.len() != 32 { - return Err("Hex32 format requires 32 hex characters".to_string()); - } - if !cleaned.chars().all(|c| c.is_ascii_hexdigit()) { - return Err("Hex32 format requires hex characters only".to_string()); - } - Ok(cleaned) - } - MachineIdFormat::Unknown => Err("Cannot format unknown machine ID format".to_string()), - } - } -} - -impl std::fmt::Display for MachineIdFormat { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - match self { - MachineIdFormat::Uuid => write!(f, "uuid"), - MachineIdFormat::Hex32 => write!(f, "hex32"), - MachineIdFormat::Unknown => write!(f, "unknown"), - } - } -} - -impl std::fmt::Display for MachineIdOperation { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - match self { - MachineIdOperation::Get => write!(f, "Get"), - MachineIdOperation::Set => write!(f, "Set"), - MachineIdOperation::Generate => write!(f, "Generate"), - MachineIdOperation::Backup => write!(f, "Backup"), - MachineIdOperation::Restore => write!(f, "Restore"), - MachineIdOperation::Reset => write!(f, "Reset"), - } - } -} diff --git a/src-tauri/src/models/mcp_model.rs b/src-tauri/src/models/mcp_model.rs deleted file mode 100644 index fe257cc13..000000000 --- a/src-tauri/src/models/mcp_model.rs +++ /dev/null @@ -1,181 +0,0 @@ -use serde::{Deserialize, Serialize}; -use serde_json::Value; -use std::collections::HashMap; - -/// MCP 服务器配置(类型化) -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct McpServerConfigTyped { - /// 启动命令 - pub command: String, - /// 命令参数 - #[serde(default)] - pub args: Vec, - /// 环境变量 - #[serde(default)] - pub env: HashMap, - /// 工作目录 - #[serde(skip_serializing_if = "Option::is_none")] - pub cwd: Option, - /// 超时时间(秒) - #[serde(default = "default_timeout")] - pub timeout: u64, -} - -fn default_timeout() -> u64 { - 30 -} - -impl Default for McpServerConfigTyped { - fn default() -> Self { - Self { - command: String::new(), - args: Vec::new(), - env: HashMap::new(), - cwd: None, - timeout: 30, - } - } -} - -/// 配置验证错误 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ConfigValidationError { - pub field: String, - pub message: String, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct McpServer { - pub id: String, - pub name: String, - pub server_config: Value, - #[serde(skip_serializing_if = "Option::is_none")] - pub description: Option, - #[serde(default)] - pub enabled_proxycast: bool, - #[serde(default)] - pub enabled_claude: bool, - #[serde(default)] - pub enabled_codex: bool, - #[serde(default)] - pub enabled_gemini: bool, - #[serde(skip_serializing_if = "Option::is_none")] - pub created_at: Option, -} - -impl McpServer { - #[allow(dead_code)] - pub fn new(id: String, name: String, server_config: Value) -> Self { - Self { - id, - name, - server_config, - description: None, - enabled_proxycast: false, - enabled_claude: false, - enabled_codex: false, - enabled_gemini: false, - created_at: Some(chrono::Utc::now().timestamp()), - } - } - - /// 解析 server_config 为类型化配置 - /// - /// 将 JSON Value 解析为 McpServerConfigTyped 结构。 - /// 如果解析失败,返回默认配置并尝试提取基本字段。 - pub fn parse_config(&self) -> McpServerConfigTyped { - serde_json::from_value(self.server_config.clone()).unwrap_or_else(|_| { - // 尝试手动提取字段 - McpServerConfigTyped { - command: self - .server_config - .get("command") - .and_then(|v| v.as_str()) - .unwrap_or("") - .to_string(), - args: self - .server_config - .get("args") - .and_then(|v| v.as_array()) - .map(|arr| { - arr.iter() - .filter_map(|v| v.as_str().map(|s| s.to_string())) - .collect() - }) - .unwrap_or_default(), - env: self - .server_config - .get("env") - .and_then(|v| v.as_object()) - .map(|obj| { - obj.iter() - .filter_map(|(k, v)| v.as_str().map(|s| (k.clone(), s.to_string()))) - .collect() - }) - .unwrap_or_default(), - cwd: self - .server_config - .get("cwd") - .and_then(|v| v.as_str()) - .map(|s| s.to_string()), - timeout: self - .server_config - .get("timeout") - .and_then(|v| v.as_u64()) - .unwrap_or(30), - } - }) - } - - /// 验证服务器配置 - /// - /// 检查配置是否有效,返回验证错误列表。 - /// 空列表表示配置有效。 - pub fn validate_config(&self) -> Vec { - let mut errors = Vec::new(); - let config = self.parse_config(); - - // 验证 command 不为空 - if config.command.trim().is_empty() { - errors.push(ConfigValidationError { - field: "command".to_string(), - message: "启动命令不能为空".to_string(), - }); - } - - // 验证 name 不为空 - if self.name.trim().is_empty() { - errors.push(ConfigValidationError { - field: "name".to_string(), - message: "服务器名称不能为空".to_string(), - }); - } - - // 验证 name 不包含特殊字符(用于工具名称前缀) - if !self - .name - .chars() - .all(|c| c.is_alphanumeric() || c == '-' || c == '_') - { - errors.push(ConfigValidationError { - field: "name".to_string(), - message: "服务器名称只能包含字母、数字、连字符和下划线".to_string(), - }); - } - - // 验证 timeout 在合理范围内 - if config.timeout == 0 || config.timeout > 300 { - errors.push(ConfigValidationError { - field: "timeout".to_string(), - message: "超时时间必须在 1-300 秒之间".to_string(), - }); - } - - errors - } - - /// 检查配置是否有效 - pub fn is_valid(&self) -> bool { - self.validate_config().is_empty() - } -} diff --git a/src-tauri/src/models/mod.rs b/src-tauri/src/models/mod.rs index 3f76fae09..f2c1fcfee 100644 --- a/src-tauri/src/models/mod.rs +++ b/src-tauri/src/models/mod.rs @@ -1,30 +1,45 @@ -pub mod anthropic; -pub mod app_type; -pub mod codewhisperer; -pub mod kiro_fingerprint; -pub mod machine_id; -pub mod mcp_model; -pub mod model_registry; -pub mod openai; -pub mod project_model; -pub mod prompt_model; -pub mod provider_model; -pub mod provider_pool_model; -pub mod route_model; -pub mod skill_model; +//! 数据模型模块 +//! +//! 从 proxycast-core crate 重新导出所有模型类型。 +//! 仅保留依赖主 crate 业务模块的类型在本地定义。 +// 从 core crate 重新导出所有模型 +pub use proxycast_core::models::anthropic; #[allow(unused_imports)] -pub use anthropic::*; -pub use app_type::AppType; +pub use proxycast_core::models::app_type; #[allow(unused_imports)] -pub use codewhisperer::*; -pub use mcp_model::McpServer; +pub use proxycast_core::models::codewhisperer; +pub use proxycast_core::models::kiro_fingerprint; +pub use proxycast_core::models::machine_id; +pub use proxycast_core::models::mcp_model; +pub use proxycast_core::models::model_registry; +pub use proxycast_core::models::openai; #[allow(unused_imports)] -pub use openai::*; +pub use proxycast_core::models::prompt_model; +#[allow(unused_imports)] +pub use proxycast_core::models::provider_model; +pub use proxycast_core::models::provider_pool_model; +pub use proxycast_core::models::route_model; +pub use proxycast_core::models::skill_model; + +// ProjectContext 及相关类型依赖 workspace::Workspace,保留在主 crate +pub mod project_model; #[allow(unused_imports)] pub use project_model::*; -pub use prompt_model::Prompt; -pub use provider_model::Provider; + +// 重新导出常用类型(保持向后兼容) #[allow(unused_imports)] -pub use provider_pool_model::*; -pub use skill_model::{Skill, SkillMetadata, SkillRepo, SkillState, SkillStates}; +pub use proxycast_core::models::anthropic::*; +pub use proxycast_core::models::app_type::AppType; +#[allow(unused_imports)] +pub use proxycast_core::models::codewhisperer::*; +pub use proxycast_core::models::mcp_model::McpServer; +#[allow(unused_imports)] +pub use proxycast_core::models::openai::*; +pub use proxycast_core::models::prompt_model::Prompt; +pub use proxycast_core::models::provider_model::Provider; +#[allow(unused_imports)] +pub use proxycast_core::models::provider_pool_model::*; +pub use proxycast_core::models::skill_model::{ + Skill, SkillMetadata, SkillRepo, SkillState, SkillStates, +}; diff --git a/src-tauri/src/models/model_registry.rs b/src-tauri/src/models/model_registry.rs deleted file mode 100644 index e2418904e..000000000 --- a/src-tauri/src/models/model_registry.rs +++ /dev/null @@ -1,665 +0,0 @@ -//! 模型注册表数据结构 -//! -//! 借鉴 opencode 的模型管理方式,定义增强的模型元数据结构 - -use serde::{Deserialize, Serialize}; - -/// 模型能力 -#[derive(Debug, Clone, Serialize, Deserialize, Default)] -pub struct ModelCapabilities { - /// 是否支持视觉输入 - pub vision: bool, - /// 是否支持工具调用 - pub tools: bool, - /// 是否支持流式输出 - pub streaming: bool, - /// 是否支持 JSON 模式 - pub json_mode: bool, - /// 是否支持函数调用 - pub function_calling: bool, - /// 是否支持推理/思考 - pub reasoning: bool, -} - -/// 模型定价 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ModelPricing { - /// 输入价格(每百万 token) - pub input_per_million: Option, - /// 输出价格(每百万 token) - pub output_per_million: Option, - /// 缓存读取价格(每百万 token) - pub cache_read_per_million: Option, - /// 缓存写入价格(每百万 token) - pub cache_write_per_million: Option, - /// 货币单位 ("USD" | "CNY") - pub currency: String, -} - -impl Default for ModelPricing { - fn default() -> Self { - Self { - input_per_million: None, - output_per_million: None, - cache_read_per_million: None, - cache_write_per_million: None, - currency: "USD".to_string(), - } - } -} - -/// 模型限制 -#[derive(Debug, Clone, Serialize, Deserialize, Default)] -pub struct ModelLimits { - /// 上下文长度 - pub context_length: Option, - /// 最大输出 token 数 - pub max_output_tokens: Option, - /// 每分钟请求数限制 - pub requests_per_minute: Option, - /// 每分钟 token 数限制 - pub tokens_per_minute: Option, -} - -/// 模型状态 -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] -#[serde(rename_all = "lowercase")] -pub enum ModelStatus { - /// 活跃可用 - Active, - /// 预览版 - Preview, - /// Alpha 测试 - Alpha, - /// Beta 测试 - Beta, - /// 已弃用 - Deprecated, - /// 旧版本 - Legacy, -} - -impl Default for ModelStatus { - fn default() -> Self { - Self::Active - } -} - -impl std::fmt::Display for ModelStatus { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - match self { - Self::Active => write!(f, "active"), - Self::Preview => write!(f, "preview"), - Self::Alpha => write!(f, "alpha"), - Self::Beta => write!(f, "beta"), - Self::Deprecated => write!(f, "deprecated"), - Self::Legacy => write!(f, "legacy"), - } - } -} - -impl std::str::FromStr for ModelStatus { - type Err = String; - - fn from_str(s: &str) -> Result { - match s.to_lowercase().as_str() { - "active" => Ok(Self::Active), - "preview" => Ok(Self::Preview), - "alpha" => Ok(Self::Alpha), - "beta" => Ok(Self::Beta), - "deprecated" => Ok(Self::Deprecated), - "legacy" => Ok(Self::Legacy), - _ => Err(format!("Unknown model status: {s}")), - } - } -} - -/// 模型服务等级 -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] -#[serde(rename_all = "lowercase")] -pub enum ModelTier { - /// 快速响应,适合简单任务 - Mini, - /// 均衡性能,适合大多数任务 - Pro, - /// 最强能力,适合复杂任务 - Max, -} - -impl Default for ModelTier { - fn default() -> Self { - Self::Pro - } -} - -impl std::fmt::Display for ModelTier { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - match self { - Self::Mini => write!(f, "mini"), - Self::Pro => write!(f, "pro"), - Self::Max => write!(f, "max"), - } - } -} - -impl std::str::FromStr for ModelTier { - type Err = String; - - fn from_str(s: &str) -> Result { - match s.to_lowercase().as_str() { - "mini" => Ok(Self::Mini), - "pro" => Ok(Self::Pro), - "max" => Ok(Self::Max), - _ => Err(format!("Unknown model tier: {s}")), - } - } -} - -/// 模型数据来源 -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] -#[serde(rename_all = "lowercase")] -pub enum ModelSource { - /// 从内嵌资源加载(构建时打包) - Embedded, - /// 从 models.dev API 获取(已弃用) - ModelsDev, - /// 本地硬编码(国内模型等) - Local, - /// 用户自定义 - Custom, - /// 从 Provider API 获取 - Api, -} - -impl Default for ModelSource { - fn default() -> Self { - Self::Local - } -} - -impl std::fmt::Display for ModelSource { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - match self { - Self::Embedded => write!(f, "embedded"), - Self::ModelsDev => write!(f, "models.dev"), - Self::Local => write!(f, "local"), - Self::Custom => write!(f, "custom"), - Self::Api => write!(f, "api"), - } - } -} - -impl std::str::FromStr for ModelSource { - type Err = String; - - fn from_str(s: &str) -> Result { - match s.to_lowercase().as_str() { - "embedded" => Ok(Self::Embedded), - "models.dev" | "modelsdev" => Ok(Self::ModelsDev), - "local" => Ok(Self::Local), - "custom" => Ok(Self::Custom), - "api" => Ok(Self::Api), - _ => Err(format!("Unknown model source: {s}")), - } - } -} - -/// 增强的模型元数据 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct EnhancedModelMetadata { - /// 模型 ID (如 "claude-sonnet-4-5-20250514") - pub id: String, - /// 显示名称 (如 "Claude Sonnet 4.5") - pub display_name: String, - /// Provider ID (如 "anthropic", "openai", "dashscope") - pub provider_id: String, - /// Provider 显示名称 - pub provider_name: String, - /// 模型家族 (如 "sonnet", "gpt-4", "qwen") - pub family: Option, - /// 服务等级 - pub tier: ModelTier, - /// 模型能力 - pub capabilities: ModelCapabilities, - /// 定价信息 - pub pricing: Option, - /// 限制信息 - pub limits: ModelLimits, - /// 模型状态 - pub status: ModelStatus, - /// 发布日期 - pub release_date: Option, - /// 是否为最新版本 - pub is_latest: bool, - /// 描述 - pub description: Option, - /// 数据来源 - pub source: ModelSource, - /// 创建时间 (Unix 时间戳) - pub created_at: i64, - /// 最后更新时间 (Unix 时间戳) - pub updated_at: i64, -} - -impl EnhancedModelMetadata { - /// 创建新的模型元数据 - pub fn new( - id: String, - display_name: String, - provider_id: String, - provider_name: String, - ) -> Self { - let now = chrono::Utc::now().timestamp(); - Self { - id, - display_name, - provider_id, - provider_name, - family: None, - tier: ModelTier::Pro, - capabilities: ModelCapabilities::default(), - pricing: None, - limits: ModelLimits::default(), - status: ModelStatus::Active, - release_date: None, - is_latest: false, - description: None, - source: ModelSource::Local, - created_at: now, - updated_at: now, - } - } - - /// 设置模型家族 - pub fn with_family(mut self, family: impl Into) -> Self { - self.family = Some(family.into()); - self - } - - /// 设置服务等级 - pub fn with_tier(mut self, tier: ModelTier) -> Self { - self.tier = tier; - self - } - - /// 设置模型能力 - pub fn with_capabilities(mut self, capabilities: ModelCapabilities) -> Self { - self.capabilities = capabilities; - self - } - - /// 设置定价信息 - pub fn with_pricing(mut self, pricing: ModelPricing) -> Self { - self.pricing = Some(pricing); - self - } - - /// 设置限制信息 - pub fn with_limits(mut self, limits: ModelLimits) -> Self { - self.limits = limits; - self - } - - /// 设置模型状态 - pub fn with_status(mut self, status: ModelStatus) -> Self { - self.status = status; - self - } - - /// 设置发布日期 - pub fn with_release_date(mut self, date: impl Into) -> Self { - self.release_date = Some(date.into()); - self - } - - /// 设置是否为最新版本 - pub fn with_is_latest(mut self, is_latest: bool) -> Self { - self.is_latest = is_latest; - self - } - - /// 设置描述 - pub fn with_description(mut self, description: impl Into) -> Self { - self.description = Some(description.into()); - self - } - - /// 设置数据来源 - pub fn with_source(mut self, source: ModelSource) -> Self { - self.source = source; - self - } -} - -/// 用户模型偏好 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct UserModelPreference { - /// 模型 ID - pub model_id: String, - /// 是否收藏 - pub is_favorite: bool, - /// 是否隐藏 - pub is_hidden: bool, - /// 自定义别名 - pub custom_alias: Option, - /// 使用次数 - pub usage_count: u32, - /// 最后使用时间 (Unix 时间戳) - pub last_used_at: Option, - /// 创建时间 (Unix 时间戳) - pub created_at: i64, - /// 更新时间 (Unix 时间戳) - pub updated_at: i64, -} - -impl UserModelPreference { - /// 创建新的用户偏好 - pub fn new(model_id: String) -> Self { - let now = chrono::Utc::now().timestamp(); - Self { - model_id, - is_favorite: false, - is_hidden: false, - custom_alias: None, - usage_count: 0, - last_used_at: None, - created_at: now, - updated_at: now, - } - } -} - -/// 模型同步状态 -#[derive(Debug, Clone, Serialize, Deserialize, Default)] -pub struct ModelSyncState { - /// 最后同步时间 (Unix 时间戳) - pub last_sync_at: Option, - /// 同步的模型数量 - pub model_count: u32, - /// 是否正在同步 - pub is_syncing: bool, - /// 最后同步错误 - pub last_error: Option, -} - -// ============================================================================ -// Provider Alias 相关类型(用于 Kiro、Antigravity 等中转服务) -// ============================================================================ - -/// 单个模型别名映射 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ModelAlias { - /// 实际模型 ID(如 "claude-sonnet-4-5-20250929") - pub actual: String, - /// 内部 API 名称(如 "CLAUDE_SONNET_4_5_20250929_V1_0") - pub internal_name: Option, - /// 原始 Provider(如 "anthropic") - pub provider: Option, - /// 描述 - pub description: Option, -} - -/// Provider 的别名配置 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ProviderAliasConfig { - /// Provider ID(如 "kiro"、"antigravity") - pub provider: String, - /// 描述 - pub description: Option, - /// 支持的模型列表 - #[serde(default)] - pub models: Vec, - /// 别名映射(模型名 -> 别名配置) - pub aliases: std::collections::HashMap, - /// 更新时间 - pub updated_at: Option, -} - -impl ProviderAliasConfig { - /// 检查是否支持指定模型 - pub fn supports_model(&self, model: &str) -> bool { - self.models.contains(&model.to_string()) || self.aliases.contains_key(model) - } - - /// 获取模型的内部名称 - pub fn get_internal_name(&self, model: &str) -> Option<&str> { - self.aliases - .get(model) - .and_then(|a| a.internal_name.as_deref()) - } - - /// 获取模型的实际 ID - pub fn get_actual_model(&self, model: &str) -> Option<&str> { - self.aliases.get(model).map(|a| a.actual.as_str()) - } -} - -/// models.dev API 响应中的 Provider 结构 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[allow(dead_code)] -pub struct ModelsDevProvider { - pub id: String, - pub name: String, - #[serde(default)] - pub api: Option, - #[serde(default)] - pub npm: Option, - #[serde(default)] - pub models: std::collections::HashMap, -} - -/// models.dev API 响应中的 Model 结构 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ModelsDevModel { - pub id: String, - pub name: String, - #[serde(default)] - pub family: Option, - #[serde(default)] - pub release_date: Option, - #[serde(default)] - pub attachment: bool, - #[serde(default)] - pub reasoning: bool, - #[serde(default)] - pub temperature: bool, - #[serde(default)] - pub tool_call: bool, - #[serde(default)] - pub cost: Option, - #[serde(default)] - pub limit: Option, - #[serde(default)] - pub modalities: Option, - #[serde(default)] - pub experimental: Option, - #[serde(default)] - pub status: Option, -} - -/// models.dev API 响应中的 Cost 结构 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ModelsDevCost { - #[serde(default)] - pub input: Option, - #[serde(default)] - pub output: Option, - #[serde(default)] - pub cache_read: Option, - #[serde(default)] - pub cache_write: Option, -} - -/// models.dev API 响应中的 Limit 结构 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ModelsDevLimit { - #[serde(default)] - pub context: Option, - #[serde(default)] - pub output: Option, -} - -/// models.dev API 响应中的 Modalities 结构 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ModelsDevModalities { - #[serde(default)] - pub input: Vec, - #[serde(default)] - pub output: Vec, -} - -impl ModelsDevModel { - /// 转换为 EnhancedModelMetadata - /// 预留:用于从 models.dev API 导入模型数据 - #[allow(dead_code)] - pub fn to_enhanced_metadata( - &self, - provider_id: &str, - provider_name: &str, - ) -> EnhancedModelMetadata { - let now = chrono::Utc::now().timestamp(); - - // 判断是否支持视觉 - let supports_vision = self - .modalities - .as_ref() - .map(|m| m.input.iter().any(|i| i == "image" || i == "video")) - .unwrap_or(false) - || self.attachment; - - // 根据模型名称推断服务等级 - let tier = infer_model_tier(&self.id, &self.name); - - // 解析状态 - let status = self - .status - .as_ref() - .and_then(|s| s.parse().ok()) - .unwrap_or(ModelStatus::Active); - - // 判断是否为最新版本 - let is_latest = self.id.contains("latest"); - - EnhancedModelMetadata { - id: self.id.clone(), - display_name: self.name.clone(), - provider_id: provider_id.to_string(), - provider_name: provider_name.to_string(), - family: self.family.clone(), - tier, - capabilities: ModelCapabilities { - vision: supports_vision, - tools: self.tool_call, - streaming: true, // 大多数模型都支持流式 - json_mode: true, // 大多数模型都支持 JSON 模式 - function_calling: self.tool_call, - reasoning: self.reasoning, - }, - pricing: self.cost.as_ref().map(|c| ModelPricing { - input_per_million: c.input, - output_per_million: c.output, - cache_read_per_million: c.cache_read, - cache_write_per_million: c.cache_write, - currency: "USD".to_string(), - }), - limits: ModelLimits { - context_length: self.limit.as_ref().and_then(|l| l.context), - max_output_tokens: self.limit.as_ref().and_then(|l| l.output), - requests_per_minute: None, - tokens_per_minute: None, - }, - status, - release_date: self.release_date.clone(), - is_latest, - description: None, - source: ModelSource::ModelsDev, - created_at: now, - updated_at: now, - } - } -} - -/// 根据模型 ID 和名称推断服务等级 -/// 用于 to_enhanced_metadata 和测试 -#[allow(dead_code)] -fn infer_model_tier(model_id: &str, model_name: &str) -> ModelTier { - let id_lower = model_id.to_lowercase(); - let name_lower = model_name.to_lowercase(); - - // Mini 等级模型(优先检查,因为 gpt-4o-mini 包含 gpt-4o) - let mini_patterns = [ - "mini", - "nano", - "lite", - "flash", - "haiku", - "gpt-4o-mini", - "gemini-flash", - "qwen-turbo", - "glm-4-flash", - ]; - for pattern in mini_patterns { - if id_lower.contains(pattern) || name_lower.contains(pattern) { - return ModelTier::Mini; - } - } - - // Max 等级模型 - let max_patterns = [ - "opus", - "gpt-4o", - "gpt-4-turbo", - "gemini-2.5-pro", - "gemini-ultra", - "claude-3-opus", - "qwen-max", - "glm-4-plus", - "deepseek-v3", - ]; - for pattern in max_patterns { - if id_lower.contains(pattern) || name_lower.contains(pattern) { - return ModelTier::Max; - } - } - - // 默认为 Pro 等级 - ModelTier::Pro -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_model_tier_inference() { - assert_eq!( - infer_model_tier("claude-opus-4-5-20250514", "Claude Opus 4.5"), - ModelTier::Max - ); - assert_eq!( - infer_model_tier("gpt-4o-mini", "GPT-4o Mini"), - ModelTier::Mini - ); - assert_eq!( - infer_model_tier("claude-sonnet-4-5", "Claude Sonnet 4.5"), - ModelTier::Pro - ); - assert_eq!( - infer_model_tier("gemini-2.5-flash", "Gemini 2.5 Flash"), - ModelTier::Mini - ); - } - - #[test] - fn test_model_status_parsing() { - assert_eq!( - "active".parse::().unwrap(), - ModelStatus::Active - ); - assert_eq!( - "deprecated".parse::().unwrap(), - ModelStatus::Deprecated - ); - assert_eq!("beta".parse::().unwrap(), ModelStatus::Beta); - } -} diff --git a/src-tauri/src/models/openai.rs b/src-tauri/src/models/openai.rs deleted file mode 100644 index 2119cb667..000000000 --- a/src-tauri/src/models/openai.rs +++ /dev/null @@ -1,309 +0,0 @@ -//! OpenAI API 数据模型 -//! -//! 支持标准 OpenAI 格式以及扩展的工具类型(如 web_search)。 -//! -//! # 工具类型支持 -//! -//! - `function`: 标准函数调用工具 -//! - `web_search`: 联网搜索工具(Claude Code 使用 `web_search_20250305`) -//! -//! # 更新日志 -//! -//! - 2025-12-27: 添加 web_search 工具支持,修复 Issue #49 -use serde::{Deserialize, Serialize}; - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ImageUrl { - pub url: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub detail: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(tag = "type")] -pub enum ContentPart { - #[serde(rename = "text")] - Text { text: String }, - #[serde(rename = "image_url")] - ImageUrl { image_url: ImageUrl }, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ToolCall { - pub id: String, - #[serde(rename = "type")] - pub call_type: String, - pub function: FunctionCall, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct FunctionCall { - pub name: String, - pub arguments: String, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(untagged)] -pub enum MessageContent { - Text(String), - Parts(Vec), -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ChatMessage { - pub role: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub content: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub tool_calls: Option>, - #[serde(skip_serializing_if = "Option::is_none")] - pub tool_call_id: Option, - /// 推理内容(DeepSeek R1 等模型的思维链内容) - /// DeepSeek Reasoner 在 Tool Calls 场景下要求此字段 - #[serde(skip_serializing_if = "Option::is_none")] - pub reasoning_content: Option, -} - -impl ChatMessage { - pub fn get_content_text(&self) -> String { - match &self.content { - Some(MessageContent::Text(s)) => s.clone(), - Some(MessageContent::Parts(parts)) => parts - .iter() - .filter_map(|p| { - if let ContentPart::Text { text } = p { - Some(text.clone()) - } else { - None - } - }) - .collect::>() - .join(""), - None => String::new(), - } - } - - /// 提取消息中的图片 URL 列表 - /// 返回 (format, base64_data) 元组列表 - pub fn get_images(&self) -> Vec<(String, String)> { - match &self.content { - Some(MessageContent::Parts(parts)) => parts - .iter() - .filter_map(|p| { - if let ContentPart::ImageUrl { image_url } = p { - // 解析 data URL: data:image/jpeg;base64,xxxxx - if image_url.url.starts_with("data:") { - let parts: Vec<&str> = image_url.url.splitn(2, ',').collect(); - if parts.len() == 2 { - // 提取 media_type: data:image/jpeg;base64 -> image/jpeg - let header = parts[0]; - let data = parts[1]; - let media_type = header - .strip_prefix("data:") - .and_then(|s| s.split(';').next()) - .unwrap_or("image/jpeg"); - // 提取格式: image/jpeg -> jpeg - let format = - media_type.split('/').nth(1).unwrap_or("jpeg").to_string(); - return Some((format, data.to_string())); - } - } - None - } else { - None - } - }) - .collect(), - _ => Vec::new(), - } - } -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct FunctionDef { - pub name: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub description: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub parameters: Option, -} - -/// 工具定义 -/// -/// 支持多种工具类型: -/// - `function`: 标准函数调用工具,包含 function 字段 -/// - `web_search`: 联网搜索工具,无需额外字段 -/// - `web_search_20250305`: Claude Code 的联网搜索工具类型 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(tag = "type")] -pub enum Tool { - /// 标准函数调用工具 - #[serde(rename = "function")] - Function { function: FunctionDef }, - /// 联网搜索工具(Codex/Kiro 格式) - #[serde(rename = "web_search")] - WebSearch, - /// 联网搜索工具(Claude Code 格式) - #[serde(rename = "web_search_20250305")] - WebSearch20250305, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ChatCompletionRequest { - pub model: String, - pub messages: Vec, - #[serde(skip_serializing_if = "Option::is_none")] - pub temperature: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub max_tokens: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub top_p: Option, - #[serde(default)] - pub stream: bool, - #[serde(skip_serializing_if = "Option::is_none")] - pub tools: Option>, - #[serde(skip_serializing_if = "Option::is_none")] - pub tool_choice: Option, - /// 思维链强度:none, low, medium, high - #[serde(skip_serializing_if = "Option::is_none")] - pub reasoning_effort: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct Usage { - pub prompt_tokens: u32, - pub completion_tokens: u32, - pub total_tokens: u32, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ResponseMessage { - pub role: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub content: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub tool_calls: Option>, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct Choice { - pub index: u32, - pub message: ResponseMessage, - pub finish_reason: String, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ChatCompletionResponse { - pub id: String, - pub object: String, - pub created: u64, - pub model: String, - pub choices: Vec, - pub usage: Usage, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct StreamDelta { - #[serde(skip_serializing_if = "Option::is_none")] - pub role: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub content: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub tool_calls: Option>, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct StreamChoice { - pub index: u32, - pub delta: StreamDelta, - #[serde(skip_serializing_if = "Option::is_none")] - pub finish_reason: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ChatCompletionChunk { - pub id: String, - pub object: String, - pub created: u64, - pub model: String, - pub choices: Vec, -} - -// ============================================================================ -// 图像生成 API 数据模型 -// ============================================================================ - -/// OpenAI 图像生成请求 -/// -/// 兼容 OpenAI Images API,支持通过 Antigravity 生成图像。 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ImageGenerationRequest { - /// 图像生成提示词 - pub prompt: String, - - /// 模型名称 (默认: gemini-3-pro-image-preview) - #[serde(default = "default_image_model")] - pub model: String, - - /// 生成图像数量 (默认: 1) - #[serde(default = "default_n")] - pub n: u32, - - /// 图像尺寸 (可选,Antigravity 可能忽略) - #[serde(skip_serializing_if = "Option::is_none")] - pub size: Option, - - /// 响应格式: "url" 或 "b64_json" (默认: "url") - #[serde(default = "default_response_format")] - pub response_format: String, - - /// 图像质量 (可选,Antigravity 可能忽略) - #[serde(skip_serializing_if = "Option::is_none")] - pub quality: Option, - - /// 图像风格 (可选,Antigravity 可能忽略) - #[serde(skip_serializing_if = "Option::is_none")] - pub style: Option, - - /// 用户标识 (可选) - #[serde(skip_serializing_if = "Option::is_none")] - pub user: Option, -} - -fn default_image_model() -> String { - "gemini-3-pro-image-preview".to_string() -} - -fn default_n() -> u32 { - 1 -} - -fn default_response_format() -> String { - "url".to_string() -} - -/// OpenAI 图像生成响应 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ImageGenerationResponse { - /// 创建时间戳 (Unix epoch seconds) - pub created: i64, - - /// 生成的图像数组 - pub data: Vec, -} - -/// 单个图像数据 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ImageData { - /// Base64 编码的图像数据 (当 response_format="b64_json") - #[serde(skip_serializing_if = "Option::is_none")] - pub b64_json: Option, - - /// 图像 URL (当 response_format="url",返回 data URL) - #[serde(skip_serializing_if = "Option::is_none")] - pub url: Option, - - /// 修订后的提示词 (如果 Antigravity 返回了文本) - #[serde(skip_serializing_if = "Option::is_none")] - pub revised_prompt: Option, -} diff --git a/src-tauri/src/models/prompt_model.rs b/src-tauri/src/models/prompt_model.rs deleted file mode 100644 index b77d529b9..000000000 --- a/src-tauri/src/models/prompt_model.rs +++ /dev/null @@ -1,35 +0,0 @@ -use serde::{Deserialize, Serialize}; - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct Prompt { - pub id: String, - pub app_type: String, - pub name: String, - pub content: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub description: Option, - /// Whether this prompt is currently enabled (synced to live file) - #[serde(default)] - pub enabled: bool, - #[serde(rename = "createdAt", skip_serializing_if = "Option::is_none")] - pub created_at: Option, - #[serde(rename = "updatedAt", skip_serializing_if = "Option::is_none")] - pub updated_at: Option, -} - -impl Prompt { - #[allow(dead_code)] - pub fn new(id: String, app_type: String, name: String, content: String) -> Self { - let now = chrono::Utc::now().timestamp(); - Self { - id, - app_type, - name, - content, - description: None, - enabled: false, - created_at: Some(now), - updated_at: Some(now), - } - } -} diff --git a/src-tauri/src/models/provider_model.rs b/src-tauri/src/models/provider_model.rs deleted file mode 100644 index 354cee392..000000000 --- a/src-tauri/src/models/provider_model.rs +++ /dev/null @@ -1,43 +0,0 @@ -use serde::{Deserialize, Serialize}; -use serde_json::Value; - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct Provider { - pub id: String, - pub app_type: String, - pub name: String, - pub settings_config: Value, - #[serde(skip_serializing_if = "Option::is_none")] - pub category: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub icon: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub icon_color: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub notes: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub created_at: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub sort_index: Option, - #[serde(default)] - pub is_current: bool, -} - -impl Provider { - #[allow(dead_code)] - pub fn new(id: String, app_type: String, name: String, settings_config: Value) -> Self { - Self { - id, - app_type, - name, - settings_config, - category: None, - icon: None, - icon_color: None, - notes: None, - created_at: Some(chrono::Utc::now().timestamp()), - sort_index: None, - is_current: false, - } - } -} diff --git a/src-tauri/src/models/provider_pool_model.rs b/src-tauri/src/models/provider_pool_model.rs deleted file mode 100644 index 4143702b7..000000000 --- a/src-tauri/src/models/provider_pool_model.rs +++ /dev/null @@ -1,1140 +0,0 @@ -//! Provider Pool 数据模型 -//! -//! 支持多凭证池管理,包括健康检测、负载均衡、故障转移等功能。 - -use chrono::{DateTime, Utc}; -use serde::{Deserialize, Serialize}; -use std::collections::HashMap; -use uuid::Uuid; - -use crate::providers::ANTIGRAVITY_MODELS_FALLBACK; - -/// 凭证来源枚举 -/// 用于标识凭证是如何添加到凭证池的 -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize, Default)] -#[serde(rename_all = "snake_case")] -pub enum CredentialSource { - /// 手动添加(通过 UI 添加) - #[default] - Manual, - /// 导入(从文件导入) - Imported, - /// 私有凭证(从高级设置迁移) - Private, -} - -/// Provider 类型别名 -/// -/// 为了向后兼容,PoolProviderType 是 crate::ProviderType 的类型别名。 -/// 所有 Provider 类型定义已统一到 lib.rs 中的 ProviderType。 -pub type PoolProviderType = crate::ProviderType; - -/// 凭证数据,根据 Provider 类型不同而不同 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(tag = "type", rename_all = "snake_case")] -pub enum CredentialData { - /// Kiro OAuth 凭证(文件路径) - KiroOAuth { creds_file_path: String }, - /// Gemini OAuth 凭证(文件路径) - GeminiOAuth { - creds_file_path: String, - project_id: Option, - }, - - /// Antigravity OAuth 凭证(文件路径)- Google 内部 Gemini 3 Pro - AntigravityOAuth { - creds_file_path: String, - project_id: Option, - }, - /// OpenAI API Key 凭证 - OpenAIKey { - api_key: String, - base_url: Option, - }, - /// Claude API Key 凭证 - ClaudeKey { - api_key: String, - base_url: Option, - }, - /// Vertex AI API Key 凭证 - VertexKey { - api_key: String, - base_url: Option, - /// Model alias mappings (alias -> upstream model name) - #[serde(default)] - model_aliases: std::collections::HashMap, - }, - /// Gemini API Key 凭证(多账号负载均衡) - GeminiApiKey { - api_key: String, - base_url: Option, - /// 排除的模型列表(支持通配符) - #[serde(default)] - excluded_models: Vec, - }, - /// Codex OAuth 凭证(OpenAI Codex) - CodexOAuth { - creds_file_path: String, - /// API Base URL(可选,默认使用凭证文件中的配置) - #[serde(default)] - api_base_url: Option, - }, - /// Claude OAuth 凭证(Anthropic OAuth) - ClaudeOAuth { creds_file_path: String }, - - /// Anthropic API Key 凭证(直接使用 Anthropic API) - AnthropicKey { - api_key: String, - base_url: Option, - }, -} - -impl CredentialData { - /// 获取凭证的显示名称(隐藏敏感信息) - pub fn display_name(&self) -> String { - match self { - CredentialData::KiroOAuth { creds_file_path } => { - format!("Kiro OAuth: {}", mask_path(creds_file_path)) - } - CredentialData::GeminiOAuth { - creds_file_path, .. - } => { - format!("Gemini OAuth: {}", mask_path(creds_file_path)) - } - - CredentialData::AntigravityOAuth { - creds_file_path, .. - } => { - format!("Antigravity OAuth: {}", mask_path(creds_file_path)) - } - CredentialData::OpenAIKey { api_key, .. } => { - format!("OpenAI: {}", mask_key(api_key)) - } - CredentialData::ClaudeKey { api_key, .. } => { - format!("Claude: {}", mask_key(api_key)) - } - CredentialData::VertexKey { api_key, .. } => { - format!("Vertex AI: {}", mask_key(api_key)) - } - CredentialData::GeminiApiKey { api_key, .. } => { - format!("Gemini API Key: {}", mask_key(api_key)) - } - CredentialData::CodexOAuth { - creds_file_path, .. - } => { - format!("Codex OAuth: {}", mask_path(creds_file_path)) - } - CredentialData::ClaudeOAuth { creds_file_path } => { - format!("Claude OAuth: {}", mask_path(creds_file_path)) - } - - CredentialData::AnthropicKey { api_key, .. } => { - format!("Anthropic: {}", mask_key(api_key)) - } - } - } - - /// 获取 Provider 类型 - pub fn provider_type(&self) -> PoolProviderType { - match self { - CredentialData::KiroOAuth { .. } => PoolProviderType::Kiro, - CredentialData::GeminiOAuth { .. } => PoolProviderType::Gemini, - - CredentialData::AntigravityOAuth { .. } => PoolProviderType::Antigravity, - CredentialData::OpenAIKey { .. } => PoolProviderType::OpenAI, - CredentialData::ClaudeKey { .. } => PoolProviderType::Claude, - CredentialData::VertexKey { .. } => PoolProviderType::Vertex, - CredentialData::GeminiApiKey { .. } => PoolProviderType::GeminiApiKey, - CredentialData::CodexOAuth { .. } => PoolProviderType::Codex, - CredentialData::ClaudeOAuth { .. } => PoolProviderType::ClaudeOAuth, - - CredentialData::AnthropicKey { .. } => PoolProviderType::Anthropic, - } - } -} - -/// 通配符模式匹配 -/// -/// 支持的通配符模式: -/// - 精确匹配: `claude-sonnet-4-5` -/// - 前缀匹配: `claude-*` -/// - 后缀匹配: `*-preview` -/// - 包含匹配: `*flash*` -pub fn pattern_matches(pattern: &str, model: &str) -> bool { - if !pattern.contains('*') { - return pattern == model; - } - - let parts: Vec<&str> = pattern.split('*').collect(); - - match parts.as_slice() { - [prefix, ""] => model.starts_with(prefix), - ["", suffix] => model.ends_with(suffix), - ["", middle, ""] => model.contains(middle), - [prefix, suffix] => model.starts_with(prefix) && model.ends_with(suffix), - _ => false, - } -} - -/// 单个凭证 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ProviderCredential { - /// 唯一标识符 - pub uuid: String, - /// Provider 类型 - pub provider_type: PoolProviderType, - /// 凭证数据 - pub credential: CredentialData, - /// 备注/名称 - pub name: Option, - /// 是否健康 - #[serde(default = "default_true")] - pub is_healthy: bool, - /// 是否禁用(手动禁用) - #[serde(default)] - pub is_disabled: bool, - /// 是否启用自动健康检查 - #[serde(default = "default_true")] - pub check_health: bool, - /// 自定义健康检查模型 - pub check_model_name: Option, - /// 不支持的模型列表(黑名单) - #[serde(default)] - pub not_supported_models: Vec, - /// 支持的模型列表(从 /v1/models 接口获取) - #[serde(default)] - pub supported_models: Vec, - /// 使用次数 - #[serde(default)] - pub usage_count: u64, - /// 错误次数 - #[serde(default)] - pub error_count: u32, - /// 最后使用时间 - pub last_used: Option>, - /// 最后错误时间 - pub last_error_time: Option>, - /// 最后错误消息 - pub last_error_message: Option, - /// 最后健康检查时间 - pub last_health_check_time: Option>, - /// 最后健康检查使用的模型 - pub last_health_check_model: Option, - /// 创建时间 - pub created_at: DateTime, - /// 更新时间 - pub updated_at: DateTime, - /// Token 缓存信息 - #[serde(default)] - pub cached_token: Option, - /// 凭证来源(手动添加/导入/私有) - #[serde(default)] - pub source: CredentialSource, - /// 代理 URL(可覆盖全局代理设置) - pub proxy_url: Option, -} - -fn default_true() -> bool { - true -} - -impl ProviderCredential { - /// 创建新凭证 - pub fn new(provider_type: PoolProviderType, credential: CredentialData) -> Self { - let now = Utc::now(); - Self { - uuid: Uuid::new_v4().to_string(), - provider_type, - credential, - name: None, - is_healthy: true, - is_disabled: false, - check_health: true, - check_model_name: None, - not_supported_models: Vec::new(), - supported_models: Vec::new(), - usage_count: 0, - error_count: 0, - last_used: None, - last_error_time: None, - last_error_message: None, - last_health_check_time: None, - last_health_check_model: None, - created_at: now, - updated_at: now, - cached_token: None, - source: CredentialSource::Manual, - proxy_url: None, - } - } - - /// 创建带来源的新凭证 - pub fn new_with_source( - provider_type: PoolProviderType, - credential: CredentialData, - source: CredentialSource, - ) -> Self { - let mut cred = Self::new(provider_type, credential); - cred.source = source; - cred - } - - /// 是否可用(健康且未禁用) - pub fn is_available(&self) -> bool { - self.is_healthy && !self.is_disabled - } - - /// 是否支持指定模型 - /// - /// 检查两个来源的排除列表: - /// 1. `not_supported_models` - 通用的不支持模型列表(精确匹配) - /// 2. `excluded_models` - 来自 CredentialData::GeminiApiKey 的排除列表(支持通配符) - /// 3. Antigravity 凭证只支持特定的模型列表 - pub fn supports_model(&self, model: &str) -> bool { - // 检查通用的不支持模型列表(精确匹配) - if self.not_supported_models.contains(&model.to_string()) { - return false; - } - - // 检查 GeminiApiKey 的 excluded_models(支持通配符) - if let CredentialData::GeminiApiKey { - excluded_models, .. - } = &self.credential - { - for pattern in excluded_models { - if pattern_matches(pattern, model) { - return false; - } - } - } - - // Antigravity 凭证只支持特定的模型 - // 使用 providers::antigravity 中定义的模型列表(fallback) - // 实际模型列表由 models/aliases/antigravity.json 定义 - if let CredentialData::AntigravityOAuth { .. } = &self.credential { - return ANTIGRAVITY_MODELS_FALLBACK.contains(&model); - } - - true - } - - /// 检查凭证是否适用于指定的客户端类型 - /// - /// 某些凭证可能有使用限制,例如 Claude Code 专用凭证只能用于 Claude Code 客户端 - pub fn is_compatible_with_client( - &self, - client_type: Option<&crate::server::client_detector::ClientType>, - ) -> bool { - // 检查是否是 Claude Code 专用凭证 - if let Some(error_msg) = &self.last_error_message { - if error_msg.contains("only authorized for use with Claude Code") { - // 这是 Claude Code 专用凭证,只能用于 Claude Code 客户端 - return matches!( - client_type, - Some(crate::server::client_detector::ClientType::ClaudeCode) - ); - } - } - - // 默认情况下,凭证适用于所有客户端 - true - } - - /// 标记为健康 - pub fn mark_healthy(&mut self, check_model: Option) { - self.is_healthy = true; - self.error_count = 0; - self.last_health_check_time = Some(Utc::now()); - self.last_health_check_model = check_model; - self.updated_at = Utc::now(); - } - - /// 标记为不健康 - pub fn mark_unhealthy(&mut self, error_message: Option) { - self.error_count += 1; - self.last_error_time = Some(Utc::now()); - self.last_error_message = error_message; - self.updated_at = Utc::now(); - // 错误次数达到阈值则标记为不健康 - if self.error_count >= 3 { - self.is_healthy = false; - } - } - - /// 记录使用 - pub fn record_usage(&mut self) { - self.usage_count += 1; - self.last_used = Some(Utc::now()); - self.updated_at = Utc::now(); - } - - /// 重置计数器 - pub fn reset_counters(&mut self) { - self.usage_count = 0; - self.error_count = 0; - self.is_healthy = true; - self.last_error_time = None; - self.last_error_message = None; - self.updated_at = Utc::now(); - } -} - -/// 凭证池统计信息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct PoolStats { - /// 总凭证数 - pub total_count: usize, - /// 健康凭证数 - pub healthy_count: usize, - /// 禁用凭证数 - pub disabled_count: usize, - /// 总使用次数 - pub total_usage: u64, - /// 总错误次数 - pub total_errors: u64, - /// 最后更新时间 - pub last_update: DateTime, -} - -impl PoolStats { - pub fn from_credentials(credentials: &[ProviderCredential]) -> Self { - Self { - total_count: credentials.len(), - healthy_count: credentials.iter().filter(|c| c.is_healthy).count(), - disabled_count: credentials.iter().filter(|c| c.is_disabled).count(), - total_usage: credentials.iter().map(|c| c.usage_count).sum(), - total_errors: credentials.iter().map(|c| c.error_count as u64).sum(), - last_update: Utc::now(), - } - } -} - -/// 健康检查结果 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct HealthCheckResult { - pub uuid: String, - pub success: bool, - pub model: Option, - pub message: Option, - pub duration_ms: u64, -} - -/// OAuth 凭证状态 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct OAuthStatus { - /// 是否有 access_token - pub has_access_token: bool, - /// 是否有 refresh_token - pub has_refresh_token: bool, - /// token 是否有效 - pub is_token_valid: bool, - /// 过期信息 - pub expiry_info: Option, - /// 凭证文件路径 - pub creds_path: String, -} - -/// Token 缓存状态(用于前端展示) -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct TokenCacheStatus { - /// 是否有缓存的 token - pub has_cached_token: bool, - /// Token 是否有效 - pub is_valid: bool, - /// Token 是否即将过期(5分钟内) - pub is_expiring_soon: bool, - /// 过期时间 - pub expiry_time: Option, - /// 最后刷新时间 - pub last_refresh: Option, - /// 连续刷新失败次数 - pub refresh_error_count: u32, - /// 最后刷新错误信息 - pub last_refresh_error: Option, -} - -/// Token 缓存信息 -#[derive(Debug, Clone, Default, Serialize, Deserialize)] -pub struct CachedTokenInfo { - /// 缓存的 access_token - pub access_token: Option, - /// 缓存的 refresh_token(刷新后可能变化) - pub refresh_token: Option, - /// Token 过期时间 - pub expiry_time: Option>, - /// 最后刷新时间 - pub last_refresh: Option>, - /// 连续刷新失败次数 - #[serde(default)] - pub refresh_error_count: u32, - /// 最后刷新错误信息 - pub last_refresh_error: Option, -} - -impl CachedTokenInfo { - /// 检查 token 是否有效(存在且未过期) - pub fn is_valid(&self) -> bool { - if self.access_token.is_none() { - return false; - } - match &self.expiry_time { - Some(expiry) => *expiry > Utc::now(), - None => true, // 没有过期时间,假设有效 - } - } - - /// 检查 token 是否即将过期(5分钟内) - pub fn is_expiring_soon(&self) -> bool { - self.is_expiring_within_minutes(5) - } - - /// 检查 token 是否在指定分钟数内过期 - /// - /// # 参数 - /// - `minutes`: 检查的时间阈值(分钟) - /// - /// # 返回 - /// - `true`: Token 将在指定分钟数内过期 - /// - `false`: Token 不会在指定分钟数内过期,或没有过期时间 - pub fn is_expiring_within_minutes(&self, minutes: i64) -> bool { - match &self.expiry_time { - Some(expiry) => { - let threshold = Utc::now() + chrono::Duration::minutes(minutes); - *expiry <= threshold - } - None => false, // 没有过期时间,假设不会过期 - } - } - - /// 检查 token 是否需要刷新(无效或即将过期) - pub fn needs_refresh(&self) -> bool { - !self.is_valid() || self.is_expiring_soon() - } -} - -/// 默认健康检查模型 -pub fn get_default_check_model(provider_type: PoolProviderType) -> &'static str { - match provider_type { - PoolProviderType::Kiro => "claude-haiku-4-5", - PoolProviderType::Gemini => "gemini-2.5-flash", - PoolProviderType::OpenAI => "gpt-3.5-turbo", - // 使用 claude-sonnet-4-5-20250929,兼容更多代理服务器 - PoolProviderType::Claude => "claude-sonnet-4-5-20250929", - PoolProviderType::ClaudeOAuth => "claude-sonnet-4-5-20250929", - // Anthropic 兼容格式使用相同的健康检查模型 - PoolProviderType::AnthropicCompatible => "claude-sonnet-4-5-20250929", - PoolProviderType::Antigravity => "gemini-3-pro-preview", - PoolProviderType::Vertex => "gemini-2.0-flash", - PoolProviderType::GeminiApiKey => "gemini-2.5-flash", - PoolProviderType::Codex => "gpt-4o-mini", - // API Key Provider 类型 - PoolProviderType::Anthropic => "claude-sonnet-4-5-20250929", - PoolProviderType::AzureOpenai => "gpt-4o-mini", - PoolProviderType::AwsBedrock => "claude-sonnet-4-5-20250929", - PoolProviderType::Ollama => "llama3.2", - } -} - -/// 凭证池前端展示数据 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CredentialDisplay { - pub uuid: String, - pub provider_type: String, - pub credential_type: String, - pub name: Option, - pub display_credential: String, - pub is_healthy: bool, - pub is_disabled: bool, - pub check_health: bool, - pub check_model_name: Option, - pub not_supported_models: Vec, - pub supported_models: Vec, - pub usage_count: u64, - pub error_count: u32, - pub last_used: Option, - pub last_error_time: Option, - pub last_error_message: Option, - pub last_health_check_time: Option, - pub last_health_check_model: Option, - pub oauth_status: Option, - pub token_cache_status: Option, - pub created_at: String, - pub updated_at: String, - /// 凭证来源(手动添加/导入/私有) - pub source: CredentialSource, - /// API Key 凭证的 base_url(仅用于 OpenAI/Claude API Key 类型) - pub base_url: Option, - /// API Key 凭证的完整 api_key(仅用于 OpenAI/Claude API Key 类型,用于编辑) - pub api_key: Option, - /// 凭证级代理 URL(可覆盖全局代理设置) - pub proxy_url: Option, -} - -/// 获取凭证类型字符串 -fn get_credential_type(cred: &CredentialData) -> String { - match cred { - CredentialData::KiroOAuth { .. } => "kiro_oauth".to_string(), - CredentialData::GeminiOAuth { .. } => "gemini_oauth".to_string(), - CredentialData::AntigravityOAuth { .. } => "antigravity_oauth".to_string(), - CredentialData::OpenAIKey { .. } => "openai_key".to_string(), - CredentialData::ClaudeKey { .. } => "claude_key".to_string(), - CredentialData::VertexKey { .. } => "vertex_key".to_string(), - CredentialData::GeminiApiKey { .. } => "gemini_api_key".to_string(), - CredentialData::CodexOAuth { .. } => "codex_oauth".to_string(), - CredentialData::ClaudeOAuth { .. } => "claude_oauth".to_string(), - CredentialData::AnthropicKey { .. } => "anthropic_key".to_string(), - } -} - -/// 获取 OAuth 凭证的文件路径 -pub fn get_oauth_creds_path(cred: &CredentialData) -> Option { - match cred { - CredentialData::KiroOAuth { creds_file_path } => Some(creds_file_path.clone()), - CredentialData::GeminiOAuth { - creds_file_path, .. - } => Some(creds_file_path.clone()), - CredentialData::AntigravityOAuth { - creds_file_path, .. - } => Some(creds_file_path.clone()), - CredentialData::CodexOAuth { - creds_file_path, .. - } => Some(creds_file_path.clone()), - CredentialData::ClaudeOAuth { creds_file_path } => Some(creds_file_path.clone()), - _ => None, - } -} - -/// 从 CredentialData 中提取 base_url(仅适用于 API Key 类型) -fn get_base_url(cred: &CredentialData) -> Option { - match cred { - CredentialData::OpenAIKey { base_url, .. } => base_url.clone(), - CredentialData::ClaudeKey { base_url, .. } => base_url.clone(), - CredentialData::AnthropicKey { base_url, .. } => base_url.clone(), - _ => None, - } -} - -/// 从 CredentialData 中提取 api_key(仅适用于 API Key 类型) -fn get_api_key(cred: &CredentialData) -> Option { - match cred { - CredentialData::OpenAIKey { api_key, .. } => Some(api_key.clone()), - CredentialData::ClaudeKey { api_key, .. } => Some(api_key.clone()), - CredentialData::AnthropicKey { api_key, .. } => Some(api_key.clone()), - _ => None, - } -} - -impl From<&ProviderCredential> for CredentialDisplay { - fn from(cred: &ProviderCredential) -> Self { - // 构建 token 缓存状态 - let token_cache_status = cred.cached_token.as_ref().map(|cache| TokenCacheStatus { - has_cached_token: cache.access_token.is_some(), - is_valid: cache.is_valid(), - is_expiring_soon: cache.is_expiring_soon(), - expiry_time: cache.expiry_time.map(|t| t.to_rfc3339()), - last_refresh: cache.last_refresh.map(|t| t.to_rfc3339()), - refresh_error_count: cache.refresh_error_count, - last_refresh_error: cache.last_refresh_error.clone(), - }); - - Self { - uuid: cred.uuid.clone(), - provider_type: cred.provider_type.to_string(), - credential_type: get_credential_type(&cred.credential), - name: cred.name.clone(), - display_credential: cred.credential.display_name(), - is_healthy: cred.is_healthy, - is_disabled: cred.is_disabled, - check_health: cred.check_health, - check_model_name: cred.check_model_name.clone(), - not_supported_models: cred.not_supported_models.clone(), - supported_models: cred.supported_models.clone(), - usage_count: cred.usage_count, - error_count: cred.error_count, - last_used: cred.last_used.map(|t| t.to_rfc3339()), - last_error_time: cred.last_error_time.map(|t| t.to_rfc3339()), - last_error_message: cred.last_error_message.clone(), - last_health_check_time: cred.last_health_check_time.map(|t| t.to_rfc3339()), - last_health_check_model: cred.last_health_check_model.clone(), - oauth_status: None, // 需要单独调用获取 - token_cache_status, - created_at: cred.created_at.to_rfc3339(), - updated_at: cred.updated_at.to_rfc3339(), - source: cred.source, - base_url: get_base_url(&cred.credential), - api_key: get_api_key(&cred.credential), - proxy_url: cred.proxy_url.clone(), - } - } -} - -/// Provider 池概览(按类型分组的统计) -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct ProviderPoolOverview { - pub provider_type: String, - pub stats: PoolStats, - pub credentials: Vec, -} - -// 辅助函数:隐藏路径中的用户名 -fn mask_path(path: &str) -> String { - if let Some(home) = dirs::home_dir() { - let home_str = home.to_string_lossy(); - path.replace(&*home_str, "~") - } else { - path.to_string() - } -} - -// 辅助函数:隐藏 API Key -fn mask_key(key: &str) -> String { - if key.len() <= 12 { - "****".to_string() - } else { - format!("{}...{}", &key[..6], &key[key.len() - 4..]) - } -} - -/// 添加凭证的请求结构 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct AddCredentialRequest { - pub provider_type: String, - pub credential: CredentialData, - pub name: Option, - pub check_health: Option, - pub check_model_name: Option, -} - -/// 更新凭证请求 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct UpdateCredentialRequest { - pub name: Option, - pub is_disabled: Option, - pub check_health: Option, - pub check_model_name: Option, - pub not_supported_models: Option>, - /// 新的凭证文件路径(仅适用于OAuth凭证,用于重新上传文件) - pub new_creds_file_path: Option, - /// OAuth相关:新的project_id(仅适用于Gemini) - pub new_project_id: Option, - /// API Key 相关:新的 base_url(仅适用于 API Key 凭证) - pub new_base_url: Option, - /// API Key 相关:新的 api_key(仅适用于 API Key 凭证) - pub new_api_key: Option, - /// 新的代理 URL(可覆盖全局代理设置) - pub new_proxy_url: Option, -} - -pub type ProviderPools = HashMap>; - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn test_pattern_matches_exact() { - assert!(pattern_matches("gemini-2.5-pro", "gemini-2.5-pro")); - assert!(!pattern_matches("gemini-2.5-pro", "gemini-2.5-flash")); - } - - #[test] - fn test_pattern_matches_prefix() { - assert!(pattern_matches("gemini-*", "gemini-2.5-pro")); - assert!(pattern_matches("gemini-*", "gemini-2.5-flash")); - assert!(!pattern_matches("gemini-*", "claude-sonnet")); - } - - #[test] - fn test_pattern_matches_suffix() { - assert!(pattern_matches("*-preview", "gemini-3-pro-preview")); - assert!(pattern_matches("*-preview", "claude-preview")); - assert!(!pattern_matches("*-preview", "gemini-2.5-pro")); - } - - #[test] - fn test_pattern_matches_contains() { - assert!(pattern_matches("*flash*", "gemini-2.5-flash")); - assert!(pattern_matches("*flash*", "gemini-2.5-flash-lite")); - assert!(!pattern_matches("*flash*", "gemini-2.5-pro")); - } - - #[test] - fn test_pattern_matches_prefix_and_suffix() { - assert!(pattern_matches("gemini-*-pro", "gemini-2.5-pro")); - assert!(pattern_matches("gemini-*-pro", "gemini-3-pro")); - assert!(!pattern_matches("gemini-*-pro", "gemini-2.5-flash")); - } - - #[test] - fn test_supports_model_not_supported_models() { - let cred = ProviderCredential { - uuid: "test-uuid".to_string(), - provider_type: PoolProviderType::Kiro, - credential: CredentialData::KiroOAuth { - creds_file_path: "/path/to/creds".to_string(), - }, - name: None, - is_healthy: true, - is_disabled: false, - check_health: true, - check_model_name: None, - not_supported_models: vec!["claude-opus".to_string()], - supported_models: vec![], - usage_count: 0, - error_count: 0, - last_used: None, - last_error_time: None, - last_error_message: None, - last_health_check_time: None, - last_health_check_model: None, - created_at: Utc::now(), - updated_at: Utc::now(), - cached_token: None, - source: CredentialSource::Manual, - proxy_url: None, - }; - - assert!(!cred.supports_model("claude-opus")); - assert!(cred.supports_model("claude-sonnet")); - } - - #[test] - fn test_supports_model_gemini_api_key_excluded_models_exact() { - let cred = ProviderCredential { - uuid: "test-uuid".to_string(), - provider_type: PoolProviderType::GeminiApiKey, - credential: CredentialData::GeminiApiKey { - api_key: "test-key".to_string(), - base_url: None, - excluded_models: vec!["gemini-2.5-pro".to_string()], - }, - name: None, - is_healthy: true, - is_disabled: false, - check_health: true, - check_model_name: None, - not_supported_models: vec![], - supported_models: vec![], - usage_count: 0, - error_count: 0, - last_used: None, - last_error_time: None, - last_error_message: None, - last_health_check_time: None, - last_health_check_model: None, - created_at: Utc::now(), - updated_at: Utc::now(), - cached_token: None, - source: CredentialSource::Manual, - proxy_url: None, - }; - - // Exact match exclusion - assert!(!cred.supports_model("gemini-2.5-pro")); - // Not excluded - assert!(cred.supports_model("gemini-2.5-flash")); - } - - #[test] - fn test_supports_model_gemini_api_key_excluded_models_wildcard() { - let cred = ProviderCredential { - uuid: "test-uuid".to_string(), - provider_type: PoolProviderType::GeminiApiKey, - credential: CredentialData::GeminiApiKey { - api_key: "test-key".to_string(), - base_url: None, - excluded_models: vec!["gemini-2.5-*".to_string(), "*-preview".to_string()], - }, - name: None, - is_healthy: true, - is_disabled: false, - check_health: true, - check_model_name: None, - not_supported_models: vec![], - supported_models: vec![], - usage_count: 0, - error_count: 0, - last_used: None, - last_error_time: None, - last_error_message: None, - last_health_check_time: None, - last_health_check_model: None, - created_at: Utc::now(), - updated_at: Utc::now(), - cached_token: None, - source: CredentialSource::Manual, - proxy_url: None, - }; - - // Prefix wildcard exclusion - assert!(!cred.supports_model("gemini-2.5-pro")); - assert!(!cred.supports_model("gemini-2.5-flash")); - // Suffix wildcard exclusion - assert!(!cred.supports_model("gemini-3-pro-preview")); - // Not excluded - assert!(cred.supports_model("gemini-2.0-flash")); - assert!(cred.supports_model("gemini-3-pro")); - } - - #[test] - fn test_supports_model_gemini_api_key_excluded_models_contains() { - let cred = ProviderCredential { - uuid: "test-uuid".to_string(), - provider_type: PoolProviderType::GeminiApiKey, - credential: CredentialData::GeminiApiKey { - api_key: "test-key".to_string(), - base_url: None, - excluded_models: vec!["*flash*".to_string()], - }, - name: None, - is_healthy: true, - is_disabled: false, - check_health: true, - check_model_name: None, - not_supported_models: vec![], - supported_models: vec![], - usage_count: 0, - error_count: 0, - last_used: None, - last_error_time: None, - last_error_message: None, - last_health_check_time: None, - last_health_check_model: None, - created_at: Utc::now(), - updated_at: Utc::now(), - cached_token: None, - source: CredentialSource::Manual, - proxy_url: None, - }; - - // Contains wildcard exclusion - assert!(!cred.supports_model("gemini-2.5-flash")); - assert!(!cred.supports_model("gemini-2.5-flash-lite")); - // Not excluded - assert!(cred.supports_model("gemini-2.5-pro")); - } - - #[test] - fn test_supports_model_combined_exclusions() { - let cred = ProviderCredential { - uuid: "test-uuid".to_string(), - provider_type: PoolProviderType::GeminiApiKey, - credential: CredentialData::GeminiApiKey { - api_key: "test-key".to_string(), - base_url: None, - excluded_models: vec!["gemini-2.5-*".to_string()], - }, - name: None, - is_healthy: true, - is_disabled: false, - check_health: true, - check_model_name: None, - not_supported_models: vec!["gemini-3-pro".to_string()], - supported_models: vec![], - usage_count: 0, - error_count: 0, - last_used: None, - last_error_time: None, - last_error_message: None, - last_health_check_time: None, - last_health_check_model: None, - created_at: Utc::now(), - updated_at: Utc::now(), - cached_token: None, - source: CredentialSource::Manual, - proxy_url: None, - }; - - // Excluded by not_supported_models (exact match) - assert!(!cred.supports_model("gemini-3-pro")); - // Excluded by excluded_models (wildcard) - assert!(!cred.supports_model("gemini-2.5-pro")); - assert!(!cred.supports_model("gemini-2.5-flash")); - // Not excluded - assert!(cred.supports_model("gemini-2.0-flash")); - } - - #[test] - fn test_supports_model_non_gemini_api_key_ignores_excluded_models() { - // For non-GeminiApiKey credentials, excluded_models in CredentialData is not checked - let cred = ProviderCredential { - uuid: "test-uuid".to_string(), - provider_type: PoolProviderType::Kiro, - credential: CredentialData::KiroOAuth { - creds_file_path: "/path/to/creds".to_string(), - }, - name: None, - is_healthy: true, - is_disabled: false, - check_health: true, - check_model_name: None, - not_supported_models: vec![], - supported_models: vec![], - usage_count: 0, - error_count: 0, - last_used: None, - last_error_time: None, - last_error_message: None, - last_health_check_time: None, - last_health_check_model: None, - created_at: Utc::now(), - updated_at: Utc::now(), - cached_token: None, - source: CredentialSource::Manual, - proxy_url: None, - }; - - // All models should be supported since not_supported_models is empty - assert!(cred.supports_model("claude-sonnet")); - assert!(cred.supports_model("claude-opus")); - } - - // ======================================================================== - // Property-Based Tests for Token Expiration Check - // ======================================================================== - - use proptest::prelude::*; - - /// 生成随机的过期时间偏移量(分钟) - fn expiry_offset_strategy() -> impl Strategy { - // 生成 -60 到 +120 分钟的偏移量 - -60i64..=120i64 - } - - /// 生成随机的检查阈值(分钟) - fn threshold_strategy() -> impl Strategy { - // 生成 1 到 30 分钟的阈值 - 1i64..=30i64 - } - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: kiro-streaming-fix, Property 7: Token 过期检查** - /// - /// *对于任意* 即将过期的 Token(指定分钟数内),`is_expiring_within_minutes` - /// 方法应该正确返回 true;对于不会在指定时间内过期的 Token,应该返回 false。 - /// - /// **Validates: Requirements 4.4** - #[test] - fn property_token_expiration_check( - offset_minutes in expiry_offset_strategy(), - threshold_minutes in threshold_strategy() - ) { - let now = Utc::now(); - let expiry_time = now + chrono::Duration::minutes(offset_minutes); - - let cache_info = CachedTokenInfo { - access_token: Some("test_token".to_string()), - refresh_token: None, - expiry_time: Some(expiry_time), - last_refresh: None, - refresh_error_count: 0, - last_refresh_error: None, - }; - - let is_expiring = cache_info.is_expiring_within_minutes(threshold_minutes); - - // Token 应该在 offset_minutes <= threshold_minutes 时被认为即将过期 - // 注意:由于时间精度问题,我们允许 1 秒的误差 - if offset_minutes <= threshold_minutes { - prop_assert!( - is_expiring, - "Token with {}min until expiry should be considered expiring within {}min", - offset_minutes, - threshold_minutes - ); - } else { - prop_assert!( - !is_expiring, - "Token with {}min until expiry should NOT be considered expiring within {}min", - offset_minutes, - threshold_minutes - ); - } - } - - /// **Feature: kiro-streaming-fix, Property 7.1: 无过期时间的 Token 不会被认为即将过期** - /// - /// *对于任意* 没有过期时间的 Token,`is_expiring_within_minutes` 应该返回 false。 - /// - /// **Validates: Requirements 4.4** - #[test] - fn property_no_expiry_time_not_expiring(threshold_minutes in threshold_strategy()) { - let cache_info = CachedTokenInfo { - access_token: Some("test_token".to_string()), - refresh_token: None, - expiry_time: None, // 没有过期时间 - last_refresh: None, - refresh_error_count: 0, - last_refresh_error: None, - }; - - let is_expiring = cache_info.is_expiring_within_minutes(threshold_minutes); - - prop_assert!( - !is_expiring, - "Token without expiry time should NOT be considered expiring within {}min", - threshold_minutes - ); - } - - /// **Feature: kiro-streaming-fix, Property 7.2: is_expiring_soon 等价于 is_expiring_within_minutes(5)** - /// - /// *对于任意* Token,`is_expiring_soon()` 应该等价于 `is_expiring_within_minutes(5)`。 - /// - /// **Validates: Requirements 4.4** - #[test] - fn property_expiring_soon_equivalence(offset_minutes in expiry_offset_strategy()) { - let now = Utc::now(); - let expiry_time = now + chrono::Duration::minutes(offset_minutes); - - let cache_info = CachedTokenInfo { - access_token: Some("test_token".to_string()), - refresh_token: None, - expiry_time: Some(expiry_time), - last_refresh: None, - refresh_error_count: 0, - last_refresh_error: None, - }; - - let is_expiring_soon = cache_info.is_expiring_soon(); - let is_expiring_within_5 = cache_info.is_expiring_within_minutes(5); - - prop_assert_eq!( - is_expiring_soon, - is_expiring_within_5, - "is_expiring_soon() should be equivalent to is_expiring_within_minutes(5)" - ); - } - - /// **Feature: kiro-streaming-fix, Property 7.3: 10分钟阈值检查** - /// - /// *对于任意* 在 10 分钟内过期的 Token,`is_expiring_within_minutes(10)` 应该返回 true。 - /// 这是流式请求前的预检查阈值。 - /// - /// **Validates: Requirements 4.4** - #[test] - fn property_streaming_threshold_check(offset_minutes in 0i64..=10i64) { - let now = Utc::now(); - let expiry_time = now + chrono::Duration::minutes(offset_minutes); - - let cache_info = CachedTokenInfo { - access_token: Some("test_token".to_string()), - refresh_token: None, - expiry_time: Some(expiry_time), - last_refresh: None, - refresh_error_count: 0, - last_refresh_error: None, - }; - - let is_expiring = cache_info.is_expiring_within_minutes(10); - - prop_assert!( - is_expiring, - "Token expiring in {}min should be considered expiring within 10min (streaming threshold)", - offset_minutes - ); - } - } -} diff --git a/src-tauri/src/models/route_model.rs b/src-tauri/src/models/route_model.rs deleted file mode 100644 index 5a8bfe24c..000000000 --- a/src-tauri/src/models/route_model.rs +++ /dev/null @@ -1,145 +0,0 @@ -//! 路由模型 -//! -//! 用于多供应商路由功能的数据结构定义。 - -use serde::{Deserialize, Serialize}; - -/// 单个路由信息 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct RouteInfo { - /// 路由选择器 (provider 类型或凭证名称) - pub selector: String, - /// Provider 类型 - pub provider_type: String, - /// 关联的凭证数量 - pub credential_count: usize, - /// 可用的端点列表 - pub endpoints: Vec, - /// 标签 (如 "突破限制", "官方API/三方") - pub tags: Vec, - /// 是否启用 - pub enabled: bool, -} - -/// 路由端点 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct RouteEndpoint { - /// 端点路径 - pub path: String, - /// 协议类型 - pub protocol: String, // "openai" 或 "claude" - /// 完整 URL - pub url: String, -} - -/// 路由列表响应 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct RouteListResponse { - /// 服务器基础 URL - pub base_url: String, - /// 默认 Provider - pub default_provider: String, - /// 所有可用路由 - pub routes: Vec, -} - -/// curl 示例 -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct CurlExample { - /// 描述 - pub description: String, - /// curl 命令 - pub command: String, -} - -impl RouteInfo { - /// 创建新的路由信息 - pub fn new(selector: String, provider_type: String) -> Self { - Self { - selector, - provider_type, - credential_count: 0, - endpoints: Vec::new(), - tags: Vec::new(), - enabled: true, - } - } - - /// 添加端点 - pub fn add_endpoint(&mut self, base_url: &str, protocol: &str) { - let path = match protocol { - "claude" => format!("/{}/v1/messages", self.selector), - "openai" => format!("/{}/v1/chat/completions", self.selector), - _ => return, - }; - let url = format!("{base_url}{path}"); - self.endpoints.push(RouteEndpoint { - path, - protocol: protocol.to_string(), - url, - }); - } - - /// 生成 curl 示例 - pub fn generate_curl_examples(&self, api_key: &str) -> Vec { - let mut examples = Vec::new(); - - for endpoint in &self.endpoints { - let (_model, body) = match endpoint.protocol.as_str() { - "claude" => { - let model = match self.provider_type.as_str() { - "kiro" | "claude" => "claude-sonnet-4-5", - "gemini" => "gemini-2.5-flash", - "qwen" => "qwen3-coder-plus", - "openai" => "gpt-4", - _ => "claude-sonnet-4-5", - }; - ( - model, - format!( - r#"{{ - "model": "{model}", - "max_tokens": 1024, - "messages": [{{"role": "user", "content": "Hello!"}}] -}}"# - ), - ) - } - "openai" => { - let model = match self.provider_type.as_str() { - "kiro" | "claude" => "claude-sonnet-4-5", - "gemini" => "gemini-2.5-flash", - "qwen" => "qwen3-coder-plus", - "openai" => "gpt-4", - _ => "claude-sonnet-4-5", - }; - ( - model, - format!( - r#"{{ - "model": "{model}", - "messages": [{{"role": "user", "content": "Hello!"}}] -}}"# - ), - ) - } - _ => continue, - }; - - let command = format!( - r#"curl {} \ - -H "Content-Type: application/json" \ - -H "Authorization: Bearer {}" \ - -d '{}'"#, - endpoint.url, api_key, body - ); - - examples.push(CurlExample { - description: format!("{} 协议", endpoint.protocol.to_uppercase()), - command, - }); - } - - examples - } -} diff --git a/src-tauri/src/models/skill_model.rs b/src-tauri/src/models/skill_model.rs deleted file mode 100644 index 09c69d9eb..000000000 --- a/src-tauri/src/models/skill_model.rs +++ /dev/null @@ -1,173 +0,0 @@ -use chrono::{DateTime, Utc}; -use serde::{Deserialize, Serialize}; -use std::collections::HashMap; - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct Skill { - pub key: String, - pub name: String, - pub description: String, - pub directory: String, - #[serde(rename = "readmeUrl", skip_serializing_if = "Option::is_none")] - pub readme_url: Option, - pub installed: bool, - #[serde(rename = "repoOwner", skip_serializing_if = "Option::is_none")] - pub repo_owner: Option, - #[serde(rename = "repoName", skip_serializing_if = "Option::is_none")] - pub repo_name: Option, - #[serde(rename = "repoBranch", skip_serializing_if = "Option::is_none")] - pub repo_branch: Option, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct SkillRepo { - pub owner: String, - pub name: String, - pub branch: String, - pub enabled: bool, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct SkillState { - pub installed: bool, - pub installed_at: DateTime, -} - -#[derive(Debug, Clone, Serialize, Deserialize)] -pub struct SkillMetadata { - pub name: Option, - pub description: Option, -} - -impl Default for SkillRepo { - fn default() -> Self { - Self { - owner: String::new(), - name: String::new(), - branch: "main".to_string(), - enabled: true, - } - } -} - -#[allow(dead_code)] -impl SkillRepo { - pub fn new(owner: String, name: String, branch: String) -> Self { - Self { - owner, - name, - branch, - enabled: true, - } - } - - pub fn github_url(&self) -> String { - format!("https://github.com/{}/{}", self.owner, self.name) - } - - pub fn zip_url(&self) -> String { - format!( - "https://github.com/{}/{}/archive/refs/heads/{}.zip", - self.owner, self.name, self.branch - ) - } -} - -pub fn get_default_skill_repos() -> Vec { - vec![ - // ProxyCast 官方仓库(排第一位) - SkillRepo { - owner: "proxycast".to_string(), - name: "skills".to_string(), - branch: "main".to_string(), - enabled: true, - }, - SkillRepo { - owner: "ComposioHQ".to_string(), - name: "awesome-claude-skills".to_string(), - branch: "main".to_string(), - enabled: true, - }, - SkillRepo { - owner: "anthropics".to_string(), - name: "skills".to_string(), - branch: "main".to_string(), - enabled: true, - }, - SkillRepo { - owner: "cexll".to_string(), - name: "myclaude".to_string(), - branch: "master".to_string(), - enabled: true, - }, - ] -} - -pub type SkillStates = HashMap; - -#[cfg(test)] -mod tests { - use super::*; - use proptest::prelude::*; - - /// Feature: skills-platform-mvp, Property 1: Default Repositories Include ProxyCast Official - /// Validates: Requirements 1.1, 1.2, 1.3 - #[test] - fn test_default_repos_include_proxycast_official() { - let repos = get_default_skill_repos(); - - // 验证列表非空 - assert!(!repos.is_empty(), "默认仓库列表不应为空"); - - // 验证第一个仓库是 ProxyCast 官方仓库 - let first_repo = &repos[0]; - assert_eq!( - first_repo.owner, "proxycast", - "第一个仓库的 owner 应为 proxycast" - ); - assert_eq!(first_repo.name, "skills", "第一个仓库的 name 应为 skills"); - assert_eq!(first_repo.branch, "main", "第一个仓库的 branch 应为 main"); - assert!(first_repo.enabled, "ProxyCast 官方仓库应默认启用"); - } - - // Property 1: Default Repositories Include ProxyCast Official (Property-Based Test) - // For any call to get_default_skill_repos(), the returned list SHALL contain - // a SkillRepo with owner="proxycast", name="skills", branch="main", and enabled=true, - // and this repo SHALL be the first item in the list. - // Validates: Requirements 1.1, 1.2, 1.3 - proptest! { - #[test] - fn prop_default_repos_proxycast_first(_seed in 0u64..1000) { - // 无论调用多少次,结果应该一致 - let repos = get_default_skill_repos(); - - // Property: 列表非空 - prop_assert!(!repos.is_empty()); - - // Property: 第一个仓库是 ProxyCast 官方仓库 - let first = &repos[0]; - prop_assert_eq!(&first.owner, "proxycast"); - prop_assert_eq!(&first.name, "skills"); - prop_assert_eq!(&first.branch, "main"); - prop_assert!(first.enabled); - } - } - - #[test] - fn test_proxycast_repo_exists_in_list() { - let repos = get_default_skill_repos(); - - // 验证 ProxyCast 仓库存在于列表中 - let proxycast_repo = repos - .iter() - .find(|r| r.owner == "proxycast" && r.name == "skills"); - assert!( - proxycast_repo.is_some(), - "ProxyCast 官方仓库应存在于默认列表中" - ); - - let repo = proxycast_repo.unwrap(); - assert_eq!(repo.branch, "main"); - assert!(repo.enabled); - } -} diff --git a/src-tauri/src/plugin/mod.rs b/src-tauri/src/plugin/mod.rs index a8900beb5..5d857dd77 100644 --- a/src-tauri/src/plugin/mod.rs +++ b/src-tauri/src/plugin/mod.rs @@ -1,38 +1,26 @@ //! 插件系统模块 //! -//! 提供插件扩展功能,支持: -//! - 插件加载和初始化 -//! - 请求前/响应后钩子 -//! - 插件隔离和错误处理 -//! - 插件配置管理 -//! - 二进制组件下载和管理 -//! - 声明式插件 UI 系统 -//! - 插件安装和卸载 +//! 核心逻辑从 proxycast-core 重新导出, +//! ui_events 依赖 Tauri 保留在主 crate -pub mod binary_downloader; -pub mod examples; -pub mod installer; -mod loader; -mod manager; -mod types; -pub mod ui_builder; -pub mod ui_events; -pub mod ui_trait; -pub mod ui_types; +// 从 core 重新导出所有插件类型和模块 +pub use proxycast_core::plugin::binary_downloader; +pub use proxycast_core::plugin::examples; +pub use proxycast_core::plugin::installer; +pub use proxycast_core::plugin::ui_builder; +pub use proxycast_core::plugin::ui_trait; +pub use proxycast_core::plugin::ui_types; -pub use binary_downloader::BinaryDownloader; -pub use loader::PluginLoader; -pub use manager::PluginManager; -pub use types::{ - BinaryComponentStatus, BinaryManifest, HookResult, PlatformBinaries, Plugin, PluginConfig, - PluginContext, PluginError, PluginInfo, PluginManifest, PluginState, PluginStatus, PluginType, -}; -pub use ui_events::{PluginUIEmitter, PluginUIEmitterState, PluginUIEventPayload}; -pub use ui_trait::{NoUI, PluginUI}; -pub use ui_types::{ +pub use proxycast_core::plugin::{ Action, BoundValue, ChildrenDef, ComponentDef, ComponentType, DataEntry, DataModelUpdate, SurfaceDefinition, SurfaceUpdate, UIMessage, UserAction, }; +pub use proxycast_core::plugin::{ + BinaryComponentStatus, BinaryDownloader, BinaryManifest, HookResult, NoUI, PlatformBinaries, + Plugin, PluginConfig, PluginContext, PluginError, PluginInfo, PluginLoader, PluginManager, + PluginManifest, PluginState, PluginStatus, PluginType, PluginUI, +}; -#[cfg(test)] -mod tests; +// Tauri 依赖的 UI 事件模块保留在主 crate +pub mod ui_events; +pub use ui_events::{PluginUIEmitter, PluginUIEmitterState, PluginUIEventPayload}; diff --git a/src-tauri/src/server/handlers/api.rs b/src-tauri/src/server/handlers/api.rs index e8c1d1fee..697678d06 100644 --- a/src-tauri/src/server/handlers/api.rs +++ b/src-tauri/src/server/handlers/api.rs @@ -27,11 +27,6 @@ use serde_json::json; use std::collections::HashMap; use crate::converter::anthropic_to_openai::convert_anthropic_to_openai; -use crate::flow_monitor::{ - ClientInfo, FlowError, FlowErrorType, FlowMetadata, FlowType, InterceptAction, InterceptType, - LLMFlow, LLMRequest, LLMResponse, Message, MessageContent, MessageRole, RequestParameters, - RoutingInfo, TokenUsage, -}; use crate::models::anthropic::AnthropicMessagesRequest; use crate::models::openai::ChatCompletionRequest; use crate::processor::RequestContext; @@ -46,324 +41,23 @@ use crate::ProviderType; use super::{call_provider_anthropic, call_provider_openai}; -// ============================================================================ -// Flow 捕获辅助函数 -// ============================================================================ - -/// 从 OpenAI 格式请求构建 LLMRequest -fn build_llm_request_from_openai( - request: &ChatCompletionRequest, - path: &str, - headers: &HeaderMap, -) -> LLMRequest { - // 转换消息 - let messages: Vec = request - .messages - .iter() - .map(|m| { - let role = match m.role.as_str() { - "system" => MessageRole::System, - "user" => MessageRole::User, - "assistant" => MessageRole::Assistant, - "tool" => MessageRole::Tool, - "function" => MessageRole::Function, - _ => MessageRole::User, - }; - - let content = match &m.content { - Some(c) => match c { - crate::models::openai::MessageContent::Text(s) => { - MessageContent::Text(s.clone()) - } - crate::models::openai::MessageContent::Parts(parts) => { - let flow_parts: Vec = parts - .iter() - .map(|p| match p { - crate::models::openai::ContentPart::Text { text } => { - crate::flow_monitor::ContentPart::Text { text: text.clone() } - } - crate::models::openai::ContentPart::ImageUrl { image_url } => { - crate::flow_monitor::ContentPart::ImageUrl { - image_url: crate::flow_monitor::models::ImageUrl { - url: image_url.url.clone(), - detail: image_url.detail.clone(), - }, - } - } - }) - .collect(); - MessageContent::MultiModal(flow_parts) - } - }, - None => MessageContent::Text(String::new()), - }; - - Message { - role, - content, - tool_calls: None, - tool_result: None, - name: None, - } - }) - .collect(); - - // 提取系统提示词 - let system_prompt = messages - .iter() - .find(|m| m.role == MessageRole::System) - .map(|m| m.content.get_all_text()); - - // 构建请求参数 - let parameters = RequestParameters { - temperature: request.temperature, - top_p: None, - max_tokens: request.max_tokens, - stop: None, - stream: request.stream, - extra: HashMap::new(), - }; - - // 提取请求头 - let mut header_map = HashMap::new(); - for (name, value) in headers.iter() { - if let Ok(v) = value.to_str() { - // 排除敏感头 - let name_lower = name.as_str().to_lowercase(); - if !name_lower.contains("authorization") && !name_lower.contains("api-key") { - header_map.insert(name.as_str().to_string(), v.to_string()); - } - } - } - - LLMRequest { - method: "POST".to_string(), - path: path.to_string(), - headers: header_map, - body: serde_json::to_value(request).unwrap_or_default(), - messages, - system_prompt, - tools: None, // TODO: 转换工具定义 - model: request.model.clone(), - original_model: None, - parameters, - size_bytes: 0, - timestamp: Utc::now(), - } -} - -/// 从 Anthropic 格式请求构建 LLMRequest -fn build_llm_request_from_anthropic( - request: &AnthropicMessagesRequest, - path: &str, - headers: &HeaderMap, -) -> LLMRequest { - // 转换消息 - let messages: Vec = request - .messages - .iter() - .map(|m| { - let role = match m.role.as_str() { - "user" => MessageRole::User, - "assistant" => MessageRole::Assistant, - _ => MessageRole::User, - }; - - let content = match &m.content { - serde_json::Value::String(s) => MessageContent::Text(s.clone()), - serde_json::Value::Array(arr) => { - let flow_parts: Vec = arr - .iter() - .filter_map(|p| { - let part_type = p.get("type").and_then(|t| t.as_str()).unwrap_or(""); - match part_type { - "text" => p.get("text").and_then(|t| t.as_str()).map(|text| { - crate::flow_monitor::ContentPart::Text { - text: text.to_string(), - } - }), - "image" => { - let source = p.get("source")?; - let media_type = source - .get("media_type") - .and_then(|m| m.as_str()) - .map(|s| s.to_string()); - let data = source - .get("data") - .and_then(|d| d.as_str()) - .map(|s| s.to_string()); - Some(crate::flow_monitor::ContentPart::Image { - media_type, - data, - url: None, - }) - } - _ => None, - } - }) - .collect(); - MessageContent::MultiModal(flow_parts) - } - _ => MessageContent::Text(String::new()), - }; - - Message { - role, - content, - tool_calls: None, - tool_result: None, - name: None, - } - }) - .collect(); - - // 提取系统提示词 - let system_prompt = request.system.as_ref().map(|s| match s { - serde_json::Value::String(text) => text.clone(), - serde_json::Value::Array(arr) => arr - .iter() - .filter_map(|p| p.get("text").and_then(|t| t.as_str())) - .collect::>() - .join("\n"), - _ => String::new(), - }); - - // 构建请求参数 - let parameters = RequestParameters { - temperature: request.temperature, - top_p: None, - max_tokens: request.max_tokens, - stop: None, - stream: request.stream, - extra: HashMap::new(), - }; - - // 提取请求头 - let mut header_map = HashMap::new(); - for (name, value) in headers.iter() { - if let Ok(v) = value.to_str() { - let name_lower = name.as_str().to_lowercase(); - if !name_lower.contains("authorization") && !name_lower.contains("api-key") { - header_map.insert(name.as_str().to_string(), v.to_string()); - } - } - } - - LLMRequest { - method: "POST".to_string(), - path: path.to_string(), - headers: header_map, - body: serde_json::to_value(request).unwrap_or_default(), - messages, - system_prompt, - tools: None, // TODO: 转换工具定义 - model: request.model.clone(), - original_model: None, - parameters, - size_bytes: 0, - timestamp: Utc::now(), - } -} - -/// 构建 FlowMetadata -fn build_flow_metadata( - provider: ProviderType, - provider_id: Option<&str>, - credential_id: Option<&str>, - credential_name: Option<&str>, - headers: &HeaderMap, - request_id: &str, -) -> FlowMetadata { - // 提取客户端信息 - let client_ip = headers - .get("x-forwarded-for") - .or_else(|| headers.get("x-real-ip")) - .and_then(|v| v.to_str().ok()) - .map(|s| s.split(',').next().unwrap_or("").trim().to_string()); - - let user_agent = headers - .get("user-agent") - .and_then(|v| v.to_str().ok()) - .map(|s| s.to_string()); - - FlowMetadata { - provider, - provider_id: provider_id.map(|s| s.to_string()), - credential_id: credential_id.map(|s| s.to_string()), - credential_name: credential_name.map(|s| s.to_string()), - retry_count: 0, - client_info: ClientInfo { - ip: client_ip, - user_agent, - request_id: Some(request_id.to_string()), - }, - routing_info: RoutingInfo::default(), - injected_params: None, - context_usage_percentage: None, - } -} - -/// 从响应构建 LLMResponse -fn build_llm_response(status_code: u16, content: &str, usage: Option<(u32, u32)>) -> LLMResponse { - let now = Utc::now(); - let (input_tokens, output_tokens) = usage.unwrap_or((0, 0)); - - LLMResponse { - status_code, - status_text: if status_code == 200 { "OK" } else { "Error" }.to_string(), - headers: HashMap::new(), - body: serde_json::Value::Null, - content: content.to_string(), - thinking: None, - tool_calls: Vec::new(), - usage: TokenUsage { - input_tokens, - output_tokens, - cache_read_tokens: None, - cache_write_tokens: None, - thinking_tokens: None, - total_tokens: input_tokens + output_tokens, - }, - stop_reason: None, - size_bytes: content.len(), - timestamp_start: now, - timestamp_end: now, - stream_info: None, - } -} - // ============================================================================ // Provider 选择辅助函数 // ============================================================================ /// 根据客户端类型和端点配置选择 Provider -/// -/// **Validates: Requirements 1.3, 1.4, 3.4** -/// -/// 优先级:端点 Provider 配置 > 默认 Provider -/// -/// # 参数 -/// - `headers`: HTTP 请求头,用于提取 User-Agent -/// - `state`: 应用状态,包含端点配置和默认 Provider -/// -/// # 返回 -/// 选择的 Provider 名称和检测到的客户端类型 async fn select_provider_for_client(headers: &HeaderMap, state: &AppState) -> (String, ClientType) { - // 从 User-Agent 检测客户端类型 let user_agent = headers .get("user-agent") .and_then(|v| v.to_str().ok()) .unwrap_or(""); let client_type = ClientType::from_user_agent(user_agent); - // 获取端点 Provider 配置 let endpoint_providers = state.endpoint_providers.read().await; let endpoint_provider = endpoint_providers.get_provider(client_type.config_key()); - // 获取默认 Provider let default_provider = state.default_provider.read().await.clone(); - // 选择 Provider:端点配置优先,否则使用默认 let selected_provider = match endpoint_provider { Some(provider) => provider.clone(), None => default_provider, @@ -372,175 +66,6 @@ async fn select_provider_for_client(headers: &HeaderMap, state: &AppState) -> (S (selected_provider, client_type) } -// ============================================================================ -// 拦截检查辅助函数 -// ============================================================================ - -/// 拦截检查结果 -pub enum InterceptCheckResult { - /// 继续处理(可能带有修改后的请求) - Continue(Option), - /// 请求被取消 - Cancelled, -} - -/// 检查是否需要拦截请求 -/// -/// **Validates: Requirements 2.1, 2.3, 2.5** -/// -/// 如果拦截器启用且请求匹配拦截规则,则拦截请求并等待用户操作。 -/// 返回 `InterceptCheckResult::Continue` 表示继续处理(可能带有修改后的请求), -/// 返回 `InterceptCheckResult::Cancelled` 表示请求被取消。 -async fn check_request_intercept( - state: &AppState, - flow_id: &str, - llm_request: &LLMRequest, - flow_metadata: &FlowMetadata, -) -> InterceptCheckResult { - // 创建临时 Flow 用于拦截检查 - let temp_flow = LLMFlow::new( - flow_id.to_string(), - FlowType::ChatCompletions, - llm_request.clone(), - flow_metadata.clone(), - ); - - // 检查是否需要拦截 - if !state - .flow_interceptor - .should_intercept(&temp_flow, &InterceptType::Request) - .await - { - return InterceptCheckResult::Continue(None); - } - - state - .logs - .write() - .await - .add("info", &format!("[INTERCEPT] 拦截请求: flow_id={flow_id}")); - - // 拦截请求 - let _intercepted = state - .flow_interceptor - .intercept_request(flow_id, llm_request.clone()) - .await; - - // 等待用户操作 - let action = state.flow_interceptor.wait_for_action(flow_id).await; - - match action { - InterceptAction::Continue(modified) => { - state.logs.write().await.add( - "info", - &format!( - "[INTERCEPT] 继续处理请求: flow_id={}, modified={}", - flow_id, - modified.is_some() - ), - ); - // 如果有修改,提取修改后的请求 - if let Some(crate::flow_monitor::ModifiedData::Request(req)) = modified { - InterceptCheckResult::Continue(Some(req)) - } else { - InterceptCheckResult::Continue(None) - } - } - InterceptAction::Cancel => { - state.logs.write().await.add( - "info", - &format!("[INTERCEPT] 请求被取消: flow_id={flow_id}"), - ); - InterceptCheckResult::Cancelled - } - InterceptAction::Timeout(timeout_action) => { - state.logs.write().await.add( - "warn", - &format!("[INTERCEPT] 请求超时: flow_id={flow_id}, action={timeout_action:?}"), - ); - match timeout_action { - crate::flow_monitor::TimeoutAction::Continue => { - InterceptCheckResult::Continue(None) - } - crate::flow_monitor::TimeoutAction::Cancel => InterceptCheckResult::Cancelled, - } - } - } -} - -/// 检查是否需要拦截响应 -/// -/// **Validates: Requirements 2.1, 2.5** -/// -/// 如果拦截器启用且响应匹配拦截规则,则拦截响应并等待用户操作。 -/// 返回修改后的响应(如果有)或 None。 -async fn check_response_intercept( - state: &AppState, - flow_id: &str, - llm_response: &LLMResponse, - llm_request: &LLMRequest, - flow_metadata: &FlowMetadata, -) -> Option { - // 创建临时 Flow 用于拦截检查 - let mut temp_flow = LLMFlow::new( - flow_id.to_string(), - FlowType::ChatCompletions, - llm_request.clone(), - flow_metadata.clone(), - ); - temp_flow.response = Some(llm_response.clone()); - - // 检查是否需要拦截 - if !state - .flow_interceptor - .should_intercept(&temp_flow, &InterceptType::Response) - .await - { - return None; - } - - state - .logs - .write() - .await - .add("info", &format!("[INTERCEPT] 拦截响应: flow_id={flow_id}")); - - // 拦截响应 - let _intercepted = state - .flow_interceptor - .intercept_response(flow_id, llm_response.clone()) - .await; - - // 等待用户操作 - let action = state.flow_interceptor.wait_for_action(flow_id).await; - - match action { - InterceptAction::Continue(modified) => { - state.logs.write().await.add( - "info", - &format!( - "[INTERCEPT] 继续处理响应: flow_id={}, modified={}", - flow_id, - modified.is_some() - ), - ); - // 如果有修改,提取修改后的响应 - if let Some(crate::flow_monitor::ModifiedData::Response(resp)) = modified { - Some(resp) - } else { - None - } - } - InterceptAction::Cancel | InterceptAction::Timeout(_) => { - state.logs.write().await.add( - "warn", - &format!("[INTERCEPT] 响应处理被取消或超时: flow_id={flow_id}"), - ); - None - } - } -} - // ============================================================================ // API Key 验证 // ============================================================================ @@ -966,7 +491,6 @@ pub async fn chat_completions( ); // 启动 Flow 捕获 - let llm_request = build_llm_request_from_openai(&request, "/v1/chat/completions", &headers); // 尝试将 selected_provider 解析为 ProviderType // 构建 Flow Metadata,同时保存 provider_type 和实际的 provider_id @@ -985,49 +509,11 @@ pub async fn chat_completions( } }); - let flow_metadata = build_flow_metadata( - provider_type, - provider_display_name, // 使用 Provider 显示名称(如 "DeepSeek") - Some(&cred.uuid), - cred.name.as_deref(), - &headers, - &ctx.request_id, - ); - let flow_id = state - .flow_monitor - .start_flow(llm_request.clone(), flow_metadata.clone()) - .await; - // 检查是否需要拦截请求 // **Validates: Requirements 2.1, 2.3, 2.5** - if let Some(ref fid) = flow_id { - match check_request_intercept(&state, fid, &llm_request, &flow_metadata).await { - InterceptCheckResult::Continue(modified_request) => { - // 如果有修改后的请求,更新请求 - if let Some(modified) = modified_request { - // 从修改后的 LLMRequest 更新 ChatCompletionRequest - if let Ok(updated) = serde_json::from_value(modified.body.clone()) { - request = updated; - } - } - } - InterceptCheckResult::Cancelled => { - // 请求被取消,标记 Flow 失败并返回错误 - let error = FlowError::new(FlowErrorType::Cancelled, "请求被用户取消"); - state.flow_monitor.fail_flow(fid, error).await; - return ( - StatusCode::BAD_REQUEST, - Json( - serde_json::json!({"error": {"message": "Request cancelled by user"}}), - ), - ) - .into_response(); - } - } - } eprintln!("[CHAT_COMPLETIONS] 调用 Provider: {}", cred.provider_type); - let response = call_provider_openai(&state, &cred, &request, flow_id.as_deref()).await; + let response = call_provider_openai(&state, &cred, &request, None).await; eprintln!( "[CHAT_COMPLETIONS] Provider 响应状态: {}", response.status() @@ -1045,203 +531,6 @@ pub async fn chat_completions( // 如果成功且需要 Flow 捕获,提取响应体内容和响应头 // 注意:非流式响应需要读取 body,所以必须在这里处理 - if is_success && flow_id.is_some() && !request.stream { - // 将 Response 转换为 bytes - let (parts, body) = response.into_parts(); - - // 提取响应头 - let mut response_headers = HashMap::new(); - for (name, value) in parts.headers.iter() { - if let Ok(v) = value.to_str() { - response_headers.insert(name.as_str().to_string(), v.to_string()); - } - } - - let body_bytes = match axum::body::to_bytes(body, usize::MAX).await { - Ok(bytes) => bytes, - Err(e) => { - eprintln!("[CHAT_COMPLETIONS] 读取响应体失败: {e}"); - // 如果读取失败,返回错误 - if let Some(fid) = flow_id { - let error = FlowError::new(FlowErrorType::Network, e.to_string()); - state.flow_monitor.fail_flow(&fid, error).await; - } - return ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": format!("Failed to read response body: {}", e)}})), - ) - .into_response(); - } - }; - - // 解析响应体 - let response_json: serde_json::Value = match serde_json::from_slice(&body_bytes) { - Ok(json) => json, - Err(e) => { - eprintln!("[CHAT_COMPLETIONS] 解析响应体失败: {e}"); - // 如果解析失败,仍然返回原始响应 - if let Some(fid) = flow_id { - let error = FlowError::new( - FlowErrorType::Other, - format!("Failed to parse response: {e}"), - ); - state.flow_monitor.fail_flow(&fid, error).await; - } - // 重新构建响应 - let response = Response::from_parts(parts, Body::from(body_bytes)); - return response; - } - }; - - // 提取内容和 token 使用量 - // 优先从 content 字段提取,如果为空则尝试从 tool_calls 提取 - let mut content = response_json["choices"][0]["message"]["content"] - .as_str() - .unwrap_or("") - .to_string(); - - // 如果 content 为空,检查是否有 tool_calls - if content.is_empty() { - if let Some(tool_calls) = - response_json["choices"][0]["message"]["tool_calls"].as_array() - { - if !tool_calls.is_empty() { - // 从第一个 tool_call 的 arguments 中提取内容 - if let Some(arguments) = tool_calls[0]["function"]["arguments"].as_str() { - content = arguments.to_string(); - eprintln!("[CHAT_COMPLETIONS] 从 tool_calls 中提取内容"); - } - } - } - } - - let input_tokens = response_json["usage"]["prompt_tokens"] - .as_u64() - .unwrap_or(0) as u32; - let output_tokens = response_json["usage"]["completion_tokens"] - .as_u64() - .unwrap_or(0) as u32; - - eprintln!("[CHAT_COMPLETIONS] 提取响应内容: content_len={}, input_tokens={}, output_tokens={}", - content.len(), input_tokens, output_tokens); - - // 记录 Token 使用量 - record_token_usage(&state, &ctx, Some(input_tokens), Some(output_tokens)); - - // 完成 Flow 捕获并检查响应拦截 - // **Validates: Requirements 2.1, 2.5** - if let Some(fid) = flow_id { - // 构建 LLMResponse,包含完整的响应体和响应头 - let mut llm_response = - build_llm_response(200, &content, Some((input_tokens, output_tokens))); - llm_response.body = response_json.clone(); - llm_response.headers = response_headers; // 设置响应头 - - // 检查是否需要拦截响应 - if let Some(modified_response) = check_response_intercept( - &state, - &fid, - &llm_response, - &llm_request, - &flow_metadata, - ) - .await - { - // 响应被修改,需要重新构建响应 - state - .logs - .write() - .await - .add("info", &format!("[INTERCEPT] 响应被修改: flow_id={fid}")); - - // 使用修改后的响应完成 Flow - state - .flow_monitor - .complete_flow(&fid, Some(modified_response.clone())) - .await; - - // 构建修改后的 HTTP 响应 - return ( - StatusCode::OK, - Json(serde_json::json!({ - "id": format!("chatcmpl-{}", uuid::Uuid::new_v4()), - "object": "chat.completion", - "created": std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_secs(), - "model": request.model, - "choices": [{ - "index": 0, - "message": { - "role": "assistant", - "content": modified_response.content - }, - "finish_reason": "stop" - }], - "usage": { - "prompt_tokens": modified_response.usage.input_tokens, - "completion_tokens": modified_response.usage.output_tokens, - "total_tokens": modified_response.usage.total_tokens - } - })), - ) - .into_response(); - } - - eprintln!("[FLOW_DEBUG] 准备完成 Flow: flow_id={}, content_len={}, input_tokens={}, output_tokens={}", - fid, llm_response.content.len(), llm_response.usage.input_tokens, llm_response.usage.output_tokens); - - state - .flow_monitor - .complete_flow(&fid, Some(llm_response)) - .await; - - eprintln!("[FLOW_DEBUG] Flow 已完成: flow_id={fid}"); - } - - // 重新构建响应返回给客户端 - let response = Response::from_parts(parts, Body::from(body_bytes)); - return response; - } else { - // 流式响应或没有 Flow 捕获,直接返回 - // 估算 Token 使用量(用于统计) - let estimated_input_tokens = request - .messages - .iter() - .map(|m| { - let content_len = match &m.content { - Some(c) => message_content_len(c), - None => 0, - }; - content_len / 4 - }) - .sum::() as u32; - let estimated_output_tokens = if is_success { 100u32 } else { 0u32 }; - - if is_success { - record_token_usage( - &state, - &ctx, - Some(estimated_input_tokens), - Some(estimated_output_tokens), - ); - } - - // 如果失败,标记 Flow 失败 - if let Some(fid) = flow_id { - if !is_success { - let error = FlowError::new( - FlowErrorType::from_status_code(status_code), - "Request failed", - ) - .with_status_code(status_code); - state.flow_monitor.fail_flow(&fid, error).await; - } - } - - return response; - } } // 回退到旧的单凭证模式(仅当选择的 Provider 是 Kiro 时) @@ -1273,50 +562,14 @@ pub async fn chat_completions( ); // 启动 Flow 捕获(legacy mode) - let llm_request = build_llm_request_from_openai(&request, "/v1/chat/completions", &headers); // 使用实际的 provider ID 构建 Flow Metadata let provider_type = selected_provider .parse::() .unwrap_or(ProviderType::OpenAI); - let flow_metadata = build_flow_metadata( - provider_type, - Some(&selected_provider), - None, - None, - &headers, - &ctx.request_id, - ); - let flow_id = state - .flow_monitor - .start_flow(llm_request.clone(), flow_metadata.clone()) - .await; - // 检查是否需要拦截请求(legacy mode) // **Validates: Requirements 2.1, 2.3, 2.5** - if let Some(ref fid) = flow_id { - match check_request_intercept(&state, fid, &llm_request, &flow_metadata).await { - InterceptCheckResult::Continue(modified_request) => { - // 如果有修改后的请求,更新请求 - if let Some(modified) = modified_request { - if let Ok(updated) = serde_json::from_value(modified.body.clone()) { - request = updated; - } - } - } - InterceptCheckResult::Cancelled => { - // 请求被取消,标记 Flow 失败并返回错误 - let error = FlowError::new(FlowErrorType::Cancelled, "请求被用户取消"); - state.flow_monitor.fail_flow(fid, error).await; - return ( - StatusCode::BAD_REQUEST, - Json(serde_json::json!({"error": {"message": "Request cancelled by user"}})), - ) - .into_response(); - } - } - } // 检查是否需要刷新 token(无 token 或即将过期) { @@ -1332,13 +585,6 @@ pub async fn chat_completions( .await .add("error", &format!("Token refresh failed: {e}")); // 标记 Flow 失败 - if let Some(fid) = &flow_id { - let error = FlowError::new( - FlowErrorType::Authentication, - format!("Token refresh failed: {e}"), - ); - state.flow_monitor.fail_flow(fid, error).await; - } return ( StatusCode::UNAUTHORIZED, Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})), @@ -1441,70 +687,6 @@ pub async fn chat_completions( ); // 完成 Flow 捕获并检查响应拦截 // **Validates: Requirements 2.1, 2.5** - if let Some(fid) = &flow_id { - let llm_response = build_llm_response( - 200, - &parsed.content, - Some((estimated_input_tokens, estimated_output_tokens)), - ); - - // 检查是否需要拦截响应 - if let Some(modified_response) = check_response_intercept( - &state, - fid, - &llm_response, - &llm_request, - &flow_metadata, - ) - .await - { - // 响应被修改,需要重新构建响应 - state - .logs - .write() - .await - .add("info", &format!("[INTERCEPT] 响应被修改: flow_id={fid}")); - - // 使用修改后的响应完成 Flow - state - .flow_monitor - .complete_flow(fid, Some(modified_response.clone())) - .await; - - // 构建修改后的响应 - let modified_message = serde_json::json!({ - "role": "assistant", - "content": modified_response.content - }); - - let modified_json_response = serde_json::json!({ - "id": format!("chatcmpl-{}", uuid::Uuid::new_v4()), - "object": "chat.completion", - "created": std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_secs(), - "model": request.model, - "choices": [{ - "index": 0, - "message": modified_message, - "finish_reason": "stop" - }], - "usage": { - "prompt_tokens": modified_response.usage.input_tokens, - "completion_tokens": modified_response.usage.output_tokens, - "total_tokens": modified_response.usage.total_tokens - } - }); - - return Json(modified_json_response).into_response(); - } - - state - .flow_monitor - .complete_flow(fid, Some(llm_response)) - .await; - } Json(response).into_response() } Err(e) => { @@ -1516,10 +698,6 @@ pub async fn chat_completions( Some(e.to_string()), ); // 标记 Flow 失败 - if let Some(fid) = &flow_id { - let error = FlowError::new(FlowErrorType::Network, e.to_string()); - state.flow_monitor.fail_flow(fid, error).await; - } ( StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": e.to_string()}})), @@ -1609,89 +787,10 @@ pub async fn chat_completions( }); // 完成 Flow 捕获并检查响应拦截(重试成功) // **Validates: Requirements 2.1, 2.5** - if let Some(fid) = &flow_id { - let (est_input, est_output) = - parsed.estimate_tokens(); - let llm_response = build_llm_response( - 200, - &parsed.content, - Some((est_input, est_output)), - ); - - // 检查是否需要拦截响应 - if let Some(modified_response) = - check_response_intercept( - &state, - fid, - &llm_response, - &llm_request, - &flow_metadata, - ) - .await - { - // 响应被修改,需要重新构建响应 - state.logs.write().await.add( - "info", - &format!( - "[INTERCEPT] 响应被修改: flow_id={fid}" - ), - ); - - // 使用修改后的响应完成 Flow - state - .flow_monitor - .complete_flow( - fid, - Some(modified_response.clone()), - ) - .await; - - // 构建修改后的响应 - let modified_message = serde_json::json!({ - "role": "assistant", - "content": modified_response.content - }); - - let modified_json_response = serde_json::json!({ - "id": format!("chatcmpl-{}", uuid::Uuid::new_v4()), - "object": "chat.completion", - "created": std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_secs(), - "model": request.model, - "choices": [{ - "index": 0, - "message": modified_message, - "finish_reason": "stop" - }], - "usage": { - "prompt_tokens": modified_response.usage.input_tokens, - "completion_tokens": modified_response.usage.output_tokens, - "total_tokens": modified_response.usage.total_tokens - } - }); - - return Json(modified_json_response) - .into_response(); - } - - state - .flow_monitor - .complete_flow(fid, Some(llm_response)) - .await; - } return Json(response).into_response(); } Err(e) => { // 标记 Flow 失败 - if let Some(fid) = &flow_id { - let error = FlowError::new( - FlowErrorType::Network, - e.to_string(), - ); - state.flow_monitor.fail_flow(fid, error).await; - } return ( StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": e.to_string()}})), @@ -1701,13 +800,6 @@ pub async fn chat_completions( } let body = retry_resp.text().await.unwrap_or_default(); // 标记 Flow 失败(重试失败) - if let Some(fid) = &flow_id { - let error = FlowError::new( - FlowErrorType::ServerError, - format!("Retry failed: {body}"), - ); - state.flow_monitor.fail_flow(fid, error).await; - } ( StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": format!("Retry failed: {}", body)}})), @@ -1715,11 +807,6 @@ pub async fn chat_completions( } Err(e) => { // 标记 Flow 失败 - if let Some(fid) = &flow_id { - let error = - FlowError::new(FlowErrorType::Network, e.to_string()); - state.flow_monitor.fail_flow(fid, error).await; - } ( StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": e.to_string()}})), @@ -1735,13 +822,6 @@ pub async fn chat_completions( .await .add("error", &format!("[AUTH] Token refresh failed: {e}")); // 标记 Flow 失败 - if let Some(fid) = &flow_id { - let error = FlowError::new( - FlowErrorType::Authentication, - format!("Token refresh failed: {e}"), - ); - state.flow_monitor.fail_flow(fid, error).await; - } ( StatusCode::UNAUTHORIZED, Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})), @@ -1756,12 +836,6 @@ pub async fn chat_completions( &format!("Upstream error {}: {}", status, safe_truncate(&body, 200)), ); // 标记 Flow 失败 - if let Some(fid) = &flow_id { - let error = - FlowError::new(FlowErrorType::from_status_code(status.as_u16()), &body) - .with_status_code(status.as_u16()); - state.flow_monitor.fail_flow(fid, error).await; - } ( StatusCode::from_u16(status.as_u16()).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), Json(serde_json::json!({"error": {"message": format!("Upstream error: {}", body)}})) @@ -1775,10 +849,6 @@ pub async fn chat_completions( .await .add("error", &format!("API call failed: {e}")); // 标记 Flow 失败 - if let Some(fid) = &flow_id { - let error = FlowError::new(FlowErrorType::Network, e.to_string()); - state.flow_monitor.fail_flow(fid, error).await; - } ( StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": e.to_string()}})), @@ -2040,7 +1110,6 @@ pub async fn anthropic_messages( ); // 启动 Flow 捕获 - let llm_request = build_llm_request_from_anthropic(&request, "/v1/messages", &headers); // 使用凭证的实际 provider_type(支持自定义 Provider) // 对于自定义 Provider ID,凭证的 provider_type 已通过数据库查询正确设置 @@ -2057,52 +1126,10 @@ pub async fn anthropic_messages( } }); - let flow_metadata = build_flow_metadata( - provider_type, - provider_display_name, // 使用 Provider 显示名称(如 "DeepSeek") - Some(&cred.uuid), - cred.name.as_deref(), - &headers, - &ctx.request_id, - ); - let flow_id = state - .flow_monitor - .start_flow(llm_request.clone(), flow_metadata.clone()) - .await; - // 检查是否需要拦截请求 // **Validates: Requirements 2.1, 2.3, 2.5** - if let Some(ref fid) = flow_id { - match check_request_intercept(&state, fid, &llm_request, &flow_metadata).await { - InterceptCheckResult::Continue(modified_request) => { - // 如果有修改后的请求,更新请求 - if let Some(modified) = modified_request { - // 从修改后的 LLMRequest 更新 AnthropicMessagesRequest - if let Ok(updated) = serde_json::from_value(modified.body.clone()) { - request = updated; - } - } - } - InterceptCheckResult::Cancelled => { - // 请求被取消,标记 Flow 失败并返回错误 - let error = FlowError::new(FlowErrorType::Cancelled, "请求被用户取消"); - state.flow_monitor.fail_flow(fid, error).await; - return ( - StatusCode::BAD_REQUEST, - Json(serde_json::json!({ - "type": "error", - "error": { - "type": "request_cancelled", - "message": "Request cancelled by user" - } - })), - ) - .into_response(); - } - } - } - let response = call_provider_anthropic(&state, &cred, &request, flow_id.as_deref()).await; + let response = call_provider_anthropic(&state, &cred, &request, None).await; // 记录请求统计 let is_success = response.status().is_success(); @@ -2143,73 +1170,6 @@ pub async fn anthropic_messages( // 完成 Flow 捕获并检查响应拦截 // **Validates: Requirements 2.1, 2.5** - if let Some(fid) = flow_id { - if is_success { - let llm_response = build_llm_response( - 200, - "", - Some((estimated_input_tokens, estimated_output_tokens)), - ); - - // 检查是否需要拦截响应 - if let Some(modified_response) = check_response_intercept( - &state, - &fid, - &llm_response, - &llm_request, - &flow_metadata, - ) - .await - { - // 响应被修改,需要重新构建响应 - state - .logs - .write() - .await - .add("info", &format!("[INTERCEPT] 响应被修改: flow_id={fid}")); - - // 使用修改后的响应完成 Flow - state - .flow_monitor - .complete_flow(&fid, Some(modified_response.clone())) - .await; - - // 构建修改后的 Anthropic 格式响应 - return ( - StatusCode::OK, - Json(serde_json::json!({ - "id": format!("msg_{}", uuid::Uuid::new_v4()), - "type": "message", - "role": "assistant", - "content": [{ - "type": "text", - "text": modified_response.content - }], - "model": request.model, - "stop_reason": "end_turn", - "stop_sequence": null, - "usage": { - "input_tokens": modified_response.usage.input_tokens, - "output_tokens": modified_response.usage.output_tokens - } - })), - ) - .into_response(); - } - - state - .flow_monitor - .complete_flow(&fid, Some(llm_response)) - .await; - } else { - let error = FlowError::new( - FlowErrorType::from_status_code(response.status().as_u16()), - "Request failed", - ) - .with_status_code(response.status().as_u16()); - state.flow_monitor.fail_flow(&fid, error).await; - } - } return response; } @@ -2243,56 +1203,14 @@ pub async fn anthropic_messages( ); // 启动 Flow 捕获(legacy mode) - let llm_request = build_llm_request_from_anthropic(&request, "/v1/messages", &headers); // 使用实际的 provider ID 构建 Flow Metadata let provider_type = selected_provider .parse::() .unwrap_or(ProviderType::OpenAI); - let flow_metadata = build_flow_metadata( - provider_type, - Some(&selected_provider), - None, - None, - &headers, - &ctx.request_id, - ); - let flow_id = state - .flow_monitor - .start_flow(llm_request.clone(), flow_metadata.clone()) - .await; - // 检查是否需要拦截请求(legacy mode) // **Validates: Requirements 2.1, 2.3, 2.5** - if let Some(ref fid) = flow_id { - match check_request_intercept(&state, fid, &llm_request, &flow_metadata).await { - InterceptCheckResult::Continue(modified_request) => { - // 如果有修改后的请求,更新请求 - if let Some(modified) = modified_request { - if let Ok(updated) = serde_json::from_value(modified.body.clone()) { - request = updated; - } - } - } - InterceptCheckResult::Cancelled => { - // 请求被取消,标记 Flow 失败并返回错误 - let error = FlowError::new(FlowErrorType::Cancelled, "请求被用户取消"); - state.flow_monitor.fail_flow(fid, error).await; - return ( - StatusCode::BAD_REQUEST, - Json(serde_json::json!({ - "type": "error", - "error": { - "type": "request_cancelled", - "message": "Request cancelled by user" - } - })), - ) - .into_response(); - } - } - } // 检查是否需要刷新 token(无 token 或即将过期) { @@ -2312,13 +1230,6 @@ pub async fn anthropic_messages( .await .add("error", &format!("[AUTH] Token refresh failed: {e}")); // 标记 Flow 失败 - if let Some(fid) = &flow_id { - let error = FlowError::new( - FlowErrorType::Authentication, - format!("Token refresh failed: {e}"), - ); - state.flow_monitor.fail_flow(fid, error).await; - } return ( StatusCode::UNAUTHORIZED, Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})), @@ -2415,129 +1326,11 @@ pub async fn anthropic_messages( if request.stream { // 完成 Flow 捕获并检查响应拦截(流式) // **Validates: Requirements 2.1, 2.5** - if let Some(fid) = &flow_id { - let (est_input, est_output) = parsed.estimate_tokens(); - let llm_response = build_llm_response( - 200, - &parsed.content, - Some((est_input, est_output)), - ); - - // 检查是否需要拦截响应 - if let Some(modified_response) = check_response_intercept( - &state, - fid, - &llm_response, - &llm_request, - &flow_metadata, - ) - .await - { - // 响应被修改,需要重新构建响应 - state.logs.write().await.add( - "info", - &format!("[INTERCEPT] 流式响应被修改: flow_id={fid}"), - ); - - // 使用修改后的响应完成 Flow - state - .flow_monitor - .complete_flow(fid, Some(modified_response.clone())) - .await; - - // 构建修改后的流式响应 - // 注意:这里简化处理,实际应该构建完整的流式响应 - return ( - StatusCode::OK, - Json(serde_json::json!({ - "id": format!("msg_{}", uuid::Uuid::new_v4()), - "type": "message", - "role": "assistant", - "content": [{ - "type": "text", - "text": modified_response.content - }], - "model": request.model, - "stop_reason": "end_turn", - "stop_sequence": null, - "usage": { - "input_tokens": modified_response.usage.input_tokens, - "output_tokens": modified_response.usage.output_tokens - } - })), - ) - .into_response(); - } - - state - .flow_monitor - .complete_flow(fid, Some(llm_response)) - .await; - } return build_anthropic_stream_response(&request.model, &parsed); } // 完成 Flow 捕获并检查响应拦截(非流式) // **Validates: Requirements 2.1, 2.5** - if let Some(fid) = &flow_id { - let (est_input, est_output) = parsed.estimate_tokens(); - let llm_response = build_llm_response( - 200, - &parsed.content, - Some((est_input, est_output)), - ); - - // 检查是否需要拦截响应 - if let Some(modified_response) = check_response_intercept( - &state, - fid, - &llm_response, - &llm_request, - &flow_metadata, - ) - .await - { - // 响应被修改,需要重新构建响应 - state - .logs - .write() - .await - .add("info", &format!("[INTERCEPT] 响应被修改: flow_id={fid}")); - - // 使用修改后的响应完成 Flow - state - .flow_monitor - .complete_flow(fid, Some(modified_response.clone())) - .await; - - // 构建修改后的 Anthropic 格式响应 - return ( - StatusCode::OK, - Json(serde_json::json!({ - "id": format!("msg_{}", uuid::Uuid::new_v4()), - "type": "message", - "role": "assistant", - "content": [{ - "type": "text", - "text": modified_response.content - }], - "model": request.model, - "stop_reason": "end_turn", - "stop_sequence": null, - "usage": { - "input_tokens": modified_response.usage.input_tokens, - "output_tokens": modified_response.usage.output_tokens - } - })), - ) - .into_response(); - } - - state - .flow_monitor - .complete_flow(fid, Some(llm_response)) - .await; - } // 非流式响应 build_anthropic_response(&request.model, &parsed) @@ -2549,10 +1342,6 @@ pub async fn anthropic_messages( .await .add("error", &format!("[ERROR] Response body read failed: {e}")); // 标记 Flow 失败 - if let Some(fid) = &flow_id { - let error = FlowError::new(FlowErrorType::Network, e.to_string()); - state.flow_monitor.fail_flow(fid, error).await; - } ( StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": e.to_string()}})), @@ -2610,92 +1399,6 @@ pub async fn anthropic_messages( ); // 完成 Flow 捕获并检查响应拦截(重试成功) // **Validates: Requirements 2.1, 2.5** - if let Some(fid) = &flow_id { - let (est_input, est_output) = - parsed.estimate_tokens(); - let llm_response = build_llm_response( - 200, - &parsed.content, - Some((est_input, est_output)), - ); - - // 检查是否需要拦截响应 - if let Some(modified_response) = - check_response_intercept( - &state, - fid, - &llm_response, - &llm_request, - &flow_metadata, - ) - .await - { - // 响应被修改,需要重新构建响应 - state.logs.write().await.add( - "info", - &format!("[INTERCEPT] 重试响应被修改: flow_id={fid}"), - ); - - // 使用修改后的响应完成 Flow - state - .flow_monitor - .complete_flow( - fid, - Some(modified_response.clone()), - ) - .await; - - // 构建修改后的响应 - if request.stream { - return ( - StatusCode::OK, - Json(serde_json::json!({ - "id": format!("msg_{}", uuid::Uuid::new_v4()), - "type": "message", - "role": "assistant", - "content": [{ - "type": "text", - "text": modified_response.content - }], - "model": request.model, - "stop_reason": "end_turn", - "stop_sequence": null, - "usage": { - "input_tokens": modified_response.usage.input_tokens, - "output_tokens": modified_response.usage.output_tokens - } - })), - ) - .into_response(); - } else { - return ( - StatusCode::OK, - Json(serde_json::json!({ - "id": format!("msg_{}", uuid::Uuid::new_v4()), - "type": "message", - "role": "assistant", - "content": [{ - "type": "text", - "text": modified_response.content - }], - "model": request.model, - "stop_reason": "end_turn", - "stop_sequence": null, - "usage": { - "input_tokens": modified_response.usage.input_tokens, - "output_tokens": modified_response.usage.output_tokens - } - })), - ) - .into_response(); - } - } - - state - .flow_monitor - .complete_flow(fid, Some(llm_response)) - .await; - } if request.stream { return build_anthropic_stream_response( &request.model, @@ -2713,13 +1416,6 @@ pub async fn anthropic_messages( &format!("[RETRY] Body read failed: {e}"), ); // 标记 Flow 失败 - if let Some(fid) = &flow_id { - let error = FlowError::new( - FlowErrorType::Network, - e.to_string(), - ); - state.flow_monitor.fail_flow(fid, error).await; - } return ( StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": e.to_string()}})), @@ -2741,13 +1437,6 @@ pub async fn anthropic_messages( ), ); // 标记 Flow 失败(重试失败) - if let Some(fid) = &flow_id { - let error = FlowError::new( - FlowErrorType::ServerError, - format!("Retry failed: {body}"), - ); - state.flow_monitor.fail_flow(fid, error).await; - } ( StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": format!("Retry failed: {}", body)}})), @@ -2761,11 +1450,6 @@ pub async fn anthropic_messages( .await .add("error", &format!("[RETRY] Request failed: {e}")); // 标记 Flow 失败 - if let Some(fid) = &flow_id { - let error = - FlowError::new(FlowErrorType::Network, e.to_string()); - state.flow_monitor.fail_flow(fid, error).await; - } ( StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": e.to_string()}})), @@ -2781,13 +1465,6 @@ pub async fn anthropic_messages( .await .add("error", &format!("[AUTH] Token refresh failed: {e}")); // 标记 Flow 失败 - if let Some(fid) = &flow_id { - let error = FlowError::new( - FlowErrorType::Authentication, - format!("Token refresh failed: {e}"), - ); - state.flow_monitor.fail_flow(fid, error).await; - } ( StatusCode::UNAUTHORIZED, Json(serde_json::json!({"error": {"message": format!("Token refresh failed: {e}")}})), @@ -2806,12 +1483,6 @@ pub async fn anthropic_messages( ), ); // 标记 Flow 失败 - if let Some(fid) = &flow_id { - let error = - FlowError::new(FlowErrorType::from_status_code(status.as_u16()), &body) - .with_status_code(status.as_u16()); - state.flow_monitor.fail_flow(fid, error).await; - } ( StatusCode::from_u16(status.as_u16()) .unwrap_or(StatusCode::INTERNAL_SERVER_ERROR), @@ -2835,10 +1506,6 @@ pub async fn anthropic_messages( &format!("[ERROR] Full error details: {error_details}"), ); // 标记 Flow 失败 - if let Some(fid) = &flow_id { - let error = FlowError::new(FlowErrorType::Network, e.to_string()); - state.flow_monitor.fail_flow(fid, error).await; - } ( StatusCode::INTERNAL_SERVER_ERROR, Json(serde_json::json!({"error": {"message": e.to_string()}})), diff --git a/src-tauri/src/server/handlers/provider_calls.rs b/src-tauri/src/server/handlers/provider_calls.rs index 89a6c607f..e8b3fb031 100644 --- a/src-tauri/src/server/handlers/provider_calls.rs +++ b/src-tauri/src/server/handlers/provider_calls.rs @@ -53,8 +53,6 @@ use crate::converter::anthropic_to_openai::convert_anthropic_to_openai; use crate::converter::openai_to_antigravity::{ convert_antigravity_to_openai_response, convert_openai_to_antigravity_with_context, }; -use crate::flow_monitor::models::{FlowError, FlowErrorType}; -use crate::flow_monitor::stream_rebuilder::StreamFormat; use crate::models::anthropic::AnthropicMessagesRequest; use crate::models::openai::ChatCompletionRequest; use crate::models::provider_pool_model::{CredentialData, ProviderCredential}; @@ -88,21 +86,6 @@ pub async fn call_provider_anthropic( request: &AnthropicMessagesRequest, flow_id: Option<&str>, ) -> Response { - // 如果是流式请求且有 flow_id,设置流式状态 - if request.stream { - if let Some(fid) = flow_id { - // 根据凭证类型确定流格式 - // 注意:Kiro 凭证虽然原始返回 AWS Event Stream,但 handle_kiro_stream 会将其转换为 Anthropic SSE 格式 - let format = match &credential.credential { - CredentialData::KiroOAuth { .. } => StreamFormat::Anthropic, // Kiro 流式响应被转换为 Anthropic SSE 格式 - CredentialData::ClaudeKey { .. } => StreamFormat::Anthropic, - CredentialData::AntigravityOAuth { .. } => StreamFormat::Gemini, - _ => StreamFormat::Unknown, - }; - state.flow_monitor.set_streaming(fid, format).await; - } - } - match &credential.credential { CredentialData::KiroOAuth { creds_file_path } => { // 如果是流式请求,使用真正的流式处理(需求 1.1, 6.1) @@ -2192,55 +2175,9 @@ pub async fn handle_streaming_response( ); // 获取 flow_id 的克隆用于回调 - let flow_id_for_callback = flow_id.map(|s| s.to_string()); - let flow_monitor = state.flow_monitor.clone(); - // 创建带回调的流式处理 - let managed_stream = if let Some(fid) = flow_id_for_callback { - // 使用带回调的流式处理,集成 Flow Monitor - let on_chunk = move |event: &str, _metrics: &crate::streaming::StreamMetrics| { - // 解析 SSE 事件并调用 process_chunk - // SSE 格式: "event: xxx\ndata: {...}\n\n" - let lines: Vec<&str> = event.lines().collect(); - let mut event_type: Option<&str> = None; - let mut data: Option<&str> = None; - - for line in lines { - if line.starts_with("event: ") { - event_type = Some(&line[7..]); - } else if line.starts_with("data: ") { - data = Some(&line[6..]); - } - } - - if let Some(d) = data { - // 使用 tokio::spawn 异步调用 process_chunk - let flow_monitor_clone = flow_monitor.clone(); - let fid_clone = fid.clone(); - let event_type_owned = event_type.map(|s| s.to_string()); - let data_owned = d.to_string(); - - tokio::spawn(async move { - flow_monitor_clone - .process_chunk(&fid_clone, event_type_owned.as_deref(), &data_owned) - .await; - }); - } - }; - - let stream = manager.handle_stream_with_callback(context, source_stream, on_chunk); - - // 转换为 Body 流 - let body_stream = stream.map(|result| -> Result { - match result { - Ok(event) => Ok(axum::body::Bytes::from(event)), - Err(e) => Ok(axum::body::Bytes::from(e.to_sse_error())), - } - }); - - Body::from_stream(body_stream) - } else { - // 没有 flow_id,使用普通流式处理 + // 创建流式处理 + let managed_stream = { let stream = manager.handle_stream(context, source_stream); let body_stream = stream.map(|result| -> Result { @@ -2318,45 +2255,12 @@ pub async fn handle_streaming_response_with_timeout( ); // 获取 flow_id 的克隆用于回调 - let flow_id_for_callback = flow_id.map(|s| s.to_string()); - let flow_monitor = state.flow_monitor.clone(); // 创建带超时的流式处理,使用 BoxStream 统一类型 - let timeout_stream: BoxStream<'static, Result> = - if let Some(fid) = flow_id_for_callback { - let on_chunk = move |event: &str, _metrics: &crate::streaming::StreamMetrics| { - let lines: Vec<&str> = event.lines().collect(); - let mut event_type: Option<&str> = None; - let mut data: Option<&str> = None; - - for line in lines { - if line.starts_with("event: ") { - event_type = Some(&line[7..]); - } else if line.starts_with("data: ") { - data = Some(&line[6..]); - } - } - - if let Some(d) = data { - let flow_monitor_clone = flow_monitor.clone(); - let fid_clone = fid.clone(); - let event_type_owned = event_type.map(|s| s.to_string()); - let data_owned = d.to_string(); - - tokio::spawn(async move { - flow_monitor_clone - .process_chunk(&fid_clone, event_type_owned.as_deref(), &data_owned) - .await; - }); - } - }; - - let stream = manager.handle_stream_with_callback(context, source_stream, on_chunk); - Box::pin(crate::streaming::with_timeout(stream, &config)) - } else { - let stream = manager.handle_stream(context, source_stream); - Box::pin(crate::streaming::with_timeout(stream, &config)) - }; + let timeout_stream: BoxStream<'static, Result> = { + let stream = manager.handle_stream(context, source_stream); + Box::pin(crate::streaming::with_timeout(stream, &config)) + }; // 转换为 Body 流 let body_stream = timeout_stream.map(|result| -> Result { @@ -2446,49 +2350,13 @@ pub async fn handle_streaming_with_disconnect_detection( ); // 获取 flow_id 的克隆 - let flow_id_for_callback = flow_id.map(|s| s.to_string()); let flow_id_for_cancel = flow_id.map(|s| s.to_string()); - let flow_monitor = state.flow_monitor.clone(); - let flow_monitor_for_cancel = state.flow_monitor.clone(); - // 创建带回调的流式处理 - // 使用 BoxStream 统一类型 + // 创建流式处理 let managed_stream: futures::stream::BoxStream< 'static, Result, - > = if let Some(fid) = flow_id_for_callback { - let on_chunk = move |event: &str, _metrics: &crate::streaming::StreamMetrics| { - let lines: Vec<&str> = event.lines().collect(); - let mut event_type: Option<&str> = None; - let mut data: Option<&str> = None; - - for line in lines { - if line.starts_with("event: ") { - event_type = Some(&line[7..]); - } else if line.starts_with("data: ") { - data = Some(&line[6..]); - } - } - - if let Some(d) = data { - let flow_monitor_clone = flow_monitor.clone(); - let fid_clone = fid.clone(); - let event_type_owned = event_type.map(|s| s.to_string()); - let data_owned = d.to_string(); - - tokio::spawn(async move { - flow_monitor_clone - .process_chunk(&fid_clone, event_type_owned.as_deref(), &data_owned) - .await; - }); - } - }; - - Box::pin(manager.handle_stream_with_callback(context, source_stream, on_chunk)) - } else { - // 没有 flow_id,使用普通流式处理 - Box::pin(manager.handle_stream(context, source_stream)) - }; + > = Box::pin(manager.handle_stream(context, source_stream)); // 如果有取消令牌,创建一个可取消的流 let body_stream = if let Some(token) = cancel_token { @@ -2502,7 +2370,6 @@ pub async fn handle_streaming_with_disconnect_detection( async move { token.cancelled().await; if let Some(fid) = flow_id { - flow_monitor_for_cancel.cancel_flow(&fid).await; tracing::info!("[STREAM] 客户端断开,已取消 Flow: {}", fid); } } @@ -2851,18 +2718,8 @@ pub async fn handle_kiro_stream( let config = PipelineConfig::kiro_to_anthropic(request.model.clone()); let pipeline = std::sync::Arc::new(tokio::sync::Mutex::new(StreamPipeline::new(config))); - // 获取 flow_id 的克隆用于回调 - let flow_id_owned = flow_id.map(|s| s.to_string()); - let flow_monitor = state.flow_monitor.clone(); - - // 创建转换流 let pipeline_clone = pipeline.clone(); - let flow_id_for_stream = flow_id_owned.clone(); - let flow_monitor_for_stream = flow_monitor.clone(); - let pipeline_for_finalize = pipeline.clone(); - let flow_id_for_finalize = flow_id_owned.clone(); - let flow_monitor_for_finalize = flow_monitor.clone(); let final_stream = async_stream::stream! { use futures::StreamExt; @@ -2872,89 +2729,27 @@ pub async fn handle_kiro_stream( while let Some(chunk_result) = stream_response.next().await { match chunk_result { Ok(bytes) => { - // 调试日志:记录接收到的字节数 tracing::info!( "[KIRO_STREAM] 收到 {} 字节数据", bytes.len() ); - // 使用 Pipeline 处理字节块 let sse_strings = { let mut pipeline_guard = pipeline_clone.lock().await; pipeline_guard.process_chunk(&bytes) }; - // 调试日志:记录生成的 SSE 事件数量 tracing::info!( "[KIRO_STREAM] 生成 {} 个 SSE 事件", sse_strings.len() ); for sse_str in sse_strings { - // 调用 FlowMonitor.process_chunk()(需求 3.2) - if let Some(ref fid) = flow_id_for_stream { - // 解析 SSE 事件类型和数据 - let lines: Vec<&str> = sse_str.lines().collect(); - let mut event_type: Option<&str> = None; - let mut data: Option<&str> = None; - - for line in &lines { - if line.starts_with("event: ") { - event_type = Some(&line[7..]); - } else if line.starts_with("data: ") { - data = Some(&line[6..]); - } - } - - if let Some(d) = data { - let flow_monitor_clone = flow_monitor_for_stream.clone(); - let fid_clone = fid.clone(); - let event_type_owned = event_type.map(|s| s.to_string()); - let data_owned = d.to_string(); - - tokio::spawn(async move { - flow_monitor_clone - .process_chunk( - &fid_clone, - event_type_owned.as_deref(), - &data_owned, - ) - .await; - }); - } - } - - // 立即 yield SSE 事件 yield Ok::(sse_str); } } Err(e) => { - // 需求 5.1, 5.3: 流式传输期间发生错误时,发出错误事件并以失败状态完成 flow tracing::error!("[KIRO_STREAM] 流式传输期间发生错误: {}", e); - - // 根据 StreamError 类型映射到 FlowErrorType - let flow_error_type = match &e { - StreamError::Network(_) => FlowErrorType::Network, - StreamError::Timeout => FlowErrorType::Timeout, - StreamError::ProviderError { status, .. } => { - FlowErrorType::from_status_code(*status) - } - StreamError::ParseError(_) => FlowErrorType::Other, - StreamError::ClientDisconnected => FlowErrorType::Cancelled, - StreamError::BufferOverflow => FlowErrorType::Other, - StreamError::Internal(_) => FlowErrorType::ServerError, - }; - - // 调用 FlowMonitor.fail_flow() 标记失败 - if let Some(ref fid) = flow_id_for_stream { - let flow_error = FlowError::new( - flow_error_type, - format!("流式传输错误: {e}"), - ); - flow_monitor_for_stream.fail_flow(fid, flow_error).await; - } - - // 发送 SSE 错误事件 yield Err(e); return; } @@ -2963,7 +2758,6 @@ pub async fn handle_kiro_stream( tracing::info!("[KIRO_STREAM] 流结束,生成 finalize 事件"); - // 流结束,使用 Pipeline 生成 finalize 事件 let final_events = { let mut pipeline_guard = pipeline_for_finalize.lock().await; pipeline_guard.finish() @@ -2972,34 +2766,6 @@ pub async fn handle_kiro_stream( tracing::info!("[KIRO_STREAM] finalize 生成 {} 个事件", final_events.len()); for sse_str in final_events { - // 调用 FlowMonitor.process_chunk() - if let Some(ref fid) = flow_id_for_finalize { - let lines: Vec<&str> = sse_str.lines().collect(); - let mut event_type: Option<&str> = None; - let mut data: Option<&str> = None; - - for line in &lines { - if line.starts_with("event: ") { - event_type = Some(&line[7..]); - } else if line.starts_with("data: ") { - data = Some(&line[6..]); - } - } - - if let Some(d) = data { - let flow_monitor_clone = flow_monitor_for_finalize.clone(); - let fid_clone = fid.clone(); - let event_type_owned = event_type.map(|s| s.to_string()); - let data_owned = d.to_string(); - - tokio::spawn(async move { - flow_monitor_clone - .process_chunk(&fid_clone, event_type_owned.as_deref(), &data_owned) - .await; - }); - } - } - yield Ok::(sse_str); } }; diff --git a/src-tauri/src/server/handlers/websocket.rs b/src-tauri/src/server/handlers/websocket.rs index 338c758ad..2fe842caf 100644 --- a/src-tauri/src/server/handlers/websocket.rs +++ b/src-tauri/src/server/handlers/websocket.rs @@ -30,7 +30,7 @@ use crate::providers::{ use crate::server::AppState; use crate::server_utils::parse_cw_response; use crate::websocket::{ - WsApiRequest, WsApiResponse, WsEndpoint, WsError, WsFlowEvent, WsMessage as WsProtoMessage, + WsApiRequest, WsApiResponse, WsEndpoint, WsError, WsMessage as WsProtoMessage, }; /// WebSocket 查询参数 @@ -125,60 +125,6 @@ pub async fn handle_websocket( let (sender, mut receiver) = socket.split(); let sender = Arc::new(Mutex::new(sender)); - // Flow 事件订阅状态 - let flow_subscribed = Arc::new(std::sync::atomic::AtomicBool::new(false)); - - // 启动 Flow 事件转发任务 - let flow_sender = sender.clone(); - let flow_subscribed_clone = flow_subscribed.clone(); - let flow_monitor = state.flow_monitor.clone(); - let conn_id_clone = conn_id.clone(); - let _logs_clone = state.logs.clone(); - - let flow_task = tokio::spawn(async move { - let mut flow_receiver = flow_monitor.subscribe(); - - loop { - match flow_receiver.recv().await { - Ok(event) => { - // 只有在订阅状态下才转发事件 - if !flow_subscribed_clone.load(std::sync::atomic::Ordering::Relaxed) { - continue; - } - - // 转换为 WebSocket 消息 - let ws_event: WsFlowEvent = event.into(); - let ws_msg = WsProtoMessage::FlowEvent(ws_event); - - if let Ok(msg_text) = serde_json::to_string(&ws_msg) { - let mut sender_guard = flow_sender.lock().await; - if sender_guard.send(WsMessage::Text(msg_text)).await.is_err() { - tracing::debug!( - "[WS] Flow event send failed for connection {}", - &conn_id_clone[..8] - ); - break; - } - } - } - Err(tokio::sync::broadcast::error::RecvError::Lagged(n)) => { - tracing::warn!( - "[WS] Flow event receiver lagged by {} messages for connection {}", - n, - &conn_id_clone[..8] - ); - } - Err(tokio::sync::broadcast::error::RecvError::Closed) => { - tracing::debug!( - "[WS] Flow event channel closed for connection {}", - &conn_id_clone[..8] - ); - break; - } - } - } - }); - // 消息处理循环 while let Some(msg) = receiver.next().await { match msg { @@ -188,8 +134,7 @@ pub async fn handle_websocket( match serde_json::from_str::(&text) { Ok(ws_msg) => { - let response = - handle_ws_message(&state, &conn_id, ws_msg, &flow_subscribed).await; + let response = handle_ws_message(&state, &conn_id, ws_msg).await; if let Some(resp) = response { let resp_text = serde_json::to_string(&resp).unwrap_or_default(); let mut sender_guard = sender.lock().await; @@ -252,9 +197,6 @@ pub async fn handle_websocket( } } - // 取消 Flow 事件转发任务 - flow_task.abort(); - // 清理连接 state.ws_manager.unregister(&conn_id); state.logs.write().await.add( @@ -268,56 +210,10 @@ async fn handle_ws_message( state: &AppState, conn_id: &str, msg: WsProtoMessage, - flow_subscribed: &Arc, ) -> Option { match msg { WsProtoMessage::Ping { timestamp } => Some(WsProtoMessage::Pong { timestamp }), WsProtoMessage::Pong { .. } => None, - WsProtoMessage::SubscribeFlowEvents => { - // 订阅 Flow 事件 - flow_subscribed.store(true, std::sync::atomic::Ordering::Relaxed); - state.logs.write().await.add( - "info", - &format!( - "[WS] Connection {} subscribed to flow events", - &conn_id[..8] - ), - ); - // 返回确认消息 - Some(WsProtoMessage::Response(WsApiResponse { - request_id: "subscribe_flow_events".to_string(), - payload: serde_json::json!({ - "status": "subscribed", - "message": "Successfully subscribed to flow events" - }), - })) - } - WsProtoMessage::UnsubscribeFlowEvents => { - // 取消订阅 Flow 事件 - flow_subscribed.store(false, std::sync::atomic::Ordering::Relaxed); - state.logs.write().await.add( - "info", - &format!( - "[WS] Connection {} unsubscribed from flow events", - &conn_id[..8] - ), - ); - // 返回确认消息 - Some(WsProtoMessage::Response(WsApiResponse { - request_id: "unsubscribe_flow_events".to_string(), - payload: serde_json::json!({ - "status": "unsubscribed", - "message": "Successfully unsubscribed from flow events" - }), - })) - } - WsProtoMessage::FlowEvent(_) => { - // 客户端不应该发送 FlowEvent 消息 - Some(WsProtoMessage::Error(WsError::invalid_request( - None, - "FlowEvent messages are server-to-client only", - ))) - } WsProtoMessage::Request(request) => { state.logs.write().await.add( "info", diff --git a/src-tauri/src/server/mod.rs b/src-tauri/src/server/mod.rs index f0ab9f4fe..97d95dfbf 100644 --- a/src-tauri/src/server/mod.rs +++ b/src-tauri/src/server/mod.rs @@ -10,7 +10,6 @@ use crate::converter::anthropic_to_openai::convert_anthropic_to_openai; use crate::credential::CredentialSyncService; use crate::database::dao::provider_pool::ProviderPoolDao; use crate::database::DbConnection; -use crate::flow_monitor::{FlowInterceptor, FlowMonitor, FlowMonitorConfig}; use crate::injection::Injector; use crate::logger::LogStore; use crate::models::anthropic::*; @@ -259,17 +258,14 @@ impl ServerState { shared_stats, shared_tokens, shared_logger, - None, - None, ) .await } - /// 启动服务器(使用共享的遥测实例和 Flow Monitor) + /// 启动服务器(使用共享的遥测实例) /// /// 这允许服务器与 TelemetryState 共享同一个 StatsAggregator、TokenTracker 和 RequestLogger, - /// 以及与 FlowMonitorState 共享同一个 FlowMonitor, - /// 使得请求处理过程中记录的统计数据和 Flow 数据能够在前端监控页面中显示。 + /// 使得请求处理过程中记录的统计数据能够在前端监控页面中显示。 pub async fn start_with_telemetry_and_flow_monitor( &mut self, logs: Arc>, @@ -279,8 +275,6 @@ impl ServerState { shared_stats: Option>>, shared_tokens: Option>>, shared_logger: Option>, - shared_flow_monitor: Option>, - shared_flow_interceptor: Option>, ) -> Result<(), Box> { if self.running { return Ok(()); @@ -391,8 +385,6 @@ impl ServerState { shared_stats, shared_tokens, shared_logger, - shared_flow_monitor, - shared_flow_interceptor, Some(config), Some(config_path), Some(processor), @@ -424,16 +416,6 @@ impl ServerState { } } -impl Clone for KiroProvider { - fn clone(&self) -> Self { - Self { - credentials: self.credentials.clone(), - client: reqwest::Client::new(), - creds_path: self.creds_path.clone(), - } - } -} - pub mod handlers; #[derive(Clone)] @@ -465,10 +447,6 @@ pub struct AppState { pub request_logger: Option>, /// Amp CLI 路由器 pub amp_router: Arc, - /// Flow 监控服务 - pub flow_monitor: Arc, - /// Flow 拦截器 - pub flow_interceptor: Arc, /// 端点 Provider 配置 pub endpoint_providers: Arc>, /// Kiro 事件服务 @@ -754,8 +732,6 @@ async fn run_server( shared_stats: Option>>, shared_tokens: Option>>, shared_logger: Option>, - shared_flow_monitor: Option>, - shared_flow_interceptor: Option>, config: Option, config_path: Option, processor: Option>, @@ -841,14 +817,6 @@ async fn run_server( .unwrap_or_default(), )); - // 使用共享的 Flow 监控服务,如果没有则创建新的 - let flow_monitor = shared_flow_monitor - .unwrap_or_else(|| Arc::new(FlowMonitor::new(FlowMonitorConfig::default(), None))); - - // 使用共享的 Flow 拦截器,如果没有则创建新的 - let flow_interceptor = - shared_flow_interceptor.unwrap_or_else(|| Arc::new(FlowInterceptor::default())); - // 初始化端点 Provider 配置 let endpoint_providers = Arc::new(RwLock::new( config @@ -883,8 +851,6 @@ async fn run_server( hot_reload_manager: hot_reload_manager.clone(), request_logger: shared_logger, amp_router, - flow_monitor, - flow_interceptor, endpoint_providers, kiro_event_service, api_key_service, diff --git a/src-tauri/src/services/provider_pool_service.rs b/src-tauri/src/services/provider_pool_service.rs index 31700c995..c954c6566 100644 --- a/src-tauri/src/services/provider_pool_service.rs +++ b/src-tauri/src/services/provider_pool_service.rs @@ -14,10 +14,28 @@ use crate::models::provider_pool_model::{ use crate::models::route_model::RouteInfo; use crate::providers::antigravity::TokenRefreshError; use crate::providers::kiro::KiroProvider; +use crate::server::client_detector::ClientType; use crate::services::api_key_provider_service::ApiKeyProviderService; use chrono::Utc; use reqwest::Client; use serde::{Deserialize, Serialize}; + +/// 扩展 ProviderCredential 的客户端兼容性检查 +/// (此方法依赖 server::client_detector,不适合放在 core crate) +trait ProviderCredentialClientCompat { + fn is_compatible_with_client(&self, client_type: Option<&ClientType>) -> bool; +} + +impl ProviderCredentialClientCompat for ProviderCredential { + fn is_compatible_with_client(&self, client_type: Option<&ClientType>) -> bool { + if let Some(error_msg) = &self.last_error_message { + if error_msg.contains("only authorized for use with Claude Code") { + return matches!(client_type, Some(ClientType::ClaudeCode)); + } + } + true + } +} use std::collections::HashMap; use std::sync::atomic::AtomicUsize; use std::time::Duration; diff --git a/src-tauri/src/session/mod.rs b/src-tauri/src/session/mod.rs index 652852a98..192e77d0d 100644 --- a/src-tauri/src/session/mod.rs +++ b/src-tauri/src/session/mod.rs @@ -7,19 +7,15 @@ //! - 调度模式配置 //! - 增强的限流处理(Duration 解析、指数退避) -mod rate_limit; -mod session_manager; -mod signature_store; -mod sticky_config; -mod sticky_manager; - -pub use rate_limit::{ - extract_retry_delay, parse_duration_string, RateLimitReason, RateLimitRecord, RateLimitTracker, -}; -pub use session_manager::SessionManager; -pub use signature_store::{ +// 从 providers crate 重新导出 session_manager 和 signature_store +pub use proxycast_providers::session::SessionManager; +pub use proxycast_providers::session::{ clear_thought_signature, get_thought_signature, has_valid_signature, store_thought_signature, take_thought_signature, }; -pub use sticky_config::{SchedulingMode, StickySessionConfig}; -pub use sticky_manager::{AccountInfo, StickySessionManager}; + +// 从 core crate 重新导出 rate_limit、sticky_config、sticky_manager +pub use proxycast_core::session::{ + extract_retry_delay, parse_duration_string, AccountInfo, RateLimitReason, RateLimitRecord, + RateLimitTracker, SchedulingMode, StickySessionConfig, StickySessionManager, +}; diff --git a/src-tauri/src/websocket/handler.rs b/src-tauri/src/websocket/handler.rs index 18d4a7206..87d991940 100644 --- a/src-tauri/src/websocket/handler.rs +++ b/src-tauri/src/websocket/handler.rs @@ -245,21 +245,6 @@ async fn handle_message( // 忽略客户端发送的错误消息 None } - WsMessage::SubscribeFlowEvents | WsMessage::UnsubscribeFlowEvents => { - // Flow 事件订阅在 server/handlers/websocket.rs 中处理 - // 这里的 handler 是旧的实现,暂时返回不支持的错误 - Some(WsMessage::Error(WsError::invalid_request( - None, - "Flow event subscription is not supported in this handler", - ))) - } - WsMessage::FlowEvent(_) => { - // 客户端不应发送 FlowEvent 消息 - Some(WsMessage::Error(WsError::invalid_request( - None, - "FlowEvent messages are server-to-client only", - ))) - } WsMessage::SubscribeKiroEvents => { // TODO: 实现Kiro事件订阅 None diff --git a/src-tauri/src/websocket/mod.rs b/src-tauri/src/websocket/mod.rs index 9ae2ea094..563a0db44 100644 --- a/src-tauri/src/websocket/mod.rs +++ b/src-tauri/src/websocket/mod.rs @@ -17,7 +17,7 @@ mod types; pub use processor::MessageProcessor; pub use types::{ KiroTokenInfo, WsApiRequest, WsApiResponse, WsConfig, WsConnection, WsEndpoint, WsError, - WsFlowEvent, WsKiroEvent, WsMessage, WsStats, WsStatsSnapshot, WsStreamChunk, WsStreamEnd, + WsKiroEvent, WsMessage, WsStats, WsStatsSnapshot, WsStreamChunk, WsStreamEnd, }; use dashmap::DashMap; diff --git a/src-tauri/src/websocket/types.rs b/src-tauri/src/websocket/types.rs index 9f07ffdb7..fe0f14032 100644 --- a/src-tauri/src/websocket/types.rs +++ b/src-tauri/src/websocket/types.rs @@ -6,11 +6,6 @@ use chrono::{DateTime, Utc}; use serde::{Deserialize, Serialize}; use std::sync::atomic::{AtomicU64, Ordering}; -use crate::flow_monitor::models::FlowError; -use crate::flow_monitor::monitor::{ - FlowEvent, FlowSummary, FlowUpdate, NotificationEvent, ThresholdCheckResult, -}; - /// WebSocket 连接信息 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct WsConnection { @@ -74,12 +69,6 @@ pub enum WsMessage { Ping { timestamp: i64 }, /// 心跳响应 Pong { timestamp: i64 }, - /// 订阅 Flow 事件 - SubscribeFlowEvents, - /// 取消订阅 Flow 事件 - UnsubscribeFlowEvents, - /// Flow 事件通知 - FlowEvent(WsFlowEvent), /// 订阅 Kiro 凭证状态事件 SubscribeKiroEvents, /// 取消订阅 Kiro 凭证状态事件 @@ -333,49 +322,6 @@ pub struct WsStatsSnapshot { pub total_errors: u64, } -/// WebSocket Flow 事件 -/// -/// 用于通过 WebSocket 推送 Flow 监控事件 -#[derive(Debug, Clone, Serialize, Deserialize)] -#[serde(tag = "event_type", rename_all = "snake_case")] -pub enum WsFlowEvent { - /// Flow 开始 - FlowStarted { flow: FlowSummary }, - /// Flow 更新 - FlowUpdated { id: String, update: FlowUpdate }, - /// Flow 完成 - FlowCompleted { id: String, summary: FlowSummary }, - /// Flow 失败 - FlowFailed { id: String, error: FlowError }, - /// 阈值警告 - ThresholdWarning { - id: String, - result: ThresholdCheckResult, - }, - /// 通知事件 - Notification { notification: NotificationEvent }, - /// 请求速率更新 - RequestRateUpdate { rate: f64, count: usize }, -} - -impl From for WsFlowEvent { - fn from(event: FlowEvent) -> Self { - match event { - FlowEvent::FlowStarted { flow } => WsFlowEvent::FlowStarted { flow }, - FlowEvent::FlowUpdated { id, update } => WsFlowEvent::FlowUpdated { id, update }, - FlowEvent::FlowCompleted { id, summary } => WsFlowEvent::FlowCompleted { id, summary }, - FlowEvent::FlowFailed { id, error } => WsFlowEvent::FlowFailed { id, error }, - FlowEvent::ThresholdWarning { id, result } => { - WsFlowEvent::ThresholdWarning { id, result } - } - FlowEvent::Notification { notification } => WsFlowEvent::Notification { notification }, - FlowEvent::RequestRateUpdate { rate, count } => { - WsFlowEvent::RequestRateUpdate { rate, count } - } - } - } -} - /// WebSocket Kiro 凭证事件 /// /// 用于通过 WebSocket 推送 Kiro 凭证状态变化 diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index ecd253afa..1c0c2873c 100644 --- a/src-tauri/tauri.conf.json +++ b/src-tauri/tauri.conf.json @@ -1,7 +1,7 @@ { "$schema": "https://schema.tauri.app/config/2", "productName": "ProxyCast", - "version": "0.59.0", + "version": "0.60.0", "identifier": "com.proxycast.app", "build": { "beforeDevCommand": "npm run dev", diff --git a/src-tauri/tests/end_to_end_tests.rs b/src-tauri/tests/end_to_end_tests.rs deleted file mode 100644 index cbe359707..000000000 --- a/src-tauri/tests/end_to_end_tests.rs +++ /dev/null @@ -1,469 +0,0 @@ -//! Flow Monitor Enhancement 端到端功能验证测试 -//! -//! 验证 Flow Monitor 的基础功能,包括: -//! - Flow 数据结构创建和操作 -//! - 内存存储功能 -//! - 查询服务功能 -//! - 导出功能 -//! -//! **Validates: Requirements 8.2** - -use std::sync::Arc; -use tempfile::TempDir; - -use chrono::Utc; -use proxycast_lib::flow_monitor::{ - ClientInfo, ExportFormat, ExportOptions, FlowAnnotations, FlowExporter, FlowFileStore, - FlowFilter, FlowMetadata, FlowMonitor, FlowMonitorConfig, FlowQueryService, FlowSortBy, - FlowState, FlowTimestamps, FlowType, LLMFlow, LLMRequest, LLMResponse, Message, MessageContent, - MessageRole, ProviderType, RequestParameters, RotationConfig, RoutingInfo, TokenUsage, -}; -use std::collections::HashMap; - -/// 端到端测试上下文 -#[allow(dead_code)] -struct E2ETestContext { - pub temp_dir: TempDir, - pub flow_monitor: Arc, - pub flow_query_service: Arc, -} - -impl E2ETestContext { - /// 创建端到端测试上下文 - pub async fn new() -> Result> { - // 创建临时目录 - let temp_dir = TempDir::new()?; - let temp_path = temp_dir.path().to_path_buf(); - - // 创建 Flow 文件存储 - let flows_dir = temp_path.join("flows"); - std::fs::create_dir_all(&flows_dir)?; - let rotation_config = RotationConfig::default(); - let flow_file_store = Arc::new(FlowFileStore::new(flows_dir, rotation_config)?); - - // 创建 Flow Monitor - let flow_monitor_config = FlowMonitorConfig::default(); - let flow_monitor = Arc::new(FlowMonitor::new( - flow_monitor_config, - Some(flow_file_store.clone()), - )); - - // 创建 Flow Query Service - let flow_query_service = Arc::new(FlowQueryService::new( - flow_monitor.memory_store(), - flow_file_store, - )); - - Ok(Self { - temp_dir, - flow_monitor, - flow_query_service, - }) - } - - /// 创建测试用的 Flow - pub fn create_test_flow(&self, id: &str, provider: ProviderType, model: &str) -> LLMFlow { - let now = Utc::now(); - - LLMFlow { - id: id.to_string(), - flow_type: FlowType::ChatCompletions, - request: LLMRequest { - method: "POST".to_string(), - path: "/v1/chat/completions".to_string(), - headers: HashMap::new(), - body: serde_json::json!({ - "model": model, - "messages": [{"role": "user", "content": "Hello, world!"}] - }), - timestamp: now, - system_prompt: None, - messages: vec![Message { - role: MessageRole::User, - content: MessageContent::Text("Hello, world!".to_string()), - name: None, - tool_calls: None, - tool_result: None, - }], - parameters: RequestParameters { - temperature: None, - top_p: None, - max_tokens: None, - stop: None, - stream: false, - extra: HashMap::new(), - }, - model: model.to_string(), - original_model: Some(model.to_string()), - size_bytes: 100, - tools: None, - }, - response: Some(LLMResponse { - status_code: 200, - status_text: "OK".to_string(), - headers: HashMap::new(), - body: serde_json::json!({ - "id": "chatcmpl-test", - "object": "chat.completion", - "created": now.timestamp(), - "model": model, - "choices": [] - }), - content: "Hello! How can I help you today?".to_string(), - stop_reason: None, - usage: TokenUsage { - input_tokens: 10, - output_tokens: 20, - total_tokens: 30, - cache_read_tokens: None, - cache_write_tokens: None, - thinking_tokens: None, - }, - stream_info: None, - thinking: None, - tool_calls: vec![], - size_bytes: 200, - timestamp_start: now, - timestamp_end: now, - }), - error: None, - metadata: FlowMetadata { - provider, - provider_id: None, - credential_name: Some("test-cred".to_string()), - credential_id: Some("test-cred-id".to_string()), - retry_count: 0, - injected_params: Some(HashMap::new()), - context_usage_percentage: None, - client_info: ClientInfo::default(), - routing_info: RoutingInfo::default(), - }, - timestamps: FlowTimestamps { - created: now, - request_start: now, - request_end: Some(now), - response_start: Some(now), - response_end: Some(now), - duration_ms: 500, - ttfb_ms: Some(100), - }, - state: FlowState::Completed, - annotations: FlowAnnotations::default(), - } - } - - /// 设置测试数据 - pub async fn setup_test_data(&self) -> Result<(), Box> { - // 创建多样化的测试 Flow - let test_flows = vec![ - ("flow-kiro-claude", ProviderType::Kiro, "claude-3-5-sonnet"), - ("flow-openai-gpt4", ProviderType::OpenAI, "gpt-4"), - ("flow-gemini-pro", ProviderType::Gemini, "gemini-pro"), - ( - "flow-kiro-claude-2", - ProviderType::Kiro, - "claude-3-5-sonnet", - ), - ("flow-openai-gpt35", ProviderType::OpenAI, "gpt-3.5-turbo"), - ]; - - for (id, provider, model) in test_flows { - let flow = self.create_test_flow(id, provider, model); - // 直接添加到内存存储 - let memory_store = self.flow_monitor.memory_store(); - let mut store = memory_store.write().await; - store.add(flow); - } - - Ok(()) - } -} - -#[cfg(test)] -mod tests { - use super::*; - - /// 端到端测试:基础 Flow 操作 - #[tokio::test] - async fn test_e2e_basic_flow_operations() { - let ctx = E2ETestContext::new().await.unwrap(); - ctx.setup_test_data().await.unwrap(); - - // 1. 验证 Flow 已添加到内存存储 - let memory_store = ctx.flow_monitor.memory_store(); - let store = memory_store.read().await; - let recent_flows = store.get_recent(10); - assert_eq!(recent_flows.len(), 5); - - // 2. 测试按 ID 获取 Flow - let flow_lock = store.get("flow-kiro-claude"); - assert!(flow_lock.is_some()); - let flow_lock = flow_lock.unwrap(); - let flow = flow_lock.read().unwrap(); - assert_eq!(flow.id, "flow-kiro-claude"); - - // 3. 测试过滤功能 - let filter = FlowFilter { - providers: Some(vec![ProviderType::Kiro]), - ..Default::default() - }; - let filtered_flows = store.query(&filter); - assert_eq!(filtered_flows.len(), 2); // 应该有 2 个 Kiro 的 Flow - - // 4. 测试状态过滤 - let state_filter = FlowFilter { - states: Some(vec![FlowState::Completed]), - ..Default::default() - }; - let completed_flows = store.query(&state_filter); - assert_eq!(completed_flows.len(), 5); // 所有 Flow 都是 Completed 状态 - - println!("✅ 基础 Flow 操作端到端测试通过"); - } - - /// 端到端测试:Flow 查询服务 - #[tokio::test] - async fn test_e2e_flow_query_service() { - let ctx = E2ETestContext::new().await.unwrap(); - ctx.setup_test_data().await.unwrap(); - - // 1. 测试查询功能 - let result = ctx - .flow_query_service - .query(FlowFilter::default(), FlowSortBy::CreatedAt, true, 1, 20) - .await - .unwrap(); - - assert_eq!(result.flows.len(), 5); - assert_eq!(result.total, 5); - assert_eq!(result.page, 1); - assert_eq!(result.page_size, 20); - - // 2. 测试分页 - let page_result = ctx - .flow_query_service - .query(FlowFilter::default(), FlowSortBy::CreatedAt, true, 1, 3) - .await - .unwrap(); - - assert_eq!(page_result.flows.len(), 3); - assert_eq!(page_result.total, 5); - - // 3. 测试按 ID 获取 - let flow = ctx - .flow_query_service - .get_flow("flow-kiro-claude") - .await - .unwrap(); - assert!(flow.is_some()); - let flow = flow.unwrap(); - assert_eq!(flow.id, "flow-kiro-claude"); - - // 4. 测试搜索功能 - let search_results = ctx.flow_query_service.search("claude", 10).await.unwrap(); - assert_eq!(search_results.len(), 2); // 应该找到 2 个包含 claude 的 Flow - - // 5. 测试统计功能 - let stats = ctx - .flow_query_service - .get_stats(&FlowFilter::default()) - .await; - assert_eq!(stats.total_requests, 5); - - println!("✅ Flow 查询服务端到端测试通过"); - } - - /// 端到端测试:Flow 导出功能 - #[tokio::test] - async fn test_e2e_flow_export() { - let ctx = E2ETestContext::new().await.unwrap(); - ctx.setup_test_data().await.unwrap(); - - // 获取所有 Flow - let memory_store = ctx.flow_monitor.memory_store(); - let store = memory_store.read().await; - let all_flows = store.get_recent(10); - - // 1. 测试 JSON 导出 - let json_options = ExportOptions { - format: ExportFormat::JSON, - filter: None, - include_raw: true, - include_stream_chunks: false, - redact_sensitive: false, - redaction_rules: Vec::new(), - compress: false, - }; - let json_exporter = FlowExporter::new(json_options); - let json_data = json_exporter.export_json(&all_flows); - // json_data 应该是一个数组,包含所有的 Flow - assert!(json_data.is_array()); - let flows_array = json_data.as_array().unwrap(); - assert_eq!(flows_array.len(), all_flows.len()); - - // 2. 测试 JSONL 导出 - let jsonl_options = ExportOptions { - format: ExportFormat::JSONL, - filter: None, - include_raw: true, - include_stream_chunks: false, - redact_sensitive: false, - redaction_rules: Vec::new(), - compress: false, - }; - let jsonl_exporter = FlowExporter::new(jsonl_options); - let jsonl_data = jsonl_exporter.export_jsonl(&all_flows); - let lines: Vec<&str> = jsonl_data.lines().collect(); - assert_eq!(lines.len(), 5); // 应该有 5 行 - - // 3. 测试 HAR 导出 - let har_options = ExportOptions { - format: ExportFormat::HAR, - filter: None, - include_raw: true, - include_stream_chunks: false, - redact_sensitive: false, - redaction_rules: Vec::new(), - compress: false, - }; - let har_exporter = FlowExporter::new(har_options); - let har_archive = har_exporter.export_har(&all_flows); - assert_eq!(har_archive.log.entries.len(), 5); - - // 4. 测试 Markdown 导出 - let md_options = ExportOptions { - format: ExportFormat::Markdown, - filter: None, - include_raw: false, - include_stream_chunks: false, - redact_sensitive: true, - redaction_rules: Vec::new(), - compress: false, - }; - let md_exporter = FlowExporter::new(md_options); - let md_data = md_exporter.export_markdown_multiple(&all_flows); - assert!(md_data.contains("#")); // Markdown 应该包含标题 - - // 5. 测试 CSV 导出 - let csv_options = ExportOptions { - format: ExportFormat::CSV, - filter: None, - include_raw: false, - include_stream_chunks: false, - redact_sensitive: false, - redaction_rules: Vec::new(), - compress: false, - }; - let csv_exporter = FlowExporter::new(csv_options); - let csv_data = csv_exporter.export_csv(&all_flows); - let lines: Vec<&str> = csv_data.lines().collect(); - assert!(lines.len() > 1); // 应该有标题行和数据行 - - println!("✅ Flow 导出端到端测试通过"); - } - - /// 端到端测试:Flow 标注功能 - #[tokio::test] - async fn test_e2e_flow_annotations() { - let ctx = E2ETestContext::new().await.unwrap(); - ctx.setup_test_data().await.unwrap(); - - let flow_id = "flow-kiro-claude"; - - // 1. 测试切换收藏状态 - let updated = ctx.flow_monitor.toggle_starred(flow_id).await; - assert!(updated); - - // 2. 测试添加评论 - let comment_added = ctx - .flow_monitor - .add_comment(flow_id, "这是一个测试评论".to_string()) - .await; - assert!(comment_added); - - // 3. 测试添加标签 - let tag_added = ctx.flow_monitor.add_tag(flow_id, "重要".to_string()).await; - assert!(tag_added); - - // 4. 测试设置标记 - let marker_set = ctx - .flow_monitor - .set_marker(flow_id, Some("⭐".to_string())) - .await; - assert!(marker_set); - - // 5. 验证标注已更新 - let memory_store = ctx.flow_monitor.memory_store(); - { - let store = memory_store.read().await; - let flow_lock = store.get(flow_id); - assert!(flow_lock.is_some()); - let flow_lock = flow_lock.unwrap(); - let flow = flow_lock.read().unwrap(); - assert!(flow.annotations.starred); - assert!(flow.annotations.comment.is_some()); - assert!(flow.annotations.tags.contains(&"重要".to_string())); - assert_eq!(flow.annotations.marker, Some("⭐".to_string())); - } // 确保 store 锁在这里被释放 - - // 6. 测试移除标签 - let tag_removed = ctx.flow_monitor.remove_tag(flow_id, "重要").await; - assert!(tag_removed); - - // 7. 测试清除标记 - let marker_cleared = ctx.flow_monitor.set_marker(flow_id, None).await; - assert!(marker_cleared); - - println!("✅ Flow 标注端到端测试通过"); - } - - /// 端到端测试:Flow 过滤和排序 - #[tokio::test] - async fn test_e2e_flow_filtering_and_sorting() { - let ctx = E2ETestContext::new().await.unwrap(); - ctx.setup_test_data().await.unwrap(); - - // 1. 测试按提供商过滤 - let provider_filter = FlowFilter { - providers: Some(vec![ProviderType::Kiro]), - ..Default::default() - }; - let kiro_result = ctx - .flow_query_service - .query(provider_filter, FlowSortBy::CreatedAt, true, 1, 20) - .await - .unwrap(); - assert_eq!(kiro_result.flows.len(), 2); - - // 2. 测试按状态过滤 - let state_filter = FlowFilter { - states: Some(vec![FlowState::Completed]), - ..Default::default() - }; - let completed_result = ctx - .flow_query_service - .query(state_filter, FlowSortBy::Duration, false, 1, 20) - .await - .unwrap(); - assert_eq!(completed_result.flows.len(), 5); - - // 3. 测试分页 - let page_result = ctx - .flow_query_service - .query(FlowFilter::default(), FlowSortBy::CreatedAt, true, 1, 3) - .await - .unwrap(); - assert_eq!(page_result.flows.len(), 3); - assert_eq!(page_result.total, 5); - - // 4. 测试排序 - let sorted_result = ctx - .flow_query_service - .query(FlowFilter::default(), FlowSortBy::Duration, true, 1, 20) - .await - .unwrap(); - assert_eq!(sorted_result.flows.len(), 5); - - println!("✅ Flow 过滤和排序端到端测试通过"); - } -} diff --git a/src/App.tsx b/src/App.tsx index 4d49681e3..40adb8a82 100644 --- a/src/App.tsx +++ b/src/App.tsx @@ -32,7 +32,6 @@ import { FileBrowserView, WebView, } from "./components/terminal"; -import { flowEventManager } from "./lib/flowEventManager"; import { OnboardingWizard, useOnboardingState } from "./components/onboarding"; import { ConnectConfirmDialog } from "./components/connect"; import { showRegistryLoadError } from "./lib/utils/connectError"; @@ -218,11 +217,6 @@ function AppContent() { refresh: _refreshRegistry, // 保留以供后续错误处理 UI 使用 } = useRelayRegistry(); - // 在应用启动时初始化 Flow 事件订阅 - useEffect(() => { - flowEventManager.subscribe(); - }, []); - // 处理 Registry 加载失败 // _Requirements: 7.2, 7.3_ useEffect(() => { diff --git a/src/components/flow-monitor/BatchOperations.tsx b/src/components/flow-monitor/BatchOperations.tsx deleted file mode 100644 index 0d6c5c7bd..000000000 --- a/src/components/flow-monitor/BatchOperations.tsx +++ /dev/null @@ -1,660 +0,0 @@ -/** - * 批量操作组件 - * 实现批量选择、批量操作菜单、操作进度显示 - * **Validates: Requirements 11.1-11.7** - */ - -import React, { useState, useCallback } from "react"; -import { safeInvoke } from "@/lib/dev-bridge"; -import { - CheckSquare, - Square, - Star, - StarOff, - Tag, - Download, - Trash2, - FolderPlus, - X, - Loader2, - AlertCircle, - Check, - ChevronDown, - MinusSquare, -} from "lucide-react"; -import { cn } from "@/lib/utils"; -import type { ExportFormat, LLMFlow } from "@/lib/api/flowMonitor"; - -export interface BatchResult { - total: number; - success: number; - failed: number; - errors: [string, string][]; - export_data?: string; -} - -export interface SessionInfo { - id: string; - name: string; -} - -export type BatchOperationType = - | "star" - | "unstar" - | "addTags" - | "removeTags" - | "export" - | "delete" - | "addToSession"; - -interface BatchOperationsProps { - flows: LLMFlow[]; - selectedIds: Set; - onSelectionChange: (selectedIds: Set) => void; - onOperationComplete?: ( - result: BatchResult, - operation: BatchOperationType, - ) => void; - sessions?: SessionInfo[]; - availableTags?: string[]; - onRefresh?: () => void; - className?: string; -} - -export function BatchOperations({ - flows, - selectedIds, - onSelectionChange, - onOperationComplete, - sessions = [], - availableTags = [], - onRefresh, - className, -}: BatchOperationsProps) { - const [showMenu, setShowMenu] = useState(false); - const [operating, setOperating] = useState(false); - const [currentOp, setCurrentOp] = useState(null); - const [error, setError] = useState(null); - const [showTagDialog, setShowTagDialog] = useState(false); - const [tagMode, setTagMode] = useState<"add" | "remove">("add"); - const [selectedTags, setSelectedTags] = useState([]); - const [newTag, setNewTag] = useState(""); - const [showSessionDialog, setShowSessionDialog] = useState(false); - const [selectedSessionId, setSelectedSessionId] = useState(""); - const [showExportDialog, setShowExportDialog] = useState(false); - const [exportFormat, setExportFormat] = useState("json"); - const [showDeleteConfirm, setShowDeleteConfirm] = useState(false); - - const selectedCount = selectedIds.size; - const totalCount = flows.length; - const allSelected = totalCount > 0 && selectedCount === totalCount; - const someSelected = selectedCount > 0 && selectedCount < totalCount; - - const handleSelectAll = useCallback(() => { - onSelectionChange( - allSelected ? new Set() : new Set(flows.map((f) => f.id)), - ); - }, [allSelected, flows, onSelectionChange]); - - const handleClearSelection = useCallback(() => { - onSelectionChange(new Set()); - setShowMenu(false); - }, [onSelectionChange]); - - const runBatchOp = useCallback( - async ( - op: BatchOperationType, - command: string, - request: Record, - onSuccess?: () => void, - ) => { - if (selectedCount === 0) return; - try { - setOperating(true); - setCurrentOp(op); - setError(null); - const result = await safeInvoke(command, { request }); - onOperationComplete?.(result, op); - onSuccess?.(); - onRefresh?.(); - setShowMenu(false); - } catch (e) { - setError(e instanceof Error ? e.message : `批量${op}失败`); - } finally { - setOperating(false); - setCurrentOp(null); - } - }, - [selectedCount, onOperationComplete, onRefresh], - ); - - const handleBatchStar = () => - runBatchOp("star", "batch_star_flows", { - flow_ids: Array.from(selectedIds), - }); - - const handleBatchUnstar = () => - runBatchOp("unstar", "batch_unstar_flows", { - flow_ids: Array.from(selectedIds), - }); - - const handleBatchAddTags = () => { - if (selectedTags.length === 0) return; - runBatchOp( - "addTags", - "batch_add_tags", - { flow_ids: Array.from(selectedIds), tags: selectedTags }, - () => { - setShowTagDialog(false); - setSelectedTags([]); - }, - ); - }; - - const handleBatchRemoveTags = () => { - if (selectedTags.length === 0) return; - runBatchOp( - "removeTags", - "batch_remove_tags", - { flow_ids: Array.from(selectedIds), tags: selectedTags }, - () => { - setShowTagDialog(false); - setSelectedTags([]); - }, - ); - }; - - const handleBatchExport = async () => { - if (selectedCount === 0) return; - try { - setOperating(true); - setCurrentOp("export"); - setError(null); - const result = await safeInvoke("batch_export_flows", { - request: { flow_ids: Array.from(selectedIds), format: exportFormat }, - }); - if (result.export_data) { - const blob = new Blob([result.export_data], { - type: "application/json", - }); - const url = URL.createObjectURL(blob); - const a = document.createElement("a"); - a.href = url; - a.download = `flows_${new Date().toISOString().slice(0, 10)}.${exportFormat}`; - document.body.appendChild(a); - a.click(); - document.body.removeChild(a); - URL.revokeObjectURL(url); - } - onOperationComplete?.(result, "export"); - setShowExportDialog(false); - setShowMenu(false); - } catch (e) { - setError(e instanceof Error ? e.message : "批量导出失败"); - } finally { - setOperating(false); - setCurrentOp(null); - } - }; - - const handleBatchDelete = () => - runBatchOp( - "delete", - "batch_delete_flows", - { flow_ids: Array.from(selectedIds) }, - () => { - handleClearSelection(); - setShowDeleteConfirm(false); - }, - ); - - const handleBatchAddToSession = () => { - if (!selectedSessionId) return; - runBatchOp( - "addToSession", - "batch_add_to_session", - { flow_ids: Array.from(selectedIds), session_id: selectedSessionId }, - () => { - setShowSessionDialog(false); - setSelectedSessionId(""); - }, - ); - }; - - const toggleTag = (tag: string) => { - setSelectedTags((prev) => - prev.includes(tag) ? prev.filter((t) => t !== tag) : [...prev, tag], - ); - }; - - const handleAddNewTag = () => { - const t = newTag.trim(); - if (t && !selectedTags.includes(t)) { - setSelectedTags([...selectedTags, t]); - setNewTag(""); - } - }; - - const getOpLabel = (op: BatchOperationType | null) => { - switch (op) { - case "star": - return "收藏"; - case "unstar": - return "取消收藏"; - case "addTags": - return "添加标签"; - case "removeTags": - return "移除标签"; - case "export": - return "导出"; - case "delete": - return "删除"; - case "addToSession": - return "添加到会话"; - default: - return ""; - } - }; - - if (selectedCount === 0) return null; - - return ( -
- {/* 批量操作工具栏 */} -
-
- - -
- - {/* 操作按钮 */} -
- - -
- - {showMenu && ( -
- - - {sessions.length > 0 && ( - - )} - -
- -
- )} -
-
-
- - {/* 操作进度/错误提示 */} - {operating && ( -
- - 正在执行批量{getOpLabel(currentOp)}... -
- )} - {error && ( -
- - {error} - -
- )} - - {/* 标签对话框 */} - {showTagDialog && ( - { - setShowTagDialog(false); - setSelectedTags([]); - }} - > -
- {tagMode === "add" && ( -
- setNewTag(e.target.value)} - onKeyDown={(e) => e.key === "Enter" && handleAddNewTag()} - placeholder="输入新标签" - className="flex-1 rounded-lg border bg-background px-3 py-2 text-sm" - /> - -
- )} - {availableTags.length > 0 && ( -
-

可用标签:

-
- {availableTags.map((tag) => ( - - ))} -
-
- )} - - {selectedTags.length > 0 && ( -
-

已选标签:

-
- {selectedTags.map((tag) => ( - - {tag} - - - ))} -
-
- )} -
-
- - -
-
- )} - - {/* 会话对话框 */} - {showSessionDialog && ( - { - setShowSessionDialog(false); - setSelectedSessionId(""); - }} - > -
-

选择要添加到的会话:

- -
-
- - -
-
- )} - - {/* 导出对话框 */} - {showExportDialog && ( - setShowExportDialog(false)}> -
-

选择导出格式:

- -

- 将导出 {selectedCount} 个 Flow -

-
-
- - -
-
- )} - - {/* 删除确认对话框 */} - {showDeleteConfirm && ( - setShowDeleteConfirm(false)}> -
-
- -
-

- 此操作不可撤销 -

-

- 确定要删除选中的 {selectedCount} 个 Flow 吗? -

-
-
-
-
- - -
-
- )} -
- ); -} - -// ============================================================================ -// 对话框组件 -// ============================================================================ - -interface DialogProps { - title: string; - children: React.ReactNode; - onClose: () => void; -} - -function Dialog({ title, children, onClose }: DialogProps) { - return ( -
-
-
-
-

{title}

- -
-
{children}
-
-
- ); -} - -export default BatchOperations; diff --git a/src/components/flow-monitor/BookmarkPanel.tsx b/src/components/flow-monitor/BookmarkPanel.tsx deleted file mode 100644 index 70f2844bc..000000000 --- a/src/components/flow-monitor/BookmarkPanel.tsx +++ /dev/null @@ -1,730 +0,0 @@ -/** - * 书签管理面板组件 - * - * 实现书签列表、书签导航功能 - * **Validates: Requirements 8.1-8.6** - */ - -import { useState, useEffect, useCallback } from "react"; -import { safeInvoke } from "@/lib/dev-bridge"; -import { - Bookmark, - Trash2, - Edit2, - Download, - Upload, - ChevronDown, - ChevronUp, - Loader2, - AlertCircle, - X, - Search, - MoreVertical, - Check, - FolderOpen, - Navigation, -} from "lucide-react"; -import { cn } from "@/lib/utils"; - -// ============================================================================ -// 类型定义 -// ============================================================================ - -/** - * Flow 书签 - */ -export interface FlowBookmark { - id: string; - flow_id: string; - name?: string; - group?: string; - created_at: string; -} - -/** - * 更新书签请求 - */ -interface UpdateBookmarkRequest { - bookmark_id: string; - name?: string | null; - group?: string | null; -} - -/** - * 导入书签请求 - */ -interface ImportBookmarksRequest { - data: string; - overwrite: boolean; -} - -// ============================================================================ -// 组件属性 -// ============================================================================ - -interface BookmarkPanelProps { - className?: string; - /** 导航到 Flow 回调 */ - onNavigateToFlow?: (flowId: string) => void; - /** 当前选中的 Flow ID */ - currentFlowId?: string; -} - -// ============================================================================ -// 主组件 -// ============================================================================ - -export function BookmarkPanel({ - className, - onNavigateToFlow, - currentFlowId, -}: BookmarkPanelProps) { - // 状态 - const [bookmarks, setBookmarks] = useState([]); - const [groups, setGroups] = useState([]); - const [loading, setLoading] = useState(true); - const [error, setError] = useState(null); - const [expanded, setExpanded] = useState(true); - const [searchQuery, setSearchQuery] = useState(""); - const [expandedGroups, setExpandedGroups] = useState>(new Set()); - - // 编辑书签对话框状态 - const [editingBookmark, setEditingBookmark] = useState( - null, - ); - const [editName, setEditName] = useState(""); - const [editGroup, setEditGroup] = useState(""); - const [saving, setSaving] = useState(false); - - // 操作菜单状态 - const [menuBookmarkId, setMenuBookmarkId] = useState(null); - - // 导入/导出状态 - const [importing, setImporting] = useState(false); - const [exporting, setExporting] = useState(false); - - // 加载书签列表 - const loadBookmarks = useCallback(async () => { - try { - setLoading(true); - setError(null); - const [bookmarkList, groupList] = await Promise.all([ - safeInvoke("list_bookmarks", { group: null }), - safeInvoke("list_bookmark_groups"), - ]); - setBookmarks(bookmarkList); - setGroups(groupList); - // 默认展开所有分组 - setExpandedGroups(new Set(groupList)); - } catch (e) { - console.error("加载书签失败:", e); - setError(e instanceof Error ? e.message : "加载书签失败"); - } finally { - setLoading(false); - } - }, []); - - // 初始化加载 - useEffect(() => { - loadBookmarks(); - }, [loadBookmarks]); - - // 更新书签 - const handleUpdate = useCallback(async () => { - if (!editingBookmark) return; - - try { - setSaving(true); - const updated = await safeInvoke("update_bookmark", { - request: { - bookmark_id: editingBookmark.id, - name: editName.trim() || null, - group: editGroup.trim() || null, - } as UpdateBookmarkRequest, - }); - setBookmarks((prev) => - prev.map((b) => (b.id === editingBookmark.id ? updated : b)), - ); - if (updated.group && !groups.includes(updated.group)) { - setGroups((prev) => [...prev, updated.group!].sort()); - } - setEditingBookmark(null); - } catch (e) { - console.error("更新书签失败:", e); - setError(e instanceof Error ? e.message : "更新书签失败"); - } finally { - setSaving(false); - } - }, [editingBookmark, editName, editGroup, groups]); - - // 删除书签 - const handleDelete = useCallback(async (bookmarkId: string) => { - if (!confirm("确定要删除此书签吗?")) return; - - try { - await safeInvoke("remove_bookmark", { bookmarkId }); - setBookmarks((prev) => prev.filter((b) => b.id !== bookmarkId)); - setMenuBookmarkId(null); - } catch (e) { - console.error("删除书签失败:", e); - setError(e instanceof Error ? e.message : "删除书签失败"); - } - }, []); - - // 导出书签 - const handleExport = useCallback(async () => { - try { - setExporting(true); - const data = await safeInvoke("export_bookmarks"); - - // 下载文件 - const blob = new Blob([data], { type: "application/json" }); - const url = URL.createObjectURL(blob); - const a = document.createElement("a"); - a.href = url; - a.download = `bookmarks_${new Date().toISOString().slice(0, 10)}.json`; - document.body.appendChild(a); - a.click(); - document.body.removeChild(a); - URL.revokeObjectURL(url); - } catch (e) { - console.error("导出书签失败:", e); - setError(e instanceof Error ? e.message : "导出书签失败"); - } finally { - setExporting(false); - } - }, []); - - // 导入书签 - const handleImport = useCallback(async () => { - try { - // 创建文件输入 - const input = document.createElement("input"); - input.type = "file"; - input.accept = ".json"; - input.onchange = async (e) => { - const file = (e.target as HTMLInputElement).files?.[0]; - if (!file) return; - - try { - setImporting(true); - const data = await file.text(); - const imported = await safeInvoke( - "import_bookmarks", - { - request: { - data, - overwrite: false, - } as ImportBookmarksRequest, - }, - ); - if (imported.length > 0) { - await loadBookmarks(); - setError(null); - } else { - setError("没有导入任何书签(可能已存在相同的书签)"); - } - } catch (err) { - console.error("导入书签失败:", err); - setError(err instanceof Error ? err.message : "导入书签失败"); - } finally { - setImporting(false); - } - }; - input.click(); - } catch (e) { - console.error("导入书签失败:", e); - setError(e instanceof Error ? e.message : "导入书签失败"); - } - }, [loadBookmarks]); - - // 导航到 Flow - const handleNavigate = useCallback( - (bookmark: FlowBookmark) => { - onNavigateToFlow?.(bookmark.flow_id); - setMenuBookmarkId(null); - }, - [onNavigateToFlow], - ); - - // 切换分组展开状态 - const toggleGroup = useCallback((group: string) => { - setExpandedGroups((prev) => { - const next = new Set(prev); - if (next.has(group)) { - next.delete(group); - } else { - next.add(group); - } - return next; - }); - }, []); - - // 过滤书签 - const filteredBookmarks = bookmarks.filter((bookmark) => { - if (searchQuery) { - const query = searchQuery.toLowerCase(); - return ( - bookmark.name?.toLowerCase().includes(query) || - bookmark.flow_id.toLowerCase().includes(query) || - bookmark.group?.toLowerCase().includes(query) - ); - } - return true; - }); - - // 按分组组织书签 - const bookmarksByGroup = filteredBookmarks.reduce( - (acc, bookmark) => { - const group = bookmark.group || "未分组"; - if (!acc[group]) { - acc[group] = []; - } - acc[group].push(bookmark); - return acc; - }, - {} as Record, - ); - - // 排序分组(未分组在最后) - const sortedGroups = Object.keys(bookmarksByGroup).sort((a, b) => { - if (a === "未分组") return 1; - if (b === "未分组") return -1; - return a.localeCompare(b); - }); - - return ( -
- {/* 头部 */} -
setExpanded(!expanded)} - > -
- - 书签 - - ({bookmarks.length}) - -
-
- {expanded ? ( - - ) : ( - - )} -
-
- - {/* 展开内容 */} - {expanded && ( -
- {/* 错误提示 */} - {error && ( -
- - {error} - -
- )} - - {/* 搜索和操作 */} -
-
- - setSearchQuery(e.target.value)} - placeholder="搜索书签..." - className="w-full pl-9 pr-3 py-2 text-sm rounded-lg border bg-background" - /> -
- - -
- - {/* 加载状态 */} - {loading && ( -
- -
- )} - - {/* 书签列表 */} - {!loading && ( -
- {sortedGroups.length > 0 ? ( - sortedGroups.map((group) => ( - toggleGroup(group)} - onNavigate={handleNavigate} - onEdit={(bookmark) => { - setEditingBookmark(bookmark); - setEditName(bookmark.name || ""); - setEditGroup(bookmark.group || ""); - }} - onDelete={handleDelete} - onMenuToggle={(id) => - setMenuBookmarkId(menuBookmarkId === id ? null : id) - } - /> - )) - ) : ( -
- -

暂无书签

-

在 Flow 详情中点击书签图标添加

-
- )} -
- )} -
- )} - - {/* 编辑书签对话框 */} - {editingBookmark && ( - setEditingBookmark(null)} - /> - )} -
- ); -} - -// ============================================================================ -// 子组件 -// ============================================================================ - -interface BookmarkGroupProps { - group: string; - bookmarks: FlowBookmark[]; - expanded: boolean; - currentFlowId?: string; - menuBookmarkId: string | null; - onToggle: () => void; - onNavigate: (bookmark: FlowBookmark) => void; - onEdit: (bookmark: FlowBookmark) => void; - onDelete: (bookmarkId: string) => void; - onMenuToggle: (bookmarkId: string) => void; -} - -function BookmarkGroup({ - group, - bookmarks, - expanded, - currentFlowId, - menuBookmarkId, - onToggle, - onNavigate, - onEdit, - onDelete, - onMenuToggle, -}: BookmarkGroupProps) { - return ( -
- {/* 分组头部 */} - - - {/* 书签列表 */} - {expanded && ( -
- {bookmarks.map((bookmark) => ( - onNavigate(bookmark)} - onEdit={() => onEdit(bookmark)} - onDelete={() => onDelete(bookmark.id)} - onMenuToggle={() => onMenuToggle(bookmark.id)} - /> - ))} -
- )} -
- ); -} - -interface BookmarkItemProps { - bookmark: FlowBookmark; - active: boolean; - menuOpen: boolean; - onNavigate: () => void; - onEdit: () => void; - onDelete: () => void; - onMenuToggle: () => void; -} - -function BookmarkItem({ - bookmark, - active, - menuOpen, - onNavigate, - onEdit, - onDelete, - onMenuToggle, -}: BookmarkItemProps) { - const formatDate = (dateStr: string) => { - const date = new Date(dateStr); - return date.toLocaleDateString("zh-CN", { - month: "short", - day: "numeric", - hour: "2-digit", - minute: "2-digit", - }); - }; - - return ( -
-
-
- - - {bookmark.name || `Flow ${bookmark.flow_id.slice(0, 8)}...`} - -
-
- - {bookmark.flow_id.slice(0, 12)}... - - - {formatDate(bookmark.created_at)} - -
-
- - {/* 操作按钮 */} -
- - - {/* 下拉菜单 */} - {menuOpen && ( -
e.stopPropagation()} - > - - - -
- )} -
-
- ); -} - -// ============================================================================ -// 编辑书签对话框 -// ============================================================================ - -interface EditBookmarkDialogProps { - bookmark: FlowBookmark; - name: string; - group: string; - groups: string[]; - saving: boolean; - onNameChange: (name: string) => void; - onGroupChange: (group: string) => void; - onSave: () => void; - onClose: () => void; -} - -function EditBookmarkDialog({ - bookmark, - name, - group, - groups, - saving, - onNameChange, - onGroupChange, - onSave, - onClose, -}: EditBookmarkDialogProps) { - return ( -
-
-
- {/* 头部 */} -
-
- -

编辑书签

-
- -
- - {/* 内容 */} -
-
- - onNameChange(e.target.value)} - placeholder="输入书签名称(可选)" - className="w-full rounded-lg border bg-background px-3 py-2 text-sm" - autoFocus - /> -
-
- - onGroupChange(e.target.value)} - placeholder="输入或选择分组(可选)" - list="bookmark-groups" - className="w-full rounded-lg border bg-background px-3 py-2 text-sm" - /> - - {groups.map((g) => ( - -
-
-

书签 ID: {bookmark.id.slice(0, 8)}...

-

Flow ID: {bookmark.flow_id.slice(0, 12)}...

-

- 创建时间: {new Date(bookmark.created_at).toLocaleString("zh-CN")} -

-
-
- - {/* 底部 */} -
- - -
-
-
- ); -} - -export default BookmarkPanel; diff --git a/src/components/flow-monitor/CleanupDialog.tsx b/src/components/flow-monitor/CleanupDialog.tsx deleted file mode 100644 index 4ffe64e09..000000000 --- a/src/components/flow-monitor/CleanupDialog.tsx +++ /dev/null @@ -1,592 +0,0 @@ -/** - * 清理日志对话框 - * 提供多种清理选项:删除所有、按时间、按数量、按状态、按Provider、按大小 - */ - -import { useState } from "react"; -import { - Trash2, - AlertTriangle, - Clock, - Hash, - Activity, - Server, - HardDrive, -} from "lucide-react"; -import { Modal } from "@/components/Modal"; -import { - flowMonitorApi, - type CleanupFlowsRequest, - type CleanupType, - type ProviderType, -} from "@/lib/api/flowMonitor"; - -interface CleanupDialogProps { - isOpen: boolean; - onClose: () => void; - onSuccess: () => void; -} - -export function CleanupDialog({ - isOpen, - onClose, - onSuccess, -}: CleanupDialogProps) { - const [cleanupType, setCleanupType] = useState("ByTime"); - const [loading, setLoading] = useState(false); - const [error, setError] = useState(null); - const [showSuccess, setShowSuccess] = useState(false); - const [cleanupResult, setCleanupResult] = useState<{ - cleaned_count: number; - cleaned_files: number; - freed_bytes: number; - } | null>(null); - - // 时间清理选项 - const [retentionDays, setRetentionDays] = useState(7); - const [retentionHours, setRetentionHours] = useState(24); - const [useHours, setUseHours] = useState(false); - - // 数量清理选项 - const [maxRecords, setMaxRecords] = useState(1000); - - // 状态清理选项 - const [targetStates, setTargetStates] = useState(["Failed"]); - - // Provider清理选项 - const [targetProviders, setTargetProviders] = useState([]); - - // 大小清理选项 - const [maxStorageGB, setMaxStorageGB] = useState(1); - - const handleCleanup = async () => { - setLoading(true); - setError(null); - - try { - const request: CleanupFlowsRequest = { - cleanup_type: cleanupType, - }; - - // 根据清理类型设置相应参数 - switch (cleanupType) { - case "All": - // 删除所有,无需额外参数 - break; - - case "ByTime": - if (useHours) { - request.retention_hours = retentionHours; - } else { - request.retention_days = retentionDays; - } - break; - - case "ByCount": - request.max_records = maxRecords; - break; - - case "ByStatus": - if (targetStates.length === 0) { - setError("请至少选择一个状态"); - setLoading(false); - return; - } - request.target_states = targetStates; - break; - - case "ByProvider": - if (targetProviders.length === 0) { - setError("请至少选择一个Provider"); - setLoading(false); - return; - } - request.target_providers = targetProviders; - break; - - case "BySize": - request.max_storage_bytes = maxStorageGB * 1024 * 1024 * 1024; // GB转字节 - break; - } - - const result = await flowMonitorApi.cleanupFlows(request); - - // 保存清理结果并显示成功提示 - setCleanupResult(result); - setShowSuccess(true); - } catch (e) { - setError(e instanceof Error ? e.message : String(e)); - } finally { - setLoading(false); - } - }; - - const handleSuccessClose = () => { - setShowSuccess(false); - setCleanupResult(null); - onSuccess(); - onClose(); - }; - - const formatBytes = (bytes: number): string => { - if (bytes === 0) return "0 B"; - const k = 1024; - const sizes = ["B", "KB", "MB", "GB", "TB"]; - const i = Math.floor(Math.log(bytes) / Math.log(k)); - return parseFloat((bytes / Math.pow(k, i)).toFixed(2)) + " " + sizes[i]; - }; - - const availableStates = [ - "Pending", - "Streaming", - "Completed", - "Failed", - "Cancelled", - ]; - const availableProviders: ProviderType[] = [ - "Kiro", - "Gemini", - "Qwen", - "Antigravity", - "OpenAI", - "Claude", - "Vertex", - "GeminiApiKey", - "Codex", - "ClaudeOAuth", - "IFlow", - ]; - - // 如果显示成功提示,渲染成功对话框 - if (showSuccess && cleanupResult) { - return ( - -
- {/* 成功图标 */} -
-
-
- - - -
- {/* 动画圆环 */} -
-
-
- - {/* 标题 */} -

清理完成!

-

- 已成功清理日志数据 -

- - {/* 清理统计 */} -
-
-
-
- -
- 删除记录 -
- - {cleanupResult.cleaned_count} 条 - -
- -
-
-
- - - -
- 清理文件 -
- - {cleanupResult.cleaned_files} 个 - -
- -
-
-
- -
- - 释放空间 - -
- - {formatBytes(cleanupResult.freed_bytes)} - -
-
- - {/* 确认按钮 */} - -
-
- ); - } - - return ( - - {/* 标题 */} -
-

清理日志

-
- -
- {/* 警告提示 */} -
- -
-

- 注意:清理操作不可撤销 -

-

- 请确认清理条件,删除的日志数据无法恢复。 -

-
-
- - {/* 清理类型选择 */} -
- -
- - - - - - - - - - - -
-
- - {/* 清理选项配置 */} -
- {cleanupType === "ByTime" && ( -
-
- - -
- - {useHours ? ( -
- - setRetentionHours(Number(e.target.value))} - className="w-full" - /> -
- 1小时 - 7天 -
-
- ) : ( -
- - setRetentionDays(Number(e.target.value))} - className="w-full" - /> -
- 1天 - 1年 -
-
- )} -
- )} - - {cleanupType === "ByCount" && ( -
- - setMaxRecords(Number(e.target.value))} - className="w-full" - /> -
- 100条 - 10,000条 -
-
- )} - - {cleanupType === "ByStatus" && ( -
- -
- {availableStates.map((state) => ( - - ))} -
-
- )} - - {cleanupType === "ByProvider" && ( -
- -
- {availableProviders.map((provider) => ( - - ))} -
-
- )} - - {cleanupType === "BySize" && ( -
- - setMaxStorageGB(Number(e.target.value))} - className="w-full" - /> -
- 100MB - 100GB -
-

- 当存储超过此大小时,将删除最旧的数据 -

-
- )} -
- - {/* 错误显示 */} - {error && ( -
- {error} -
- )} - - {/* 操作按钮 */} -
- - -
-
-
- ); -} diff --git a/src/components/flow-monitor/ExportDialog.tsx b/src/components/flow-monitor/ExportDialog.tsx deleted file mode 100644 index 1ed61988e..000000000 --- a/src/components/flow-monitor/ExportDialog.tsx +++ /dev/null @@ -1,484 +0,0 @@ -import React, { useState, useCallback } from "react"; -import { - X, - Download, - FileJson, - FileText, - FileSpreadsheet, - FileCode, - Loader2, - Check, - AlertCircle, - Shield, - Settings, - ChevronDown, - ChevronUp, -} from "lucide-react"; -import { - flowMonitorApi, - type ExportFormat, - type ExportOptions, - type FlowFilter, - type RedactionRule, -} from "@/lib/api/flowMonitor"; -import { cn } from "@/lib/utils"; - -interface ExportDialogProps { - /** 是否显示对话框 */ - open: boolean; - /** 关闭对话框回调 */ - onClose: () => void; - /** 要导出的 Flow ID 列表(批量导出) */ - flowIds?: string[]; - /** 过滤条件(按条件导出) */ - filter?: FlowFilter; - /** 导出成功回调 */ - onExportSuccess?: (filename: string) => void; -} - -interface FormatOption { - value: ExportFormat; - label: string; - description: string; - icon: React.ReactNode; -} - -const FORMAT_OPTIONS: FormatOption[] = [ - { - value: "json", - label: "JSON", - description: "完整的 JSON 格式,适合程序处理", - icon: , - }, - { - value: "jsonl", - label: "JSONL", - description: "每行一个 JSON 对象,适合大数据处理", - icon: , - }, - { - value: "har", - label: "HAR", - description: "HTTP Archive 格式,可在浏览器开发工具中查看", - icon: , - }, - { - value: "markdown", - label: "Markdown", - description: "可读性强的文档格式,适合分享和文档", - icon: , - }, - { - value: "csv", - label: "CSV", - description: "表格格式,仅包含元数据,适合 Excel 分析", - icon: , - }, -]; - -const DEFAULT_REDACTION_RULES: RedactionRule[] = [ - { - name: "API 密钥", - pattern: - "(sk-[a-zA-Z0-9]{20,}|api[_-]?key[\"']?\\s*[:=]\\s*[\"']?[a-zA-Z0-9-_]{20,})", - replacement: "[REDACTED_API_KEY]", - enabled: true, - }, - { - name: "邮箱地址", - pattern: "[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\\.[a-zA-Z]{2,}", - replacement: "[REDACTED_EMAIL]", - enabled: true, - }, - { - name: "手机号码", - pattern: "1[3-9]\\d{9}", - replacement: "[REDACTED_PHONE]", - enabled: true, - }, - { - name: "Bearer Token", - pattern: "Bearer\\s+[a-zA-Z0-9._-]+", - replacement: "Bearer [REDACTED_TOKEN]", - enabled: true, - }, -]; - -export function ExportDialog({ - open, - onClose, - flowIds, - filter, - onExportSuccess, -}: ExportDialogProps) { - const [format, setFormat] = useState("json"); - const [includeRaw, setIncludeRaw] = useState(true); - const [includeStreamChunks, setIncludeStreamChunks] = useState(false); - const [redactSensitive, setRedactSensitive] = useState(false); - const [redactionRules, setRedactionRules] = useState( - DEFAULT_REDACTION_RULES, - ); - const [showAdvanced, setShowAdvanced] = useState(false); - const [exporting, setExporting] = useState(false); - const [error, setError] = useState(null); - const [success, setSuccess] = useState(false); - - const exportCount = flowIds?.length || 0; - const isFilterExport = !flowIds || flowIds.length === 0; - - const handleExport = useCallback(async () => { - setExporting(true); - setError(null); - setSuccess(false); - - try { - const options: ExportOptions = { - format, - include_raw: includeRaw, - include_stream_chunks: includeStreamChunks, - redact_sensitive: redactSensitive, - redaction_rules: redactSensitive - ? redactionRules.filter((r) => r.enabled) - : undefined, - }; - - let result; - if (flowIds && flowIds.length > 0) { - // 批量导出指定 ID - result = await flowMonitorApi.exportFlowsByIds(flowIds, options); - } else { - // 按过滤条件导出 - result = await flowMonitorApi.exportFlows({ - ...options, - filter: filter || {}, - }); - } - - // 下载文件 - downloadFile(result.data, result.filename, result.mime_type); - setSuccess(true); - onExportSuccess?.(result.filename); - - // 延迟关闭对话框 - setTimeout(() => { - onClose(); - setSuccess(false); - }, 1500); - } catch (e) { - console.error("Export failed:", e); - setError(e instanceof Error ? e.message : "导出失败"); - } finally { - setExporting(false); - } - }, [ - format, - includeRaw, - includeStreamChunks, - redactSensitive, - redactionRules, - flowIds, - filter, - onClose, - onExportSuccess, - ]); - - const toggleRedactionRule = (index: number) => { - setRedactionRules((prev) => - prev.map((rule, i) => - i === index ? { ...rule, enabled: !rule.enabled } : rule, - ), - ); - }; - - if (!open) return null; - - return ( -
- {/* 背景遮罩 */} -
- - {/* 对话框 */} -
- {/* 头部 */} -
-
- -

导出 Flow

-
- -
- - {/* 内容 */} -
- {/* 导出数量提示 */} -
-
- {isFilterExport ? ( - 将导出符合当前过滤条件的所有 Flow - ) : ( - - 已选择 {exportCount} 个 Flow - - )} -
-
- - {/* 格式选择 */} -
- -
- {FORMAT_OPTIONS.map((option) => ( - setFormat(option.value)} - /> - ))} -
-
- - {/* 基本选项 */} -
- -
- - -
-
- - {/* 隐私选项 */} -
-
- - -
- - - {/* 脱敏规则 */} - {redactSensitive && ( -
-
- 脱敏规则: -
- {redactionRules.map((rule, index) => ( - - ))} -
- )} -
- - {/* 高级选项 */} -
- - - {showAdvanced && ( -
-
-

• JSON/JSONL 格式适合程序处理和数据分析

-

• HAR 格式可在 Chrome DevTools 中导入查看

-

• Markdown 格式适合生成文档和分享

-

• CSV 格式仅包含元数据,不含消息内容

-
-
- )} -
- - {/* 错误提示 */} - {error && ( -
-
- - {error} -
-
- )} - - {/* 成功提示 */} - {success && ( -
-
- - 导出成功! -
-
- )} -
- - {/* 底部按钮 */} -
- - -
-
-
- ); -} - -// ============================================================================ -// 子组件 -// ============================================================================ - -interface FormatCardProps { - option: FormatOption; - selected: boolean; - onClick: () => void; -} - -function FormatCard({ option, selected, onClick }: FormatCardProps) { - return ( - - ); -} - -interface OptionCheckboxProps { - checked: boolean; - onChange: (checked: boolean) => void; - label: string; - description?: string; -} - -function OptionCheckbox({ - checked, - onChange, - label, - description, -}: OptionCheckboxProps) { - return ( - - ); -} - -// ============================================================================ -// 辅助函数 -// ============================================================================ - -/** - * 下载文件 - */ -function downloadFile(data: string, filename: string, mimeType: string) { - const blob = new Blob([data], { type: mimeType }); - const url = URL.createObjectURL(blob); - const a = document.createElement("a"); - a.href = url; - a.download = filename; - document.body.appendChild(a); - a.click(); - document.body.removeChild(a); - URL.revokeObjectURL(url); -} - -export default ExportDialog; diff --git a/src/components/flow-monitor/FilterExpressionInput.tsx b/src/components/flow-monitor/FilterExpressionInput.tsx deleted file mode 100644 index 83eee3592..000000000 --- a/src/components/flow-monitor/FilterExpressionInput.tsx +++ /dev/null @@ -1,552 +0,0 @@ -import React, { useState, useEffect, useRef, useCallback } from "react"; -import { safeInvoke } from "@/lib/dev-bridge"; -import { Search, AlertCircle, CheckCircle2, HelpCircle, X } from "lucide-react"; -import { cn } from "@/lib/utils"; - -// ============================================================================ -// 类型定义 -// ============================================================================ - -/** - * 过滤表达式解析结果 - */ -interface ParseFilterResult { - valid: boolean; - error: string | null; - expr: unknown | null; -} - -/** - * 过滤表达式帮助项 - */ -export interface FilterHelpItem { - syntax: string; - description: string; -} - -/** - * 自动补全建议 - */ -interface AutocompleteSuggestion { - text: string; - description: string; - type: "filter" | "operator" | "value"; -} - -interface FilterExpressionInputProps { - value: string; - onChange: (value: string) => void; - onSubmit: (expression: string) => void; - onValidationChange?: (valid: boolean, error: string | null) => void; - placeholder?: string; - className?: string; - showHelp?: boolean; - onHelpToggle?: () => void; -} - -// ============================================================================ -// 过滤器定义(用于语法高亮和自动补全) -// ============================================================================ - -const FILTER_KEYWORDS = [ - { prefix: "~m", name: "model", hasArg: true, description: "模型名称匹配" }, - { prefix: "~p", name: "provider", hasArg: true, description: "提供商匹配" }, - { prefix: "~s", name: "state", hasArg: true, description: "状态匹配" }, - { prefix: "~e", name: "error", hasArg: false, description: "有错误" }, - { prefix: "~t", name: "toolcalls", hasArg: false, description: "有工具调用" }, - { prefix: "~k", name: "thinking", hasArg: false, description: "有思维链" }, - { prefix: "~starred", name: "starred", hasArg: false, description: "已收藏" }, - { prefix: "~tag", name: "tag", hasArg: true, description: "包含标签" }, - { prefix: "~b", name: "body", hasArg: true, description: "内容匹配" }, - { - prefix: "~bq", - name: "bodyrequest", - hasArg: true, - description: "请求内容匹配", - }, - { - prefix: "~bs", - name: "bodyresponse", - hasArg: true, - description: "响应内容匹配", - }, - { - prefix: "~tokens", - name: "tokens", - hasArg: true, - description: "Token 数量比较", - }, - { - prefix: "~latency", - name: "latency", - hasArg: true, - description: "延迟比较", - }, -]; - -const OPERATORS = [ - { symbol: "&", description: "AND 逻辑" }, - { symbol: "|", description: "OR 逻辑" }, - { symbol: "!", description: "NOT 逻辑" }, - { symbol: "(", description: "左括号" }, - { symbol: ")", description: "右括号" }, -]; - -const STATE_VALUES = [ - "pending", - "streaming", - "completed", - "failed", - "cancelled", -]; - -const COMPARISON_OPS = [">", ">=", "<", "<=", "="]; - -// ============================================================================ -// 组件实现 -// ============================================================================ - -export function FilterExpressionInput({ - value, - onChange, - onSubmit, - onValidationChange, - placeholder = "输入过滤表达式,如 ~m claude & ~p kiro", - className, - showHelp = false, - onHelpToggle, -}: FilterExpressionInputProps) { - const [isValid, setIsValid] = useState(null); - const [error, setError] = useState(null); - const [suggestions, setSuggestions] = useState([]); - const [showSuggestions, setShowSuggestions] = useState(false); - const [selectedSuggestionIndex, setSelectedSuggestionIndex] = useState(0); - const validationTimeoutRef = useRef | null>( - null, - ); - - const inputRef = useRef(null); - const suggestionsRef = useRef(null); - - // 验证表达式 - const validateExpression = useCallback( - async (expr: string) => { - if (!expr.trim()) { - setIsValid(null); - setError(null); - onValidationChange?.(true, null); - return; - } - - try { - const result = await safeInvoke("parse_filter", { - expression: expr, - }); - - setIsValid(result.valid); - setError(result.error); - onValidationChange?.(result.valid, result.error); - } catch (e) { - setIsValid(false); - const errorMsg = e instanceof Error ? e.message : "验证失败"; - setError(errorMsg); - onValidationChange?.(false, errorMsg); - } - }, - [onValidationChange], - ); - - // 防抖验证 - useEffect(() => { - if (validationTimeoutRef.current) { - clearTimeout(validationTimeoutRef.current); - } - - const timeout = setTimeout(() => { - validateExpression(value); - }, 300); - - validationTimeoutRef.current = timeout; - - return () => { - if (timeout) { - clearTimeout(timeout); - } - }; - }, [value, validateExpression]); - - // 生成自动补全建议 - const generateSuggestions = useCallback( - (input: string, cursorPos: number) => { - const textBeforeCursor = input.slice(0, cursorPos); - const lastToken = textBeforeCursor.split(/[\s&|!()]+/).pop() || ""; - - const newSuggestions: AutocompleteSuggestion[] = []; - - // 如果以 ~ 开头,建议过滤器 - if (lastToken.startsWith("~")) { - const filterPrefix = lastToken.toLowerCase(); - FILTER_KEYWORDS.forEach((filter) => { - if (filter.prefix.toLowerCase().startsWith(filterPrefix)) { - newSuggestions.push({ - text: filter.prefix, - description: filter.description, - type: "filter", - }); - } - }); - } - // 如果刚输入了 ~s,建议状态值 - else if (/~s\s*$/.test(textBeforeCursor)) { - STATE_VALUES.forEach((state) => { - newSuggestions.push({ - text: state, - description: `状态: ${state}`, - type: "value", - }); - }); - } - // 如果刚输入了 ~tokens 或 ~latency,建议比较运算符 - else if (/~(tokens|latency)\s*$/.test(textBeforeCursor)) { - COMPARISON_OPS.forEach((op) => { - newSuggestions.push({ - text: op, - description: `比较运算符: ${op}`, - type: "operator", - }); - }); - } - // 如果输入为空或刚输入了运算符,建议过滤器 - else if (!lastToken || /[&|!()]$/.test(textBeforeCursor.trim())) { - FILTER_KEYWORDS.slice(0, 6).forEach((filter) => { - newSuggestions.push({ - text: filter.prefix, - description: filter.description, - type: "filter", - }); - }); - } - // 如果刚输入了过滤器值,建议逻辑运算符 - else if (lastToken && !lastToken.startsWith("~")) { - OPERATORS.slice(0, 3).forEach((op) => { - newSuggestions.push({ - text: op.symbol, - description: op.description, - type: "operator", - }); - }); - } - - setSuggestions(newSuggestions); - setShowSuggestions(newSuggestions.length > 0); - setSelectedSuggestionIndex(0); - }, - [], - ); - - // 处理输入变化 - const handleInputChange = (e: React.ChangeEvent) => { - const newValue = e.target.value; - onChange(newValue); - generateSuggestions(newValue, e.target.selectionStart || newValue.length); - }; - - // 处理键盘事件 - const handleKeyDown = (e: React.KeyboardEvent) => { - if (showSuggestions && suggestions.length > 0) { - switch (e.key) { - case "ArrowDown": - e.preventDefault(); - setSelectedSuggestionIndex((prev) => - prev < suggestions.length - 1 ? prev + 1 : 0, - ); - break; - case "ArrowUp": - e.preventDefault(); - setSelectedSuggestionIndex((prev) => - prev > 0 ? prev - 1 : suggestions.length - 1, - ); - break; - case "Tab": - case "Enter": - if (showSuggestions && suggestions[selectedSuggestionIndex]) { - e.preventDefault(); - applySuggestion(suggestions[selectedSuggestionIndex]); - } else if (e.key === "Enter" && isValid !== false) { - e.preventDefault(); - onSubmit(value); - } - break; - case "Escape": - e.preventDefault(); - setShowSuggestions(false); - break; - } - } else if (e.key === "Enter" && isValid !== false) { - e.preventDefault(); - onSubmit(value); - } - }; - - // 应用建议 - const applySuggestion = (suggestion: AutocompleteSuggestion) => { - const input = inputRef.current; - if (!input) return; - - const cursorPos = input.selectionStart || value.length; - const textBeforeCursor = value.slice(0, cursorPos); - const textAfterCursor = value.slice(cursorPos); - - // 找到最后一个 token 的开始位置 - const lastTokenMatch = textBeforeCursor.match(/[~\w\-.*]+$/); - const lastTokenStart = lastTokenMatch - ? cursorPos - lastTokenMatch[0].length - : cursorPos; - - // 构建新值 - const newValue = - value.slice(0, lastTokenStart) + - suggestion.text + - (suggestion.type === "filter" && - FILTER_KEYWORDS.find((f) => f.prefix === suggestion.text)?.hasArg - ? " " - : "") + - textAfterCursor; - - onChange(newValue); - setShowSuggestions(false); - - // 设置光标位置 - setTimeout(() => { - const newCursorPos = - lastTokenStart + - suggestion.text.length + - (suggestion.type === "filter" && - FILTER_KEYWORDS.find((f) => f.prefix === suggestion.text)?.hasArg - ? 1 - : 0); - input.setSelectionRange(newCursorPos, newCursorPos); - input.focus(); - }, 0); - }; - - // 点击外部关闭建议 - useEffect(() => { - const handleClickOutside = (e: MouseEvent) => { - if ( - suggestionsRef.current && - !suggestionsRef.current.contains(e.target as Node) && - inputRef.current && - !inputRef.current.contains(e.target as Node) - ) { - setShowSuggestions(false); - } - }; - - document.addEventListener("mousedown", handleClickOutside); - return () => document.removeEventListener("mousedown", handleClickOutside); - }, []); - - // 渲染语法高亮的文本 - const renderHighlightedText = () => { - if (!value) return null; - - const parts: React.ReactNode[] = []; - let remaining = value; - let key = 0; - - while (remaining.length > 0) { - let matched = false; - - // 匹配过滤器 - for (const filter of FILTER_KEYWORDS) { - const regex = new RegExp( - `^(${filter.prefix.replace(/[.*+?^${}()|[\]\\]/g, "\\$&")})(?:\\s+([^&|!()]+))?`, - "i", - ); - const match = remaining.match(regex); - if (match) { - parts.push( - - {match[1]} - , - ); - if (match[2]) { - parts.push( - - {" "} - {match[2]} - , - ); - } - remaining = remaining.slice(match[0].length); - matched = true; - break; - } - } - - if (!matched) { - // 匹配运算符 - const opMatch = remaining.match(/^([&|!()])/); - if (opMatch) { - parts.push( - - {opMatch[1]} - , - ); - remaining = remaining.slice(1); - matched = true; - } - } - - if (!matched) { - // 匹配空白 - const wsMatch = remaining.match(/^(\s+)/); - if (wsMatch) { - parts.push({wsMatch[1]}); - remaining = remaining.slice(wsMatch[1].length); - matched = true; - } - } - - if (!matched) { - // 其他字符 - parts.push({remaining[0]}); - remaining = remaining.slice(1); - } - } - - return parts; - }; - - return ( -
- {/* 输入框容器 */} -
- {/* 语法高亮层 */} - - - {/* 实际输入框 */} -
- - generateSuggestions(value, value.length)} - placeholder={placeholder} - className={cn( - "w-full rounded-lg border bg-transparent pl-9 pr-20 py-2 text-sm text-transparent caret-foreground", - "focus:outline-none focus:ring-2 focus:ring-primary", - isValid === false && "border-red-500 focus:ring-red-500", - isValid === true && - value && - "border-green-500 focus:ring-green-500", - )} - spellCheck={false} - autoComplete="off" - /> - - {/* 右侧图标 */} -
- {/* 验证状态图标 */} - {value && isValid === true && ( - - )} - {value && isValid === false && ( - - )} - - {/* 清除按钮 */} - {value && ( - - )} - - {/* 帮助按钮 */} - {onHelpToggle && ( - - )} -
-
-
- - {/* 错误提示 */} - {error && ( -
- - {error} -
- )} - - {/* 自动补全建议 */} - {showSuggestions && suggestions.length > 0 && ( -
- {suggestions.map((suggestion, index) => ( - - ))} -
- )} -
- ); -} - -export default FilterExpressionInput; diff --git a/src/components/flow-monitor/FilterHelp.tsx b/src/components/flow-monitor/FilterHelp.tsx deleted file mode 100644 index 02cb5e0b4..000000000 --- a/src/components/flow-monitor/FilterHelp.tsx +++ /dev/null @@ -1,440 +0,0 @@ -import React, { useState } from "react"; -import { - HelpCircle, - X, - Search, - Filter, - Zap, - Hash, - Tag, - ChevronDown, - ChevronRight, - Copy, - Check, - FileText, -} from "lucide-react"; -import { cn } from "@/lib/utils"; - -interface FilterHelpProps { - onClose?: () => void; - onInsertExample?: (example: string) => void; - className?: string; -} - -// ============================================================================ -// 过滤器分类 -// ============================================================================ - -interface FilterCategory { - name: string; - icon: React.ReactNode; - description: string; - filters: FilterInfo[]; -} - -interface FilterInfo { - syntax: string; - description: string; - example?: string; - hasArg: boolean; -} - -const FILTER_CATEGORIES: FilterCategory[] = [ - { - name: "基础过滤器", - icon: , - description: "按模型、提供商、状态等基本属性过滤", - filters: [ - { - syntax: "~m ", - description: "模型名称匹配(支持 * 通配符)", - example: "~m claude*", - hasArg: true, - }, - { - syntax: "~p ", - description: "提供商匹配", - example: "~p kiro", - hasArg: true, - }, - { - syntax: "~s ", - description: "状态匹配 (pending/streaming/completed/failed/cancelled)", - example: "~s completed", - hasArg: true, - }, - ], - }, - { - name: "特性过滤器", - icon: , - description: "按 Flow 特性过滤", - filters: [ - { syntax: "~e", description: "有错误", example: "~e", hasArg: false }, - { syntax: "~t", description: "有工具调用", example: "~t", hasArg: false }, - { syntax: "~k", description: "有思维链", example: "~k", hasArg: false }, - { - syntax: "~starred", - description: "已收藏", - example: "~starred", - hasArg: false, - }, - ], - }, - { - name: "标签过滤器", - icon: , - description: "按标签过滤", - filters: [ - { - syntax: "~tag ", - description: "包含指定标签", - example: "~tag important", - hasArg: true, - }, - ], - }, - { - name: "内容搜索", - icon: , - description: "搜索请求或响应内容", - filters: [ - { - syntax: "~b ", - description: "请求或响应内容匹配(正则表达式)", - example: '~b "hello"', - hasArg: true, - }, - { - syntax: "~bq ", - description: "仅请求内容匹配", - example: "~bq user", - hasArg: true, - }, - { - syntax: "~bs ", - description: "仅响应内容匹配", - example: "~bs assistant", - hasArg: true, - }, - ], - }, - { - name: "数值比较", - icon: , - description: "按 Token 数量或延迟过滤", - filters: [ - { - syntax: "~tokens ", - description: "Token 数量比较 (>, >=, <, <=, =)", - example: "~tokens >1000", - hasArg: true, - }, - { - syntax: "~latency ", - description: "延迟比较(支持 s/ms 后缀)", - example: "~latency >5s", - hasArg: true, - }, - ], - }, -]; - -const OPERATORS = [ - { - symbol: "&", - description: "AND 逻辑 - 同时满足两个条件", - example: "~p kiro & ~m claude", - }, - { - symbol: "|", - description: "OR 逻辑 - 满足任一条件", - example: "~p kiro | ~p gemini", - }, - { symbol: "!", description: "NOT 逻辑 - 取反", example: "!~e" }, - { - symbol: "()", - description: "分组 - 控制优先级", - example: "(~p kiro | ~p gemini) & ~m claude", - }, -]; - -const EXAMPLES = [ - { name: "Claude 模型", expr: "~m claude" }, - { name: "Kiro 提供商的 Claude 模型", expr: "~p kiro & ~m claude" }, - { name: "有错误或高延迟", expr: "~e | ~latency >5s" }, - { name: "没有错误", expr: "!~e" }, - { name: "大 Token 请求", expr: "~tokens >10000" }, - { name: "有工具调用的已完成请求", expr: "~t & ~s completed" }, - { name: "已收藏的有思维链请求", expr: "~starred & ~k" }, - { - name: "多提供商的高 Token 请求", - expr: "(~p kiro | ~p gemini) & ~tokens >1000", - }, -]; - -// ============================================================================ -// 组件实现 -// ============================================================================ - -export function FilterHelp({ - onClose, - onInsertExample, - className, -}: FilterHelpProps) { - const [expandedCategories, setExpandedCategories] = useState>( - new Set(FILTER_CATEGORIES.map((c) => c.name)), - ); - const [copiedExample, setCopiedExample] = useState(null); - const [searchQuery, setSearchQuery] = useState(""); - - const toggleCategory = (name: string) => { - setExpandedCategories((prev) => { - const next = new Set(prev); - if (next.has(name)) { - next.delete(name); - } else { - next.add(name); - } - return next; - }); - }; - - const handleCopyExample = async (example: string) => { - try { - await navigator.clipboard.writeText(example); - setCopiedExample(example); - setTimeout(() => setCopiedExample(null), 2000); - } catch (e) { - console.error("复制失败:", e); - } - }; - - const handleInsertExample = (example: string) => { - onInsertExample?.(example); - }; - - // 过滤搜索结果 - const filteredCategories = FILTER_CATEGORIES.map((category) => ({ - ...category, - filters: category.filters.filter( - (f) => - !searchQuery || - f.syntax.toLowerCase().includes(searchQuery.toLowerCase()) || - f.description.toLowerCase().includes(searchQuery.toLowerCase()), - ), - })).filter((c) => c.filters.length > 0); - - const filteredExamples = EXAMPLES.filter( - (e) => - !searchQuery || - e.name.toLowerCase().includes(searchQuery.toLowerCase()) || - e.expr.toLowerCase().includes(searchQuery.toLowerCase()), - ); - - return ( -
- {/* 头部 */} -
-
- - 过滤表达式帮助 -
- {onClose && ( - - )} -
- - {/* 搜索框 */} -
-
- - setSearchQuery(e.target.value)} - className="w-full rounded-lg border bg-background pl-9 pr-4 py-2 text-sm focus:outline-none focus:ring-2 focus:ring-primary" - /> -
-
- - {/* 内容区域 */} -
- {/* 过滤器分类 */} -
- {filteredCategories.map((category) => ( -
- - - {expandedCategories.has(category.name) && ( -
-

- {category.description} -

-
- {category.filters.map((filter) => ( - - ))} -
-
- )} -
- ))} -
- - {/* 逻辑运算符 */} -
-
-
- - 逻辑运算符 -
-
- {OPERATORS.map((op) => ( -
- - {op.symbol} - -
-

{op.description}

- -
-
- ))} -
-
-
- - {/* 示例 */} -
-
-
- - 常用示例 -
-
- {filteredExamples.map((example) => ( -
-
-

- {example.name} -

- - {example.expr} - -
-
- - {onInsertExample && ( - - )} -
-
- ))} -
-
-
-
-
- ); -} - -// ============================================================================ -// 子组件 -// ============================================================================ - -interface FilterItemProps { - filter: FilterInfo; - onCopy: (example: string) => void; - onInsert: (example: string) => void; - copied: boolean; -} - -function FilterItem({ filter, onCopy, onInsert, copied }: FilterItemProps) { - return ( -
- - {filter.syntax} - -
-

{filter.description}

- {filter.example && ( -
- - {filter.example} - - - -
- )} -
-
- ); -} - -export default FilterHelp; diff --git a/src/components/flow-monitor/FlowDetail.tsx b/src/components/flow-monitor/FlowDetail.tsx deleted file mode 100644 index 82d241895..000000000 --- a/src/components/flow-monitor/FlowDetail.tsx +++ /dev/null @@ -1,1283 +0,0 @@ -import React, { useState, useEffect } from "react"; -import { - ArrowLeft, - Copy, - Download, - Star, - StarOff, - Clock, - CheckCircle2, - XCircle, - Loader2, - ChevronDown, - ChevronRight, - Wrench, - Brain, - MessageSquare, - Tag, - AlertCircle, - FileJson, - Code, - User, - Bot, - Settings, - Zap, -} from "lucide-react"; -import { - flowMonitorApi, - type LLMFlow, - type Message, - type ToolCall, - type FlowState, - type ExportFormat, - formatFlowState, - formatFlowType, - formatErrorType, - formatLatency, - formatTokenCount, - formatBytes, - getMessageText, -} from "@/lib/api/flowMonitor"; -import { useFlowActions } from "@/hooks/useFlowActions"; -import { FlowTimeline } from "./FlowTimeline"; -import { cn } from "@/lib/utils"; - -interface FlowDetailProps { - flowId: string; - onBack?: () => void; - onExport?: (flowId: string, format: ExportFormat) => void; -} - -export function FlowDetail({ flowId, onBack, onExport }: FlowDetailProps) { - const [flow, setFlow] = useState(null); - const [loading, setLoading] = useState(true); - const [error, setError] = useState(null); - const [activeTab, setActiveTab] = useState< - "request" | "response" | "metadata" | "timeline" - >("request"); - const [expandedSections, setExpandedSections] = useState>( - new Set(["messages", "content", "toolCalls"]), - ); - // 代码模式:显示原始 JSON - const [codeMode, setCodeMode] = useState(false); - - // 使用 Flow 操作 Hook - const { copyText, copyFlowContent, exportFlow, exporting } = useFlowActions(); - - useEffect(() => { - const loadFlowDetail = async () => { - try { - setLoading(true); - setError(null); - const detail = await flowMonitorApi.getFlowDetail(flowId); - if (detail) { - setFlow(detail); - } else { - setError("Flow 不存在"); - } - } catch (e) { - console.error("Failed to load flow detail:", e); - setError(e instanceof Error ? e.message : "加载失败"); - } finally { - setLoading(false); - } - }; - loadFlowDetail(); - }, [flowId]); - - const handleToggleStar = async () => { - if (!flow) return; - try { - await flowMonitorApi.toggleFlowStar(flow.id); - setFlow({ - ...flow, - annotations: { - ...flow.annotations, - starred: !flow.annotations.starred, - }, - }); - } catch (e) { - console.error("Failed to toggle star:", e); - } - }; - - const handleCopyContent = async (content: string, _label?: string) => { - await copyText(content); - }; - - const handleExport = async (format: ExportFormat) => { - if (onExport) { - onExport(flowId, format); - } else { - await exportFlow(flowId, format); - } - }; - - const toggleSection = (section: string) => { - setExpandedSections((prev) => { - const next = new Set(prev); - if (next.has(section)) { - next.delete(section); - } else { - next.add(section); - } - return next; - }); - }; - - const getStateIcon = (state: FlowState) => { - switch (state) { - case "Completed": - return ; - case "Failed": - return ; - case "Streaming": - return ; - case "Pending": - return ; - case "Cancelled": - return ; - default: - return ; - } - }; - - if (loading) { - return ( -
- -
- ); - } - - if (error || !flow) { - return ( -
-
- - {error || "Flow 不存在"} -
- {onBack && ( - - )} -
- ); - } - - return ( -
- {/* 头部 */} - copyFlowContent(flow)} - getStateIcon={getStateIcon} - codeMode={codeMode} - onToggleCodeMode={() => setCodeMode(!codeMode)} - /> - - {/* 代码模式:显示原始 JSON */} - {codeMode ? ( -
-
- 原始 JSON - -
-
-            {JSON.stringify(flow, null, 2)}
-          
-
- ) : ( - <> - {/* 标签页 */} -
- setActiveTab("request")} - > - 请求 - - setActiveTab("response")} - > - 响应 - - setActiveTab("metadata")} - > - 元数据 - - setActiveTab("timeline")} - > - 时间线 - -
- - {/* 内容区域 */} -
- {activeTab === "request" && ( - - )} - {activeTab === "response" && ( - - )} - {activeTab === "metadata" && ( - - )} - {activeTab === "timeline" && } -
- - )} - - {/* 导出状态提示 */} - {exporting && ( -
- - 正在导出... -
- )} -
- ); -} - -// ============================================================================ -// 子组件 -// ============================================================================ - -interface FlowDetailHeaderProps { - flow: LLMFlow; - onBack?: () => void; - onToggleStar: () => void; - onExport: (format: ExportFormat) => void; - onCopyAll: () => void; - getStateIcon: (state: FlowState) => React.ReactNode; - codeMode: boolean; - onToggleCodeMode: () => void; -} - -function FlowDetailHeader({ - flow, - onBack, - onToggleStar, - onExport, - onCopyAll, - getStateIcon, - codeMode, - onToggleCodeMode, -}: FlowDetailHeaderProps) { - const [showExportMenu, setShowExportMenu] = useState(false); - - const formatTime = (timestamp: string) => { - return new Date(timestamp).toLocaleString("zh-CN"); - }; - - return ( -
- {/* 顶部操作栏 */} -
-
- {onBack && ( - - )} -
- {getStateIcon(flow.state)} - {formatFlowState(flow.state)} -
-
- -
- {/* 代码模式切换 */} - - - -
- - {showExportMenu && ( -
- {(["json", "markdown", "har"] as ExportFormat[]).map( - (format) => ( - - ), - )} -
- )} -
-
-
- - {/* 基本信息卡片 */} -
-
- - - - - - - - -
- - {/* 标签和标记 */} - {(flow.annotations.tags.length > 0 || - flow.annotations.marker || - flow.annotations.comment) && ( -
- {flow.annotations.marker && ( -
- {flow.annotations.marker} -
- )} - {flow.annotations.tags.length > 0 && ( -
- - {flow.annotations.tags.map((tag) => ( - - {tag} - - ))} -
- )} - {flow.annotations.comment && ( -
- - - {flow.annotations.comment} - -
- )} -
- )} -
- - {/* 错误信息 */} - {flow.error && ( -
-
- - 错误: {formatErrorType(flow.error.error_type)} -
-

- {flow.error.message} -

- {flow.error.status_code && ( -

- 状态码: {flow.error.status_code} -

- )} -
- )} -
- ); -} - -interface InfoItemProps { - label: string; - value: string; -} - -function InfoItem({ label, value }: InfoItemProps) { - return ( -
-
{label}
-
- {value} -
-
- ); -} - -interface TabButtonProps { - active: boolean; - onClick: () => void; - children: React.ReactNode; -} - -function TabButton({ active, onClick, children }: TabButtonProps) { - return ( - - ); -} - -// ============================================================================ -// 请求标签页 -// ============================================================================ - -interface RequestTabProps { - flow: LLMFlow; - expandedSections: Set; - toggleSection: (section: string) => void; - onCopy: (content: string, label: string) => void; -} - -function RequestTab({ - flow, - expandedSections, - toggleSection, - onCopy, -}: RequestTabProps) { - const { request } = flow; - - return ( -
- {/* 请求基本信息 */} - } - expanded={expandedSections.has("requestInfo")} - onToggle={() => toggleSection("requestInfo")} - > -
-
- - {request.method} - - - {request.path} - -
-
-
- 请求大小:{" "} - {formatBytes(request.size_bytes)} -
-
- 流式:{" "} - {request.parameters.stream ? "是" : "否"} -
- {request.parameters.temperature !== undefined && ( -
- Temperature:{" "} - {request.parameters.temperature} -
- )} - {request.parameters.max_tokens !== undefined && ( -
- Max Tokens:{" "} - {request.parameters.max_tokens} -
- )} -
-
-
- - {/* 系统提示词 */} - {request.system_prompt && ( - } - expanded={expandedSections.has("systemPrompt")} - onToggle={() => toggleSection("systemPrompt")} - onCopy={() => onCopy(request.system_prompt!, "系统提示词")} - > -
-            {request.system_prompt}
-          
-
- )} - - {/* 消息列表 */} - } - expanded={expandedSections.has("messages")} - onToggle={() => toggleSection("messages")} - > -
- {request.messages.map((message, index) => ( - onCopy(content, `消息 ${index + 1}`)} - /> - ))} -
-
- - {/* 工具定义 */} - {request.tools && request.tools.length > 0 && ( - } - expanded={expandedSections.has("tools")} - onToggle={() => toggleSection("tools")} - > -
- {request.tools.map((tool, index) => ( -
-
{tool.function.name}
- {tool.function.description && ( -
- {tool.function.description} -
- )} -
- ))} -
-
- )} - - {/* 请求头 */} - } - expanded={expandedSections.has("requestHeaders")} - onToggle={() => toggleSection("requestHeaders")} - > -
- {Object.entries(request.headers).map(([key, value]) => ( -
- - {key}: - - - {key.toLowerCase().includes("authorization") || - key.toLowerCase().includes("api-key") - ? "***" - : value} - -
- ))} -
-
- - {/* 原始请求体 */} - } - expanded={expandedSections.has("requestBody")} - onToggle={() => toggleSection("requestBody")} - onCopy={() => onCopy(JSON.stringify(request.body, null, 2), "请求体")} - > -
-          {JSON.stringify(request.body, null, 2)}
-        
-
-
- ); -} - -interface MessageItemProps { - message: Message; - onCopy: (content: string) => void; -} - -function MessageItem({ message, onCopy }: MessageItemProps) { - const [expanded, setExpanded] = useState(true); - const content = getMessageText(message.content); - - const getRoleIcon = (role: string) => { - switch (role) { - case "user": - return ; - case "assistant": - return ; - case "system": - return ; - case "tool": - case "function": - return ; - default: - return ; - } - }; - - const getRoleLabel = (role: string) => { - const labels: Record = { - user: "用户", - assistant: "助手", - system: "系统", - tool: "工具", - function: "函数", - }; - return labels[role] || role; - }; - - return ( -
-
setExpanded(!expanded)} - > - {expanded ? ( - - ) : ( - - )} - {getRoleIcon(message.role)} - - {getRoleLabel(message.role)} - - {message.name && ( - - ({message.name}) - - )} - -
- {expanded && ( -
-
-            {content}
-          
- {/* 工具调用 */} - {message.tool_calls && message.tool_calls.length > 0 && ( -
-
工具调用:
- {message.tool_calls.map((tc, i) => ( - - ))} -
- )} - {/* 工具结果 */} - {message.tool_result && ( -
-
工具结果:
-
-                {message.tool_result.content}
-              
-
- )} -
- )} -
- ); -} - -// ============================================================================ -// 响应标签页 -// ============================================================================ - -interface ResponseTabProps { - flow: LLMFlow; - expandedSections: Set; - toggleSection: (section: string) => void; - onCopy: (content: string, label: string) => void; -} - -function ResponseTab({ - flow, - expandedSections, - toggleSection, - onCopy, -}: ResponseTabProps) { - const { response } = flow; - - // 调试日志 - console.log("[FlowDetail] ResponseTab - response:", response); - console.log("[FlowDetail] ResponseTab - content:", response?.content); - console.log( - "[FlowDetail] ResponseTab - content length:", - response?.content?.length, - ); - console.log( - "[FlowDetail] ResponseTab - expandedSections:", - Array.from(expandedSections), - ); - - if (!response) { - return ( -
- 暂无响应数据 -
- ); - } - - return ( -
- {/* 响应基本信息 */} - } - expanded={expandedSections.has("responseInfo")} - onToggle={() => toggleSection("responseInfo")} - > -
-
- 状态码:{" "} - = 200 && response.status_code < 300 - ? "text-green-600" - : "text-red-600", - )} - > - {response.status_code} {response.status_text} - -
-
- 响应大小:{" "} - {formatBytes(response.size_bytes)} -
- {response.stop_reason && ( -
- 停止原因:{" "} - {typeof response.stop_reason === "string" - ? response.stop_reason - : response.stop_reason.other} -
- )} - {response.stream_info && ( - <> -
- Chunk 数:{" "} - {response.stream_info.chunk_count} -
-
- 首 Chunk 延迟:{" "} - {formatLatency(response.stream_info.first_chunk_latency_ms)} -
-
- 平均 Chunk 间隔:{" "} - {response.stream_info.avg_chunk_interval_ms.toFixed(1)}ms -
- - )} -
-
- - {/* Token 使用统计 */} - } - expanded={expandedSections.has("tokenUsage")} - onToggle={() => toggleSection("tokenUsage")} - > -
-
- 输入 Token:{" "} - {formatTokenCount(response.usage.input_tokens)} -
-
- 输出 Token:{" "} - {formatTokenCount(response.usage.output_tokens)} -
-
- 总 Token:{" "} - {formatTokenCount(response.usage.total_tokens)} -
- {response.usage.cache_read_tokens !== undefined && ( -
- 缓存读取:{" "} - {formatTokenCount(response.usage.cache_read_tokens)} -
- )} - {response.usage.cache_write_tokens !== undefined && ( -
- 缓存写入:{" "} - {formatTokenCount(response.usage.cache_write_tokens)} -
- )} - {response.usage.thinking_tokens !== undefined && ( -
- 思维链 Token:{" "} - {formatTokenCount(response.usage.thinking_tokens)} -
- )} -
-
- - {/* 响应内容 */} - } - expanded={expandedSections.has("content")} - onToggle={() => toggleSection("content")} - onCopy={() => onCopy(response.content, "响应内容")} - > -
-          {response.content || "(空)"}
-        
-
- - {/* 思维链内容 */} - {response.thinking && ( - } - expanded={expandedSections.has("thinking")} - onToggle={() => toggleSection("thinking")} - onCopy={() => onCopy(response.thinking!.text, "思维链")} - > -
- {response.thinking.tokens && ( -
- Token 数: {formatTokenCount(response.thinking.tokens)} -
- )} -
-              {response.thinking.text}
-            
-
-
- )} - - {/* 工具调用 */} - {response.tool_calls.length > 0 && ( - } - expanded={expandedSections.has("toolCalls")} - onToggle={() => toggleSection("toolCalls")} - > -
- {response.tool_calls.map((tc, index) => ( - - ))} -
-
- )} - - {/* 响应头 */} - } - expanded={expandedSections.has("responseHeaders")} - onToggle={() => toggleSection("responseHeaders")} - > -
- {Object.entries(response.headers).map(([key, value]) => ( -
- - {key}: - - {value} -
- ))} -
-
- - {/* 原始响应体 */} - } - expanded={expandedSections.has("responseBody")} - onToggle={() => toggleSection("responseBody")} - onCopy={() => onCopy(JSON.stringify(response.body, null, 2), "响应体")} - > -
-          {JSON.stringify(response.body, null, 2)}
-        
-
-
- ); -} - -interface ToolCallItemProps { - toolCall: ToolCall; -} - -function ToolCallItem({ toolCall }: ToolCallItemProps) { - const [expanded, setExpanded] = useState(false); - - let parsedArgs: unknown = null; - try { - parsedArgs = JSON.parse(toolCall.function.arguments); - } catch (_e) { - // 保持原始字符串 - } - - return ( -
-
setExpanded(!expanded)} - > - {expanded ? ( - - ) : ( - - )} - - {toolCall.function.name} - - ID: {toolCall.id.slice(0, 8)}... - -
- {expanded && ( -
-
-            {parsedArgs
-              ? JSON.stringify(parsedArgs, null, 2)
-              : toolCall.function.arguments}
-          
-
- )} -
- ); -} - -// ============================================================================ -// 元数据标签页 -// ============================================================================ - -interface MetadataTabProps { - flow: LLMFlow; - onCopy: (content: string, label: string) => void; -} - -function MetadataTab({ flow, onCopy }: MetadataTabProps) { - const { metadata, timestamps } = flow; - - const formatTime = (timestamp: string | undefined) => { - if (!timestamp) return "-"; - return new Date(timestamp).toLocaleString("zh-CN"); - }; - - return ( -
- {/* 时间戳 */} -
-

- - 时间戳 -

-
-
- 创建时间:{" "} - {formatTime(timestamps.created)} -
-
- 请求开始:{" "} - {formatTime(timestamps.request_start)} -
-
- 请求结束:{" "} - {formatTime(timestamps.request_end)} -
-
- 响应开始:{" "} - {formatTime(timestamps.response_start)} -
-
- 响应结束:{" "} - {formatTime(timestamps.response_end)} -
-
- 总耗时:{" "} - {formatLatency(timestamps.duration_ms)} -
- {timestamps.ttfb_ms && ( -
- TTFB:{" "} - {formatLatency(timestamps.ttfb_ms)} -
- )} -
-
- - {/* 提供商信息 */} -
-

- - 提供商信息 -

-
-
- 提供商:{" "} - {metadata.provider} -
- {metadata.credential_id && ( -
- 凭证 ID:{" "} - {metadata.credential_id.slice(0, 8)}... -
- )} - {metadata.credential_name && ( -
- 凭证名称:{" "} - {metadata.credential_name} -
- )} -
- 重试次数:{" "} - {metadata.retry_count} -
- {metadata.context_usage_percentage !== undefined && ( -
- 上下文使用率:{" "} - {(metadata.context_usage_percentage * 100).toFixed(1)}% -
- )} -
-
- - {/* 客户端信息 */} - {(metadata.client_info.ip || - metadata.client_info.user_agent || - metadata.client_info.request_id) && ( -
-

- - 客户端信息 -

-
- {metadata.client_info.ip && ( -
- IP:{" "} - {metadata.client_info.ip} -
- )} - {metadata.client_info.user_agent && ( -
- User-Agent:{" "} - - {metadata.client_info.user_agent} - -
- )} - {metadata.client_info.request_id && ( -
- Request ID:{" "} - {metadata.client_info.request_id} -
- )} -
-
- )} - - {/* 路由信息 */} - {(metadata.routing_info.target_url || - metadata.routing_info.route_rule || - metadata.routing_info.load_balance_strategy) && ( -
-

- - 路由信息 -

-
- {metadata.routing_info.target_url && ( -
- 目标 URL:{" "} - - {metadata.routing_info.target_url} - -
- )} - {metadata.routing_info.route_rule && ( -
- 路由规则:{" "} - {metadata.routing_info.route_rule} -
- )} - {metadata.routing_info.load_balance_strategy && ( -
- 负载均衡策略:{" "} - {metadata.routing_info.load_balance_strategy} -
- )} -
-
- )} - - {/* 注入参数 */} - {metadata.injected_params && - Object.keys(metadata.injected_params).length > 0 && ( -
-

- - 注入参数 -

-
-              {JSON.stringify(metadata.injected_params, null, 2)}
-            
-
- )} - - {/* Flow ID */} -
-

Flow ID

-
- - {flow.id} - - -
-
-
- ); -} - -// ============================================================================ -// 通用组件 -// ============================================================================ - -interface CollapsibleSectionProps { - title: string; - icon?: React.ReactNode; - expanded: boolean; - onToggle: () => void; - onCopy?: () => void; - children: React.ReactNode; -} - -function CollapsibleSection({ - title, - icon, - expanded, - onToggle, - onCopy, - children, -}: CollapsibleSectionProps) { - return ( -
-
- {expanded ? ( - - ) : ( - - )} - {icon} - {title} - {onCopy && ( - - )} -
- {expanded &&
{children}
} -
- ); -} - -export default FlowDetail; diff --git a/src/components/flow-monitor/FlowDiffView.tsx b/src/components/flow-monitor/FlowDiffView.tsx deleted file mode 100644 index 9a6bf03b2..000000000 --- a/src/components/flow-monitor/FlowDiffView.tsx +++ /dev/null @@ -1,1084 +0,0 @@ -/** - * Flow 差异对比视图组件 - * - * 实现差异对比视图和并排/统一视图切换 - * **Validates: Requirements 4.1-4.7** - */ - -import React, { useState, useEffect, useCallback } from "react"; -import { safeInvoke } from "@/lib/dev-bridge"; -import { - X, - Loader2, - AlertCircle, - ArrowLeftRight, - Columns, - Rows, - ChevronDown, - ChevronRight, - Plus, - Minus, - Edit3, - Settings, - MessageSquare, - Zap, - FileJson, -} from "lucide-react"; -import { cn } from "@/lib/utils"; -import type { LLMFlow } from "@/lib/api/flowMonitor"; - -// ============================================================================ -// 类型定义 -// ============================================================================ - -/** - * 差异类型 - */ -export type DiffType = "Added" | "Removed" | "Modified" | "Unchanged"; - -/** - * 差异项 - */ -export interface DiffItem { - path: string; - diff_type: DiffType; - left_value: unknown; - right_value: unknown; -} - -/** - * 消息差异项 - */ -export interface MessageDiffItem { - index: number; - diff_type: DiffType; - left_message: unknown; - right_message: unknown; - content_diffs: DiffItem[]; -} - -/** - * Token 差异 - */ -export interface TokenDiff { - input_diff: number; - output_diff: number; - total_diff: number; -} - -/** - * 差异配置 - */ -export interface DiffConfig { - ignore_fields: string[]; - ignore_timestamps: boolean; - ignore_ids: boolean; -} - -/** - * Flow 差异结果 - */ -export interface FlowDiffResult { - left_flow_id: string; - right_flow_id: string; - request_diffs: DiffItem[]; - response_diffs: DiffItem[]; - metadata_diffs: DiffItem[]; - message_diffs: MessageDiffItem[]; - token_diff: TokenDiff; -} - -/** - * 视图模式 - */ -export type ViewMode = "side-by-side" | "unified"; - -// ============================================================================ -// 组件属性 -// ============================================================================ - -interface FlowDiffViewProps { - /** 左侧 Flow ID */ - leftFlowId: string; - /** 右侧 Flow ID */ - rightFlowId: string; - /** 左侧 Flow(可选,如果提供则不需要加载) */ - leftFlow?: LLMFlow; - /** 右侧 Flow(可选,如果提供则不需要加载) */ - rightFlow?: LLMFlow; - /** 关闭回调 */ - onClose?: () => void; - /** 自定义类名 */ - className?: string; -} - -// ============================================================================ -// 主组件 -// ============================================================================ - -export function FlowDiffView({ - leftFlowId, - rightFlowId, - leftFlow: initialLeftFlow, - rightFlow: initialRightFlow, - onClose, - className, -}: FlowDiffViewProps) { - // 状态 - const [leftFlow, setLeftFlow] = useState( - initialLeftFlow || null, - ); - const [rightFlow, setRightFlow] = useState( - initialRightFlow || null, - ); - const [diffResult, setDiffResult] = useState(null); - const [loading, setLoading] = useState(true); - const [error, setError] = useState(null); - const [viewMode, setViewMode] = useState("side-by-side"); - const [config, setConfig] = useState({ - ignore_fields: [], - ignore_timestamps: true, - ignore_ids: true, - }); - const [showConfig, setShowConfig] = useState(false); - const [activeSection, setActiveSection] = useState("request"); - const [expandedPaths, setExpandedPaths] = useState>(new Set()); - - // 加载 Flow 和计算差异 - const loadDiff = useCallback(async () => { - try { - setLoading(true); - setError(null); - - // 调用后端计算差异 - const result = await safeInvoke("diff_flows", { - request: { - left_flow_id: leftFlowId, - right_flow_id: rightFlowId, - config, - }, - }); - - setDiffResult(result); - - // 如果没有提供 Flow,加载它们用于显示 - if (!initialLeftFlow) { - const left = await safeInvoke("get_flow_detail", { - flowId: leftFlowId, - }); - setLeftFlow(left); - } - if (!initialRightFlow) { - const right = await safeInvoke("get_flow_detail", { - flowId: rightFlowId, - }); - setRightFlow(right); - } - } catch (e) { - console.error("加载差异失败:", e); - setError(e instanceof Error ? e.message : "加载差异失败"); - } finally { - setLoading(false); - } - }, [leftFlowId, rightFlowId, config, initialLeftFlow, initialRightFlow]); - - useEffect(() => { - loadDiff(); - }, [loadDiff]); - - // 切换路径展开状态 - const togglePath = (path: string) => { - setExpandedPaths((prev) => { - const next = new Set(prev); - if (next.has(path)) { - next.delete(path); - } else { - next.add(path); - } - return next; - }); - }; - - if (loading) { - return ( -
- -
- ); - } - - if (error) { - return ( -
-
- - {error} -
- {onClose && ( - - )} -
- ); - } - - if (!diffResult) { - return null; - } - - return ( -
- {/* 头部 */} - setShowConfig(!showConfig)} - onClose={onClose} - /> - - {/* 配置面板 */} - {showConfig && } - - {/* Token 差异摘要 */} - - - {/* 标签页 */} -
- setActiveSection("request")} - count={ - diffResult.request_diffs.filter((d) => d.diff_type !== "Unchanged") - .length - } - > - 请求 - - setActiveSection("response")} - count={ - diffResult.response_diffs.filter((d) => d.diff_type !== "Unchanged") - .length - } - > - 响应 - - setActiveSection("messages")} - count={ - diffResult.message_diffs.filter((d) => d.diff_type !== "Unchanged") - .length - } - > - 消息 - - setActiveSection("metadata")} - count={ - diffResult.metadata_diffs.filter((d) => d.diff_type !== "Unchanged") - .length - } - > - 元数据 - -
- - {/* 差异内容 */} -
- {activeSection === "request" && ( - - )} - {activeSection === "response" && ( - - )} - {activeSection === "messages" && ( - - )} - {activeSection === "metadata" && ( - - )} -
-
- ); -} - -// ============================================================================ -// 头部组件 -// ============================================================================ - -interface DiffHeaderProps { - leftFlow: LLMFlow | null; - rightFlow: LLMFlow | null; - viewMode: ViewMode; - onViewModeChange: (mode: ViewMode) => void; - showConfig: boolean; - onToggleConfig: () => void; - onClose?: () => void; -} - -function DiffHeader({ - leftFlow, - rightFlow, - viewMode, - onViewModeChange, - showConfig, - onToggleConfig, - onClose, -}: DiffHeaderProps) { - return ( -
-
- -
- - {leftFlow?.id.slice(0, 8) || "..."} - - vs - - {rightFlow?.id.slice(0, 8) || "..."} - -
-
- -
- {/* 视图模式切换 */} -
- - -
- - {/* 配置按钮 */} - - - {/* 关闭按钮 */} - {onClose && ( - - )} -
-
- ); -} - -// ============================================================================ -// 配置面板 -// ============================================================================ - -interface DiffConfigPanelProps { - config: DiffConfig; - onChange: (config: DiffConfig) => void; -} - -function DiffConfigPanel({ config, onChange }: DiffConfigPanelProps) { - return ( -
-
差异配置
-
- - -
-
- ); -} - -// ============================================================================ -// Token 差异摘要 -// ============================================================================ - -interface TokenDiffSummaryProps { - tokenDiff: TokenDiff; -} - -function TokenDiffSummary({ tokenDiff }: TokenDiffSummaryProps) { - const hasDiff = - tokenDiff.input_diff !== 0 || - tokenDiff.output_diff !== 0 || - tokenDiff.total_diff !== 0; - - if (!hasDiff) return null; - - const formatDiff = (diff: number) => { - if (diff > 0) return `+${diff}`; - return diff.toString(); - }; - - const getDiffColor = (diff: number) => { - if (diff > 0) return "text-green-600"; - if (diff < 0) return "text-red-600"; - return "text-muted-foreground"; - }; - - return ( -
-
- - Token 差异: - - 输入 {formatDiff(tokenDiff.input_diff)} - - - 输出 {formatDiff(tokenDiff.output_diff)} - - - 总计 {formatDiff(tokenDiff.total_diff)} - -
-
- ); -} - -// ============================================================================ -// 标签页按钮 -// ============================================================================ - -interface DiffTabButtonProps { - active: boolean; - onClick: () => void; - count: number; - children: React.ReactNode; -} - -function DiffTabButton({ - active, - onClick, - count, - children, -}: DiffTabButtonProps) { - return ( - - ); -} - -// ============================================================================ -// 差异区域组件 -// ============================================================================ - -interface DiffSectionProps { - diffs: DiffItem[]; - viewMode: ViewMode; - expandedPaths: Set; - onTogglePath: (path: string) => void; -} - -function DiffSection({ - diffs, - viewMode, - expandedPaths, - onTogglePath, -}: DiffSectionProps) { - // 过滤掉未变化的项 - const changedDiffs = diffs.filter((d) => d.diff_type !== "Unchanged"); - - if (changedDiffs.length === 0) { - return ( -
- -

没有差异

-
- ); - } - - if (viewMode === "side-by-side") { - return ( -
- {changedDiffs.map((diff, idx) => ( - onTogglePath(diff.path)} - /> - ))} -
- ); - } - - return ( -
- {changedDiffs.map((diff, idx) => ( - onTogglePath(diff.path)} - /> - ))} -
- ); -} - -// ============================================================================ -// 并排差异项 -// ============================================================================ - -interface SideBySideDiffItemProps { - diff: DiffItem; - expanded: boolean; - onToggle: () => void; -} - -function SideBySideDiffItem({ - diff, - expanded, - onToggle, -}: SideBySideDiffItemProps) { - const isLongValue = - JSON.stringify(diff.left_value || diff.right_value).length > 100; - - return ( -
- {/* 路径头部 */} -
- {isLongValue ? ( - expanded ? ( - - ) : ( - - ) - ) : ( - - )} - - {diff.path} -
- - {/* 值对比 */} - {(!isLongValue || expanded) && ( -
- {/* 左侧值 */} -
- {diff.left_value !== null && diff.left_value !== undefined ? ( -
-                {formatValue(diff.left_value)}
-              
- ) : ( - (无) - )} -
- - {/* 右侧值 */} -
- {diff.right_value !== null && diff.right_value !== undefined ? ( -
-                {formatValue(diff.right_value)}
-              
- ) : ( - (无) - )} -
-
- )} -
- ); -} - -// ============================================================================ -// 统一差异项 -// ============================================================================ - -interface UnifiedDiffItemProps { - diff: DiffItem; - expanded: boolean; - onToggle: () => void; -} - -function UnifiedDiffItem({ diff, expanded, onToggle }: UnifiedDiffItemProps) { - const isLongValue = - JSON.stringify(diff.left_value || diff.right_value).length > 100; - - return ( -
- {/* 路径头部 */} -
- {isLongValue ? ( - expanded ? ( - - ) : ( - - ) - ) : ( - - )} - - {diff.path} -
- - {/* 值显示 */} - {(!isLongValue || expanded) && ( -
- {/* 删除的值 */} - {(diff.diff_type === "Removed" || diff.diff_type === "Modified") && - diff.left_value !== null && - diff.left_value !== undefined && ( -
-
- -
-
-                  {formatValue(diff.left_value)}
-                
-
- )} - - {/* 新增的值 */} - {(diff.diff_type === "Added" || diff.diff_type === "Modified") && - diff.right_value !== null && - diff.right_value !== undefined && ( -
-
- -
-
-                  {formatValue(diff.right_value)}
-                
-
- )} -
- )} -
- ); -} - -// ============================================================================ -// 消息差异区域 -// ============================================================================ - -interface MessageDiffSectionProps { - diffs: MessageDiffItem[]; - viewMode: ViewMode; -} - -function MessageDiffSection({ diffs, viewMode }: MessageDiffSectionProps) { - const changedDiffs = diffs.filter((d) => d.diff_type !== "Unchanged"); - - if (changedDiffs.length === 0) { - return ( -
- -

消息列表没有差异

-
- ); - } - - return ( -
- {diffs.map((diff, idx) => ( - - ))} -
- ); -} - -// ============================================================================ -// 消息差异项视图 -// ============================================================================ - -interface MessageDiffItemViewProps { - diff: MessageDiffItem; - viewMode: ViewMode; -} - -function MessageDiffItemView({ diff, viewMode }: MessageDiffItemViewProps) { - const [expanded, setExpanded] = useState(diff.diff_type !== "Unchanged"); - - const leftMsg = diff.left_message as { - role?: string; - content?: string; - } | null; - const rightMsg = diff.right_message as { - role?: string; - content?: string; - } | null; - - return ( -
- {/* 头部 */} -
setExpanded(!expanded)} - > - {expanded ? ( - - ) : ( - - )} - - 消息 #{diff.index + 1} - {leftMsg?.role && ( - - ({leftMsg.role}) - - )} - {!leftMsg?.role && rightMsg?.role && ( - - ({rightMsg.role}) - - )} -
- - {/* 内容 */} - {expanded && ( -
- {viewMode === "side-by-side" ? ( - <> - {/* 左侧消息 */} -
- {leftMsg ? ( -
-
- {leftMsg.role} -
-
-                      {typeof leftMsg.content === "string"
-                        ? leftMsg.content
-                        : JSON.stringify(leftMsg.content, null, 2)}
-                    
-
- ) : ( - - (无) - - )} -
- - {/* 右侧消息 */} -
- {rightMsg ? ( -
-
- {rightMsg.role} -
-
-                      {typeof rightMsg.content === "string"
-                        ? rightMsg.content
-                        : JSON.stringify(rightMsg.content, null, 2)}
-                    
-
- ) : ( - - (无) - - )} -
- - ) : ( -
- {/* 删除的消息 */} - {(diff.diff_type === "Removed" || - diff.diff_type === "Modified") && - leftMsg && ( -
-
- -
-
-
- {leftMsg.role} -
-
-                        {typeof leftMsg.content === "string"
-                          ? leftMsg.content
-                          : JSON.stringify(leftMsg.content, null, 2)}
-                      
-
-
- )} - - {/* 新增的消息 */} - {(diff.diff_type === "Added" || diff.diff_type === "Modified") && - rightMsg && ( -
-
- -
-
-
- {rightMsg.role} -
-
-                        {typeof rightMsg.content === "string"
-                          ? rightMsg.content
-                          : JSON.stringify(rightMsg.content, null, 2)}
-                      
-
-
- )} -
- )} -
- )} -
- ); -} - -// ============================================================================ -// 辅助组件和函数 -// ============================================================================ - -/** - * 差异类型图标 - */ -function DiffTypeIcon({ type }: { type: DiffType }) { - switch (type) { - case "Added": - return ; - case "Removed": - return ; - case "Modified": - return ; - default: - return null; - } -} - -/** - * 获取差异背景颜色 - */ -function getDiffBgColor(type: DiffType, isHeader: boolean = false): string { - const opacity = isHeader ? "30" : "50"; - switch (type) { - case "Added": - return `bg-green-50/${opacity} dark:bg-green-950/10`; - case "Removed": - return `bg-red-50/${opacity} dark:bg-red-950/10`; - case "Modified": - return `bg-yellow-50/${opacity} dark:bg-yellow-950/10`; - default: - return "bg-muted/30"; - } -} - -/** - * 获取差异边框颜色 - */ -function getDiffBorderColor(type: DiffType): string { - switch (type) { - case "Added": - return "border-green-200 dark:border-green-800"; - case "Removed": - return "border-red-200 dark:border-red-800"; - case "Modified": - return "border-yellow-200 dark:border-yellow-800"; - default: - return ""; - } -} - -/** - * 格式化值为字符串 - */ -function formatValue(value: unknown): string { - if (value === null) return "null"; - if (value === undefined) return "undefined"; - if (typeof value === "string") return value; - if (typeof value === "number" || typeof value === "boolean") { - return String(value); - } - return JSON.stringify(value, null, 2); -} - -// ============================================================================ -// 对话框包装组件 -// ============================================================================ - -interface FlowDiffDialogProps { - /** 是否显示对话框 */ - open: boolean; - /** 关闭对话框回调 */ - onClose: () => void; - /** 左侧 Flow ID */ - leftFlowId: string; - /** 右侧 Flow ID */ - rightFlowId: string; -} - -export function FlowDiffDialog({ - open, - onClose, - leftFlowId, - rightFlowId, -}: FlowDiffDialogProps) { - if (!open) return null; - - return ( -
- {/* 背景遮罩 */} -
- - {/* 对话框 */} -
- -
-
- ); -} - -export default FlowDiffView; diff --git a/src/components/flow-monitor/FlowFilters.tsx b/src/components/flow-monitor/FlowFilters.tsx deleted file mode 100644 index 67f1d1f26..000000000 --- a/src/components/flow-monitor/FlowFilters.tsx +++ /dev/null @@ -1,717 +0,0 @@ -import React, { useState, useEffect } from "react"; -import { - Filter, - X, - Search, - Star, - Clock, - Tag, - ChevronDown, - ChevronUp, - Code2, - HelpCircle, -} from "lucide-react"; -import { - flowMonitorApi, - type FlowFilter, - type FlowState, - type ProviderType, -} from "@/lib/api/flowMonitor"; -import { FilterExpressionInput } from "./FilterExpressionInput"; -import { FilterHelp } from "./FilterHelp"; -import { cn } from "@/lib/utils"; - -interface FlowFiltersProps { - filter: FlowFilter; - onChange: (filter: FlowFilter) => void; -} - -// 过滤模式 -type FilterMode = "simple" | "expression"; - -const PROVIDERS: ProviderType[] = [ - "Kiro", - "Gemini", - "Qwen", - "Antigravity", - "OpenAI", - "Claude", - "Vertex", - "GeminiApiKey", - "Codex", - "ClaudeOAuth", - "IFlow", -]; - -const STATES: FlowState[] = [ - "Pending", - "Streaming", - "Completed", - "Failed", - "Cancelled", -]; - -const TIME_PRESETS = [ - { label: "最近 1 小时", hours: 1 }, - { label: "最近 6 小时", hours: 6 }, - { label: "最近 24 小时", hours: 24 }, - { label: "最近 7 天", hours: 168 }, - { label: "全部", hours: 0 }, -]; - -export function FlowFilters({ filter, onChange }: FlowFiltersProps) { - const [searchQuery, setSearchQuery] = useState(""); - const [expanded, setExpanded] = useState(false); - const [availableTags, setAvailableTags] = useState([]); - const [filterMode, setFilterMode] = useState("simple"); - const [expressionValue, setExpressionValue] = useState(""); - const [showHelp, setShowHelp] = useState(false); - const [expressionValid, setExpressionValid] = useState(true); - - // 加载可用标签 - useEffect(() => { - flowMonitorApi.getAllTags().then(setAvailableTags).catch(console.error); - }, []); - - const handleSearchSubmit = (e: React.FormEvent) => { - e.preventDefault(); - // 直接更新 filter 的 content_search 字段 - onChange({ - ...filter, - content_search: searchQuery.trim() || undefined, - }); - }; - - // 当搜索框内容改变时,如果为空则清除搜索 - const handleSearchChange = (value: string) => { - setSearchQuery(value); - // 如果清空搜索框,立即清除搜索过滤 - if (!value.trim() && filter.content_search) { - onChange({ - ...filter, - content_search: undefined, - }); - } - }; - - const handleTimePreset = (hours: number) => { - if (hours === 0) { - onChange({ ...filter, time_range: undefined }); - } else { - const end = new Date(); - const start = new Date(end.getTime() - hours * 60 * 60 * 1000); - onChange({ - ...filter, - time_range: { - start: start.toISOString(), - end: end.toISOString(), - }, - }); - } - }; - - const handleProviderToggle = (provider: ProviderType) => { - const current = filter.providers || []; - const updated = current.includes(provider) - ? current.filter((p) => p !== provider) - : [...current, provider]; - onChange({ - ...filter, - providers: updated.length > 0 ? updated : undefined, - }); - }; - - const handleStateToggle = (state: FlowState) => { - const current = filter.states || []; - const updated = current.includes(state) - ? current.filter((s) => s !== state) - : [...current, state]; - onChange({ - ...filter, - states: updated.length > 0 ? updated : undefined, - }); - }; - - const handleTagToggle = (tag: string) => { - const current = filter.tags || []; - const updated = current.includes(tag) - ? current.filter((t) => t !== tag) - : [...current, tag]; - onChange({ - ...filter, - tags: updated.length > 0 ? updated : undefined, - }); - }; - - const handleClearFilters = () => { - onChange({}); - setSearchQuery(""); - setExpressionValue(""); - }; - - const handleModeToggle = () => { - const newMode = filterMode === "simple" ? "expression" : "simple"; - setFilterMode(newMode); - - // 切换模式时清除当前过滤器 - if (newMode === "expression") { - // 切换到表达式模式时,尝试将当前过滤器转换为表达式 - const expr = convertFilterToExpression(filter); - setExpressionValue(expr); - onChange({}); - } else { - // 切换到简单模式时,清除表达式 - setExpressionValue(""); - onChange({}); - } - }; - - const handleExpressionSubmit = (expression: string) => { - if (expressionValid && expression.trim()) { - // 使用表达式查询(这里需要后端支持) - // 暂时将表达式存储在 content_search 字段中作为标记 - onChange({ - filter_expression: expression.trim(), - }); - } - }; - - const handleExpressionValidation = ( - valid: boolean, - _error: string | null, - ) => { - setExpressionValid(valid); - }; - - const handleInsertExample = (example: string) => { - setExpressionValue(example); - setShowHelp(false); - }; - - // 将当前过滤器转换为表达式(简单实现) - const convertFilterToExpression = (currentFilter: FlowFilter): string => { - const parts: string[] = []; - - if (currentFilter.providers?.length) { - const providerExprs = currentFilter.providers.map((p) => `~p ${p}`); - if (providerExprs.length === 1) { - parts.push(providerExprs[0]); - } else { - parts.push(`(${providerExprs.join(" | ")})`); - } - } - - if (currentFilter.states?.length) { - const stateExprs = currentFilter.states.map( - (s) => `~s ${s.toLowerCase()}`, - ); - if (stateExprs.length === 1) { - parts.push(stateExprs[0]); - } else { - parts.push(`(${stateExprs.join(" | ")})`); - } - } - - if (currentFilter.has_error === true) { - parts.push("~e"); - } - - if (currentFilter.has_tool_calls === true) { - parts.push("~t"); - } - - if (currentFilter.has_thinking === true) { - parts.push("~k"); - } - - if (currentFilter.starred_only) { - parts.push("~starred"); - } - - if (currentFilter.content_search) { - parts.push(`~b "${currentFilter.content_search}"`); - } - - if (currentFilter.tags?.length) { - const tagExprs = currentFilter.tags.map((t) => `~tag ${t}`); - parts.push(...tagExprs); - } - - return parts.join(" & "); - }; - - const hasActiveFilters = - filter.providers?.length || - filter.states?.length || - filter.tags?.length || - filter.time_range || - filter.has_error !== undefined || - filter.has_tool_calls !== undefined || - filter.has_thinking !== undefined || - filter.starred_only || - filter.content_search || - filter.models?.length || - filter.filter_expression; - - const activeFilterCount = [ - filter.providers?.length, - filter.states?.length, - filter.tags?.length, - filter.time_range ? 1 : 0, - filter.has_error !== undefined ? 1 : 0, - filter.has_tool_calls !== undefined ? 1 : 0, - filter.has_thinking !== undefined ? 1 : 0, - filter.starred_only ? 1 : 0, - filter.content_search ? 1 : 0, - filter.models?.length, - filter.filter_expression ? 1 : 0, - ].reduce((sum: number, val) => sum + (val || 0), 0); - - return ( -
- {/* 过滤模式切换 */} -
-
- - -
- - {/* 帮助按钮(仅在表达式模式显示) */} - {filterMode === "expression" && ( - - )} -
- - {/* 表达式模式 */} - {filterMode === "expression" ? ( -
-
-
- setShowHelp(!showHelp)} - /> -
- -
- - {/* 帮助面板 */} - {showHelp && ( - setShowHelp(false)} - onInsertExample={handleInsertExample} - /> - )} - - {/* 当前表达式状态 */} - {filter.filter_expression && ( -
- 当前表达式: - - {filter.filter_expression} - -
- )} -
- ) : ( - /* 简单模式 - 原有的搜索栏和过滤器 */ - <> - {/* 搜索栏 */} -
-
- - handleSearchChange(e.target.value)} - className="w-full rounded-lg border bg-background pl-9 pr-4 py-2 text-sm focus:outline-none focus:ring-2 focus:ring-primary" - /> -
- -
- - )} - - {/* 快捷过滤器(两种模式都显示) */} -
- {/* 时间预设 */} -
- - {TIME_PRESETS.map((preset) => ( - - ))} -
- - {/* 收藏过滤 */} - - - {/* 展开/收起高级过滤器(仅简单模式) */} - {filterMode === "simple" && ( - - )} - - {/* 清除过滤器 */} - {hasActiveFilters && ( - - )} -
- - {/* 高级过滤器面板(仅简单模式且展开时显示) */} - {filterMode === "simple" && expanded && ( -
- {/* 提供商过滤 */} - -
- {PROVIDERS.map((provider) => ( - handleProviderToggle(provider)} - /> - ))} -
-
- - {/* 状态过滤 */} - -
- {STATES.map((state) => ( - handleStateToggle(state)} - /> - ))} -
-
- - {/* 特性过滤 */} - -
- - onChange({ - ...filter, - has_error: filter.has_error === true ? undefined : true, - }) - } - /> - - onChange({ - ...filter, - has_tool_calls: - filter.has_tool_calls === true ? undefined : true, - }) - } - /> - - onChange({ - ...filter, - has_thinking: - filter.has_thinking === true ? undefined : true, - }) - } - /> - - onChange({ - ...filter, - is_streaming: - filter.is_streaming === true ? undefined : true, - }) - } - /> -
-
- - {/* 标签过滤 */} - {availableTags.length > 0 && ( - -
- {availableTags.map((tag) => ( - handleTagToggle(tag)} - icon={} - /> - ))} -
-
- )} - - {/* Token 范围 */} - -
- - onChange({ - ...filter, - token_range: { - ...filter.token_range, - min: e.target.value - ? parseInt(e.target.value) - : undefined, - }, - }) - } - className="w-24 rounded border bg-background px-2 py-1 text-sm" - /> - - - - onChange({ - ...filter, - token_range: { - ...filter.token_range, - max: e.target.value - ? parseInt(e.target.value) - : undefined, - }, - }) - } - className="w-24 rounded border bg-background px-2 py-1 text-sm" - /> -
-
- - {/* 延迟范围 */} - -
- - onChange({ - ...filter, - latency_range: { - ...filter.latency_range, - min_ms: e.target.value - ? parseInt(e.target.value) - : undefined, - }, - }) - } - className="w-24 rounded border bg-background px-2 py-1 text-sm" - /> - - - - onChange({ - ...filter, - latency_range: { - ...filter.latency_range, - max_ms: e.target.value - ? parseInt(e.target.value) - : undefined, - }, - }) - } - className="w-24 rounded border bg-background px-2 py-1 text-sm" - /> -
-
- - {/* 模型过滤 */} - - - onChange({ - ...filter, - models: e.target.value ? [e.target.value] : undefined, - }) - } - className="w-full rounded border bg-background px-3 py-1.5 text-sm" - /> - -
- )} -
- ); -} - -interface FilterSectionProps { - title: string; - children: React.ReactNode; -} - -function FilterSection({ title, children }: FilterSectionProps) { - return ( -
-
- {title} -
- {children} -
- ); -} - -interface FilterChipProps { - label: string; - active?: boolean; - onClick: () => void; - icon?: React.ReactNode; -} - -function FilterChip({ label, active, onClick, icon }: FilterChipProps) { - return ( - - ); -} - -function getStateLabel(state: FlowState): string { - const labels: Record = { - Pending: "等待中", - Streaming: "流式传输中", - Completed: "已完成", - Failed: "失败", - Cancelled: "已取消", - }; - return labels[state] || state; -} - -function isTimeRangeMatch( - timeRange: { start?: string; end?: string }, - hours: number, -): boolean { - if (!timeRange.start || !timeRange.end) return false; - const start = new Date(timeRange.start); - const end = new Date(timeRange.end); - const diff = (end.getTime() - start.getTime()) / (1000 * 60 * 60); - // 允许 5% 的误差 - return Math.abs(diff - hours) < hours * 0.05; -} - -export default FlowFilters; diff --git a/src/components/flow-monitor/FlowList.tsx b/src/components/flow-monitor/FlowList.tsx deleted file mode 100644 index db28516c4..000000000 --- a/src/components/flow-monitor/FlowList.tsx +++ /dev/null @@ -1,1044 +0,0 @@ -import React, { useState, useEffect, useCallback } from "react"; -import * as Select from "@radix-ui/react-select"; -import { - CheckCircle2, - XCircle, - Clock, - Loader2, - ChevronDown, - ChevronRight, - Star, - StarOff, - Wrench, - Brain, - RefreshCw, - Copy, - ExternalLink, - Wifi, - WifiOff, - Pause, - Play, - AlertTriangle, - Activity, - Bell, - BellOff, - Check, -} from "lucide-react"; -import { - flowMonitorApi, - realtimeMonitorApi, - type LLMFlow, - type FlowState, - type FlowFilter, - type FlowSortBy, - type FlowQueryResult, - type ThresholdCheckResult, - type RequestRateResponse, - formatFlowState, - formatLatency, - formatTokenCount, - truncateText, -} from "@/lib/api/flowMonitor"; -import { useFlowEvents } from "@/hooks/useFlowEvents"; -import { useFlowNotifications } from "@/hooks/useFlowNotifications"; -import { NotificationSettings } from "./NotificationSettings"; -import { FlowRecordContextMenu } from "./FlowRecordContextMenu"; -import { cn } from "@/lib/utils"; - -interface FlowListProps { - filter?: FlowFilter; - onFlowSelect?: (flow: LLMFlow) => void; - selectedFlowId?: string; - onRefresh?: () => void; - /** 是否启用实时更新 */ - enableRealtime?: boolean; -} - -export function FlowList({ - filter = {}, - onFlowSelect, - selectedFlowId, - onRefresh, - enableRealtime = true, -}: FlowListProps) { - const [flows, setFlows] = useState([]); - const [loading, setLoading] = useState(true); - const [error, setError] = useState(null); - const [page, setPage] = useState(1); - const [pageSize] = useState(20); - const [totalPages, setTotalPages] = useState(1); - const [total, setTotal] = useState(0); - const [sortBy, setSortBy] = useState("created_at"); - const [sortDesc, setSortDesc] = useState(true); - const [expandedId, setExpandedId] = useState(null); - - // 暂停/恢复实时更新状态 - const [isPaused, setIsPaused] = useState(false); - - // 请求速率状态 - const [requestRate, setRequestRate] = useState( - null, - ); - - // 阈值警告状态 - const [thresholdWarnings, setThresholdWarnings] = useState< - Map - >(new Map()); - - // 通知设置面板状态 - const [showNotificationSettings, setShowNotificationSettings] = - useState(false); - - // 通知功能 - const { - notificationService, - requestPermission, - permissionStatus, - canNotify, - } = useFlowNotifications({ - enabled: enableRealtime && !isPaused, - autoRequestPermission: false, // 不自动请求权限,让用户手动控制 - }); - - // 实时更新 Hook - const { - connected: wsConnected, - connecting: wsConnecting, - activeFlows, - lastThresholdWarning, - } = useFlowEvents({ - autoConnect: enableRealtime && !isPaused, - onFlowStarted: (flow) => { - // 如果暂停了,不更新列表 - if (isPaused) return; - - // 新 Flow 开始时,添加到列表顶部 - if (page === 1 && sortBy === "created_at" && sortDesc) { - setFlows((prev) => { - // 将 FlowSummary 转换为 LLMFlow 的简化版本 - const newFlow: LLMFlow = { - id: flow.id, - flow_type: flow.flow_type, - state: flow.state, - request: { - method: "POST", - path: "", - headers: {}, - body: {}, - messages: [], - model: flow.model, - parameters: { stream: false }, - size_bytes: 0, - timestamp: flow.created_at, - }, - metadata: { - provider: flow.provider, - retry_count: 0, - client_info: {}, - routing_info: {}, - }, - timestamps: { - created: flow.created_at, - request_start: flow.created_at, - duration_ms: flow.duration_ms, - }, - annotations: { - tags: [], - starred: false, - }, - }; - // 避免重复 - if (prev.some((f) => f.id === flow.id)) { - return prev; - } - return [newFlow, ...prev.slice(0, pageSize - 1)]; - }); - setTotal((prev) => prev + 1); - } - }, - onFlowCompleted: (id, summary) => { - // Flow 完成时,更新状态 - setFlows((prev) => - prev.map((f) => - f.id === id - ? { - ...f, - state: "Completed" as FlowState, - timestamps: { - ...f.timestamps, - duration_ms: summary.duration_ms, - }, - response: f.response - ? { - ...f.response, - usage: { - ...f.response.usage, - input_tokens: summary.input_tokens || 0, - output_tokens: summary.output_tokens || 0, - total_tokens: - (summary.input_tokens || 0) + - (summary.output_tokens || 0), - }, - } - : undefined, - } - : f, - ), - ); - }, - onFlowFailed: (id) => { - // Flow 失败时,更新状态 - setFlows((prev) => - prev.map((f) => - f.id === id ? { ...f, state: "Failed" as FlowState } : f, - ), - ); - }, - onFlowUpdated: (id, update) => { - // Flow 更新时,更新状态 - if (update.state) { - setFlows((prev) => - prev.map((f) => (f.id === id ? { ...f, state: update.state! } : f)), - ); - } - }, - onThresholdWarning: (id, result) => { - // 阈值警告时,记录警告 - setThresholdWarnings((prev) => { - const next = new Map(prev); - next.set(id, result); - return next; - }); - }, - }); - - // 处理阈值警告 - useEffect(() => { - if (lastThresholdWarning) { - setThresholdWarnings((prev) => { - const next = new Map(prev); - next.set(lastThresholdWarning.id, lastThresholdWarning.result); - return next; - }); - } - }, [lastThresholdWarning]); - - // 定期获取请求速率 - useEffect(() => { - const fetchRequestRate = async () => { - try { - const rate = await realtimeMonitorApi.getRequestRate(); - setRequestRate(rate); - } catch (e) { - console.error("Failed to fetch request rate:", e); - } - }; - - // 初始获取 - fetchRequestRate(); - - // 每 5 秒更新一次 - const interval = setInterval(fetchRequestRate, 5000); - - return () => clearInterval(interval); - }, []); - - const fetchFlows = useCallback(async () => { - try { - setLoading(true); - setError(null); - console.log("查询 Flow,过滤条件:", JSON.stringify(filter, null, 2)); - const result: FlowQueryResult = await flowMonitorApi.queryFlows( - filter, - sortBy, - sortDesc, - page, - pageSize, - ); - console.log("查询结果:", result.total, "条记录"); - setFlows(result.flows); - setTotalPages(result.total_pages); - setTotal(result.total); - } catch (e) { - console.error("Failed to fetch flows:", e); - setError(e instanceof Error ? e.message : "加载 Flow 列表失败"); - } finally { - setLoading(false); - } - }, [filter, sortBy, sortDesc, page, pageSize]); - - // 当 filter 改变时,重置到第一页 - useEffect(() => { - setPage(1); - }, [filter]); - - useEffect(() => { - fetchFlows(); - }, [fetchFlows]); - - const handleRefresh = () => { - fetchFlows(); - onRefresh?.(); - }; - - const handleToggleStar = async (e: React.MouseEvent, flowId: string) => { - e.stopPropagation(); - try { - await flowMonitorApi.toggleFlowStar(flowId); - // 更新本地状态 - setFlows((prev) => - prev.map((f) => - f.id === flowId - ? { - ...f, - annotations: { - ...f.annotations, - starred: !f.annotations.starred, - }, - } - : f, - ), - ); - } catch (e) { - console.error("Failed to toggle star:", e); - } - }; - - const handleCopyId = async (e: React.MouseEvent, flowId: string) => { - e.stopPropagation(); - try { - await navigator.clipboard.writeText(flowId); - } catch (e) { - console.error("Failed to copy:", e); - } - }; - - const getStateIcon = (state: FlowState) => { - switch (state) { - case "Completed": - return ; - case "Failed": - return ; - case "Streaming": - return ; - case "Pending": - return ; - case "Cancelled": - return ; - default: - return ; - } - }; - - // 从模型名推断实际的模型提供商 - const inferProviderFromModel = (model: string): string | null => { - const modelLower = model.toLowerCase(); - - // DeepSeek 模型 - if (modelLower.includes("deepseek")) { - return "DeepSeek"; - } - // Claude 模型 - if ( - modelLower.includes("claude") || - modelLower.includes("anthropic") || - modelLower.includes("sonnet") || - modelLower.includes("opus") || - modelLower.includes("haiku") - ) { - return "Claude"; - } - // GPT/OpenAI 模型 - if ( - modelLower.includes("gpt") || - modelLower.includes("o1") || - modelLower.includes("o3") || - modelLower.includes("chatgpt") - ) { - return "OpenAI"; - } - // Gemini 模型 - if (modelLower.includes("gemini")) { - return "Gemini"; - } - // Qwen 模型 - if (modelLower.includes("qwen") || modelLower.includes("qwq")) { - return "Qwen"; - } - // Llama 模型 - if (modelLower.includes("llama")) { - return "Llama"; - } - // Mistral 模型 - if (modelLower.includes("mistral") || modelLower.includes("mixtral")) { - return "Mistral"; - } - - return null; - }; - - // 获取显示的提供商名称(优先使用模型推断的提供商) - const getDisplayProvider = ( - metadataProvider: string, - model: string, - ): string => { - const inferredProvider = inferProviderFromModel(model); - // 如果能从模型名推断出提供商,且与凭证池不同,则显示推断的提供商 - if (inferredProvider && inferredProvider !== metadataProvider) { - return inferredProvider; - } - return metadataProvider; - }; - - const getProviderColor = (provider: string) => { - const colors: Record = { - Kiro: "bg-purple-100 text-purple-700 dark:bg-purple-900/30 dark:text-purple-300", - Gemini: - "bg-blue-100 text-blue-700 dark:bg-blue-900/30 dark:text-blue-300", - OpenAI: - "bg-green-100 text-green-700 dark:bg-green-900/30 dark:text-green-300", - Claude: - "bg-orange-100 text-orange-700 dark:bg-orange-900/30 dark:text-orange-300", - Qwen: "bg-cyan-100 text-cyan-700 dark:bg-cyan-900/30 dark:text-cyan-300", - Antigravity: - "bg-pink-100 text-pink-700 dark:bg-pink-900/30 dark:text-pink-300", - DeepSeek: - "bg-indigo-100 text-indigo-700 dark:bg-indigo-900/30 dark:text-indigo-300", - Llama: - "bg-yellow-100 text-yellow-700 dark:bg-yellow-900/30 dark:text-yellow-300", - Mistral: - "bg-teal-100 text-teal-700 dark:bg-teal-900/30 dark:text-teal-300", - }; - return ( - colors[provider] || - "bg-gray-100 text-gray-700 dark:bg-gray-900/30 dark:text-gray-300" - ); - }; - - const formatTime = (timestamp: string) => { - const date = new Date(timestamp); - return date.toLocaleString("zh-CN", { - month: "2-digit", - day: "2-digit", - hour: "2-digit", - minute: "2-digit", - second: "2-digit", - }); - }; - - if (loading && flows.length === 0) { - return ( -
- -
- ); - } - - if (error) { - return ( -
- {error} - -
- ); - } - - return ( -
- {/* 工具栏 */} -
-
- - 共 {total} 条记录 - - {/* 实时连接状态 */} - {enableRealtime && ( -
- {wsConnecting ? ( - - ) : wsConnected && !isPaused ? ( - - ) : ( - - )} - - {wsConnecting - ? "连接中..." - : wsConnected && !isPaused - ? "实时更新" - : isPaused - ? "已暂停" - : "离线"} - -
- )} - {/* 活跃 Flow 数量 */} - {activeFlows.size > 0 && ( - - {activeFlows.size} 进行中 - - )} - {/* 请求速率显示 */} - {requestRate && ( -
- - {requestRate.rate.toFixed(2)} req/s - - ({requestRate.count} / {requestRate.window_seconds}s) - -
- )} - {/* 阈值警告数量 */} - {thresholdWarnings.size > 0 && ( - - - {thresholdWarnings.size} 警告 - - )} -
-
- {/* 通知设置按钮 */} - {enableRealtime && ( - - )} - {/* 暂停/恢复按钮 */} - {enableRealtime && ( - - )} - setSortBy(value as FlowSortBy)} - > - - - - - - - - - - - - - - 按时间 - - - - - - 按耗时 - - - - - - 按 Token - - - - - - 按模型 - - - - - - - -
-
- - {/* Flow 列表 */} -
- {flows.length === 0 ? ( -
- 暂无 Flow 记录 -
- ) : ( -
- {flows.map((flow) => ( - onFlowSelect?.(flow)} - > -
- - setExpandedId(expandedId === flow.id ? null : flow.id) - } - onSelect={() => onFlowSelect?.(flow)} - onToggleStar={(e) => handleToggleStar(e, flow.id)} - onCopyId={(e) => handleCopyId(e, flow.id)} - getStateIcon={getStateIcon} - getProviderColor={getProviderColor} - getDisplayProvider={getDisplayProvider} - formatTime={formatTime} - /> -
-
- ))} -
- )} -
- - {/* 分页 */} - {totalPages > 1 && ( -
- - - {page} / {totalPages} - - -
- )} - - {/* 通知设置面板 */} - setShowNotificationSettings(false)} - permissionStatus={permissionStatus} - onRequestPermission={requestPermission} - /> -
- ); -} - -interface FlowListItemProps { - flow: LLMFlow; - expanded: boolean; - selected: boolean; - thresholdWarning?: ThresholdCheckResult; - onToggleExpand: () => void; - onSelect: () => void; - onToggleStar: (e: React.MouseEvent) => void; - onCopyId: (e: React.MouseEvent) => void; - getStateIcon: (state: FlowState) => React.ReactNode; - getProviderColor: (provider: string) => string; - getDisplayProvider: (metadataProvider: string, model: string) => string; - formatTime: (timestamp: string) => string; -} - -function FlowListItem({ - flow, - expanded, - selected, - thresholdWarning, - onToggleExpand, - onSelect, - onToggleStar, - onCopyId, - getStateIcon, - getProviderColor, - getDisplayProvider, - formatTime, -}: FlowListItemProps) { - const hasToolCalls = - flow.response?.tool_calls && flow.response.tool_calls.length > 0; - const hasThinking = !!flow.response?.thinking; - const hasError = !!flow.error; - const hasThresholdWarning = - thresholdWarning && - (thresholdWarning.latency_exceeded || - thresholdWarning.token_exceeded || - thresholdWarning.input_token_exceeded || - thresholdWarning.output_token_exceeded); - - return ( -
- {/* 主行 */} -
- {/* 展开按钮 */} - - - {/* 状态图标 */} - {getStateIcon(flow.state)} - - {/* 时间 */} - - {formatTime(flow.timestamps.created)} - - - {/* 提供商 */} - {(() => { - const displayProvider = getDisplayProvider( - flow.metadata.provider, - flow.request.model, - ); - return ( - - {displayProvider} - - ); - })()} - - {/* 模型 */} - - {flow.request.model} - - - {/* 特性标记 */} -
- {hasToolCalls && ( - - - - )} - {hasThinking && ( - - - - )} - {hasError && ( - - - - )} - {hasThresholdWarning && ( - - - - )} -
- - {/* Token 数 */} - - {flow.response?.usage - ? formatTokenCount(flow.response.usage.total_tokens) - : "-"}{" "} - tokens - - - {/* 耗时 */} - - {formatLatency(flow.timestamps.duration_ms)} - - - {/* 收藏按钮 */} - -
- - {/* 阈值警告详情 */} - {hasThresholdWarning && expanded && ( -
-
-
- - 阈值警告 -
-
- {thresholdWarning.latency_exceeded && ( -
- 延迟: {formatLatency(thresholdWarning.actual_latency_ms)}{" "} - (超限) -
- )} - {thresholdWarning.token_exceeded && ( -
- Token: {formatTokenCount(thresholdWarning.actual_tokens)}{" "} - (超限) -
- )} - {thresholdWarning.input_token_exceeded && ( -
- 输入 Token:{" "} - {formatTokenCount(thresholdWarning.actual_input_tokens)}{" "} - (超限) -
- )} - {thresholdWarning.output_token_exceeded && ( -
- 输出 Token:{" "} - {formatTokenCount(thresholdWarning.actual_output_tokens)}{" "} - (超限) -
- )} -
-
-
- )} - - {/* 展开详情 */} - {expanded && ( -
- {/* 基本信息 */} -
-
- 状态:{" "} - - {formatFlowState(flow.state)} - -
-
- 流式:{" "} - {flow.request.parameters.stream ? "是" : "否"} -
-
- TTFB:{" "} - {flow.timestamps.ttfb_ms - ? formatLatency(flow.timestamps.ttfb_ms) - : "-"} -
- {flow.response?.usage && ( - <> -
- 输入 Token:{" "} - {formatTokenCount(flow.response.usage.input_tokens)} -
-
- 输出 Token:{" "} - {formatTokenCount(flow.response.usage.output_tokens)} -
- {flow.response.usage.cache_read_tokens && ( -
- 缓存读取:{" "} - {formatTokenCount(flow.response.usage.cache_read_tokens)} -
- )} - - )} - {flow.metadata.credential_name && ( -
- 凭证:{" "} - {flow.metadata.credential_name} -
- )} - {flow.metadata.retry_count > 0 && ( -
- 重试次数:{" "} - {flow.metadata.retry_count} -
- )} -
- - {/* 内容预览 */} - {flow.response?.content && ( -
-
- 响应内容预览: -
-
- {truncateText(flow.response.content, 300)} -
-
- )} - - {/* 错误信息 */} - {flow.error && ( -
-
错误: {flow.error.error_type}
-
{flow.error.message}
-
- )} - - {/* 标签 */} - {flow.annotations.tags.length > 0 && ( -
- 标签: - {flow.annotations.tags.map((tag) => ( - - {tag} - - ))} -
- )} - - {/* 操作按钮 */} -
- - -
-
- )} -
- ); -} - -export default FlowList; diff --git a/src/components/flow-monitor/FlowRecordContextMenu.tsx b/src/components/flow-monitor/FlowRecordContextMenu.tsx deleted file mode 100644 index fa68ec1a8..000000000 --- a/src/components/flow-monitor/FlowRecordContextMenu.tsx +++ /dev/null @@ -1,160 +0,0 @@ -/** - * Flow 记录右键菜单组件 - * - * 为 Flow Monitor 的请求记录提供右键菜单功能 - * 支持查看详情、复制 ID、复制为 cURL、导出 JSON 等操作 - * - * @module components/flow-monitor/FlowRecordContextMenu - */ - -import React from "react"; -import { ExternalLink, Copy, Terminal, FileJson } from "lucide-react"; -import { - ContextMenu, - ContextMenuContent, - ContextMenuItem, - ContextMenuSeparator, - ContextMenuShortcut, - ContextMenuTrigger, -} from "@/components/ui/context-menu"; -import { toast } from "sonner"; -import type { LLMFlow } from "@/lib/api/flowMonitor"; - -interface FlowRecordContextMenuProps { - /** Flow 记录数据 */ - flow: LLMFlow; - /** 子元素 */ - children: React.ReactNode; - /** 查看详情回调 */ - onViewDetail: () => void; - /** 导出 JSON 回调 */ - onExportJson?: (flowId: string) => void; -} - -/** - * 生成 cURL 命令 - */ -function generateCurlCommand(flow: LLMFlow): string { - const { request, metadata } = flow; - - // 基础 URL(根据 provider 推断) - const baseUrls: Record = { - Kiro: "https://codewhisperer.us-east-1.amazonaws.com", - OpenAI: "https://api.openai.com/v1/chat/completions", - Claude: "https://api.anthropic.com/v1/messages", - Gemini: "https://generativelanguage.googleapis.com/v1beta/models", - Qwen: "https://dashscope.aliyuncs.com/api/v1/services/aigc/text-generation/generation", - }; - - const url = - baseUrls[metadata.provider] || - "https://api.example.com/v1/chat/completions"; - - // 构建请求体(避免 stream 重复) - const { stream: _stream, ...otherParams } = request.parameters; - const body = { - model: request.model, - messages: request.messages, - stream: request.parameters.stream, - ...otherParams, - }; - - // 构建 cURL 命令 - const parts = [ - "curl", - `-X ${request.method || "POST"}`, - `'${url}'`, - "-H 'Content-Type: application/json'", - "-H 'Authorization: Bearer YOUR_API_KEY'", - `-d '${JSON.stringify(body, null, 2)}'`, - ]; - - return parts.join(" \\\n "); -} - -export function FlowRecordContextMenu({ - flow, - children, - onViewDetail, - onExportJson, -}: FlowRecordContextMenuProps) { - // 复制请求 ID - const handleCopyId = async () => { - try { - await navigator.clipboard.writeText(flow.id); - toast.success("已复制请求 ID"); - } catch (error) { - console.error("复制失败:", error); - toast.error("复制失败"); - } - }; - - // 复制为 cURL - const handleCopyAsCurl = async () => { - try { - const curlCommand = generateCurlCommand(flow); - await navigator.clipboard.writeText(curlCommand); - toast.success("已复制 cURL 命令"); - } catch (error) { - console.error("复制失败:", error); - toast.error("复制失败"); - } - }; - - // 导出为 JSON - const handleExportJson = () => { - if (onExportJson) { - onExportJson(flow.id); - } else { - // 默认导出行为:下载 JSON 文件 - const jsonStr = JSON.stringify(flow, null, 2); - const blob = new Blob([jsonStr], { type: "application/json" }); - const url = URL.createObjectURL(blob); - const a = document.createElement("a"); - a.href = url; - a.download = `flow-${flow.id.slice(0, 8)}.json`; - document.body.appendChild(a); - a.click(); - document.body.removeChild(a); - URL.revokeObjectURL(url); - toast.success("已导出 JSON 文件"); - } - }; - - return ( - - {children} - - {/* 查看详情 */} - - - 查看详情 - ↵ - - - {/* 复制请求 ID */} - - - 复制请求 ID - C - - - {/* 复制为 cURL */} - - - 复制为 cURL - ⇧C - - - - - {/* 导出为 JSON */} - - - 导出为 JSON - E - - - - ); -} diff --git a/src/components/flow-monitor/FlowStats.tsx b/src/components/flow-monitor/FlowStats.tsx deleted file mode 100644 index 7c2cc5422..000000000 --- a/src/components/flow-monitor/FlowStats.tsx +++ /dev/null @@ -1,1212 +0,0 @@ -import React, { useState, useEffect, useCallback } from "react"; -import { - Activity, - CheckCircle2, - XCircle, - Clock, - Zap, - TrendingUp, - TrendingDown, - RefreshCw, - BarChart3, - PieChart, - Loader2, - AlertCircle, - LineChart, -} from "lucide-react"; -import { - flowMonitorApi, - enhancedStatsApi, - type FlowStats as FlowStatsType, - type FlowFilter, - type ProviderStats, - type ModelStats, - type EnhancedStats, - type TrendData, - type Distribution, - type StatsTimeRange, - formatLatency, - formatTokenCount, -} from "@/lib/api/flowMonitor"; -import { cn } from "@/lib/utils"; - -interface FlowStatsProps { - /** 过滤条件 */ - filter?: FlowFilter; - /** 自动刷新间隔(毫秒),0 表示不自动刷新 */ - autoRefreshInterval?: number; - /** 刷新回调 */ - onRefresh?: () => void; - /** 是否显示紧凑模式 */ - compact?: boolean; - /** 是否显示增强统计 */ - showEnhanced?: boolean; -} - -export function FlowStats({ - filter = {}, - autoRefreshInterval = 0, - onRefresh, - compact = false, - showEnhanced = true, -}: FlowStatsProps) { - const [stats, setStats] = useState(null); - const [enhancedStats, setEnhancedStats] = useState( - null, - ); - const [loading, setLoading] = useState(true); - const [error, setError] = useState(null); - const [lastUpdated, setLastUpdated] = useState(null); - const [timeRangeHours, setTimeRangeHours] = useState(24); - const [activeTab, setActiveTab] = useState< - "overview" | "trends" | "distribution" - >("overview"); - - const getTimeRange = useCallback((): StatsTimeRange => { - const now = new Date(); - return { - start: new Date( - now.getTime() - timeRangeHours * 60 * 60 * 1000, - ).toISOString(), - end: now.toISOString(), - }; - }, [timeRangeHours]); - - const fetchStats = useCallback(async () => { - try { - setLoading(true); - setError(null); - - console.log("正在获取统计数据,过滤条件:", filter); - - const [basicStats, enhanced] = await Promise.all([ - flowMonitorApi.getFlowStats(filter), - showEnhanced - ? enhancedStatsApi.getEnhancedStats(filter, getTimeRange()) - : Promise.resolve(null), - ]); - - console.log("获取到的基础统计数据:", basicStats); - console.log("获取到的增强统计数据:", enhanced); - - setStats(basicStats); - setEnhancedStats(enhanced); - setLastUpdated(new Date()); - } catch (e) { - console.error("Failed to fetch flow stats:", e); - setError(e instanceof Error ? e.message : "加载统计数据失败"); - } finally { - setLoading(false); - } - }, [filter, showEnhanced, getTimeRange]); - - useEffect(() => { - fetchStats(); - }, [fetchStats]); - - // 自动刷新 - useEffect(() => { - if (autoRefreshInterval > 0) { - const interval = setInterval(fetchStats, autoRefreshInterval); - return () => clearInterval(interval); - } - }, [autoRefreshInterval, fetchStats]); - - const handleRefresh = () => { - fetchStats(); - onRefresh?.(); - }; - - if (loading && !stats) { - return ( -
- -
- ); - } - - if (error) { - return ( -
-
- - {error} -
- -
- ); - } - - if (!stats) { - return null; - } - - if (compact) { - return ( - - ); - } - - return ( -
- {/* 头部工具栏 */} -
-

- - 统计仪表板 -

-
- {/* 调试按钮 */} - - - {/* 时间范围选择 */} - - {lastUpdated && ( - - 更新于 {lastUpdated.toLocaleTimeString("zh-CN")} - - )} - -
-
- - {/* 标签页切换 */} - {showEnhanced && ( -
- - - -
- )} - - {/* 概览标签页 */} - {activeTab === "overview" && ( - - )} - - {/* 趋势标签页 */} - {activeTab === "trends" && enhancedStats && ( - - )} - - {/* 分布标签页 */} - {activeTab === "distribution" && enhancedStats && ( - - )} -
- ); -} - -// ============================================================================ -// 概览标签页 -// ============================================================================ - -interface OverviewTabProps { - stats: FlowStatsType; - enhancedStats: EnhancedStats | null; -} - -function OverviewTab({ stats, enhancedStats }: OverviewTabProps) { - return ( -
- {/* 核心指标卡片 */} -
- } - trend={null} - /> - = 0.95 ? ( - - ) : stats.success_rate >= 0.8 ? ( - - ) : ( - - ) - } - trend={ - stats.success_rate >= 0.95 - ? "up" - : stats.success_rate < 0.8 - ? "down" - : null - } - valueColor={ - stats.success_rate >= 0.95 - ? "text-green-600" - : stats.success_rate >= 0.8 - ? "text-yellow-600" - : "text-red-600" - } - /> - } - subtitle={`${formatLatency(stats.min_latency_ms)} - ${formatLatency(stats.max_latency_ms)}`} - /> - } - subtitle={`输入 ${formatTokenCount(stats.total_input_tokens)} / 输出 ${formatTokenCount(stats.total_output_tokens)}`} - /> -
- - {/* 请求速率(如果有增强统计) */} - {enhancedStats && ( -
-
-

- - 请求速率 -

- - {enhancedStats.request_rate.toFixed(2)}{" "} - - 请求/秒 - - -
-
- )} - - {/* 成功/失败统计 */} -
-
-

- - 请求状态 -

-
- - -
-
- -
-

- - Token 统计 -

-
-
-
平均输入
-
- {formatTokenCount(stats.avg_input_tokens)} -
-
-
-
平均输出
-
- {formatTokenCount(stats.avg_output_tokens)} -
-
-
-
总输入
-
- {formatTokenCount(stats.total_input_tokens)} -
-
-
-
总输出
-
- {formatTokenCount(stats.total_output_tokens)} -
-
-
-
-
- - {/* 按提供商分布 */} - {stats.by_provider.length > 0 && ( -
-

- - 按提供商分布 -

- -
- )} - - {/* 按模型分布 */} - {stats.by_model.length > 0 && ( -
-

- - 按模型分布 -

- -
- )} - - {/* 按状态分布 */} - {stats.by_state.length > 0 && ( -
-

- - 按状态分布 -

- -
- )} -
- ); -} - -// ============================================================================ -// 趋势标签页 -// ============================================================================ - -interface TrendsTabProps { - enhancedStats: EnhancedStats; -} - -function TrendsTab({ enhancedStats }: TrendsTabProps) { - return ( -
- {/* 请求趋势图 */} -
-

- - 请求趋势 -

- -
- - {/* 成功率趋势(按提供商) */} - {enhancedStats.success_by_provider.length > 0 && ( -
-

- - 按提供商成功率 -

- -
- )} -
- ); -} - -// ============================================================================ -// 分布标签页 -// ============================================================================ - -interface DistributionTabProps { - enhancedStats: EnhancedStats; -} - -function DistributionTab({ enhancedStats }: DistributionTabProps) { - return ( -
- {/* Token 分布(按模型) */} - {enhancedStats.token_by_model.buckets.length > 0 && ( -
-

- - Token 分布(按模型) -

- -
- )} - - {/* 延迟直方图 */} - {enhancedStats.latency_histogram.buckets.length > 0 && ( -
-

- - 延迟分布 -

- -
- )} - - {/* 错误分布 */} - {enhancedStats.error_distribution.buckets.length > 0 && ( -
-

- - 错误分布 -

- -
- )} -
- ); -} - -// ============================================================================ -// 趋势图组件 -// ============================================================================ - -interface TrendChartProps { - data: TrendData; -} - -function TrendChart({ data }: TrendChartProps) { - if (data.points.length === 0) { - return ( -
- 暂无数据 -
- ); - } - - const maxValue = Math.max(...data.points.map((p) => p.value), 1); - - return ( -
-
- {/* Y 轴标签 */} -
- {maxValue} - {Math.round(maxValue / 2)} - 0 -
- - {/* 图表区域 */} -
- {/* 网格线 */} -
-
-
-
-
- - {/* 数据条 */} -
- {data.points.map((point, index) => { - const height = (point.value / maxValue) * 100; - return ( -
- ); - })} -
-
-
- - {/* X 轴标签 */} -
- {data.points.length > 0 && ( - <> - - {new Date(data.points[0].timestamp).toLocaleTimeString("zh-CN", { - hour: "2-digit", - minute: "2-digit", - })} - - {data.points.length > 1 && ( - - {new Date( - data.points[data.points.length - 1].timestamp, - ).toLocaleTimeString("zh-CN", { - hour: "2-digit", - minute: "2-digit", - })} - - )} - - )} -
- -
- 时间间隔: {data.interval} -
-
- ); -} - -// ============================================================================ -// 成功率图表组件 -// ============================================================================ - -interface SuccessRateChartProps { - data: [string, number][]; -} - -function SuccessRateChart({ data }: SuccessRateChartProps) { - if (data.length === 0) { - return ( -
- 暂无数据 -
- ); - } - - return ( -
- {data.map(([provider, rate]) => { - const percentage = rate * 100; - return ( -
-
- {provider} - = 95 - ? "text-green-600" - : percentage >= 80 - ? "text-yellow-600" - : "text-red-600", - )} - > - {percentage.toFixed(1)}% - -
-
-
= 95 - ? "bg-green-500" - : percentage >= 80 - ? "bg-yellow-500" - : "bg-red-500", - )} - style={{ width: `${percentage}%` }} - /> -
-
- ); - })} -
- ); -} - -// ============================================================================ -// 分布图组件 -// ============================================================================ - -interface DistributionChartProps { - data: Distribution; - formatValue?: (value: number) => string; - color?: string; -} - -function DistributionChart({ - data, - formatValue = (v) => v.toString(), - color = "bg-blue-500", -}: DistributionChartProps) { - if (data.buckets.length === 0) { - return ( -
- 暂无数据 -
- ); - } - - const maxValue = Math.max(...data.buckets.map(([, v]) => v), 1); - - return ( -
- {data.buckets.slice(0, 10).map(([label, value]) => { - const percentage = (value / maxValue) * 100; - const totalPercentage = data.total > 0 ? (value / data.total) * 100 : 0; - return ( -
-
- - {label} - - - {formatValue(value)} ({totalPercentage.toFixed(1)}%) - -
-
-
-
-
- ); - })} - {data.buckets.length > 10 && ( -
- 还有 {data.buckets.length - 10} 项未显示 -
- )} -
- 总计: {formatValue(data.total)} -
-
- ); -} - -// ============================================================================ -// 直方图组件 -// ============================================================================ - -interface HistogramChartProps { - data: Distribution; - color?: string; -} - -function HistogramChart({ - data, - color = "bg-purple-500", -}: HistogramChartProps) { - if (data.buckets.length === 0) { - return ( -
- 暂无数据 -
- ); - } - - const maxValue = Math.max(...data.buckets.map(([, v]) => v), 1); - - return ( -
- {/* 直方图 */} -
- {data.buckets.map(([label, value], index) => { - const height = (value / maxValue) * 100; - const percentage = data.total > 0 ? (value / data.total) * 100 : 0; - return ( -
-
0 ? "4px" : "0", - }} - title={`${label}: ${value} (${percentage.toFixed(1)}%)`} - /> -
- ); - })} -
- - {/* X 轴标签 */} -
- {data.buckets.map(([label], index) => ( -
- {label} -
- ))} -
- - {/* 总计 */} -
- 总计: {data.total} 请求 -
-
- ); -} - -// ============================================================================ -// 紧凑模式组件 -// ============================================================================ - -interface CompactStatsProps { - stats: FlowStatsType; - loading: boolean; - onRefresh: () => void; - lastUpdated: Date | null; -} - -function CompactStats({ - stats, - loading, - onRefresh, - lastUpdated, -}: CompactStatsProps) { - return ( -
-
- 统计概览 - -
-
-
-
{stats.total_requests}
-
请求
-
-
-
= 0.95 - ? "text-green-600" - : stats.success_rate >= 0.8 - ? "text-yellow-600" - : "text-red-600", - )} - > - {(stats.success_rate * 100).toFixed(0)}% -
-
成功率
-
-
-
- {formatLatency(stats.avg_latency_ms)} -
-
平均延迟
-
-
-
- {formatTokenCount( - stats.total_input_tokens + stats.total_output_tokens, - )} -
-
Token
-
-
-
- ); -} - -// ============================================================================ -// 统计卡片组件 -// ============================================================================ - -interface StatCardProps { - title: string; - value: string; - icon: React.ReactNode; - trend?: "up" | "down" | null; - subtitle?: string; - valueColor?: string; -} - -function StatCard({ - title, - value, - icon, - trend, - subtitle, - valueColor, -}: StatCardProps) { - return ( -
-
- {title} - {icon} -
-
- {value} - {trend === "up" && } - {trend === "down" && } -
- {subtitle && ( -
{subtitle}
- )} -
- ); -} - -// ============================================================================ -// 状态条组件 -// ============================================================================ - -interface StatusBarProps { - label: string; - value: number; - total: number; - color: string; -} - -function StatusBar({ label, value, total, color }: StatusBarProps) { - const percentage = total > 0 ? (value / total) * 100 : 0; - - return ( -
-
- {label} - - {value} ({percentage.toFixed(1)}%) - -
-
-
-
-
- ); -} - -// ============================================================================ -// 提供商分布组件 -// ============================================================================ - -interface ProviderDistributionProps { - providers: ProviderStats[]; - total: number; -} - -function ProviderDistribution({ providers, total }: ProviderDistributionProps) { - const providerColors: Record = { - Kiro: "bg-purple-500", - Gemini: "bg-blue-500", - OpenAI: "bg-green-500", - Claude: "bg-orange-500", - Qwen: "bg-cyan-500", - Antigravity: "bg-pink-500", - Vertex: "bg-indigo-500", - GeminiApiKey: "bg-blue-400", - Codex: "bg-emerald-500", - ClaudeOAuth: "bg-amber-500", - IFlow: "bg-rose-500", - }; - - const sortedProviders = [...providers].sort((a, b) => b.count - a.count); - - return ( -
- {/* 分布条 */} -
- {sortedProviders.map((provider) => { - const percentage = total > 0 ? (provider.count / total) * 100 : 0; - if (percentage < 1) return null; - return ( -
- ); - })} -
- - {/* 详细列表 */} -
- {sortedProviders.map((provider) => { - const percentage = total > 0 ? (provider.count / total) * 100 : 0; - return ( -
-
-
-
- {provider.provider} -
-
- {provider.count} 次 ({percentage.toFixed(1)}%) -
-
-
-
= 0.95 - ? "text-green-600" - : provider.success_rate >= 0.8 - ? "text-yellow-600" - : "text-red-600", - )} - > - {(provider.success_rate * 100).toFixed(0)}% -
-
- {formatLatency(provider.avg_latency_ms)} -
-
-
- ); - })} -
-
- ); -} - -// ============================================================================ -// 模型分布组件 -// ============================================================================ - -interface ModelDistributionProps { - models: ModelStats[]; - total: number; -} - -function ModelDistribution({ models, total }: ModelDistributionProps) { - const sortedModels = [...models].sort((a, b) => b.count - a.count); - const topModels = sortedModels.slice(0, 10); // 只显示前 10 个模型 - - return ( -
- {topModels.map((model, index) => { - const percentage = total > 0 ? (model.count / total) * 100 : 0; - return ( -
-
-
- - {index + 1}. - - - {model.model} - -
-
- = 0.95 - ? "text-green-600" - : model.success_rate >= 0.8 - ? "text-yellow-600" - : "text-red-600", - )} - > - {(model.success_rate * 100).toFixed(0)}% - - - {formatLatency(model.avg_latency_ms)} - - - {model.count} ({percentage.toFixed(1)}%) - -
-
-
-
-
-
- ); - })} - {sortedModels.length > 10 && ( -
- 还有 {sortedModels.length - 10} 个模型未显示 -
- )} -
- ); -} - -// ============================================================================ -// 状态分布组件 -// ============================================================================ - -interface StateDistributionProps { - states: { state: string; count: number }[]; - total: number; -} - -function StateDistribution({ states, total }: StateDistributionProps) { - const stateColors: Record = { - Completed: "bg-green-500", - Failed: "bg-red-500", - Streaming: "bg-blue-500", - Pending: "bg-yellow-500", - Cancelled: "bg-gray-500", - }; - - const stateLabels: Record = { - Completed: "已完成", - Failed: "失败", - Streaming: "流式传输中", - Pending: "等待中", - Cancelled: "已取消", - }; - - const sortedStates = [...states].sort((a, b) => b.count - a.count); - - return ( -
- {/* 分布条 */} -
- {sortedStates.map((state) => { - const percentage = total > 0 ? (state.count / total) * 100 : 0; - if (percentage < 1) return null; - return ( -
- ); - })} -
- - {/* 图例 */} -
- {sortedStates.map((state) => { - const percentage = total > 0 ? (state.count / total) * 100 : 0; - return ( -
-
- - {stateLabels[state.state] || state.state} - - - {state.count} ({percentage.toFixed(1)}%) - -
- ); - })} -
-
- ); -} - -export default FlowStats; diff --git a/src/components/flow-monitor/FlowTimeline.tsx b/src/components/flow-monitor/FlowTimeline.tsx deleted file mode 100644 index 1b666f90c..000000000 --- a/src/components/flow-monitor/FlowTimeline.tsx +++ /dev/null @@ -1,435 +0,0 @@ -import React from "react"; -import { - Clock, - ArrowRight, - CheckCircle2, - XCircle, - Loader2, - Zap, - Send, - Download, -} from "lucide-react"; -import type { LLMFlow, FlowTimestamps } from "@/lib/api/flowMonitor"; -import { formatLatency } from "@/lib/api/flowMonitor"; -import { cn } from "@/lib/utils"; - -interface FlowTimelineProps { - flow: LLMFlow; - className?: string; -} - -interface TimelineEvent { - id: string; - label: string; - timestamp: Date | null; - icon: React.ReactNode; - color: string; - duration?: number; - durationLabel?: string; -} - -export function FlowTimeline({ flow, className }: FlowTimelineProps) { - const { timestamps, state, response } = flow; - - // 构建时间线事件 - const events = buildTimelineEvents(timestamps, state, response?.stream_info); - - // 计算时间范围 - const validTimestamps = events - .filter((e) => e.timestamp !== null) - .map((e) => e.timestamp!.getTime()); - - if (validTimestamps.length === 0) { - return ( -
-
暂无时间线数据
-
- ); - } - - const minTime = Math.min(...validTimestamps); - const maxTime = Math.max(...validTimestamps); - const totalDuration = maxTime - minTime; - - return ( -
-

- - 请求时间线 -

- - {/* 总耗时 */} -
- 总耗时 - - {formatLatency(timestamps.duration_ms)} - -
- - {/* 时间线可视化 */} -
- {/* 时间轴背景 */} -
- - {/* 事件列表 */} -
- {events.map((event) => ( - - ))} -
-
- - {/* 时间分布条 */} - -
- ); -} - -function buildTimelineEvents( - timestamps: FlowTimestamps, - state: string, - streamInfo?: { first_chunk_latency_ms: number; chunk_count: number }, -): TimelineEvent[] { - const events: TimelineEvent[] = []; - - // 创建时间 - events.push({ - id: "created", - label: "Flow 创建", - timestamp: new Date(timestamps.created), - icon: , - color: "text-gray-500", - }); - - // 请求开始 - events.push({ - id: "request_start", - label: "请求开始", - timestamp: new Date(timestamps.request_start), - icon: , - color: "text-blue-500", - }); - - // 请求结束 - if (timestamps.request_end) { - const requestDuration = - new Date(timestamps.request_end).getTime() - - new Date(timestamps.request_start).getTime(); - events.push({ - id: "request_end", - label: "请求发送完成", - timestamp: new Date(timestamps.request_end), - icon: , - color: "text-blue-500", - duration: requestDuration, - durationLabel: `请求耗时 ${formatLatency(requestDuration)}`, - }); - } - - // 响应开始 (TTFB) - if (timestamps.response_start) { - events.push({ - id: "response_start", - label: "首字节到达 (TTFB)", - timestamp: new Date(timestamps.response_start), - icon: , - color: "text-green-500", - duration: timestamps.ttfb_ms, - durationLabel: timestamps.ttfb_ms - ? `TTFB ${formatLatency(timestamps.ttfb_ms)}` - : undefined, - }); - } - - // 流式响应信息 - if (streamInfo && streamInfo.chunk_count > 0) { - events.push({ - id: "streaming", - label: `流式传输 (${streamInfo.chunk_count} chunks)`, - timestamp: timestamps.response_start - ? new Date(timestamps.response_start) - : null, - icon: , - color: "text-purple-500", - durationLabel: `首 chunk ${formatLatency(streamInfo.first_chunk_latency_ms)}`, - }); - } - - // 响应结束 - if (timestamps.response_end) { - const isSuccess = state === "Completed"; - const isFailed = state === "Failed"; - - events.push({ - id: "response_end", - label: isSuccess ? "响应完成" : isFailed ? "请求失败" : "响应结束", - timestamp: new Date(timestamps.response_end), - icon: isSuccess ? ( - - ) : isFailed ? ( - - ) : ( - - ), - color: isSuccess - ? "text-green-500" - : isFailed - ? "text-red-500" - : "text-gray-500", - duration: timestamps.duration_ms, - durationLabel: `总耗时 ${formatLatency(timestamps.duration_ms)}`, - }); - } - - return events; -} - -interface TimelineEventItemProps { - event: TimelineEvent; - totalDuration: number; - minTime: number; -} - -function TimelineEventItem({ - event, - totalDuration, - minTime, -}: TimelineEventItemProps) { - const formatTime = (date: Date | null) => { - if (!date) return "-"; - // 格式化时间,包含毫秒 - const hours = date.getHours().toString().padStart(2, "0"); - const minutes = date.getMinutes().toString().padStart(2, "0"); - const seconds = date.getSeconds().toString().padStart(2, "0"); - const ms = date.getMilliseconds().toString().padStart(3, "0"); - return `${hours}:${minutes}:${seconds}.${ms}`; - }; - - // 计算相对位置百分比(保留用于未来可能的动画效果) - void (event.timestamp && totalDuration > 0 - ? ((event.timestamp.getTime() - minTime) / totalDuration) * 100 - : 0); - - return ( -
- {/* 时间点标记 */} -
- {event.icon} -
- - {/* 事件内容 */} -
-
- {event.label} - - {formatTime(event.timestamp)} - -
- {event.durationLabel && ( -
- {event.durationLabel} -
- )} -
-
- ); -} - -// ============================================================================ -// 时间分布条组件 -// ============================================================================ - -interface TimelineBarProps { - timestamps: FlowTimestamps; - streamInfo?: { first_chunk_latency_ms: number; chunk_count: number }; - className?: string; -} - -function TimelineBar({ timestamps, streamInfo, className }: TimelineBarProps) { - const totalDuration = timestamps.duration_ms; - - if (totalDuration === 0) { - return null; - } - - // 计算各阶段占比 - const phases: { - id: string; - label: string; - duration: number; - color: string; - percentage: number; - }[] = []; - - // 请求发送阶段 - if (timestamps.request_end) { - const requestDuration = - new Date(timestamps.request_end).getTime() - - new Date(timestamps.request_start).getTime(); - if (requestDuration > 0) { - phases.push({ - id: "request", - label: "请求发送", - duration: requestDuration, - color: "bg-blue-500", - percentage: (requestDuration / totalDuration) * 100, - }); - } - } - - // 等待响应阶段 (TTFB) - if (timestamps.ttfb_ms && timestamps.request_end) { - const waitDuration = - timestamps.ttfb_ms - - (new Date(timestamps.request_end).getTime() - - new Date(timestamps.request_start).getTime()); - if (waitDuration > 0) { - phases.push({ - id: "wait", - label: "等待响应", - duration: waitDuration, - color: "bg-yellow-500", - percentage: (waitDuration / totalDuration) * 100, - }); - } - } else if (timestamps.ttfb_ms) { - phases.push({ - id: "ttfb", - label: "TTFB", - duration: timestamps.ttfb_ms, - color: "bg-yellow-500", - percentage: (timestamps.ttfb_ms / totalDuration) * 100, - }); - } - - // 响应接收阶段 - if (timestamps.response_start && timestamps.response_end) { - const responseDuration = - new Date(timestamps.response_end).getTime() - - new Date(timestamps.response_start).getTime(); - if (responseDuration > 0) { - phases.push({ - id: "response", - label: streamInfo ? "流式接收" : "响应接收", - duration: responseDuration, - color: streamInfo ? "bg-purple-500" : "bg-green-500", - percentage: (responseDuration / totalDuration) * 100, - }); - } - } - - // 如果没有详细阶段,显示总时间 - if (phases.length === 0) { - phases.push({ - id: "total", - label: "总耗时", - duration: totalDuration, - color: "bg-gray-500", - percentage: 100, - }); - } - - return ( -
-
时间分布
- - {/* 进度条 */} -
- {phases.map((phase) => ( -
- ))} -
- - {/* 图例 */} -
- {phases.map((phase) => ( -
-
- {phase.label} - {formatLatency(phase.duration)} - - ({phase.percentage.toFixed(1)}%) - -
- ))} -
-
- ); -} - -// ============================================================================ -// 简化版时间线(用于列表预览) -// ============================================================================ - -interface FlowTimelineCompactProps { - timestamps: FlowTimestamps; - state: string; - className?: string; -} - -export function FlowTimelineCompact({ - timestamps, - state, - className, -}: FlowTimelineCompactProps) { - const totalDuration = timestamps.duration_ms; - const ttfb = timestamps.ttfb_ms || 0; - - // 计算 TTFB 占比 - const ttfbPercentage = totalDuration > 0 ? (ttfb / totalDuration) * 100 : 0; - const responsePercentage = 100 - ttfbPercentage; - - const isSuccess = state === "Completed"; - const isFailed = state === "Failed"; - - return ( -
-
- {ttfb > 0 && ( -
- )} -
-
-
- {ttfb > 0 ? `TTFB ${formatLatency(ttfb)}` : ""} - {formatLatency(totalDuration)} -
-
- ); -} - -export default FlowTimeline; diff --git a/src/components/flow-monitor/InterceptEditor.tsx b/src/components/flow-monitor/InterceptEditor.tsx deleted file mode 100644 index 5ca19508b..000000000 --- a/src/components/flow-monitor/InterceptEditor.tsx +++ /dev/null @@ -1,678 +0,0 @@ -/** - * 拦截编辑器组件 - * - * 实现请求/响应编辑器和继续/取消按钮 - * **Validates: Requirements 2.2, 2.3, 2.4, 2.5** - */ - -import { useState, useEffect, useCallback } from "react"; -import { safeInvoke } from "@/lib/dev-bridge"; -import { - X, - Play, - Save, - RotateCcw, - AlertCircle, - Loader2, - ArrowRight, - ArrowLeft, - Clock, - Copy, - Check, -} from "lucide-react"; -import { cn } from "@/lib/utils"; -import type { - InterceptedFlow, - InterceptType, - InterceptState, -} from "./InterceptPanel"; -import type { LLMRequest, LLMResponse } from "@/lib/api/flowMonitor"; - -// ============================================================================ -// 组件属性 -// ============================================================================ - -interface InterceptEditorProps { - flowId: string; - onClose?: () => void; - onContinue?: () => void; - onCancel?: () => void; - className?: string; -} - -// ============================================================================ -// 主组件 -// ============================================================================ - -export function InterceptEditor({ - flowId, - onClose, - onContinue, - onCancel, - className, -}: InterceptEditorProps) { - // 状态 - const [interceptedFlow, setInterceptedFlow] = - useState(null); - const [loading, setLoading] = useState(true); - const [error, setError] = useState(null); - const [continuing, setContinuing] = useState(false); - const [cancelling, setCancelling] = useState(false); - - // 编辑状态 - const [editedContent, setEditedContent] = useState(""); - const [isModified, setIsModified] = useState(false); - const [parseError, setParseError] = useState(null); - - // 视图模式 - const [viewMode, setViewMode] = useState<"formatted" | "raw">("formatted"); - const [copied, setCopied] = useState(false); - - // 加载被拦截的 Flow - const loadInterceptedFlow = useCallback(async () => { - try { - setLoading(true); - setError(null); - const flow = await safeInvoke( - "intercept_get_flow", - { - flowId, - }, - ); - if (flow) { - setInterceptedFlow(flow); - // 初始化编辑内容 - const content = getOriginalContent(flow); - setEditedContent(JSON.stringify(content, null, 2)); - setIsModified(false); - } else { - setError("拦截的 Flow 不存在或已处理"); - } - } catch (e) { - console.error("加载拦截 Flow 失败:", e); - setError(e instanceof Error ? e.message : "加载失败"); - } finally { - setLoading(false); - } - }, [flowId]); - - // 获取原始内容 - const getOriginalContent = (flow: InterceptedFlow): unknown => { - if (flow.intercept_type === "request") { - return flow.original_request; - } else { - return flow.original_response; - } - }; - - // 处理内容变更 - const handleContentChange = (value: string) => { - setEditedContent(value); - setIsModified(true); - - // 验证 JSON - try { - JSON.parse(value); - setParseError(null); - } catch (e) { - setParseError(e instanceof Error ? e.message : "JSON 解析错误"); - } - }; - - // 重置内容 - const handleReset = () => { - if (interceptedFlow) { - const content = getOriginalContent(interceptedFlow); - setEditedContent(JSON.stringify(content, null, 2)); - setIsModified(false); - setParseError(null); - } - }; - - // 继续处理 - const handleContinue = async () => { - if (!interceptedFlow) return; - - try { - setContinuing(true); - setError(null); - - let modifiedRequest: LLMRequest | null = null; - let modifiedResponse: LLMResponse | null = null; - - // 如果有修改,解析修改后的内容 - if (isModified && !parseError) { - try { - const parsed = JSON.parse(editedContent); - if (interceptedFlow.intercept_type === "request") { - modifiedRequest = parsed as LLMRequest; - } else { - modifiedResponse = parsed as LLMResponse; - } - } catch (_e) { - setError("JSON 解析失败,请检查格式"); - return; - } - } - - await safeInvoke("intercept_continue", { - flowId: interceptedFlow.flow_id, - modifiedRequest, - modifiedResponse, - }); - - onContinue?.(); - onClose?.(); - } catch (e) { - console.error("继续 Flow 失败:", e); - setError(e instanceof Error ? e.message : "操作失败"); - } finally { - setContinuing(false); - } - }; - - // 取消处理 - const handleCancel = async () => { - if (!interceptedFlow) return; - - try { - setCancelling(true); - setError(null); - - await safeInvoke("intercept_cancel", { - flowId: interceptedFlow.flow_id, - }); - - onCancel?.(); - onClose?.(); - } catch (e) { - console.error("取消 Flow 失败:", e); - setError(e instanceof Error ? e.message : "操作失败"); - } finally { - setCancelling(false); - } - }; - - // 复制内容 - const handleCopy = async () => { - try { - await navigator.clipboard.writeText(editedContent); - setCopied(true); - setTimeout(() => setCopied(false), 2000); - } catch (e) { - console.error("复制失败:", e); - } - }; - - // 初始化加载 - useEffect(() => { - loadInterceptedFlow(); - }, [loadInterceptedFlow]); - - // 格式化时间 - const formatTime = (timestamp: string) => { - return new Date(timestamp).toLocaleString("zh-CN"); - }; - - // 获取状态标签 - const getStateLabel = (state: InterceptState) => { - const labels: Record = { - pending: "等待处理", - editing: "编辑中", - continued: "已继续", - cancelled: "已取消", - timedout: "已超时", - }; - return labels[state] || state; - }; - - if (loading) { - return ( -
- -
- ); - } - - if (error && !interceptedFlow) { - return ( -
-
- - {error} -
- {onClose && ( - - )} -
- ); - } - - if (!interceptedFlow) { - return null; - } - - return ( -
- {/* 头部 */} -
-
- {interceptedFlow.intercept_type === "request" ? ( - - ) : ( - - )} -
-
- {interceptedFlow.intercept_type === "request" - ? "拦截请求" - : "拦截响应"} -
-
- - {interceptedFlow.flow_id.slice(0, 12)}... - - • - {getStateLabel(interceptedFlow.state)} -
-
-
-
- {/* 视图模式切换 */} -
- - -
- {onClose && ( - - )} -
-
- - {/* 信息栏 */} -
-
- - - {formatTime(interceptedFlow.intercepted_at)} - - {isModified && ( - 已修改 - )} -
-
- - {isModified && ( - - )} -
-
- - {/* 错误提示 */} - {error && ( -
- - {error} -
- )} - - {/* 解析错误提示 */} - {parseError && ( -
- - JSON 格式错误: {parseError} -
- )} - - {/* 编辑区域 */} -
- {viewMode === "formatted" ? ( - - ) : ( -