diff --git a/docs/LLM_FLOW_MONITOR_SPEC.md b/docs/LLM_FLOW_MONITOR_SPEC.md new file mode 100644 index 000000000..2c31eba92 --- /dev/null +++ b/docs/LLM_FLOW_MONITOR_SPEC.md @@ -0,0 +1,1565 @@ +# LLM Flow Monitor - 详细设计方案 + +> 参考 mitmproxy 的 Flow 模型,为 ProxyCast 设计一套完整的 LLM API 流量监控系统, +> 用于捕获、存储、分析和回放 AI Agent 与大模型之间的完整交互数据。 + +## 一、背景与目标 + +### 1.1 当前问题 + +1. **日志信息不完整**:当前 `RequestLog` 只记录元数据(id、provider、model、duration、tokens),不保存完整的请求和响应内容 +2. **流式响应丢失**:SSE 流式响应的 chunks 分散,无法重建完整的响应内容 +3. **无法调试 Agent**:开发 AI Agent 时,需要查看完整的 prompt 和 response 来调优 +4. **缺乏历史回放**:无法回放历史请求,难以复现问题 +5. **数据不可导出**:无法导出为标准格式(如 HAR)供其他工具分析 + +### 1.2 设计目标 + +1. **完整捕获**:记录每个请求的完整 headers、body、响应内容 +2. **流式重建**:自动将 SSE chunks 合并为完整响应 +3. **高效存储**:内存 + 文件双层存储,支持大量请求 +4. **灵活查询**:按时间、模型、provider、内容等多维度过滤 +5. **标准导出**:支持 HAR、JSON、Markdown 等格式导出 +6. **实时监控**:前端实时展示请求列表和详情 +7. **隐私保护**:敏感信息脱敏,可配置存储策略 + +--- + +## 二、数据模型设计 + +### 2.1 核心数据结构 + +```rust +/// 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, +} + +/// 流类型 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub enum FlowType { + /// OpenAI Chat Completions + ChatCompletions, + /// Anthropic Messages + AnthropicMessages, + /// Gemini Generate Content + GeminiGenerateContent, + /// Embeddings + Embeddings, + /// 其他 + Other(String), +} + +/// 流状态 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub enum FlowState { + /// 等待响应 + Pending, + /// 正在流式传输 + Streaming, + /// 已完成 + Completed, + /// 失败 + Failed, + /// 已取消 + Cancelled, + /// 已拦截(用于调试) + Intercepted, +} +``` + +### 2.2 请求数据结构 + +```rust +/// 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, +} + +/// 消息结构 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct Message { + /// 角色 + pub role: MessageRole, + + /// 内容(可以是文本或多模态) + pub content: MessageContent, + + /// 工具调用(assistant 消息) + pub tool_calls: Option>, + + /// 工具结果(tool 消息) + pub tool_result: Option, + + /// 消息名称(function/tool 消息) + pub name: Option, +} + +/// 消息角色 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub enum MessageRole { + System, + User, + Assistant, + Tool, + Function, +} + +/// 消息内容(支持多模态) +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(untagged)] +pub enum MessageContent { + /// 纯文本 + Text(String), + + /// 多模态内容 + MultiModal(Vec), +} + +/// 内容部分 +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "type")] +pub enum ContentPart { + /// 文本 + #[serde(rename = "text")] + Text { text: String }, + + /// 图片 + #[serde(rename = "image_url")] + Image { + image_url: ImageUrl, + /// 图片摘要(用于显示,不存储完整 base64) + #[serde(skip_serializing_if = "Option::is_none")] + thumbnail: Option, + }, + + /// 音频 + #[serde(rename = "audio")] + Audio { + audio: AudioData, + }, + + /// 文件 + #[serde(rename = "file")] + File { + file: FileData, + }, +} + +/// 请求参数 +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +pub struct RequestParameters { + /// 温度 + pub temperature: Option, + /// Top P + pub top_p: Option, + /// 最大 tokens + pub max_tokens: Option, + /// 停止序列 + pub stop: Option>, + /// 是否流式 + pub stream: bool, + /// 其他参数 + #[serde(flatten)] + pub extra: HashMap, +} +``` + +### 2.3 响应数据结构 + +```rust +/// 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, +} + +/// 思维链内容 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ThinkingContent { + /// 思维链文本 + pub text: String, + /// 思维链 token 数 + pub tokens: Option, + /// 思维链签名(用于验证) + pub signature: Option, +} + +/// 工具调用 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ToolCall { + /// 调用 ID + pub id: String, + /// 工具类型 + pub call_type: String, + /// 函数信息 + pub function: FunctionCall, +} + +/// 函数调用 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct FunctionCall { + /// 函数名 + pub name: String, + /// 参数(JSON 字符串) + pub arguments: String, + /// 解析后的参数(方便查看) + pub parsed_arguments: Option, +} + +/// Token 使用统计 +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +pub struct TokenUsage { + /// 输入 tokens + pub input_tokens: u32, + /// 输出 tokens + pub output_tokens: u32, + /// 缓存读取 tokens + pub cache_read_tokens: Option, + /// 缓存写入 tokens + pub cache_write_tokens: Option, + /// 思维链 tokens + pub thinking_tokens: Option, + /// 总 tokens + pub total_tokens: u32, +} + +/// 停止原因 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub enum StopReason { + /// 正常结束 + Stop, + /// 达到长度限制 + Length, + /// 工具调用 + ToolUse, + /// 内容过滤 + ContentFilter, + /// 其他 + 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 { + /// 序号 + pub index: u32, + /// 时间戳 + pub timestamp: DateTime, + /// 原始数据 + pub data: String, + /// 增量内容 + pub delta_content: Option, +} +``` + +### 2.4 元数据结构 + +```rust +/// 流元数据 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct FlowMetadata { + /// Provider 类型 + pub provider: ProviderType, + + /// 使用的凭证 ID + pub credential_id: Option, + + /// 凭证名称(用于显示) + pub credential_name: Option, + + /// 重试次数 + pub retry_count: u32, + + /// 客户端信息 + pub client_info: ClientInfo, + + /// 路由信息 + pub routing_info: RoutingInfo, + + /// 注入的参数 + pub injected_params: Option>, + + /// 上下文使用率(%) + pub context_usage_percentage: Option, +} + +/// 客户端信息 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ClientInfo { + /// 客户端 IP + pub ip: Option, + /// User-Agent + pub user_agent: Option, + /// 客户端 SDK + pub sdk: Option, + /// 客户端版本 + pub sdk_version: Option, +} + +/// 路由信息 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct RoutingInfo { + /// 原始模型(别名) + pub original_model: String, + /// 解析后的模型 + pub resolved_model: String, + /// 路由到的 Provider + pub routed_provider: ProviderType, + /// 匹配的路由规则 + pub matched_rule: Option, +} + +/// 时间戳集合 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct FlowTimestamps { + /// 请求创建时间 + pub created: DateTime, + /// 请求发送时间 + pub request_start: DateTime, + /// 请求发送完成时间 + pub request_end: Option>, + /// 响应开始时间(收到第一个字节) + pub response_start: Option>, + /// 响应结束时间 + pub response_end: Option>, + /// 总耗时(毫秒) + pub duration_ms: u64, + /// TTFB(Time To First Byte,毫秒) + pub ttfb_ms: Option, +} + +/// 用户标注 +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +pub struct FlowAnnotations { + /// 用户标记(如 ⭐、🔴、🟢) + pub marker: Option, + /// 用户备注 + pub comment: Option, + /// 标签 + pub tags: Vec, + /// 是否已收藏 + pub starred: bool, +} +``` + +### 2.5 错误结构 + +```rust +/// 流错误 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct FlowError { + /// 错误类型 + pub error_type: FlowErrorType, + /// 错误消息 + pub message: String, + /// HTTP 状态码(如果有) + pub status_code: Option, + /// 原始错误响应 + pub raw_response: Option, + /// 错误发生时间 + pub timestamp: DateTime, + /// 是否可重试 + pub retryable: bool, +} + +/// 错误类型 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub enum FlowErrorType { + /// 网络错误 + Network, + /// 超时 + Timeout, + /// 认证失败 + Authentication, + /// 限流 + RateLimit, + /// 内容过滤 + ContentFilter, + /// 服务端错误 + ServerError, + /// 请求格式错误 + BadRequest, + /// 模型不可用 + ModelUnavailable, + /// Token 超限 + TokenLimitExceeded, + /// 其他 + Other, +} +``` + +--- + +## 三、流式响应重建 + +### 3.1 SSE 解析器 + +```rust +/// SSE 流重建器 +pub struct StreamRebuilder { + /// 累积的 chunks + chunks: Vec, + /// 累积的内容 + content_buffer: String, + /// 累积的 tool calls + tool_calls_buffer: HashMap, + /// 累积的 thinking + thinking_buffer: Option, + /// 第一个 chunk 时间 + first_chunk_time: Option>, + /// 上一个 chunk 时间 + last_chunk_time: Option>, + /// 流格式 + format: StreamFormat, +} + +/// 流格式 +pub enum StreamFormat { + /// OpenAI 格式 + OpenAI, + /// Anthropic 格式 + Anthropic, + /// Gemini 格式 + Gemini, + /// 未知格式 + Unknown, +} + +impl StreamRebuilder { + /// 处理一个 SSE 事件 + pub fn process_event(&mut self, event: &str, data: &str) -> Result<(), Error> { + let chunk = StreamChunk { + index: self.chunks.len() as u32, + timestamp: Utc::now(), + data: data.to_string(), + delta_content: None, + }; + + // 根据格式解析增量内容 + 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)?, + _ => {}, + } + + self.chunks.push(chunk); + Ok(()) + } + + /// 完成重建,返回完整响应 + pub fn finish(self) -> LLMResponse { + // 构建完整的响应对象 + LLMResponse { + content: self.content_buffer, + tool_calls: self.tool_calls_buffer.into_values().map(|b| b.build()).collect(), + thinking: self.thinking_buffer.map(|t| ThinkingContent { text: t, tokens: None, signature: None }), + stream_info: Some(StreamInfo { + chunk_count: self.chunks.len() as u32, + first_chunk_latency_ms: self.calculate_first_chunk_latency(), + avg_chunk_interval_ms: self.calculate_avg_interval(), + raw_chunks: if self.should_save_raw_chunks() { Some(self.chunks) } else { None }, + }), + // ... 其他字段 + } + } +} +``` + +### 3.2 不同格式处理 + +```rust +impl StreamRebuilder { + /// 处理 OpenAI 格式的 chunk + fn process_openai_chunk(&mut self, data: &str, chunk: &mut StreamChunk) -> Result<(), Error> { + if data == "[DONE]" { + return Ok(()); + } + + let parsed: OpenAIStreamChunk = serde_json::from_str(data)?; + + for choice in &parsed.choices { + if let Some(delta) = &choice.delta { + // 文本内容 + if let Some(content) = &delta.content { + self.content_buffer.push_str(content); + chunk.delta_content = Some(content.clone()); + } + + // 工具调用 + if let Some(tool_calls) = &delta.tool_calls { + for tc in tool_calls { + self.process_tool_call_delta(tc); + } + } + } + } + + Ok(()) + } + + /// 处理 Anthropic 格式的 chunk + fn process_anthropic_chunk(&mut self, event: &str, data: &str, chunk: &mut StreamChunk) -> Result<(), Error> { + match event { + "content_block_delta" => { + let parsed: AnthropicDelta = serde_json::from_str(data)?; + match &parsed.delta { + Delta::TextDelta { text } => { + self.content_buffer.push_str(text); + chunk.delta_content = Some(text.clone()); + }, + Delta::ThinkingDelta { thinking } => { + self.thinking_buffer.get_or_insert(String::new()).push_str(thinking); + }, + Delta::InputJsonDelta { partial_json } => { + // 处理工具调用参数 + self.process_tool_call_json_delta(parsed.index, partial_json); + }, + } + }, + "content_block_start" => { + // 处理新的内容块 + }, + "message_delta" => { + // 处理消息级别的更新(stop_reason, usage 等) + }, + _ => {}, + } + + Ok(()) + } +} +``` + +--- + +## 四、存储系统设计 + +### 4.1 双层存储架构 + +``` +┌─────────────────────────────────────────────────────┐ +│ 查询层 │ +│ (按 ID / 时间 / 模型 / Provider / 内容 查询) │ +└─────────────────────────────────────────────────────┘ + │ + ┌───────────────┼───────────────┐ + ▼ ▼ ▼ +┌─────────────┐ ┌─────────────┐ ┌─────────────┐ +│ 内存缓存 │ │ 索引层 │ │ 文件层 │ +│ (热数据) │ │ (SQLite) │ │ (JSONL) │ +│ 最近 1000 │ │ 元数据索引 │ │ 完整数据 │ +└─────────────┘ └─────────────┘ └─────────────┘ +``` + +### 4.2 内存缓存 + +```rust +/// 内存 Flow 存储 +pub struct FlowMemoryStore { + /// 按 ID 索引的 flows + flows: HashMap>>, + /// 按时间排序的 flow IDs + ordered_ids: VecDeque, + /// 最大缓存数量 + max_size: usize, + /// 内存使用估算 + memory_usage: AtomicUsize, +} + +impl FlowMemoryStore { + /// 添加 flow + pub fn add(&mut self, flow: LLMFlow) { + let id = flow.id.clone(); + let size = self.estimate_size(&flow); + + self.flows.insert(id.clone(), Arc::new(RwLock::new(flow))); + self.ordered_ids.push_back(id); + self.memory_usage.fetch_add(size, Ordering::Relaxed); + + // 驱逐旧数据 + while self.ordered_ids.len() > self.max_size { + if let Some(old_id) = self.ordered_ids.pop_front() { + if let Some(old_flow) = self.flows.remove(&old_id) { + let old_size = self.estimate_size(&old_flow.read()); + self.memory_usage.fetch_sub(old_size, Ordering::Relaxed); + } + } + } + } + + /// 获取最近 N 条 + pub fn get_recent(&self, limit: usize) -> Vec>> { + self.ordered_ids + .iter() + .rev() + .take(limit) + .filter_map(|id| self.flows.get(id).cloned()) + .collect() + } +} +``` + +### 4.3 文件持久化 + +```rust +/// Flow 文件存储 +pub struct FlowFileStore { + /// 存储目录 + base_dir: PathBuf, + /// 当前写入文件 + current_file: RwLock>, + /// 轮转配置 + rotation_config: RotationConfig, +} + +/// 轮转配置 +pub struct RotationConfig { + /// 按日期轮转 + pub rotate_daily: bool, + /// 单文件最大大小 + pub max_file_size: u64, + /// 保留天数 + pub retention_days: u32, + /// 是否压缩旧文件 + pub compress_old: bool, +} + +impl FlowFileStore { + /// 存储文件结构: + /// ~/.proxycast/flows/ + /// ├── 2024-01-15/ + /// │ ├── flows_001.jsonl + /// │ ├── flows_002.jsonl + /// │ └── index.sqlite (当日索引) + /// ├── 2024-01-14/ + /// │ ├── flows.jsonl.gz (压缩后) + /// │ └── index.sqlite + /// └── global_index.sqlite (全局索引) + + /// 写入 flow + pub async fn write(&self, flow: &LLMFlow) -> Result<(), Error> { + let mut writer = self.get_or_create_writer().await?; + + // 写入 JSONL + let json = serde_json::to_string(flow)?; + writer.write_line(&json).await?; + + // 更新索引 + self.update_index(flow).await?; + + // 检查是否需要轮转 + if writer.size() > self.rotation_config.max_file_size { + self.rotate().await?; + } + + Ok(()) + } + + /// 按条件查询 + pub async fn query(&self, filter: &FlowFilter) -> Result, Error> { + // 先查询索引获取文件位置 + let locations = self.query_index(filter).await?; + + // 从文件读取 + let mut flows = Vec::new(); + for loc in locations { + let flow = self.read_flow(&loc).await?; + if filter.matches(&flow) { + flows.push(flow); + } + } + + Ok(flows) + } +} +``` + +### 4.4 SQLite 索引 + +```sql +-- 全局索引表 +CREATE TABLE flow_index ( + id TEXT PRIMARY KEY, + created_at DATETIME 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 BOOLEAN DEFAULT FALSE, + has_tool_calls BOOLEAN DEFAULT FALSE, + has_thinking BOOLEAN DEFAULT FALSE, + file_path TEXT NOT NULL, + file_offset INTEGER NOT NULL, + -- 用于全文搜索 + content_preview TEXT, + request_preview TEXT +); + +CREATE INDEX idx_created_at ON flow_index(created_at); +CREATE INDEX idx_provider ON flow_index(provider); +CREATE INDEX idx_model ON flow_index(model); +CREATE INDEX idx_status ON flow_index(status); + +-- 全文搜索表(可选,使用 FTS5) +CREATE VIRTUAL TABLE flow_fts USING fts5( + id, + content, + request, + thinking, + content='flow_index' +); +``` + +--- + +## 五、查询与过滤 + +### 5.1 过滤器设计 + +```rust +/// Flow 过滤器 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct FlowFilter { + /// 时间范围 + pub time_range: Option, + + /// Provider 过滤 + pub providers: Option>, + + /// 模型过滤(支持通配符) + pub models: Option>, + + /// 状态过滤 + pub states: Option>, + + /// 是否有错误 + pub has_error: Option, + + /// 是否有工具调用 + pub has_tool_calls: Option, + + /// 是否有思维链 + pub has_thinking: Option, + + /// 是否流式 + pub is_streaming: Option, + + /// 内容搜索(全文) + pub content_search: Option, + + /// 请求内容搜索 + pub request_search: Option, + + /// Token 范围 + pub token_range: Option, + + /// 延迟范围 + pub latency_range: Option, + + /// 标签过滤 + pub tags: Option>, + + /// 只显示收藏 + pub starred_only: bool, + + /// 凭证 ID + pub credential_id: Option, +} + +/// 排序选项 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub enum FlowSortBy { + /// 创建时间(默认) + CreatedAt, + /// 耗时 + Duration, + /// Token 数 + TotalTokens, + /// 内容长度 + ContentLength, + /// 模型 + Model, +} +``` + +### 5.2 查询 API + +```rust +/// Flow 查询服务 +pub struct FlowQueryService { + memory_store: Arc, + file_store: Arc, +} + +impl FlowQueryService { + /// 查询 flows + pub async fn query(&self, + filter: FlowFilter, + sort_by: FlowSortBy, + sort_desc: bool, + page: usize, + page_size: usize, + ) -> Result { + // 优先从内存查询 + let mut flows = self.memory_store.query(&filter); + + // 如果需要更多数据,从文件查询 + if flows.len() < page * page_size { + let file_flows = self.file_store.query(&filter).await?; + flows.extend(file_flows); + } + + // 排序 + self.sort_flows(&mut flows, sort_by, sort_desc); + + // 分页 + let total = flows.len(); + let start = page * page_size; + let end = (start + page_size).min(total); + let flows = flows[start..end].to_vec(); + + Ok(FlowQueryResult { + flows, + total, + page, + page_size, + }) + } + + /// 获取统计信息 + pub async fn get_stats(&self, filter: &FlowFilter) -> FlowStats { + // 计算聚合统计 + } + + /// 全文搜索 + pub async fn search(&self, query: &str, limit: usize) -> Vec { + // 使用 FTS 搜索 + } +} +``` + +--- + +## 六、导出功能 + +### 6.1 支持的导出格式 + +```rust +/// 导出格式 +pub enum ExportFormat { + /// HAR (HTTP Archive) 格式 + HAR, + /// JSON 格式 + JSON, + /// JSONL (每行一个 JSON) + JSONL, + /// Markdown 格式(用于文档) + Markdown, + /// CSV 格式(仅元数据) + CSV, + /// OpenAI JSONL(用于 fine-tuning) + OpenAIFineTune, + /// Anthropic JSONL(用于 fine-tuning) + AnthropicFineTune, +} + +/// 导出选项 +pub struct ExportOptions { + /// 导出格式 + pub format: ExportFormat, + /// 过滤器 + pub filter: FlowFilter, + /// 是否包含原始数据 + pub include_raw: bool, + /// 是否包含流式 chunks + pub include_stream_chunks: bool, + /// 是否脱敏 + pub redact_sensitive: bool, + /// 脱敏规则 + pub redaction_rules: Vec, + /// 是否压缩 + pub compress: bool, +} +``` + +### 6.2 HAR 导出 + +```rust +impl FlowExporter { + /// 导出为 HAR 格式 + pub fn export_har(&self, flows: &[LLMFlow]) -> HarArchive { + HarArchive { + log: HarLog { + version: "1.2".to_string(), + creator: HarCreator { + name: "ProxyCast".to_string(), + version: env!("CARGO_PKG_VERSION").to_string(), + }, + entries: flows.iter().map(|f| self.flow_to_har_entry(f)).collect(), + // LLM 特定扩展 + _llm_metadata: Some(LLMHarMetadata { + total_tokens: flows.iter().map(|f| f.response.as_ref().map(|r| r.usage.total_tokens).unwrap_or(0) as u64).sum(), + models_used: flows.iter().map(|f| f.request.model.clone()).collect::>().into_iter().collect(), + providers_used: flows.iter().map(|f| f.metadata.provider.to_string()).collect::>().into_iter().collect(), + }), + }, + } + } + + fn flow_to_har_entry(&self, flow: &LLMFlow) -> HarEntry { + HarEntry { + started_date_time: flow.timestamps.created.to_rfc3339(), + time: flow.timestamps.duration_ms as f64, + request: HarRequest { + method: flow.request.method.clone(), + url: format!("https://api.provider.com{}", flow.request.path), + http_version: "HTTP/1.1".to_string(), + headers: flow.request.headers.iter() + .map(|(k, v)| HarHeader { name: k.clone(), value: v.clone() }) + .collect(), + post_data: Some(HarPostData { + mime_type: "application/json".to_string(), + text: serde_json::to_string(&flow.request.body).unwrap(), + }), + // ... + }, + response: flow.response.as_ref().map(|r| HarResponse { + status: r.status_code as i32, + status_text: r.status_text.clone(), + headers: r.headers.iter() + .map(|(k, v)| HarHeader { name: k.clone(), value: v.clone() }) + .collect(), + content: HarContent { + size: r.size_bytes as i64, + mime_type: "application/json".to_string(), + text: Some(serde_json::to_string(&r.body).unwrap()), + }, + // ... + }), + // LLM 特定扩展 + _llm: Some(LLMHarExtension { + provider: flow.metadata.provider.to_string(), + model: flow.request.model.clone(), + input_tokens: flow.response.as_ref().map(|r| r.usage.input_tokens), + output_tokens: flow.response.as_ref().map(|r| r.usage.output_tokens), + has_tool_calls: flow.response.as_ref().map(|r| !r.tool_calls.is_empty()).unwrap_or(false), + has_thinking: flow.response.as_ref().and_then(|r| r.thinking.as_ref()).is_some(), + }), + } + } +} +``` + +### 6.3 Markdown 导出(用于文档和分享) + +```rust +impl FlowExporter { + /// 导出为 Markdown(用于复制分享) + pub fn export_markdown(&self, flow: &LLMFlow) -> String { + let mut md = String::new(); + + // 标题 + writeln!(md, "# LLM Request - {}", flow.id).unwrap(); + writeln!(md, "").unwrap(); + + // 元信息 + writeln!(md, "## Metadata").unwrap(); + writeln!(md, "- **Provider**: {}", flow.metadata.provider).unwrap(); + writeln!(md, "- **Model**: {}", flow.request.model).unwrap(); + writeln!(md, "- **Time**: {}", flow.timestamps.created).unwrap(); + writeln!(md, "- **Duration**: {}ms", flow.timestamps.duration_ms).unwrap(); + writeln!(md, "").unwrap(); + + // 请求 + writeln!(md, "## Request").unwrap(); + if let Some(system) = &flow.request.system_prompt { + writeln!(md, "### System Prompt").unwrap(); + writeln!(md, "```").unwrap(); + writeln!(md, "{}", system).unwrap(); + writeln!(md, "```").unwrap(); + } + + writeln!(md, "### Messages").unwrap(); + for msg in &flow.request.messages { + writeln!(md, "**{}**:", msg.role).unwrap(); + writeln!(md, "{}", msg.content.to_string()).unwrap(); + writeln!(md, "").unwrap(); + } + + // 响应 + if let Some(resp) = &flow.response { + writeln!(md, "## Response").unwrap(); + + if let Some(thinking) = &resp.thinking { + writeln!(md, "### Thinking").unwrap(); + writeln!(md, "
Click to expand").unwrap(); + writeln!(md, "").unwrap(); + writeln!(md, "{}", thinking.text).unwrap(); + writeln!(md, "
").unwrap(); + writeln!(md, "").unwrap(); + } + + writeln!(md, "### Content").unwrap(); + writeln!(md, "{}", resp.content).unwrap(); + + if !resp.tool_calls.is_empty() { + writeln!(md, "### Tool Calls").unwrap(); + for tc in &resp.tool_calls { + writeln!(md, "- **{}**: `{}`", tc.function.name, tc.function.arguments).unwrap(); + } + } + + writeln!(md, "### Usage").unwrap(); + writeln!(md, "- Input: {} tokens", resp.usage.input_tokens).unwrap(); + writeln!(md, "- Output: {} tokens", resp.usage.output_tokens).unwrap(); + } + + md + } +} +``` + +--- + +## 七、前端界面设计 + +### 7.1 流量列表视图 + +``` +┌─────────────────────────────────────────────────────────────────────────────┐ +│ 🔍 Search... │ Provider ▾ │ Model ▾ │ Status ▾ │ Time Range ▾ │ ⚙️ Export │ +├─────────────────────────────────────────────────────────────────────────────┤ +│ │ +│ ┌─────────────────────────────────────────────────────────────────────────┐ │ +│ │ ⭐ 14:32:05 │ claude-sonnet-4-5 │ Kiro │ ✅ 2.3s │ 1.2k→3.4k │ 🔧 tool │ │ +│ │ "请帮我分析这段代码的性能问题..." │ │ +│ └─────────────────────────────────────────────────────────────────────────┘ │ +│ │ +│ ┌─────────────────────────────────────────────────────────────────────────┐ │ +│ │ 14:31:42 │ gemini-2.5-flash │ Gemini │ ✅ 0.8s │ 500→1.2k │ │ │ +│ │ "Write a Python function to..." │ │ +│ └─────────────────────────────────────────────────────────────────────────┘ │ +│ │ +│ ┌─────────────────────────────────────────────────────────────────────────┐ │ +│ │ 14:31:15 │ claude-sonnet-4-5 │ Kiro │ ❌ 5.2s │ Error: Rate limit │ │ +│ │ "Explain the difference between..." │ │ +│ └─────────────────────────────────────────────────────────────────────────┘ │ +│ │ +└─────────────────────────────────────────────────────────────────────────────┘ +``` + +### 7.2 流量详情视图 + +``` +┌─────────────────────────────────────────────────────────────────────────────┐ +│ ← Back │ Request abc123 │ ⭐ Star │ 📋 Copy │ 📤 Export │ 🔄 Replay │ +├─────────────────────────────────────────────────────────────────────────────┤ +│ │ +│ ┌─ Metadata ────────────────────────────────────────────────────────────┐ │ +│ │ Provider: Kiro Model: claude-sonnet-4-5 │ │ +│ │ Duration: 2.3s TTFB: 1.2s │ │ +│ │ Tokens: 1,234 → 3,456 Cost: $0.045 │ │ +│ │ Credential: work-account-1 │ │ +│ └───────────────────────────────────────────────────────────────────────┘ │ +│ │ +│ ┌─ Request ─────────────────────────────────────────────────────────────┐ │ +│ │ [Headers] [Body] [Messages] [Tools] │ │ +│ │ │ │ +│ │ System: You are a helpful assistant... │ │ +│ │ │ │ +│ │ User: 请帮我分析这段代码的性能问题: │ │ +│ │ ```python │ │ +│ │ def slow_function(): │ │ +│ │ for i in range(10000): │ │ +│ │ result = expensive_operation(i) │ │ +│ │ ``` │ │ +│ └───────────────────────────────────────────────────────────────────────┘ │ +│ │ +│ ┌─ Response ────────────────────────────────────────────────────────────┐ │ +│ │ [Content] [Thinking] [Tool Calls] [Raw] [Stream] │ │ +│ │ │ │ +│ │ 这段代码存在几个性能问题: │ │ +│ │ │ │ +│ │ 1. **循环中的重复计算**:`expensive_operation` 被调用 10000 次... │ │ +│ │ 2. **缺少缓存**:如果操作结果可以重用... │ │ +│ │ │ │ +│ │ [Show more...] │ │ +│ └───────────────────────────────────────────────────────────────────────┘ │ +│ │ +│ ┌─ Timeline ────────────────────────────────────────────────────────────┐ │ +│ │ Request ████░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░ 0.1s │ │ +│ │ TTFB ░░░░████████████████████░░░░░░░░░░░░░░░░░░░░░░░░░░░░ 1.2s │ │ +│ │ Stream ░░░░░░░░░░░░░░░░░░░░░░░░██████████████████████████░ 1.0s │ │ +│ │ Total ████████████████████████████████████████████████████ 2.3s │ │ +│ └───────────────────────────────────────────────────────────────────────┘ │ +│ │ +└─────────────────────────────────────────────────────────────────────────────┘ +``` + +### 7.3 统计仪表板 + +``` +┌─────────────────────────────────────────────────────────────────────────────┐ +│ 📊 Flow Statistics │ +├─────────────────────────────────────────────────────────────────────────────┤ +│ │ +│ ┌─ Overview ──────────────────────────┐ ┌─ Token Usage ─────────────────┐ │ +│ │ Total Requests │ 1,234 │ │ │ │ +│ │ Success Rate │ 98.2% │ │ ███████████ Input: 1.2M │ │ +│ │ Avg Latency │ 1.8s │ │ █████████████████ Output: 2.1M│ │ +│ │ Total Tokens │ 3.3M │ │ │ │ +│ └─────────────────────────────────────┘ └───────────────────────────────┘ │ +│ │ +│ ┌─ Requests by Provider ──────────────────────────────────────────────────┐│ +│ │ Kiro ██████████████████████████████████████████░░░░░░░░ 68% ││ +│ │ Gemini ████████████████░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░ 22% ││ +│ │ OpenAI ████████░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░ 10% ││ +│ └─────────────────────────────────────────────────────────────────────────┘│ +│ │ +│ ┌─ Latency Distribution ─────────────┐ ┌─ Requests Timeline ───────────┐ │ +│ │ ▃▅█▇▅▃▂▁ │ │ ▂▃▅▇█▇▅▃▂▁▂▃▅▇█▇▅▃▂ │ │ +│ │ 0s 1s 2s 3s 4s 5s+ │ │ 00:00 06:00 12:00 18:00│ │ +│ └────────────────────────────────────┘ └────────────────────────────────┘ │ +│ │ +└─────────────────────────────────────────────────────────────────────────────┘ +``` + +--- + +## 八、API 设计 + +### 8.1 Tauri Commands + +```rust +// 查询 flows +#[tauri::command] +async fn query_flows( + filter: FlowFilter, + sort_by: Option, + sort_desc: Option, + page: Option, + page_size: Option, + state: State<'_, FlowMonitorState>, +) -> Result; + +// 获取单个 flow 详情 +#[tauri::command] +async fn get_flow_detail( + id: String, + state: State<'_, FlowMonitorState>, +) -> Result; + +// 搜索 flows +#[tauri::command] +async fn search_flows( + query: String, + limit: Option, + state: State<'_, FlowMonitorState>, +) -> Result, String>; + +// 获取统计信息 +#[tauri::command] +async fn get_flow_stats( + filter: Option, + state: State<'_, FlowMonitorState>, +) -> Result; + +// 导出 flows +#[tauri::command] +async fn export_flows( + options: ExportOptions, + path: String, + state: State<'_, FlowMonitorState>, +) -> Result; + +// 更新 flow 标注 +#[tauri::command] +async fn update_flow_annotations( + id: String, + annotations: FlowAnnotations, + state: State<'_, FlowMonitorState>, +) -> Result<(), String>; + +// 重放请求 +#[tauri::command] +async fn replay_flow( + id: String, + modifications: Option, + state: State<'_, FlowMonitorState>, +) -> Result; + +// 清理旧数据 +#[tauri::command] +async fn cleanup_flows( + before: DateTime, + state: State<'_, FlowMonitorState>, +) -> Result; +``` + +### 8.2 WebSocket 实时推送 + +```rust +/// 实时 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 }, + /// 统计更新 + StatsUpdated { stats: FlowStats }, +} + +/// Flow 摘要(用于列表显示) +#[derive(Debug, Clone, Serialize)] +pub struct FlowSummary { + pub id: String, + pub provider: String, + pub model: String, + pub state: FlowState, + pub duration_ms: Option, + pub input_tokens: Option, + pub output_tokens: Option, + pub content_preview: String, + pub has_error: bool, + pub has_tool_calls: bool, + pub created_at: DateTime, +} +``` + +--- + +## 九、性能与隐私 + +### 9.1 性能优化 + +```rust +/// Flow 监控配置 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct FlowMonitorConfig { + /// 是否启用监控 + pub enabled: bool, + + /// 内存中最大 flow 数量 + pub max_memory_flows: usize, + + /// 是否保存到文件 + pub persist_to_file: bool, + + /// 文件保留天数 + pub retention_days: u32, + + /// 是否保存原始 stream chunks + pub save_stream_chunks: bool, + + /// 请求体大小限制(超过则截断) + pub max_request_body_size: usize, + + /// 响应体大小限制 + pub max_response_body_size: usize, + + /// 是否保存图片内容(base64) + pub save_image_content: bool, + + /// 图片缩略图大小 + pub thumbnail_size: (u32, u32), + + /// 采样率(0.0-1.0,用于高流量场景) + pub sampling_rate: f32, + + /// 排除的模型(不记录) + pub excluded_models: Vec, + + /// 排除的路径 + pub excluded_paths: Vec, +} + +impl Default for FlowMonitorConfig { + fn default() -> Self { + Self { + enabled: true, + max_memory_flows: 1000, + persist_to_file: true, + retention_days: 7, + save_stream_chunks: false, // 默认不保存原始 chunks + max_request_body_size: 1024 * 1024, // 1MB + max_response_body_size: 10 * 1024 * 1024, // 10MB + save_image_content: false, // 默认不保存图片 + thumbnail_size: (100, 100), + sampling_rate: 1.0, + excluded_models: vec![], + excluded_paths: vec!["/health".to_string()], + } + } +} +``` + +### 9.2 隐私保护 + +```rust +/// 脱敏规则 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct RedactionRule { + /// 规则名称 + pub name: String, + /// 匹配模式(正则) + pub pattern: String, + /// 替换内容 + pub replacement: String, + /// 应用位置 + pub apply_to: Vec, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub enum RedactionTarget { + /// 请求头 + RequestHeaders, + /// 请求体 + RequestBody, + /// 响应头 + ResponseHeaders, + /// 响应体 + ResponseBody, + /// 所有位置 + All, +} + +impl Default for Vec { + fn default() -> Self { + vec![ + // API Key 脱敏 + RedactionRule { + name: "api_key".to_string(), + pattern: r"(sk-[a-zA-Z0-9]{20,}|api[_-]?key[=:]\s*['\"]?)[a-zA-Z0-9\-_]+".to_string(), + replacement: "$1***REDACTED***".to_string(), + apply_to: vec![RedactionTarget::All], + }, + // Email 脱敏 + RedactionRule { + name: "email".to_string(), + pattern: r"[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}".to_string(), + replacement: "***@***.***".to_string(), + apply_to: vec![RedactionTarget::RequestBody, RedactionTarget::ResponseBody], + }, + // 手机号脱敏 + RedactionRule { + name: "phone".to_string(), + pattern: r"\b1[3-9]\d{9}\b".to_string(), + replacement: "1**********".to_string(), + apply_to: vec![RedactionTarget::RequestBody, RedactionTarget::ResponseBody], + }, + ] + } +} +``` + +--- + +## 十、实现路线图 + +### Phase 1: 基础设施(1-2 周) + +- [ ] 定义完整的数据模型(LLMFlow, LLMRequest, LLMResponse) +- [ ] 实现内存存储 FlowMemoryStore +- [ ] 实现 SSE 流重建器 StreamRebuilder +- [ ] 在现有 API handlers 中集成 flow 捕获 + +### Phase 2: 持久化与查询(1-2 周) + +- [ ] 实现文件存储 FlowFileStore +- [ ] 实现 SQLite 索引 +- [ ] 实现查询过滤器 +- [ ] 添加全文搜索支持 + +### Phase 3: 前端界面(2-3 周) + +- [ ] 实现 Flow 列表页面 +- [ ] 实现 Flow 详情页面 +- [ ] 实现统计仪表板 +- [ ] 实现实时更新(WebSocket) + +### Phase 4: 导出与高级功能(1-2 周) + +- [ ] 实现 HAR 导出 +- [ ] 实现 Markdown 导出 +- [ ] 实现请求重放 +- [ ] 实现隐私脱敏 + +### Phase 5: 优化与文档(1 周) + +- [ ] 性能优化 +- [ ] 编写用户文档 +- [ ] 添加测试用例 +- [ ] 发布 v1.0 + +--- + +## 十一、附录 + +### A. 与现有系统的集成点 + +1. **server/handlers/api.rs**: 在 `chat_completions` 和 `anthropic_messages` 函数中添加 flow 捕获 +2. **server_utils.rs**: 复用 `parse_cw_response` 用于流式响应解析 +3. **services/provider_pool_service.rs**: 获取凭证信息用于 metadata +4. **models/log_model.rs**: 将 RequestLog 与 LLMFlow 关联 + +### B. 参考实现 + +- [mitmproxy](https://github.com/mitmproxy/mitmproxy) - HTTP 流量捕获的黄金标准 +- [Charles Proxy](https://www.charlesproxy.com/) - 商业代理调试工具 +- [Fiddler](https://www.telerik.com/fiddler) - .NET 平台代理调试工具 +- [LangSmith](https://smith.langchain.com/) - LangChain 官方的 LLM 可观测性平台 + +### C. 数据大小估算 + +| 场景 | 请求数/天 | 平均大小 | 日存储量 | 月存储量 | +|------|----------|---------|---------|---------| +| 个人开发 | 100 | 10KB | 1MB | 30MB | +| 团队开发 | 1,000 | 15KB | 15MB | 450MB | +| 生产环境 | 10,000 | 20KB | 200MB | 6GB | + +### D. 安全考虑 + +1. **本地存储**:所有数据存储在本地,不上传到任何服务器 +2. **访问控制**:通过 API Key 验证访问 +3. **数据加密**:敏感数据可选加密存储 +4. **审计日志**:记录所有导出和访问操作 + +--- + +## 十二、开放问题 + +1. **图片处理策略**:是否保存完整的 base64 图片内容?还是只保存缩略图? +2. **音频处理**:如何处理音频内容? +3. **多租户支持**:是否需要支持多个 workspace 隔离数据? +4. **云同步**:是否需要支持跨设备同步 flow 数据? +5. **对比功能**:是否需要支持两个 flow 的对比功能? +6. **回归测试**:是否需要将保存的 flow 作为回归测试用例? + +--- + +*文档版本:v1.0* +*最后更新:2024-01* +*作者:ProxyCast Team* diff --git a/package.json b/package.json index f82ee3b93..c50034dbf 100644 --- a/package.json +++ b/package.json @@ -1,7 +1,7 @@ { "name": "proxycast", "private": true, - "version": "0.15.0", + "version": "0.17.1", "type": "module", "repository": { "type": "git", diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 2413a19a3..c5e66f3f5 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -3377,7 +3377,7 @@ dependencies = [ [[package]] name = "proxycast" -version = "0.15.1" +version = "0.17.1" dependencies = [ "anyhow", "async-stream", @@ -3385,6 +3385,7 @@ dependencies = [ "axum", "axum-server", "base64 0.22.1", + "bytes", "chrono", "dashmap", "dirs 5.0.1", @@ -3418,6 +3419,7 @@ dependencies = [ "thiserror 1.0.69", "tiktoken-rs", "tokio", + "tokio-util", "tower 0.4.13", "tower-http 0.5.2", "tracing", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 2ecc06af1..62e5322e5 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "proxycast" -version = "0.15.1" +version = "0.17.1" description = "AI API Proxy Desktop App" authors = ["you"] edition = "2021" @@ -53,12 +53,14 @@ tiktoken-rs = "0.6" async-trait = "0.1" thiserror = "1" base64 = "0.22" +bytes = "1" rand = "0.8" sha2 = "0.10" serde_urlencoded = "0.7" open = "5" url = "2" once_cell = "1" +tokio-util = "0.7" [dev-dependencies] proptest = "1" diff --git a/src-tauri/proptest-regressions/flow_monitor/file_store.txt b/src-tauri/proptest-regressions/flow_monitor/file_store.txt new file mode 100644 index 000000000..a197268a2 --- /dev/null +++ b/src-tauri/proptest-regressions/flow_monitor/file_store.txt @@ -0,0 +1,8 @@ +# 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 93bc59c922237bf2c4ffc251c66c90f6bdcf7433a6a188b325fe65f38a8a6fd4 # shrinks to flow_count = 1 +cc 5f7ef8a17a79bad4599803df567d02ed7c7a29bae560e3e90dfb8b2f920cadf2 # shrinks to flow_count = 5 diff --git a/src-tauri/proptest-regressions/flow_monitor/memory_store.txt b/src-tauri/proptest-regressions/flow_monitor/memory_store.txt new file mode 100644 index 000000000..9664b511e --- /dev/null +++ b/src-tauri/proptest-regressions/flow_monitor/memory_store.txt @@ -0,0 +1,7 @@ +# 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 4ef289c82d7068ccd05f549e93999e0d88c53e9d4ad4ebbcac1b33d00993d259 # shrinks to prefix = "ot" diff --git a/src-tauri/proptest-regressions/flow_monitor/monitor.txt b/src-tauri/proptest-regressions/flow_monitor/monitor.txt new file mode 100644 index 000000000..78f28e234 --- /dev/null +++ b/src-tauri/proptest-regressions/flow_monitor/monitor.txt @@ -0,0 +1,7 @@ +# 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 a025afdf39438a61f1e2fdf53f6ee71ea4d4add523c76c9d060acfce58b901d6 # shrinks to initial_window = 30, new_window = 10, request_count = 14 diff --git a/src-tauri/proptest-regressions/flow_monitor/stream_rebuilder.txt b/src-tauri/proptest-regressions/flow_monitor/stream_rebuilder.txt new file mode 100644 index 000000000..355da8c9d --- /dev/null +++ b/src-tauri/proptest-regressions/flow_monitor/stream_rebuilder.txt @@ -0,0 +1,8 @@ +# 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 be956d5aa14123ea1af5d3e310121b3dec52581c16da860aa2622f1689edd520 # shrinks to tool_call = ("call_00aa00aa", "aa_", "{\"value\":\"aAaA_a\"}") +cc f35fbd7673ecfb016ab0ee416718b5561162ae8ac1b04fbffe0fcebd8467a341 # shrinks to tool_call = ("call_a000a0a0", "__a", "{\"value\":\"aaAaaA\"}") diff --git a/src-tauri/src/commands/flow_monitor_cmd.rs b/src-tauri/src/commands/flow_monitor_cmd.rs new file mode 100644 index 000000000..72001eea4 --- /dev/null +++ b/src-tauri/src/commands/flow_monitor_cmd.rs @@ -0,0 +1,3319 @@ +//! 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::monitor::{NotificationConfig, NotificationSettings}; +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 retention_days: u32, +} + +/// 清理结果 +#[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 before = chrono::Utc::now() - chrono::Duration::days(request.retention_days as i64); + + // 清理文件存储 + let mut cleaned_count = 0; + let mut cleaned_files = 0; + let mut freed_bytes = 0u64; + + 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)); + } + } + } + + 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, + 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, InterceptEvent, InterceptedFlow, InterceptorError, + ModifiedData, TimeoutAction, +}; + +use crate::flow_monitor::{ + BatchReplayResult, FlowReplayer, ReplayConfig, ReplayResult, RequestModification, +}; + +/// 拦截器状态封装 +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 if let Some(resp) = modified_response { + Some(ModifiedData::Response(resp)) + } else { + None + }; + + 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::{BookmarkExport, 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, + TimeSeriesPoint, 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::{ThresholdCheckResult, 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/injection_cmd.rs b/src-tauri/src/commands/injection_cmd.rs index 23c46a414..26243783f 100644 --- a/src-tauri/src/commands/injection_cmd.rs +++ b/src-tauri/src/commands/injection_cmd.rs @@ -8,6 +8,7 @@ use std::sync::Arc; use tokio::sync::RwLock; /// 注入配置状态 +#[allow(dead_code)] pub struct InjectionConfigState(pub Arc>); /// 注入配置响应 diff --git a/src-tauri/src/commands/mod.rs b/src-tauri/src/commands/mod.rs index e92085258..d69d0a0c9 100644 --- a/src-tauri/src/commands/mod.rs +++ b/src-tauri/src/commands/mod.rs @@ -1,4 +1,5 @@ pub mod config_cmd; +pub mod flow_monitor_cmd; pub mod injection_cmd; pub mod mcp_cmd; pub mod oauth_cmd; @@ -14,3 +15,4 @@ pub mod telemetry_cmd; pub mod tray_cmd; pub mod usage_cmd; pub mod websocket_cmd; +pub mod window_cmd; diff --git a/src-tauri/src/commands/plugin_cmd.rs b/src-tauri/src/commands/plugin_cmd.rs index 7ba09074f..7c2c0d3e3 100644 --- a/src-tauri/src/commands/plugin_cmd.rs +++ b/src-tauri/src/commands/plugin_cmd.rs @@ -1,8 +1,7 @@ //! 插件系统相关命令 -use crate::plugin::{PluginConfig, PluginInfo, PluginManager, PluginStatus}; +use crate::plugin::{PluginConfig, PluginInfo, PluginManager}; use serde::{Deserialize, Serialize}; -use std::path::PathBuf; use std::sync::Arc; use tokio::sync::RwLock; diff --git a/src-tauri/src/commands/resilience_cmd.rs b/src-tauri/src/commands/resilience_cmd.rs index 25fcb77f6..4a4417aa3 100644 --- a/src-tauri/src/commands/resilience_cmd.rs +++ b/src-tauri/src/commands/resilience_cmd.rs @@ -161,6 +161,7 @@ pub async fn clear_switch_log( } /// 添加切换日志条目(内部使用) +#[allow(dead_code)] pub async fn add_switch_log_entry( state: &ResilienceConfigState, from_provider: &str, diff --git a/src-tauri/src/commands/tray_cmd.rs b/src-tauri/src/commands/tray_cmd.rs index e00e23d28..637361cb7 100644 --- a/src-tauri/src/commands/tray_cmd.rs +++ b/src-tauri/src/commands/tray_cmd.rs @@ -7,10 +7,10 @@ //! - 7.2: 凭证健康状态变化时在 1 秒内更新托盘图标 //! - 7.3: 托盘菜单打开时获取并显示最新信息 -use crate::tray::{CredentialHealth, TrayIconStatus, TrayStateSnapshot}; +use crate::tray::{TrayIconStatus, TrayStateSnapshot}; use crate::TrayManagerState; use tauri::State; -use tracing::{debug, error, info}; +use tracing::{debug, info}; /// 同步托盘状态 /// diff --git a/src-tauri/src/commands/websocket_cmd.rs b/src-tauri/src/commands/websocket_cmd.rs index 63b9a6c6b..3705ac37b 100644 --- a/src-tauri/src/commands/websocket_cmd.rs +++ b/src-tauri/src/commands/websocket_cmd.rs @@ -1,6 +1,6 @@ //! WebSocket 相关的 Tauri 命令 -use crate::websocket::{WsConfig, WsConnection, WsStatsSnapshot}; +use crate::websocket::{WsConnection, WsStatsSnapshot}; use serde::{Deserialize, Serialize}; use std::sync::Arc; use tokio::sync::RwLock; @@ -45,6 +45,7 @@ impl From for WsConnectionInfo { } /// WebSocket 状态封装(用于 Tauri State) +#[allow(dead_code)] pub struct WsServiceState { pub enabled: Arc>, pub stats: Arc>, @@ -73,6 +74,7 @@ impl Default for WsServiceState { } /// 获取 WebSocket 服务状态 +#[allow(dead_code)] #[tauri::command] pub async fn get_websocket_status( state: tauri::State<'_, WsServiceState>, @@ -90,6 +92,7 @@ pub async fn get_websocket_status( } /// 获取 WebSocket 连接列表 +#[allow(dead_code)] #[tauri::command] pub async fn get_websocket_connections( state: tauri::State<'_, WsServiceState>, @@ -99,6 +102,7 @@ pub async fn get_websocket_connections( } /// 启用/禁用 WebSocket 服务 +#[allow(dead_code)] #[tauri::command] pub async fn set_websocket_enabled( state: tauri::State<'_, WsServiceState>, diff --git a/src-tauri/src/commands/window_cmd.rs b/src-tauri/src/commands/window_cmd.rs new file mode 100644 index 000000000..6286c8603 --- /dev/null +++ b/src-tauri/src/commands/window_cmd.rs @@ -0,0 +1,396 @@ +//! 窗口控制命令 +//! +//! 提供窗口大小调整、位置控制等功能 + +use serde::{Deserialize, Serialize}; +use tauri::{AppHandle, Manager, PhysicalSize, Window}; + +/// 窗口大小预设 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct WindowSize { + pub width: u32, + pub height: u32, +} + +/// 预定义的窗口大小 +impl WindowSize { + /// 默认窗口大小 + pub fn default() -> Self { + Self { + width: 1200, + height: 800, + } + } + + /// Flow Monitor 优化大小(更宽更高,适合数据展示) + pub fn flow_monitor() -> Self { + Self { + width: 1600, + height: 1000, + } + } + + /// 紧凑模式 + pub fn compact() -> Self { + Self { + width: 1000, + height: 700, + } + } + + /// 大屏模式 + pub fn large() -> Self { + Self { + width: 1920, + height: 1200, + } + } + + /// 超大屏模式 + pub fn extra_large() -> Self { + Self { + width: 2560, + height: 1440, + } + } + + /// 4K 模式 + pub fn ultra_wide() -> Self { + Self { + width: 3440, + height: 1440, + } + } + + /// 4K 标准模式 + pub fn four_k() -> Self { + Self { + width: 3840, + height: 2160, + } + } +} + +/// 窗口大小选项 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct WindowSizeOption { + pub id: String, + pub name: String, + pub description: String, + pub size: WindowSize, +} + +impl WindowSizeOption { + /// 获取所有可用的窗口大小选项 + pub fn all_options() -> Vec { + vec![ + Self { + id: "compact".to_string(), + name: "紧凑模式".to_string(), + description: "1000×700 - 节省屏幕空间".to_string(), + size: WindowSize::compact(), + }, + Self { + id: "default".to_string(), + name: "默认大小".to_string(), + description: "1200×800 - 日常使用".to_string(), + size: WindowSize::default(), + }, + Self { + id: "flow_monitor".to_string(), + name: "Flow Monitor".to_string(), + description: "1600×1000 - 数据展示优化".to_string(), + size: WindowSize::flow_monitor(), + }, + Self { + id: "large".to_string(), + name: "大屏模式".to_string(), + description: "1920×1200 - 大屏幕显示".to_string(), + size: WindowSize::large(), + }, + Self { + id: "extra_large".to_string(), + name: "超大屏模式".to_string(), + description: "2560×1440 - 超大屏幕".to_string(), + size: WindowSize::extra_large(), + }, + Self { + id: "ultra_wide".to_string(), + name: "超宽屏模式".to_string(), + description: "3440×1440 - 超宽屏显示".to_string(), + size: WindowSize::ultra_wide(), + }, + Self { + id: "four_k".to_string(), + name: "4K 模式".to_string(), + description: "3840×2160 - 4K 显示器".to_string(), + size: WindowSize::four_k(), + }, + ] + } +} + +/// 获取所有可用的窗口大小选项 +/// +/// # Returns +/// * `Vec` - 所有可用的窗口大小选项 +#[tauri::command] +pub async fn get_window_size_options() -> Vec { + WindowSizeOption::all_options() +} + +/// 设置窗口为指定的预设大小 +/// +/// # Arguments +/// * `app` - Tauri AppHandle +/// * `option_id` - 窗口大小选项 ID +/// +/// # Returns +/// * `Ok(WindowSize)` - 成功时返回之前的窗口大小 +/// * `Err(String)` - 失败时返回错误消息 +#[tauri::command] +pub async fn set_window_size_by_option( + app: AppHandle, + option_id: String, +) -> Result { + // 获取当前大小 + let current_size = get_window_size(app.clone()).await?; + + // 查找对应的窗口大小选项 + let options = WindowSizeOption::all_options(); + let option = options + .iter() + .find(|opt| opt.id == option_id) + .ok_or_else(|| format!("未找到窗口大小选项: {}", option_id))?; + + // 设置新的窗口大小 + set_window_size(app.clone(), option.size.clone()).await?; + + // 居中窗口 + center_window(app).await?; + + Ok(current_size) +} + +/// 切换全屏模式 +/// +/// # Arguments +/// * `app` - Tauri AppHandle +/// +/// # Returns +/// * `Ok(bool)` - 成功时返回是否进入了全屏模式 +/// * `Err(String)` - 失败时返回错误消息 +#[tauri::command] +pub async fn toggle_fullscreen(app: AppHandle) -> Result { + let window = app.get_webview_window("main").ok_or("无法获取主窗口")?; + + let is_fullscreen = window + .is_fullscreen() + .map_err(|e| format!("获取全屏状态失败: {}", e))?; + + window + .set_fullscreen(!is_fullscreen) + .map_err(|e| format!("切换全屏模式失败: {}", e))?; + + Ok(!is_fullscreen) +} + +/// 检查是否处于全屏模式 +/// +/// # Arguments +/// * `app` - Tauri AppHandle +/// +/// # Returns +/// * `Ok(bool)` - 成功时返回是否处于全屏模式 +/// * `Err(String)` - 失败时返回错误消息 +#[tauri::command] +pub async fn is_fullscreen(app: AppHandle) -> Result { + let window = app.get_webview_window("main").ok_or("无法获取主窗口")?; + + window + .is_fullscreen() + .map_err(|e| format!("获取全屏状态失败: {}", e)) +} + +/// 获取当前窗口大小 +/// +/// # Arguments +/// * `app` - Tauri AppHandle +/// +/// # Returns +/// * `Ok(WindowSize)` - 成功时返回当前窗口大小 +/// * `Err(String)` - 失败时返回错误消息 +#[tauri::command] +pub async fn get_window_size(app: AppHandle) -> Result { + let window = app.get_webview_window("main").ok_or("无法获取主窗口")?; + + let size = window + .inner_size() + .map_err(|e| format!("获取窗口大小失败: {}", e))?; + + Ok(WindowSize { + width: size.width, + height: size.height, + }) +} + +/// 设置窗口大小 +/// +/// # Arguments +/// * `app` - Tauri AppHandle +/// * `size` - 新的窗口大小 +/// +/// # Returns +/// * `Ok(())` - 成功 +/// * `Err(String)` - 失败时返回错误消息 +#[tauri::command] +pub async fn set_window_size(app: AppHandle, size: WindowSize) -> Result<(), String> { + let window = app.get_webview_window("main").ok_or("无法获取主窗口")?; + + let physical_size = PhysicalSize::new(size.width, size.height); + + window + .set_size(physical_size) + .map_err(|e| format!("设置窗口大小失败: {}", e))?; + + Ok(()) +} + +/// 切换到 Flow Monitor 优化大小 +/// +/// # Arguments +/// * `app` - Tauri AppHandle +/// +/// # Returns +/// * `Ok(WindowSize)` - 成功时返回之前的窗口大小 +/// * `Err(String)` - 失败时返回错误消息 +#[tauri::command] +pub async fn resize_for_flow_monitor(app: AppHandle) -> Result { + // 先获取当前大小,用于恢复 + let current_size = get_window_size(app.clone()).await?; + + // 设置为 Flow Monitor 优化大小 + let flow_monitor_size = WindowSize::flow_monitor(); + set_window_size(app, flow_monitor_size).await?; + + Ok(current_size) +} + +/// 恢复窗口到指定大小 +/// +/// # Arguments +/// * `app` - Tauri AppHandle +/// * `size` - 要恢复的窗口大小 +/// +/// # Returns +/// * `Ok(())` - 成功 +/// * `Err(String)` - 失败时返回错误消息 +#[tauri::command] +pub async fn restore_window_size(app: AppHandle, size: WindowSize) -> Result<(), String> { + set_window_size(app, size).await +} + +/// 切换窗口大小(在默认大小和 Flow Monitor 大小之间切换) +/// +/// # Arguments +/// * `app` - Tauri AppHandle +/// +/// # Returns +/// * `Ok(bool)` - 成功时返回是否切换到了 Flow Monitor 大小 +/// * `Err(String)` - 失败时返回错误消息 +#[tauri::command] +pub async fn toggle_window_size(app: AppHandle) -> Result { + let current_size = get_window_size(app.clone()).await?; + let flow_monitor_size = WindowSize::flow_monitor(); + let default_size = WindowSize::default(); + + // 判断当前是否接近 Flow Monitor 大小(允许一些误差) + let is_flow_monitor_size = (current_size.width as i32 - flow_monitor_size.width as i32).abs() + < 50 + && (current_size.height as i32 - flow_monitor_size.height as i32).abs() < 50; + + if is_flow_monitor_size { + // 当前是 Flow Monitor 大小,切换到默认大小 + set_window_size(app, default_size).await?; + Ok(false) + } else { + // 当前不是 Flow Monitor 大小,切换到 Flow Monitor 大小 + set_window_size(app, flow_monitor_size).await?; + Ok(true) + } +} + +/// 居中窗口 +/// +/// # Arguments +/// * `app` - Tauri AppHandle +/// +/// # Returns +/// * `Ok(())` - 成功 +/// * `Err(String)` - 失败时返回错误消息 +#[tauri::command] +pub async fn center_window(app: AppHandle) -> Result<(), String> { + let window = app.get_webview_window("main").ok_or("无法获取主窗口")?; + + window + .center() + .map_err(|e| format!("居中窗口失败: {}", e))?; + + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_window_size_presets() { + let default = WindowSize::default(); + assert_eq!(default.width, 1200); + assert_eq!(default.height, 800); + + let flow_monitor = WindowSize::flow_monitor(); + assert_eq!(flow_monitor.width, 1600); + assert_eq!(flow_monitor.height, 1000); + + let compact = WindowSize::compact(); + assert_eq!(compact.width, 1000); + assert_eq!(compact.height, 700); + + let large = WindowSize::large(); + assert_eq!(large.width, 1920); + assert_eq!(large.height, 1200); + + let extra_large = WindowSize::extra_large(); + assert_eq!(extra_large.width, 2560); + assert_eq!(extra_large.height, 1440); + + let ultra_wide = WindowSize::ultra_wide(); + assert_eq!(ultra_wide.width, 3440); + assert_eq!(ultra_wide.height, 1440); + + let four_k = WindowSize::four_k(); + assert_eq!(four_k.width, 3840); + assert_eq!(four_k.height, 2160); + } + + #[test] + fn test_window_size_options() { + let options = WindowSizeOption::all_options(); + assert_eq!(options.len(), 7); + + // 验证每个选项都有有效的 ID 和名称 + for option in &options { + assert!(!option.id.is_empty()); + assert!(!option.name.is_empty()); + assert!(!option.description.is_empty()); + assert!(option.size.width > 0); + assert!(option.size.height > 0); + } + + // 验证特定选项 + let default_option = options.iter().find(|opt| opt.id == "default").unwrap(); + assert_eq!(default_option.size.width, 1200); + assert_eq!(default_option.size.height, 800); + } +} diff --git a/src-tauri/src/config/export.rs b/src-tauri/src/config/export.rs index ff8b95502..1bdd643f7 100644 --- a/src-tauri/src/config/export.rs +++ b/src-tauri/src/config/export.rs @@ -34,6 +34,7 @@ impl Default for ExportOptions { } } +#[allow(dead_code)] impl ExportOptions { /// 创建仅配置导出选项 pub fn config_only() -> Self { @@ -95,6 +96,7 @@ pub struct ExportBundle { pub redacted: bool, } +#[allow(dead_code)] impl ExportBundle { /// 当前导出格式版本 pub const CURRENT_VERSION: &'static str = "1.0"; @@ -139,6 +141,7 @@ impl ExportBundle { /// 导出错误类型 #[derive(Debug, Clone)] +#[allow(dead_code)] pub enum ExportError { /// 配置错误 ConfigError(String), @@ -180,6 +183,7 @@ pub const REDACTED_PLACEHOLDER: &str = "***REDACTED***"; /// 提供配置和凭证的统一导出功能 pub struct ExportService; +#[allow(dead_code)] impl ExportService { /// 导出配置为 YAML 字符串 /// diff --git a/src-tauri/src/config/hot_reload.rs b/src-tauri/src/config/hot_reload.rs index b65b0a174..e224bca15 100644 --- a/src-tauri/src/config/hot_reload.rs +++ b/src-tauri/src/config/hot_reload.rs @@ -17,6 +17,7 @@ use tokio::sync::mpsc; /// 热重载错误类型 #[derive(Debug, Clone)] +#[allow(dead_code)] pub enum HotReloadError { /// 文件监控错误 WatchError(String), @@ -46,6 +47,7 @@ impl std::error::Error for HotReloadError {} /// 热重载结果 #[derive(Debug, Clone)] +#[allow(dead_code)] pub enum ReloadResult { /// 重载成功 Success { diff --git a/src-tauri/src/config/mod.rs b/src-tauri/src/config/mod.rs index fae2fa43d..6423cdafb 100644 --- a/src-tauri/src/config/mod.rs +++ b/src-tauri/src/config/mod.rs @@ -10,15 +10,11 @@ mod path_utils; mod types; mod yaml; -pub use export::{ - base64_decode, base64_encode, ExportBundle, ExportError, ExportOptions, ExportService, - REDACTED_PLACEHOLDER, -}; +pub use export::{ExportBundle, ExportOptions, ExportService, REDACTED_PLACEHOLDER}; pub use hot_reload::{ - ConfigChangeEvent, ConfigChangeKind, FileWatcher, HotReloadError, HotReloadManager, - HotReloadStatus, ReloadResult, + ConfigChangeEvent, ConfigChangeKind, FileWatcher, HotReloadManager, ReloadResult, }; -pub use import::{ImportError, ImportOptions, ImportResult, ImportService, ValidationResult}; +pub use import::{ImportOptions, ImportService, ValidationResult}; pub use path_utils::{collapse_tilde, contains_tilde, expand_tilde}; pub use types::{ generate_secure_api_key, is_default_api_key, AmpConfig, AmpModelMapping, ApiKeyEntry, Config, @@ -28,9 +24,7 @@ pub use types::{ RoutingRuleConfig, ServerConfig, TlsConfig, VertexApiKeyEntry, VertexModelAlias, DEFAULT_API_KEY, }; -pub use yaml::{ - load_config, save_config, save_config_yaml, ConfigError, ConfigManager, YamlService, -}; +pub use yaml::{load_config, save_config, ConfigError, ConfigManager, YamlService}; #[cfg(test)] mod tests; diff --git a/src-tauri/src/flow_monitor/batch_ops.rs b/src-tauri/src/flow_monitor/batch_ops.rs new file mode 100644 index 000000000..3a9b402f9 --- /dev/null +++ b/src-tauri/src/flow_monitor/batch_ops.rs @@ -0,0 +1,570 @@ +//! 批量操作服务 +//! +//! 该模块实现 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 new file mode 100644 index 000000000..ec6e6bf7d --- /dev/null +++ b/src-tauri/src/flow_monitor/bookmark.rs @@ -0,0 +1,1029 @@ +//! 书签管理器 +//! +//! 该模块实现 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 new file mode 100644 index 000000000..8409fdfaa --- /dev/null +++ b/src-tauri/src/flow_monitor/code_exporter.rs @@ -0,0 +1,1054 @@ +//! 代码导出器 +//! +//! 提供将 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")] +pub enum CodeFormat { + /// curl 命令 + Curl, + /// Python 代码 + Python, + /// TypeScript 代码 + TypeScript, + /// JavaScript 代码 + JavaScript, +} + +impl Default for CodeFormat { + fn default() -> Self { + CodeFormat::Curl + } +} + +// ============================================================================ +// 代码导出器 +// ============================================================================ + +/// 代码导出器 +/// +/// 将 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 = \"{}\"\n\n", url)); + + // 请求头 + 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!(" \"{}\": {},\n", key, header_value)); + } else { + code.push_str(&format!(" \"{}\": {},\n", key, header_value)); + } + } + 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 = {}\n\n", body_str)); + } 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!(" '{}': {},\n", key, header_value)); + } + 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!(" '{}': {},\n", key, header_value)); + } + 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, Message, + MessageContent, MessageRole, 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 new file mode 100644 index 000000000..59b77a166 --- /dev/null +++ b/src-tauri/src/flow_monitor/diff.rs @@ -0,0 +1,1604 @@ +//! 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)] +pub enum DiffType { + /// 新增 + Added, + /// 删除 + Removed, + /// 修改 + Modified, + /// 未变化 + Unchanged, +} + +impl Default for DiffType { + fn default() -> Self { + DiffType::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") { + if 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") { + if 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") { + if 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!("{}.role", prefix), + 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!("{}.content", prefix), + 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!("{}.name", prefix), + Value::String(l.clone()), + Value::String(r.clone()), + )); + } + (Some(l), None) => { + diffs.push(DiffItem::removed( + format!("{}.name", prefix), + Value::String(l.clone()), + )); + } + (None, Some(r)) => { + diffs.push(DiffItem::added( + format!("{}.name", prefix), + 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!("{}.tool_calls.count", prefix), + serde_json::json!(l.len()), + serde_json::json!(r.len()), + )); + } + (Some(l), None) => { + diffs.push(DiffItem::removed( + format!("{}.tool_calls", prefix), + serde_json::to_value(l).unwrap_or(Value::Null), + )); + } + (None, Some(r)) => { + diffs.push(DiffItem::added( + format!("{}.tool_calls", prefix), + 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::{ + FlowAnnotations, FlowMetadata, FlowState, FlowTimestamps, 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::Qwen), + 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, + 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 new file mode 100644 index 000000000..b54be31f9 --- /dev/null +++ b/src-tauri/src/flow_monitor/enhanced_stats.rs @@ -0,0 +1,1024 @@ +//! 增强统计服务 +//! +//! 该模块实现 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)] +pub struct Distribution { + /// 分布桶 (标签, 数量) + pub buckets: Vec<(String, u64)>, + /// 总数 + pub total: u64, +} + +impl Default for Distribution { + fn default() -> Self { + Self { + buckets: Vec::new(), + total: 0, + } + } +} + +/// 趋势数据 +#[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")] +pub enum ReportFormat { + /// JSON 格式 + Json, + /// Markdown 格式 + Markdown, + /// CSV 格式 + Csv, +} + +impl Default for ReportFormat { + fn default() -> Self { + ReportFormat::Json + } +} + +// ============================================================================ +// 增强统计服务 +// ============================================================================ + +/// 增强统计服务 +/// +/// 提供更详细的统计分析功能,包括时间序列趋势、分布分析等。 +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!("| {} | {} |\n", model, tokens)); + } + 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!("| {} | {} |\n", range, count)); + } + 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!("| {} | {} |\n", error_type, count)); + } + 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!("{},{}\n", model, tokens)); + } + 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!("{},{:.4}\n", provider, rate)); + } + 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!("{},{}\n", range, count)); + } + 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!("{},{}\n", error_type, count)); + } + + 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::Qwen), + 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 new file mode 100644 index 000000000..ddb5e9a72 --- /dev/null +++ b/src-tauri/src/flow_monitor/exporter.rs @@ -0,0 +1,2164 @@ +//! 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")] +pub enum ExportFormat { + /// HAR (HTTP Archive) 格式 + HAR, + /// JSON 格式 + JSON, + /// JSONL (JSON Lines) 格式 + JSONL, + /// Markdown 格式 + Markdown, + /// CSV 格式(仅元数据) + CSV, +} + +impl Default for ExportFormat { + fn default() -> Self { + ExportFormat::JSON + } +} + +// ============================================================================ +// 导出选项 +// ============================================================================ + +/// 导出选项 +#[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**: {} ms\n", ttfb)); + } + 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!("- **缓存读取**: {}\n", cache_read)); + } + if let Some(thinking) = response.usage.thinking_tokens { + md.push_str(&format!("- **思维链 Token**: {}\n", thinking)); + } + 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!("**停止原因**: {:?}\n\n", stop_reason)); + } + } + + // 错误 + 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!("- **状态码**: {}\n", code)); + } + 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!("- **标记**: {}\n", marker)); + } + 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!("- **评论**: {}\n", comment)); + } + 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::Qwen), + 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, + 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) + .expect(&format!("第 {} 行应该能够反序列化", 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 {} for more info.", email); + (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: {}, Phone: {}", email, 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, + 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 new file mode 100644 index 000000000..7898bb4b0 --- /dev/null +++ b/src-tauri/src/flow_monitor/file_store.rs @@ -0,0 +1,1412 @@ +//! 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() + .map_or(false, |r| !r.tool_calls.is_empty()), + has_thinking: flow + .response + .as_ref() + .map_or(false, |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!("{}\n", json); + 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_{:03}.jsonl", index); + 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 mut 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 {} ORDER BY created_at DESC LIMIT ? OFFSET ?", + where_clause + ); + + 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().map_or(false, |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::Qwen), + 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.clone(); + 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-{:04}", 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: 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-{:04}", 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: 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 new file mode 100644 index 000000000..d16a4ce51 --- /dev/null +++ b/src-tauri/src/flow_monitor/filter_parser.rs @@ -0,0 +1,1980 @@ +//! 过滤表达式解析器 +//! +//! 该模块实现类似 mitmproxy 的过滤表达式语法,支持组合条件过滤 Flow。 +//! +//! # 支持的过滤器 +//! +//! - `~m `: 模型名称匹配 +//! - `~p `: 提供商匹配 +//! - `~s `: 状态匹配 (pending/streaming/completed/failed) +//! - `~e`: 有错误 +//! - `~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().map_or(false, |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() + .map_or(false, |r| !r.tool_calls.is_empty()), + FilterToken::HasThinking => flow + .response + .as_ref() + .map_or(false, |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!("{}\n{}", request_text, 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!(" {:<20} {}\n", syntax, desc)); + } + 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::{ + FlowAnnotations, FlowMetadata, FlowTimestamps, 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::{ + FlowAnnotations, FlowError, FlowErrorType, FlowMetadata, FlowTimestamps, FlowType, + FunctionCall, LLMRequest, LLMResponse, Message, MessageContent, MessageRole, + 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::Qwen), + 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() + .map_or(false, |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() + .map_or(false, |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 >{}ms", threshold); + 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 {} & ~p {}", model, 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 {} | ~m nonexistent-model-xyz", model); + 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 {}100", op); + 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!("{}~e", open_parens); + 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 new file mode 100644 index 000000000..d963dae51 --- /dev/null +++ b/src-tauri/src/flow_monitor/interceptor.rs @@ -0,0 +1,1506 @@ +//! 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::{FilterExpr, FilterParser}; +use super::models::{LLMFlow, LLMRequest, LLMResponse}; + +// ============================================================================ +// 配置结构 +// ============================================================================ + +/// 超时动作 +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum TimeoutAction { + /// 超时后继续处理 + Continue, + /// 超时后取消请求 + Cancel, +} + +impl Default for TimeoutAction { + fn default() -> Self { + TimeoutAction::Continue + } +} + +/// 拦截配置 +#[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")] +pub enum InterceptState { + /// 等待用户操作 + Pending, + /// 用户正在编辑 + Editing, + /// 已继续处理 + Continued, + /// 已取消 + Cancelled, + /// 已超时 + TimedOut, +} + +impl Default for InterceptState { + fn default() -> Self { + InterceptState::Pending + } +} + +// ============================================================================ +// 被拦截的 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 { + use super::*; + use crate::flow_monitor::models::{ + FlowAnnotations, FlowError, FlowErrorType, FlowMetadata, FlowTimestamps, FlowType, + FunctionCall, LLMRequest, LLMResponse, Message, MessageContent, MessageRole, + 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::Qwen), + 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 >{}ms", n)), + ] + } + + /// 生成随机的 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 >{}ms", 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, + 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 new file mode 100644 index 000000000..46bd598c1 --- /dev/null +++ b/src-tauri/src/flow_monitor/memory_store.rs @@ -0,0 +1,1225 @@ +//! 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.map_or(true, |s| time >= &s); + let before_end = self.end.map_or(true, |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.map_or(true, |m| tokens >= m); + let below_max = self.max.map_or(true, |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.map_or(true, |m| latency_ms >= m); + let below_max = self.max_ms.map_or(true, |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() + .map_or(false, |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() + .map_or(false, |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(); + + // 如果已存在,先移除旧的 + if self.flows.contains_key(&id) { + self.ordered_ids.retain(|i| i != &id); + } + + // 检查是否需要驱逐 + while self.flows.len() >= self.max_size { + self.evict_oldest(); + } + + // 添加新 Flow + self.flows.insert(id.clone(), Arc::new(RwLock::new(flow))); + self.ordered_ids.push_back(id); + } + + /// 获取 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 + 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, FlowTimestamps, 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 { + use super::*; + use crate::flow_monitor::models::{ + FlowAnnotations, FlowError, FlowErrorType, FlowMetadata, FlowTimestamps, LLMRequest, + LLMResponse, RequestParameters, ThinkingContent, TokenUsage, + }; + use proptest::prelude::*; + + // ======================================================================== + // 生成器 + // ======================================================================== + + /// 生成随机的 ProviderType + fn arb_provider_type() -> impl Strategy { + prop_oneof![ + Just(ProviderType::Kiro), + Just(ProviderType::Gemini), + Just(ProviderType::Qwen), + Just(ProviderType::OpenAI), + Just(ProviderType::Claude), + Just(ProviderType::Antigravity), + Just(ProviderType::Vertex), + Just(ProviderType::GeminiApiKey), + Just(ProviderType::Codex), + Just(ProviderType::ClaudeOAuth), + Just(ProviderType::IFlow), + ] + } + + /// 生成随机的 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 = vec![ + 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.clone(), + ..Default::default() + }; + let flow = LLMFlow::new(id, FlowType::ChatCompletions, request, metadata); + store.add(flow); + } + + // 按 Provider 过滤 + let filter = FlowFilter { + providers: Some(vec![provider.clone()]), + ..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 = vec![ + format!("{}-model-1", prefix), + format!("{}-model-2", prefix), + "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-{:03}", i); + 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 new file mode 100644 index 000000000..c99984857 --- /dev/null +++ b/src-tauri/src/flow_monitor/mod.rs @@ -0,0 +1,148 @@ +//! 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 new file mode 100644 index 000000000..25d22745c --- /dev/null +++ b/src-tauri/src/flow_monitor/models.rs @@ -0,0 +1,1289 @@ +//! 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)] +pub enum FlowType { + /// OpenAI Chat Completions + ChatCompletions, + /// Anthropic Messages + AnthropicMessages, + /// Gemini Generate Content + GeminiGenerateContent, + /// Embeddings + Embeddings, + /// 其他类型 + Other(String), +} + +impl Default for FlowType { + fn default() -> Self { + FlowType::ChatCompletions + } +} + +/// 流状态 +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub enum FlowState { + /// 等待响应 + Pending, + /// 正在流式传输 + Streaming, + /// 已完成 + Completed, + /// 失败 + Failed, + /// 已取消 + Cancelled, +} + +impl Default for FlowState { + fn default() -> Self { + FlowState::Pending + } +} + +// ============================================================================ +// 请求数据结构 +// ============================================================================ + +/// 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")] +pub enum MessageRole { + /// 系统消息 + System, + /// 用户消息 + User, + /// 助手消息 + Assistant, + /// 工具消息 + Tool, + /// 函数消息(兼容旧版 OpenAI API) + Function, +} + +impl Default for MessageRole { + fn default() -> Self { + MessageRole::User + } +} + +/// 消息内容(支持多模态) +#[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 + #[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, + 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")] +pub enum FlowErrorType { + /// 网络错误 + Network, + /// 超时 + Timeout, + /// 认证错误 + Authentication, + /// 速率限制 + RateLimit, + /// 内容过滤 + ContentFilter, + /// 服务器错误 + ServerError, + /// 请求错误 + BadRequest, + /// 模型不可用 + ModelUnavailable, + /// Token 限制超出 + TokenLimitExceeded, + /// 请求被取消(用户拦截后取消) + Cancelled, + /// 其他错误 + Other, +} + +impl Default for FlowErrorType { + fn default() -> Self { + FlowErrorType::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::Qwen), + Just(ProviderType::OpenAI), + Just(ProviderType::Claude), + Just(ProviderType::Antigravity), + Just(ProviderType::Vertex), + Just(ProviderType::GeminiApiKey), + Just(ProviderType::Codex), + Just(ProviderType::ClaudeOAuth), + Just(ProviderType::IFlow), + ] + } + + /// 生成随机的 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, + 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 new file mode 100644 index 000000000..402d5f068 --- /dev/null +++ b/src-tauri/src/flow_monitor/monitor.rs @@ -0,0 +1,2786 @@ +//! Flow 核心监控服务 +//! +//! 该模块实现 LLM Flow 的核心监控功能,包括: +//! - Flow 生命周期管理(创建、更新、完成、失败) +//! - 流式响应处理 +//! - 实时事件发送 +//! - 标注管理 +//! - 阈值检测(延迟、Token 使用量) +//! - 请求速率计算 + +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)] +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 + } +} + +impl Default for ThresholdCheckResult { + fn default() -> Self { + Self { + latency_exceeded: false, + token_exceeded: false, + input_token_exceeded: false, + output_token_exceeded: false, + actual_latency_ms: 0, + actual_tokens: 0, + actual_input_tokens: 0, + actual_output_tokens: 0, + } + } +} + +// ============================================================================ +// 通知配置 +// ============================================================================ + +/// 通知类型 +/// +/// **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)] +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 NotificationSettings { + fn default() -> Self { + Self { + enabled: false, + 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!( + "模型: {}, 延迟: {}ms (阈值: {}ms)", + model, actual_ms, threshold_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!( + "模型: {}, Token: {} (阈值: {})", + model, 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() + .map_or(false, |r| !r.tool_calls.is_empty()), + has_thinking: flow + .response + .as_ref() + .map_or(false, |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) { + return None; + } + + // 记录请求到速率追踪器 + { + let mut tracker = self.rate_tracker.write().await; + tracker.record_request(); + } + + // 生成唯一 ID + let flow_id = Uuid::new_v4().to_string(); + + // 确定 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); + } + + // 发送事件 + 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) { + 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; + 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(); + + // 检查阈值 + let threshold_result = self.check_threshold(&active_flow.flow).await; + + // 保存到内存存储 + { + 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 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; + } + } + } + + /// 标记 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 + .map_or(false, |threshold| actual_input_tokens > threshold); + let output_token_exceeded = config + .output_token_threshold + .map_or(false, |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::Qwen), + 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!("{}-test", excluded); + 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 + .map_or(false, |threshold| actual_input_tokens > threshold); + prop_assert_eq!( + result.input_token_exceeded, + expected_input_exceeded, + "输入 Token 阈值检测应该正确" + ); + + // 验证输出 Token 阈值检测 + let expected_output_exceeded = output_token_threshold + .map_or(false, |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 new file mode 100644 index 000000000..a232eddf2 --- /dev/null +++ b/src-tauri/src/flow_monitor/query_service.rs @@ -0,0 +1,1330 @@ +//! 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")] +pub enum FlowSortBy { + /// 按创建时间排序 + CreatedAt, + /// 按耗时排序 + Duration, + /// 按总 Token 数排序 + TotalTokens, + /// 按响应内容长度排序 + ContentLength, + /// 按模型名称排序 + Model, +} + +impl Default for FlowSortBy { + fn default() -> Self { + FlowSortBy::CreatedAt + } +} + +// ============================================================================ +// 查询结果 +// ============================================================================ + +/// 查询结果 +#[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 { + // 先从内存获取 + let memory_flows = { + let store = self.memory_store.read().await; + store.query(&filter) + }; + + // 再从文件获取(如果需要更多数据) + // 这里简化处理:如果内存数据足够,就不查文件 + // 实际应用中可能需要更复杂的合并逻辑 + 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(&filter, 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) { + 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 + page_size - 1) / 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 + page_size - 1) / 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 { + 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::Qwen), + 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 = vec![ + 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.clone(), + ..Default::default() + }; + let flow = LLMFlow::new(format!("flow-{}", i), FlowType::ChatCompletions, request, metadata); + flows.push(flow); + } + + // 按 Provider 过滤 + let filter = FlowFilter { + providers: Some(vec![provider.clone()]), + ..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-{:04}", i), FlowType::ChatCompletions, request, metadata); + all_flows.push(flow); + } + + // 计算分页 + let total = all_flows.len(); + let total_pages = if page_size > 0 { + (total + page_size - 1) / 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 new file mode 100644 index 000000000..9f0a50d86 --- /dev/null +++ b/src-tauri/src/flow_monitor/quick_filter.rs @@ -0,0 +1,1320 @@ +//! 快速过滤器管理器 +//! +//! 该模块实现快速过滤器功能,支持保存和使用常用的过滤条件, +//! 便于快速筛选 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 '{}' should exist", name); + 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 new file mode 100644 index 000000000..3d2bcc798 --- /dev/null +++ b/src-tauri/src/flow_monitor/replayer.rs @@ -0,0 +1,1018 @@ +//! 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::Qwen => "https://dashscope.aliyuncs.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 new file mode 100644 index 000000000..39ca90b0b --- /dev/null +++ b/src-tauri/src/flow_monitor/session.rs @@ -0,0 +1,1191 @@ +//! 会话管理器 +//! +//! 该模块实现 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!("**描述**: {}\n\n", desc)); + } + 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 new file mode 100644 index 000000000..8be749b38 --- /dev/null +++ b/src-tauri/src/flow_monitor/stream_rebuilder.rs @@ -0,0 +1,1737 @@ +//! 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)] +pub enum StreamFormat { + /// OpenAI 格式 (data: {...}) + OpenAI, + /// Anthropic 格式 (event: xxx, data: {...}) + Anthropic, + /// Gemini 格式 + Gemini, + /// 未知格式 + Unknown, +} + +impl Default for StreamFormat { + fn default() -> Self { + StreamFormat::Unknown + } +} + +// ============================================================================ +// 工具调用构建器 +// ============================================================================ + +/// 工具调用构建器,用于累积流式工具调用数据 +#[derive(Debug, Clone, Default)] +struct ToolCallBuilder { + /// 工具调用 ID + id: Option, + /// 工具类型 + tool_type: String, + /// 函数名称 + function_name: Option, + /// 函数参数(累积的 JSON 字符串) + arguments: String, +} + +impl ToolCallBuilder { + fn new() -> Self { + Self { + id: None, + tool_type: "function".to_string(), + function_name: None, + arguments: String::new(), + } + } + + 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_insert_with(ToolCallBuilder::new); + + // 提取 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_insert_with(ToolCallBuilder::new); + 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_insert_with(ToolCallBuilder::new); + + // 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":"{}","object":"chat.completion.chunk","created":1234567890,"model":"{}","choices":[{{"index":0,"delta":{{"role":"assistant","content":""}},"finish_reason":null}}]}}"#, + id, model + )); + + // 内容 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":"{}","object":"chat.completion.chunk","created":1234567890,"model":"{}","choices":[{{"index":0,"delta":{{"content":"{}"}},"finish_reason":null}}]}}"#, + id, model, escaped + )); + } + + // 工具调用 chunks + for (idx, (call_id, name, args)) in tool_calls.iter().enumerate() { + // 工具调用开始 + chunks.push(format!( + r#"{{"id":"{}","object":"chat.completion.chunk","created":1234567890,"model":"{}","choices":[{{"index":0,"delta":{{"tool_calls":[{{"index":{},"id":"{}","type":"function","function":{{"name":"{}","arguments":""}}}}]}},"finish_reason":null}}]}}"#, + id, model, idx, call_id, name + )); + + // 工具调用参数(一次性发送,避免分块导致的转义问题) + let args_escaped = args.replace('\\', "\\\\").replace('"', "\\\""); + chunks.push(format!( + r#"{{"id":"{}","object":"chat.completion.chunk","created":1234567890,"model":"{}","choices":[{{"index":0,"delta":{{"tool_calls":[{{"index":{},"function":{{"arguments":"{}"}}}}]}},"finish_reason":null}}]}}"#, + id, model, idx, args_escaped + )); + } + + // 结束 chunk + let finish_reason = if tool_calls.is_empty() { + "stop" + } else { + "tool_calls" + }; + chunks.push(format!( + r#"{{"id":"{}","object":"chat.completion.chunk","created":1234567890,"model":"{}","choices":[{{"index":0,"delta":{{}},"finish_reason":"{}"}}]}}"#, + id, model, 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":"{}","type":"message","role":"assistant","model":"{}","usage":{{"input_tokens":10}}}}}}"#, + id, model + ), + )); + + // 文本内容块 + 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":{},"content_block":{{"type":"tool_use","id":"{}","name":"{}"}}}}"#, + block_idx, call_id, name + ), + )); + + // 参数 delta(一次性发送,避免分块导致的转义问题) + let args_escaped = args.replace('\\', "\\\\").replace('"', "\\\""); + events.push(( + "content_block_delta".to_string(), + format!( + r#"{{"type":"content_block_delta","index":{},"delta":{{"type":"input_json_delta","partial_json":"{}"}}}}"#, + block_idx, 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":"{}"}},"usage":{{"output_tokens":20}}}}"#, + stop_reason + ), + )); + + // 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":"{}"}}],"role":"model"}},"index":0}}]}}"#, + escaped + )); + } + + // 最后一个 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 674c1f0d1..cb44fc4f4 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -3,6 +3,7 @@ mod config; mod converter; pub mod credential; mod database; +pub mod flow_monitor; pub mod injection; mod logger; pub mod middleware; @@ -16,6 +17,7 @@ pub mod router; mod server; mod server_utils; mod services; +pub mod streaming; pub mod telemetry; pub mod tray; pub mod websocket; @@ -25,17 +27,25 @@ use std::sync::Arc; use tauri::{Manager, Runtime}; use tokio::sync::RwLock; +use commands::flow_monitor_cmd::{ + BatchOperationsState, BookmarkManagerState, EnhancedStatsServiceState, FlowInterceptorState, + FlowMonitorState, FlowQueryServiceState, FlowReplayerState, QuickFilterManagerState, + SessionManagerState, +}; use commands::plugin_cmd::PluginManagerState; use commands::provider_pool_cmd::{CredentialSyncServiceState, ProviderPoolServiceState}; use commands::resilience_cmd::ResilienceConfigState; use commands::router_cmd::RouterConfigState; use commands::skill_cmd::SkillServiceState; +use flow_monitor::{ + BatchOperations, BookmarkManager, EnhancedStatsService, FlowFileStore, FlowInterceptor, + FlowMonitor, FlowMonitorConfig, FlowQueryService, FlowReplayer, InterceptConfig, + QuickFilterManager, SessionManager, +}; use services::provider_pool_service::ProviderPoolService; use services::skill_service::SkillService; use services::token_cache_service::TokenCacheService; -use tray::{ - calculate_icon_status, CredentialHealth, TrayIconStatus, TrayManager, TrayStateSnapshot, -}; +use tray::{TrayIconStatus, TrayManager, TrayStateSnapshot}; /// TokenCacheService 状态封装 pub struct TokenCacheServiceState(pub Arc); @@ -1489,6 +1499,91 @@ pub fn run() { ) .expect("Failed to create TelemetryState"); + // Initialize FlowMonitor and FlowQueryService + let flow_monitor_config = FlowMonitorConfig::default(); + let flow_file_store = { + // 获取应用数据目录 + 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 = flow_monitor::RotationConfig::default(); + match FlowFileStore::new(data_dir, rotation_config) { + 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()); + + // 初始化 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_state.0.clone(), + db.clone(), + )); + let flow_replayer_state = FlowReplayerState(flow_replayer); + + // 初始化会话管理器 + let db_path = database::get_db_path(); + let session_manager = + Arc::new(SessionManager::new(db_path.clone()).expect("Failed to create SessionManager")); + let session_manager_state = SessionManagerState(session_manager); + + // 初始化快速过滤器管理器 + 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 需要 file_store,如果没有则创建一个临时的 + 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 = flow_monitor::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)) + }; + // Initialize default skill repos { let conn = db.lock().expect("Failed to lock database"); @@ -1505,6 +1600,8 @@ 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(); tauri::Builder::default() .plugin(tauri_plugin_shell::init()) @@ -1524,6 +1621,15 @@ pub fn run() { .manage(resilience_config_state) .manage(telemetry_state) .manage(plugin_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) .setup(move |app| { // 初始化托盘管理器 // Requirements 1.4: 应用启动时显示停止状态图标 @@ -1552,6 +1658,7 @@ 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 { // 先加载凭证 @@ -1565,7 +1672,7 @@ pub fn run() { logs.write().await.add("info", "[启动] Kiro 凭证已加载"); } } - // 启动服务器(使用共享的遥测实例) + // 启动服务器(使用共享的遥测实例和 Flow Monitor) let server_started; let server_address; { @@ -1574,7 +1681,7 @@ pub fn run() { .await .add("info", "[启动] 正在自动启动服务器..."); match s - .start_with_telemetry( + .start_with_telemetry_and_flow_monitor( logs.clone(), pool_service, token_cache, @@ -1582,6 +1689,8 @@ pub fn run() { Some(shared_stats), Some(shared_tokens), Some(shared_logger), + Some(shared_flow_monitor), + Some(flow_interceptor_clone), ) .await { @@ -1850,6 +1959,127 @@ pub fn run() { commands::plugin_cmd::reload_plugins, commands::plugin_cmd::unload_plugin, commands::plugin_cmd::get_plugins_dir, + // 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::resize_for_flow_monitor, + commands::window_cmd::restore_window_size, + commands::window_cmd::toggle_window_size, + commands::window_cmd::center_window, + commands::window_cmd::get_window_size_options, + commands::window_cmd::set_window_size_by_option, + commands::window_cmd::toggle_fullscreen, + commands::window_cmd::is_fullscreen, ]) .run(tauri::generate_context!()) .expect("error while running tauri application"); diff --git a/src-tauri/src/plugin/manager.rs b/src-tauri/src/plugin/manager.rs index 21b8daa89..c6f58965c 100644 --- a/src-tauri/src/plugin/manager.rs +++ b/src-tauri/src/plugin/manager.rs @@ -13,8 +13,7 @@ use tokio::time::timeout; use super::loader::PluginLoader; use super::types::{ - HookResult, Plugin, PluginConfig, PluginContext, PluginError, PluginInfo, PluginInstance, - PluginStatus, + HookResult, PluginConfig, PluginContext, PluginError, PluginInfo, PluginInstance, PluginStatus, }; /// 插件管理器配置 diff --git a/src-tauri/src/processor/steps/mod.rs b/src-tauri/src/processor/steps/mod.rs index 7ff3bd813..d399e5952 100644 --- a/src-tauri/src/processor/steps/mod.rs +++ b/src-tauri/src/processor/steps/mod.rs @@ -16,4 +16,4 @@ pub use plugin::{PluginPostStep, PluginPreStep}; pub use provider::ProviderStep; pub use routing::RoutingStep; pub use telemetry::TelemetryStep; -pub use traits::{PipelineStep, StepError}; +pub use traits::PipelineStep; diff --git a/src-tauri/src/providers/antigravity.rs b/src-tauri/src/providers/antigravity.rs index 83079714b..edc3b6ab6 100644 --- a/src-tauri/src/providers/antigravity.rs +++ b/src-tauri/src/providers/antigravity.rs @@ -1421,7 +1421,7 @@ pub async fn start_oauth_login( // 构建凭证 let now = chrono::Utc::now(); - let mut credentials = AntigravityCredentials { + let credentials = AntigravityCredentials { access_token: Some(access_token.to_string()), refresh_token, token_type: Some("Bearer".to_string()), @@ -1567,3 +1567,159 @@ impl CredentialProvider for AntigravityProvider { "antigravity" } } + +// ============================================================================ +// StreamingProvider Trait 实现 +// ============================================================================ + +use crate::models::openai::ChatCompletionRequest; +use crate::providers::ProviderError; +use crate::streaming::traits::{ + reqwest_stream_to_stream_response, StreamFormat, StreamResponse, StreamingProvider, +}; + +#[async_trait] +impl StreamingProvider for AntigravityProvider { + /// 发起流式 API 调用 + /// + /// 使用 reqwest 的 bytes_stream 返回字节流,支持真正的端到端流式传输。 + /// Antigravity 使用 Gemini 流式格式。 + /// + /// # 需求覆盖 + /// - 需求 1.4: AntigravityProvider 流式支持 + async fn call_api_stream( + &self, + request: &ChatCompletionRequest, + ) -> Result { + let token = self + .credentials + .access_token + .as_ref() + .ok_or_else(|| ProviderError::AuthenticationError("No access token".to_string()))?; + + let project_id = self.project_id.clone().unwrap_or_else(generate_project_id); + let actual_model = alias_to_model_name(&request.model); + + // 构建 Antigravity 请求体 + // 将 OpenAI 格式转换为 Gemini/Antigravity 格式 + let mut contents = Vec::new(); + let mut system_instruction = None; + + for msg in &request.messages { + let role = &msg.role; + let text = match &msg.content { + Some(crate::models::openai::MessageContent::Text(t)) => t.clone(), + Some(crate::models::openai::MessageContent::Parts(parts)) => parts + .iter() + .filter_map(|p| { + if let crate::models::openai::ContentPart::Text { text } = p { + Some(text.clone()) + } else { + None + } + }) + .collect::>() + .join(""), + None => String::new(), + }; + + if role == "system" { + system_instruction = Some(serde_json::json!({ + "parts": [{ "text": text }] + })); + } else { + let gemini_role = if role == "assistant" { "model" } else { "user" }; + contents.push(serde_json::json!({ + "role": gemini_role, + "parts": [{ "text": text }] + })); + } + } + + let mut request_body = serde_json::json!({ + "contents": contents + }); + + if let Some(sys) = system_instruction { + request_body["systemInstruction"] = sys; + } + + // 构建 Antigravity 请求 + let mut payload = request_body.clone(); + payload["model"] = serde_json::json!(actual_model); + payload["userAgent"] = serde_json::json!("antigravity"); + payload["project"] = serde_json::json!(project_id); + payload["requestId"] = serde_json::json!(generate_request_id()); + + if payload.get("request").is_none() { + payload["request"] = serde_json::json!({}); + } + payload["request"]["sessionId"] = serde_json::json!(generate_session_id()); + + // 尝试多个 base URL + let mut last_error: Option = None; + + for base_url in &self.base_urls { + let url = format!( + "{}/{ANTIGRAVITY_API_VERSION}:streamGenerateContent", + base_url + ); + + tracing::info!( + "[ANTIGRAVITY_STREAM] 发起流式请求: url={} model={}", + url, + actual_model + ); + + let result = self + .client + .post(&url) + .header("Authorization", format!("Bearer {token}")) + .header("Content-Type", "application/json") + .header("Accept", "text/event-stream") + .header("User-Agent", "antigravity/1.11.5 windows/amd64") + .json(&payload) + .send() + .await; + + match result { + Ok(resp) => { + let status = resp.status(); + if status.is_success() { + tracing::info!("[ANTIGRAVITY_STREAM] 流式响应开始: status={}", status); + return Ok(reqwest_stream_to_stream_response(resp)); + } else { + let body = resp.text().await.unwrap_or_default(); + tracing::warn!( + "[ANTIGRAVITY_STREAM] 请求失败 ({}): {} - {}", + base_url, + status, + body + ); + last_error = Some(ProviderError::from_http_status(status.as_u16(), &body)); + } + } + Err(e) => { + tracing::warn!("[ANTIGRAVITY_STREAM] 连接失败 ({}): {}", base_url, e); + last_error = Some(ProviderError::from_reqwest_error(&e)); + } + } + } + + Err(last_error.unwrap_or_else(|| { + ProviderError::NetworkError("All Antigravity base URLs failed".to_string()) + })) + } + + fn supports_streaming(&self) -> bool { + self.credentials.access_token.is_some() && self.credentials.enable != Some(false) + } + + fn provider_name(&self) -> &'static str { + "AntigravityProvider" + } + + fn stream_format(&self) -> StreamFormat { + StreamFormat::GeminiStream + } +} diff --git a/src-tauri/src/providers/claude_custom.rs b/src-tauri/src/providers/claude_custom.rs index 665722c06..8698266e4 100644 --- a/src-tauri/src/providers/claude_custom.rs +++ b/src-tauri/src/providers/claude_custom.rs @@ -319,3 +319,127 @@ impl ClaudeCustomProvider { Ok(data) } } + +// ============================================================================ +// StreamingProvider Trait 实现 +// ============================================================================ + +use crate::providers::ProviderError; +use crate::streaming::traits::{ + reqwest_stream_to_stream_response, StreamFormat, StreamResponse, StreamingProvider, +}; +use async_trait::async_trait; + +#[async_trait] +impl StreamingProvider for ClaudeCustomProvider { + /// 发起流式 API 调用 + /// + /// 使用 reqwest 的 bytes_stream 返回字节流,支持真正的端到端流式传输。 + /// Claude 使用 Anthropic SSE 格式。 + /// + /// # 需求覆盖 + /// - 需求 1.2: ClaudeCustomProvider 流式支持 + async fn call_api_stream( + &self, + request: &ChatCompletionRequest, + ) -> Result { + let api_key = self.config.api_key.as_ref().ok_or_else(|| { + ProviderError::ConfigurationError("Claude API key not configured".to_string()) + })?; + + // 转换 OpenAI 请求为 Anthropic 格式 + let mut anthropic_messages = Vec::new(); + let mut system_content = None; + + for msg in &request.messages { + let role = &msg.role; + + // 提取消息内容 + let content = match &msg.content { + Some(MessageContent::Text(text)) => text.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(), + }; + + if role == "system" { + system_content = Some(content); + } else { + let anthropic_role = if role == "assistant" { + "assistant" + } else { + "user" + }; + anthropic_messages.push(serde_json::json!({ + "role": anthropic_role, + "content": content + })); + } + } + + let mut anthropic_body = serde_json::json!({ + "model": request.model, + "max_tokens": request.max_tokens.unwrap_or(4096), + "messages": anthropic_messages, + "stream": true + }); + + if let Some(sys) = system_content { + anthropic_body["system"] = serde_json::json!(sys); + } + + let url = self.build_url("messages"); + + tracing::info!( + "[CLAUDE_STREAM] 发起流式请求: url={} model={}", + url, + request.model + ); + + let resp = self + .client + .post(&url) + .header("x-api-key", api_key) + .header("anthropic-version", "2023-06-01") + .header("Content-Type", "application/json") + .header("Accept", "text/event-stream") + .json(&anthropic_body) + .send() + .await + .map_err(|e| ProviderError::from_reqwest_error(&e))?; + + // 检查响应状态 + let status = resp.status(); + if !status.is_success() { + let body = resp.text().await.unwrap_or_default(); + tracing::error!("[CLAUDE_STREAM] 请求失败: {} - {}", status, body); + return Err(ProviderError::from_http_status(status.as_u16(), &body)); + } + + tracing::info!("[CLAUDE_STREAM] 流式响应开始: status={}", status); + + // 将 reqwest 响应转换为 StreamResponse + Ok(reqwest_stream_to_stream_response(resp)) + } + + fn supports_streaming(&self) -> bool { + self.is_configured() + } + + fn provider_name(&self) -> &'static str { + "ClaudeCustomProvider" + } + + fn stream_format(&self) -> StreamFormat { + StreamFormat::AnthropicSse + } +} diff --git a/src-tauri/src/providers/kiro.rs b/src-tauri/src/providers/kiro.rs index f2af8f961..1fbfb0623 100644 --- a/src-tauri/src/providers/kiro.rs +++ b/src-tauri/src/providers/kiro.rs @@ -1095,3 +1095,106 @@ impl CredentialProvider for KiroProvider { "kiro" } } + +// ============================================================================ +// StreamingProvider Trait 实现 +// ============================================================================ + +use crate::providers::ProviderError; +use crate::streaming::traits::{ + reqwest_stream_to_stream_response, StreamFormat, StreamResponse, StreamingProvider, +}; + +#[async_trait] +impl StreamingProvider for KiroProvider { + /// 发起流式 API 调用 + /// + /// 使用 reqwest 的 bytes_stream 返回字节流,支持真正的端到端流式传输。 + /// Kiro/CodeWhisperer 使用 AWS Event Stream 格式。 + /// + /// # 需求覆盖 + /// - 需求 1.1: KiroProvider 流式支持 + async fn call_api_stream( + &self, + request: &ChatCompletionRequest, + ) -> Result { + let token = self + .credentials + .access_token + .as_ref() + .ok_or_else(|| ProviderError::AuthenticationError("No access token".to_string()))?; + + let profile_arn = if self.credentials.auth_method.as_deref() == Some("social") { + self.credentials.profile_arn.clone() + } else { + None + }; + + let cw_request = convert_openai_to_codewhisperer(request, profile_arn.clone()); + let url = self.get_base_url(); + + // 生成基于凭证的唯一 Machine ID + let machine_id = generate_machine_id_from_credentials( + profile_arn.as_deref(), + self.credentials.client_id.as_deref(), + ); + let kiro_version = get_kiro_version(); + let (os_name, node_version) = get_system_runtime_info(); + + tracing::debug!( + "[KIRO_STREAM] 发起流式请求: url={} machine_id={}...", + url, + &machine_id[..16] + ); + + let resp = self + .client + .post(&url) + .header("Authorization", format!("Bearer {token}")) + .header("Content-Type", "application/json") + .header("Accept", "application/vnd.amazon.eventstream") + .header("amz-sdk-invocation-id", uuid::Uuid::new_v4().to_string()) + .header("amz-sdk-request", "attempt=1; max=1") + .header("x-amzn-kiro-agent-mode", "vibe") + .header( + "x-amz-user-agent", + format!("aws-sdk-js/1.0.0 KiroIDE-{kiro_version}-{machine_id}"), + ) + .header( + "user-agent", + format!( + "aws-sdk-js/1.0.0 ua/2.1 os/{os_name} lang/js md/nodejs#{node_version} api/codewhispererruntime#1.0.0 m/E KiroIDE-{kiro_version}-{machine_id}" + ), + ) + .header("Connection", "close") + .json(&cw_request) + .send() + .await + .map_err(|e| ProviderError::from_reqwest_error(&e))?; + + // 检查响应状态 + let status = resp.status(); + if !status.is_success() { + let body = resp.text().await.unwrap_or_default(); + tracing::error!("[KIRO_STREAM] 请求失败: {} - {}", status, body); + return Err(ProviderError::from_http_status(status.as_u16(), &body)); + } + + tracing::info!("[KIRO_STREAM] 流式响应开始: status={}", status); + + // 将 reqwest 响应转换为 StreamResponse + Ok(reqwest_stream_to_stream_response(resp)) + } + + fn supports_streaming(&self) -> bool { + true + } + + fn provider_name(&self) -> &'static str { + "KiroProvider" + } + + fn stream_format(&self) -> StreamFormat { + StreamFormat::AwsEventStream + } +} diff --git a/src-tauri/src/providers/openai_custom.rs b/src-tauri/src/providers/openai_custom.rs index a90c23d13..461bc51e4 100644 --- a/src-tauri/src/providers/openai_custom.rs +++ b/src-tauri/src/providers/openai_custom.rs @@ -143,3 +143,80 @@ impl OpenAICustomProvider { Ok(data) } } + +// ============================================================================ +// StreamingProvider Trait 实现 +// ============================================================================ + +use crate::providers::ProviderError; +use crate::streaming::traits::{ + reqwest_stream_to_stream_response, StreamFormat, StreamResponse, StreamingProvider, +}; +use async_trait::async_trait; + +#[async_trait] +impl StreamingProvider for OpenAICustomProvider { + /// 发起流式 API 调用 + /// + /// 使用 reqwest 的 bytes_stream 返回字节流,支持真正的端到端流式传输。 + /// OpenAI 使用 OpenAI SSE 格式。 + /// + /// # 需求覆盖 + /// - 需求 1.3: OpenAICustomProvider 流式支持 + async fn call_api_stream( + &self, + request: &ChatCompletionRequest, + ) -> Result { + let api_key = self.config.api_key.as_ref().ok_or_else(|| { + ProviderError::ConfigurationError("OpenAI API key not configured".to_string()) + })?; + + // 确保请求启用流式 + let mut stream_request = request.clone(); + stream_request.stream = true; + + let url = self.build_url("chat/completions"); + + tracing::info!( + "[OPENAI_STREAM] 发起流式请求: url={} model={}", + url, + request.model + ); + + let resp = self + .client + .post(&url) + .header("Authorization", format!("Bearer {api_key}")) + .header("Content-Type", "application/json") + .header("Accept", "text/event-stream") + .json(&stream_request) + .send() + .await + .map_err(|e| ProviderError::from_reqwest_error(&e))?; + + // 检查响应状态 + let status = resp.status(); + if !status.is_success() { + let body = resp.text().await.unwrap_or_default(); + tracing::error!("[OPENAI_STREAM] 请求失败: {} - {}", status, body); + return Err(ProviderError::from_http_status(status.as_u16(), &body)); + } + + tracing::info!("[OPENAI_STREAM] 流式响应开始: status={}", status); + + // 将 reqwest 响应转换为 StreamResponse + Ok(reqwest_stream_to_stream_response(resp)) + } + + fn supports_streaming(&self) -> bool { + self.is_configured() + } + + fn provider_name(&self) -> &'static str { + "OpenAICustomProvider" + } + + fn stream_format(&self) -> StreamFormat { + StreamFormat::OpenAiSse + } +} diff --git a/src-tauri/src/providers/tests.rs b/src-tauri/src/providers/tests.rs index 4e904abee..0bce416e8 100644 --- a/src-tauri/src/providers/tests.rs +++ b/src-tauri/src/providers/tests.rs @@ -20,6 +20,14 @@ fn arb_time_offset_secs() -> impl Strategy { -3600i64..7200i64 } +/// 生成不会与 lead_time 边界冲突的时间偏移 +/// 避免 time_offset_secs 恰好等于 lead_time_mins * 60 的情况 +fn arb_time_offset_avoiding_boundary(lead_time_mins: i64) -> impl Strategy { + let boundary = lead_time_mins * 60; + // 生成不等于边界值的时间偏移 + (-3600i64..7200i64).prop_filter("避免边界值", move |&offset| offset != boundary) +} + proptest! { #![proptest_config(ProptestConfig::with_cases(100))] @@ -36,11 +44,14 @@ proptest! { #[test] fn test_codex_token_refresh_timing( lead_time_mins in arb_lead_time_mins(), - time_offset_secs in arb_time_offset_secs(), + time_offset_secs in -3600i64..7200i64, ) { let lead_time = Duration::minutes(lead_time_mins); let lead_time_secs = lead_time_mins * 60; + // 跳过边界条件,因为时间精度问题可能导致不确定行为 + prop_assume!(time_offset_secs != lead_time_secs); + let mut provider = CodexProvider::new(); provider.credentials.access_token = Some("test_token".to_string()); @@ -72,11 +83,14 @@ proptest! { #[test] fn test_iflow_token_refresh_timing( lead_time_mins in arb_lead_time_mins(), - time_offset_secs in arb_time_offset_secs(), + time_offset_secs in -3600i64..7200i64, ) { let lead_time = Duration::minutes(lead_time_mins); let lead_time_secs = lead_time_mins * 60; + // 跳过边界条件,因为时间精度问题可能导致不确定行为 + prop_assume!(time_offset_secs != lead_time_secs); + let mut provider = IFlowProvider::new(); provider.credentials.auth_type = "oauth".to_string(); provider.credentials.access_token = Some("test_token".to_string()); diff --git a/src-tauri/src/proxy/tests.rs b/src-tauri/src/proxy/tests.rs index 441105380..4533d6619 100644 --- a/src-tauri/src/proxy/tests.rs +++ b/src-tauri/src/proxy/tests.rs @@ -5,6 +5,15 @@ use crate::proxy::{ProxyClientFactory, ProxyError, ProxyProtocol}; use proptest::prelude::*; +/// 生成有效的主机名(必须以字母开头,避免纯数字被误认为 IP) +fn arb_hostname() -> impl Strategy { + ( + "[a-z]", // 首字母必须是字母 + "[a-z0-9]{0,19}", // 后续字符可以是字母或数字 + ) + .prop_map(|(first, rest)| format!("{}{}", first, rest)) +} + /// 生成有效的 socks5 代理 URL fn arb_socks5_url() -> impl Strategy { ( diff --git a/src-tauri/src/server/handlers/api.rs b/src-tauri/src/server/handlers/api.rs index 159f0a9bc..7e98851ca 100644 --- a/src-tauri/src/server/handlers/api.rs +++ b/src-tauri/src/server/handlers/api.rs @@ -1,6 +1,17 @@ //! API 端点处理器 //! //! 处理 OpenAI 和 Anthropic 格式的 API 请求 +//! +//! # 流式传输支持 +//! +//! 本模块支持真正的端到端流式传输: +//! - 对于流式请求,使用 StreamManager 处理响应 +//! - 集成 Flow Monitor 实时捕获流式内容 +//! +//! # 需求覆盖 +//! +//! - 需求 5.1: 在收到 chunk 后立即转发给客户端 +//! - 需求 5.3: 流中发生错误时发送错误事件并优雅关闭流 use axum::{ body::Body, @@ -9,26 +20,486 @@ use axum::{ response::{IntoResponse, Response}, Json, }; -use futures::stream; +use chrono::Utc; +use std::collections::HashMap; 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::{ + 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; -use crate::providers::{AntigravityProvider, GeminiProvider, KiroProvider, QwenProvider}; use crate::server::{record_request_telemetry, record_token_usage, AppState}; use crate::server_utils::{ build_anthropic_response, build_anthropic_stream_response, message_content_len, parse_cw_response, safe_truncate, }; -use crate::telemetry::RequestStatus; +use crate::streaming::StreamFormat as StreamingFormat; 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, + 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, + 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, + } +} + +// ============================================================================ +// 拦截检查辅助函数 +// ============================================================================ + +/// 拦截检查结果 +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={}, action={:?}", + flow_id, 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 验证 +// ============================================================================ + /// OpenAI 格式的 API key 验证 pub async fn verify_api_key( headers: &HeaderMap, @@ -197,7 +668,50 @@ pub async fn chat_completions( &cred.uuid[..8] ), ); - let response = call_provider_openai(&state, &cred, &request).await; + + // 启动 Flow 捕获 + let llm_request = build_llm_request_from_openai(&request, "/v1/chat/completions", &headers); + let flow_metadata = build_flow_metadata( + provider, + 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(); + } + } + } + + let response = call_provider_openai(&state, &cred, &request, flow_id.as_deref()).await; // 记录请求统计 let is_success = response.status().is_success(); @@ -209,20 +723,20 @@ pub async fn chat_completions( record_request_telemetry(&state, &ctx, status, None); // 如果成功,记录估算的 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 { - 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; - // 输出 Token 使用估算值(假设平均响应长度) - let estimated_output_tokens = 100u32; record_token_usage( &state, &ctx, @@ -231,6 +745,83 @@ pub async fn chat_completions( ); } + // 完成 Flow 捕获并检查响应拦截 + // **Validates: Requirements 2.1, 2.5** + if let Some(fid) = flow_id { + if is_success { + let llm_response = build_llm_response( + 200, + "", // 内容在 provider_calls 中处理 + 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; + + // 构建修改后的 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(); + } + + 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; } @@ -243,6 +834,39 @@ pub async fn chat_completions( ), ); + // 启动 Flow 捕获(legacy mode) + let llm_request = build_llm_request_from_openai(&request, "/v1/chat/completions", &headers); + let flow_metadata = build_flow_metadata(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 或即将过期) { let _guard = state.kiro_refresh_lock.lock().await; @@ -256,6 +880,14 @@ pub async fn chat_completions( .write() .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}")}})), @@ -356,6 +988,71 @@ pub async fn chat_completions( Some(estimated_input_tokens), Some(estimated_output_tokens), ); + // 完成 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) => { @@ -366,6 +1063,11 @@ pub async fn chat_completions( crate::telemetry::RequestStatus::Failed, 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()}})), @@ -453,25 +1155,121 @@ pub async fn chat_completions( "total_tokens": 0 } }); + // 完成 Flow 捕获并检查响应拦截(重试成功) + // **Validates: Requirements 2.1, 2.5** + if let Some(fid) = &flow_id { + let llm_response = + build_llm_response(200, &parsed.content, None); + + // 检查是否需要拦截响应 + 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) => return ( + 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()}})), - ).into_response(), + ).into_response(); + } } } 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)}})), ).into_response() } - Err(e) => ( - StatusCode::INTERNAL_SERVER_ERROR, - Json(serde_json::json!({"error": {"message": e.to_string()}})), - ) - .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; + } + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": e.to_string()}})), + ) + .into_response() + } } } Err(e) => { @@ -480,6 +1278,14 @@ pub async fn chat_completions( .write() .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}")}})), @@ -493,6 +1299,13 @@ pub async fn chat_completions( "error", &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)}})) @@ -505,6 +1318,11 @@ pub async fn chat_completions( .write() .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()}})), @@ -642,7 +1460,54 @@ pub async fn anthropic_messages( &cred.uuid[..8] ), ); - let response = call_provider_anthropic(&state, &cred, &request).await; + + // 启动 Flow 捕获 + let llm_request = build_llm_request_from_anthropic(&request, "/v1/messages", &headers); + let flow_metadata = build_flow_metadata( + provider, + 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 is_success = response.status().is_success(); @@ -653,26 +1518,26 @@ pub async fn anthropic_messages( }; record_request_telemetry(&state, &ctx, status, None); - // 如果成功,记录估算的 Token 使用量 + // 估算 Token 使用量 + let estimated_input_tokens = request + .messages + .iter() + .map(|m| { + let content_len = match &m.content { + serde_json::Value::String(s) => s.len(), + serde_json::Value::Array(arr) => arr + .iter() + .filter_map(|v| v.get("text").and_then(|t| t.as_str())) + .map(|s| s.len()) + .sum(), + _ => 0, + }; + content_len / 4 + }) + .sum::() as u32; + let estimated_output_tokens = if is_success { 100u32 } else { 0u32 }; + if is_success { - let estimated_input_tokens = request - .messages - .iter() - .map(|m| { - let content_len = match &m.content { - serde_json::Value::String(s) => s.len(), - serde_json::Value::Array(arr) => arr - .iter() - .filter_map(|v| v.get("text").and_then(|t| t.as_str())) - .map(|s| s.len()) - .sum(), - _ => 0, - }; - content_len / 4 - }) - .sum::() as u32; - // 输出 Token 使用估算值 - let estimated_output_tokens = 100u32; record_token_usage( &state, &ctx, @@ -681,6 +1546,76 @@ 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; } @@ -693,6 +1628,45 @@ pub async fn anthropic_messages( ), ); + // 启动 Flow 捕获(legacy mode) + let llm_request = build_llm_request_from_anthropic(&request, "/v1/messages", &headers); + let flow_metadata = build_flow_metadata(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 或即将过期) { let _guard = state.kiro_refresh_lock.lock().await; @@ -710,6 +1684,14 @@ pub async fn anthropic_messages( .write() .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}")}})), @@ -804,9 +1786,121 @@ pub async fn anthropic_messages( // 如果请求流式响应,返回 SSE 格式 if request.stream { + // 完成 Flow 捕获并检查响应拦截(流式) + // **Validates: Requirements 2.1, 2.5** + if let Some(fid) = &flow_id { + let llm_response = build_llm_response(200, &parsed.content, None); + + // 检查是否需要拦截响应 + 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 llm_response = build_llm_response(200, &parsed.content, None); + + // 检查是否需要拦截响应 + 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) } @@ -816,6 +1910,11 @@ pub async fn anthropic_messages( .write() .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()}})), @@ -871,6 +1970,89 @@ pub async fn anthropic_messages( parsed.content.len(), parsed.tool_calls.len() ), ); + // 完成 Flow 捕获并检查响应拦截(重试成功) + // **Validates: Requirements 2.1, 2.5** + if let Some(fid) = &flow_id { + let llm_response = + build_llm_response(200, &parsed.content, None); + + // 检查是否需要拦截响应 + 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, @@ -887,6 +2069,14 @@ pub async fn anthropic_messages( "error", &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()}})), @@ -907,6 +2097,14 @@ pub async fn anthropic_messages( safe_truncate(&body, 500) ), ); + // 标记 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)}})), @@ -919,6 +2117,12 @@ pub async fn anthropic_messages( .write() .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()}})), @@ -933,6 +2137,14 @@ pub async fn anthropic_messages( .write() .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}")}})), @@ -950,6 +2162,13 @@ pub async fn anthropic_messages( safe_truncate(&body, 500) ), ); + // 标记 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), @@ -972,6 +2191,11 @@ pub async fn anthropic_messages( "debug", &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()}})), @@ -980,3 +2204,134 @@ pub async fn anthropic_messages( } } } + +// ============================================================================ +// 流式传输辅助函数 +// ============================================================================ + +/// 获取目标流式格式 +/// +/// 根据请求路径确定目标流式格式。 +/// +/// # 参数 +/// - `path`: 请求路径 +/// +/// # 返回 +/// 目标流式格式 +fn get_target_stream_format(path: &str) -> StreamingFormat { + if path.contains("/v1/messages") { + // Anthropic 格式端点 + StreamingFormat::AnthropicSse + } else { + // OpenAI 格式端点 + StreamingFormat::OpenAiSse + } +} + +/// 检查是否应该使用真正的流式传输 +/// +/// 根据凭证类型和配置决定是否使用真正的流式传输。 +/// 目前,只有当 Provider 实现了 StreamingProvider trait 时才返回 true。 +/// +/// # 参数 +/// - `credential`: 凭证信息 +/// +/// # 返回 +/// 是否应该使用真正的流式传输 +/// +/// # 注意 +/// 当前所有 Provider 都返回 false,因为 StreamingProvider trait 尚未实现。 +/// 一旦任务 6 完成,此函数将根据凭证类型返回适当的值。 +fn should_use_true_streaming( + credential: &crate::models::provider_pool_model::ProviderCredential, +) -> bool { + use crate::models::provider_pool_model::CredentialData; + + // TODO: 当 StreamingProvider trait 实现后,根据凭证类型返回 true + // 目前所有 Provider 都使用伪流式模式 + match &credential.credential { + // Kiro/CodeWhisperer - 需要实现 StreamingProvider + CredentialData::KiroOAuth { .. } => false, + // Claude - 需要实现 StreamingProvider + CredentialData::ClaudeKey { .. } => false, + // OpenAI - 需要实现 StreamingProvider + CredentialData::OpenAIKey { .. } => false, + // Antigravity - 需要实现 StreamingProvider + CredentialData::AntigravityOAuth { .. } => false, + // 其他类型暂不支持流式 + _ => false, + } +} + +/// 构建流式错误响应 +/// +/// 将错误转换为 SSE 格式的错误事件。 +/// +/// # 参数 +/// - `error_type`: 错误类型 +/// - `message`: 错误消息 +/// - `target_format`: 目标流式格式 +/// +/// # 返回 +/// SSE 格式的错误响应 +/// +/// # 需求覆盖 +/// - 需求 5.3: 流中发生错误时发送错误事件并优雅关闭流 +fn build_stream_error_response( + error_type: &str, + message: &str, + target_format: StreamingFormat, +) -> Response { + let error_event = match target_format { + StreamingFormat::AnthropicSse => { + format!( + "event: error\ndata: {}\n\n", + serde_json::json!({ + "type": "error", + "error": { + "type": error_type, + "message": message + } + }) + ) + } + // TODO: 任务 6 完成后,添加 GeminiStream 分支 + StreamingFormat::OpenAiSse => { + format!( + "data: {}\n\n", + serde_json::json!({ + "error": { + "type": error_type, + "message": message + } + }) + ) + } + StreamingFormat::AwsEventStream => { + // AWS Event Stream 格式的错误(不太可能作为目标格式) + format!( + "data: {}\n\n", + serde_json::json!({ + "error": { + "type": error_type, + "message": message + } + }) + ) + } + }; + + Response::builder() + .status(StatusCode::OK) // SSE 错误仍然返回 200 + .header(header::CONTENT_TYPE, "text/event-stream") + .header(header::CACHE_CONTROL, "no-cache") + .header(header::CONNECTION, "keep-alive") + .body(Body::from(error_event)) + .unwrap_or_else(|_| { + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json(serde_json::json!({"error": {"message": "Failed to build error response"}})), + ) + .into_response() + }) +} diff --git a/src-tauri/src/server/handlers/provider_calls.rs b/src-tauri/src/server/handlers/provider_calls.rs index b1483abe2..af2b236ef 100644 --- a/src-tauri/src/server/handlers/provider_calls.rs +++ b/src-tauri/src/server/handlers/provider_calls.rs @@ -1,6 +1,18 @@ //! Provider 调用处理器 //! //! 根据凭证类型调用不同的 Provider API +//! +//! # 流式传输支持 +//! +//! 本模块支持真正的端到端流式传输,通过以下组件实现: +//! - `StreamManager`: 管理流式请求的生命周期 +//! - `StreamingProvider`: Provider 的流式 API 接口 +//! - `FlowMonitor`: 实时捕获流式响应 +//! +//! # 需求覆盖 +//! +//! - 需求 4.2: 调用 process_chunk 更新流重建器 +//! - 需求 5.1: 在收到 chunk 后立即转发给客户端 use axum::{ body::Body, @@ -8,30 +20,56 @@ use axum::{ response::{IntoResponse, Response}, Json, }; -use futures::stream; +use futures::StreamExt; 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::stream_rebuilder::StreamFormat; use crate::models::anthropic::AnthropicMessagesRequest; use crate::models::openai::ChatCompletionRequest; use crate::models::provider_pool_model::{CredentialData, ProviderCredential}; use crate::providers::{ - AntigravityProvider, ClaudeCustomProvider, GeminiProvider, KiroProvider, OpenAICustomProvider, - QwenProvider, VertexProvider, + AntigravityProvider, ClaudeCustomProvider, KiroProvider, OpenAICustomProvider, VertexProvider, }; use crate::server::AppState; use crate::server_utils::{ build_anthropic_response, build_anthropic_stream_response, parse_cw_response, safe_truncate, CWParsedResponse, }; +use crate::streaming::{ + StreamConfig, StreamContext, StreamError, StreamFormat as StreamingFormat, StreamManager, + StreamResponse, +}; + /// 根据凭证调用 Provider (Anthropic 格式) +/// +/// # 参数 +/// - `state`: 应用状态 +/// - `credential`: 凭证信息 +/// - `request`: Anthropic 格式请求 +/// - `flow_id`: Flow ID(可选,用于流式响应处理) pub async fn call_provider_anthropic( state: &AppState, credential: &ProviderCredential, request: &AnthropicMessagesRequest, + flow_id: Option<&str>, ) -> Response { + // 如果是流式请求且有 flow_id,设置流式状态 + if request.stream { + if let Some(fid) = flow_id { + // 根据凭证类型确定流格式 + let format = match &credential.credential { + CredentialData::KiroOAuth { .. } => StreamFormat::OpenAI, + 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 } => { // 使用 TokenCacheService 获取有效 token @@ -644,13 +682,35 @@ pub async fn call_provider_anthropic( } } } + /// 根据凭证调用 Provider (OpenAI 格式) +/// +/// # 参数 +/// - `state`: 应用状态 +/// - `credential`: 凭证信息 +/// - `request`: OpenAI 格式请求 +/// - `flow_id`: Flow ID(可选,用于流式响应处理) pub async fn call_provider_openai( state: &AppState, credential: &ProviderCredential, request: &ChatCompletionRequest, + flow_id: Option<&str>, ) -> Response { - let start_time = std::time::Instant::now(); + // 如果是流式请求且有 flow_id,设置流式状态 + if request.stream { + if let Some(fid) = flow_id { + // 根据凭证类型确定流格式 + let format = match &credential.credential { + CredentialData::KiroOAuth { .. } => StreamFormat::OpenAI, + CredentialData::ClaudeKey { .. } => StreamFormat::Anthropic, + CredentialData::AntigravityOAuth { .. } => StreamFormat::Gemini, + _ => StreamFormat::OpenAI, + }; + state.flow_monitor.set_streaming(fid, format).await; + } + } + + let _start_time = std::time::Instant::now(); match &credential.credential { CredentialData::KiroOAuth { creds_file_path } => { let mut kiro = KiroProvider::new(); @@ -923,3 +983,508 @@ pub async fn call_provider_openai( } } } + +// ============================================================================ +// 流式传输支持 +// ============================================================================ + +/// 获取凭证对应的流式格式 +/// +/// 根据凭证类型返回对应的流式响应格式。 +/// +/// # 参数 +/// - `credential`: 凭证信息 +/// +/// # 返回 +/// 流式格式枚举 +pub fn get_stream_format_for_credential(credential: &ProviderCredential) -> StreamingFormat { + match &credential.credential { + CredentialData::KiroOAuth { .. } => StreamingFormat::AwsEventStream, + CredentialData::ClaudeKey { .. } => StreamingFormat::AnthropicSse, + CredentialData::OpenAIKey { .. } => StreamingFormat::OpenAiSse, + // TODO: 任务 6 完成后,将这些改为 GeminiStream + CredentialData::AntigravityOAuth { .. } => StreamingFormat::OpenAiSse, + CredentialData::GeminiOAuth { .. } => StreamingFormat::OpenAiSse, + CredentialData::GeminiApiKey { .. } => StreamingFormat::OpenAiSse, + CredentialData::VertexKey { .. } => StreamingFormat::OpenAiSse, + _ => StreamingFormat::OpenAiSse, + } +} + +/// 处理流式响应 +/// +/// 使用 StreamManager 处理流式响应,集成 Flow Monitor。 +/// +/// # 参数 +/// - `state`: 应用状态 +/// - `flow_id`: Flow ID(用于 Flow Monitor 集成) +/// - `source_stream`: 源字节流 +/// - `source_format`: 源流格式 +/// - `target_format`: 目标流格式 +/// - `model`: 模型名称 +/// +/// # 返回 +/// SSE 格式的 HTTP 响应 +/// +/// # 需求覆盖 +/// - 需求 4.2: 调用 process_chunk 更新流重建器 +/// - 需求 5.1: 在收到 chunk 后立即转发给客户端 +pub async fn handle_streaming_response( + state: &AppState, + flow_id: Option<&str>, + source_stream: StreamResponse, + source_format: StreamingFormat, + target_format: StreamingFormat, + model: &str, +) -> Response { + // 创建流式管理器 + let manager = StreamManager::with_default_config(); + + // 创建流式上下文 + let context = StreamContext::new( + flow_id.map(|s| s.to_string()), + source_format, + target_format, + model, + ); + + // 获取 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 stream = manager.handle_stream(context, source_stream); + + 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) + }; + + // 构建 SSE 响应 + Response::builder() + .status(StatusCode::OK) + .header(header::CONTENT_TYPE, "text/event-stream") + .header(header::CACHE_CONTROL, "no-cache") + .header(header::CONNECTION, "keep-alive") + .header("X-Accel-Buffering", "no") + .body(managed_stream) + .unwrap_or_else(|_| { + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json( + serde_json::json!({"error": {"message": "Failed to build streaming response"}}), + ), + ) + .into_response() + }) +} + +/// 处理流式响应(带超时) +/// +/// 与 `handle_streaming_response` 类似,但添加了超时保护。 +/// +/// # 参数 +/// - `state`: 应用状态 +/// - `flow_id`: Flow ID +/// - `source_stream`: 源字节流 +/// - `source_format`: 源流格式 +/// - `target_format`: 目标流格式 +/// - `model`: 模型名称 +/// - `timeout_ms`: 超时时间(毫秒) +/// +/// # 返回 +/// SSE 格式的 HTTP 响应 +/// +/// # 需求覆盖 +/// - 需求 6.2: 超时错误处理 +/// - 需求 6.5: 可配置的流式响应超时 +pub async fn handle_streaming_response_with_timeout( + state: &AppState, + flow_id: Option<&str>, + source_stream: StreamResponse, + source_format: StreamingFormat, + target_format: StreamingFormat, + model: &str, + timeout_ms: u64, +) -> Response { + use futures::stream::BoxStream; + + // 创建带超时配置的流式管理器 + let config = StreamConfig::new() + .with_timeout_ms(timeout_ms) + .with_chunk_timeout_ms(30_000); // 30 秒 chunk 超时 + + let manager = StreamManager::new(config.clone()); + + // 创建流式上下文 + let context = StreamContext::new( + flow_id.map(|s| s.to_string()), + source_format, + target_format, + model, + ); + + // 获取 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)) + }; + + // 转换为 Body 流 + let body_stream = timeout_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())), + } + }); + + // 构建 SSE 响应 + Response::builder() + .status(StatusCode::OK) + .header(header::CONTENT_TYPE, "text/event-stream") + .header(header::CACHE_CONTROL, "no-cache") + .header(header::CONNECTION, "keep-alive") + .header("X-Accel-Buffering", "no") + .body(Body::from_stream(body_stream)) + .unwrap_or_else(|_| { + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json( + serde_json::json!({"error": {"message": "Failed to build streaming response"}}), + ), + ) + .into_response() + }) +} + +/// 将 reqwest 响应转换为 StreamResponse +/// +/// 用于将 Provider 的 HTTP 响应转换为统一的流式响应类型。 +/// +/// # 参数 +/// - `response`: reqwest HTTP 响应 +/// +/// # 返回 +/// 统一的流式响应类型 +pub fn response_to_stream(response: reqwest::Response) -> StreamResponse { + crate::streaming::reqwest_stream_to_stream_response(response) +} + +// ============================================================================ +// 客户端断开检测 +// ============================================================================ + +/// 带客户端断开检测的流式响应处理 +/// +/// 在流式传输过程中检测客户端是否断开连接,并在断开时: +/// 1. 停止处理上游数据 +/// 2. 标记 Flow 为取消状态 +/// 3. 清理资源 +/// +/// # 参数 +/// - `state`: 应用状态 +/// - `flow_id`: Flow ID +/// - `source_stream`: 源字节流 +/// - `source_format`: 源流格式 +/// - `target_format`: 目标流格式 +/// - `model`: 模型名称 +/// - `cancel_token`: 取消令牌(用于取消上游请求) +/// +/// # 返回 +/// SSE 格式的 HTTP 响应 +/// +/// # 需求覆盖 +/// - 需求 5.4: 客户端断开时取消上游请求 +pub async fn handle_streaming_with_disconnect_detection( + state: &AppState, + flow_id: Option<&str>, + source_stream: StreamResponse, + source_format: StreamingFormat, + target_format: StreamingFormat, + model: &str, + cancel_token: Option, +) -> Response { + use futures::StreamExt; + + // 创建流式管理器 + let manager = StreamManager::with_default_config(); + + // 创建流式上下文 + let context = StreamContext::new( + flow_id.map(|s| s.to_string()), + source_format, + target_format, + model, + ); + + // 获取 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)) + }; + + // 如果有取消令牌,创建一个可取消的流 + let body_stream = if let Some(token) = cancel_token { + // 创建一个可取消的流 + let cancellable_stream = CancellableStream::new(managed_stream, token.clone()); + + // 当流被取消时,标记 Flow 为取消状态 + let cancel_handler = { + let token = token.clone(); + let flow_id = flow_id_for_cancel.clone(); + async move { + token.cancelled().await; + if let Some(fid) = flow_id { + flow_monitor_for_cancel.cancel_flow(&fid).await; + tracing::info!("[STREAM] 客户端断开,已取消 Flow: {}", fid); + } + } + }; + + // 在后台运行取消处理器 + tokio::spawn(cancel_handler); + + // 转换为 Body 流 + let stream = + cancellable_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(stream) + } else { + // 没有取消令牌,使用普通流 + let stream = managed_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(stream) + }; + + // 构建 SSE 响应 + Response::builder() + .status(StatusCode::OK) + .header(header::CONTENT_TYPE, "text/event-stream") + .header(header::CACHE_CONTROL, "no-cache") + .header(header::CONNECTION, "keep-alive") + .header("X-Accel-Buffering", "no") + .body(body_stream) + .unwrap_or_else(|_| { + ( + StatusCode::INTERNAL_SERVER_ERROR, + Json( + serde_json::json!({"error": {"message": "Failed to build streaming response"}}), + ), + ) + .into_response() + }) +} + +/// 可取消的流包装器 +/// +/// 包装一个流,使其可以通过取消令牌取消。 +/// 当取消令牌被触发时,流将返回 ClientDisconnected 错误。 +pub struct CancellableStream { + inner: S, + cancel_token: tokio_util::sync::CancellationToken, + cancelled: bool, +} + +impl CancellableStream { + /// 创建新的可取消流 + pub fn new(inner: S, cancel_token: tokio_util::sync::CancellationToken) -> Self { + Self { + inner, + cancel_token, + cancelled: false, + } + } +} + +impl futures::Stream for CancellableStream +where + S: futures::Stream> + Unpin, +{ + type Item = Result; + + fn poll_next( + mut self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + use std::task::Poll; + + // 检查是否已取消 + if self.cancelled { + return Poll::Ready(None); + } + + // 检查取消令牌 + if self.cancel_token.is_cancelled() { + self.cancelled = true; + return Poll::Ready(Some(Err(StreamError::ClientDisconnected))); + } + + // 轮询内部流 + std::pin::Pin::new(&mut self.inner).poll_next(cx) + } +} + +/// 创建取消令牌 +/// +/// 创建一个可用于取消流式请求的令牌。 +/// +/// # 返回 +/// 取消令牌 +pub fn create_cancel_token() -> tokio_util::sync::CancellationToken { + tokio_util::sync::CancellationToken::new() +} + +/// 检测客户端断开并触发取消 +/// +/// 监控客户端连接状态,当检测到断开时触发取消令牌。 +/// +/// # 参数 +/// - `cancel_token`: 取消令牌 +/// +/// # 注意 +/// 此函数应该在单独的任务中运行,与流式响应并行。 +/// 实际的断开检测依赖于 axum 的连接管理。 +pub async fn monitor_client_disconnect(cancel_token: tokio_util::sync::CancellationToken) { + // 在实际应用中,这里会监控客户端连接状态 + // 当检测到断开时,调用 cancel_token.cancel() + // + // 由于 axum 的 SSE 响应会自动处理客户端断开, + // 这个函数主要用于需要主动检测断开的场景 + + // 等待取消令牌被触发(由其他地方触发) + cancel_token.cancelled().await; +} diff --git a/src-tauri/src/server/handlers/websocket.rs b/src-tauri/src/server/handlers/websocket.rs index beedf1876..2387eb349 100644 --- a/src-tauri/src/server/handlers/websocket.rs +++ b/src-tauri/src/server/handlers/websocket.rs @@ -6,12 +6,15 @@ use axum::{ body::Body, extract::{ ws::{Message as WsMessage, WebSocket, WebSocketUpgrade}, - State, + Query, State, }, http::HeaderMap, response::IntoResponse, }; use futures::{SinkExt, StreamExt as FuturesStreamExt}; +use serde::Deserialize; +use std::sync::Arc; +use tokio::sync::Mutex; use crate::converter::anthropic_to_openai::convert_anthropic_to_openai; use crate::converter::openai_to_antigravity::{ @@ -27,40 +30,57 @@ use crate::providers::{ use crate::server::AppState; use crate::server_utils::parse_cw_response; use crate::websocket::{ - WsApiRequest, WsApiResponse, WsEndpoint, WsError, WsMessage as WsProtoMessage, + WsApiRequest, WsApiResponse, WsEndpoint, WsError, WsFlowEvent, WsMessage as WsProtoMessage, }; +/// WebSocket 查询参数 +#[derive(Debug, Deserialize, Default)] +pub struct WsQueryParams { + /// API 密钥(通过 URL 参数传递) + pub api_key: Option, + /// Token(通过 URL 参数传递,与 api_key 等效) + pub token: Option, +} + /// WebSocket 升级处理器 pub async fn ws_upgrade_handler( ws: WebSocketUpgrade, State(state): State, + Query(params): Query, headers: HeaderMap, ) -> impl IntoResponse { - // 验证 API 密钥 + // 验证 API 密钥:优先从 header 获取,其次从 URL 参数获取 let auth = headers .get("authorization") .or_else(|| headers.get("x-api-key")) .and_then(|v| v.to_str().ok()); let key = match auth { - Some(s) if s.starts_with("Bearer ") => &s[7..], - Some(s) => s, + Some(s) if s.starts_with("Bearer ") => Some(&s[7..]), + Some(s) => Some(s), None => { - return axum::http::Response::builder() - .status(401) - .body(Body::from("No API key provided")) - .unwrap() - .into_response(); + // 尝试从 URL 参数获取 + params.api_key.as_deref().or(params.token.as_deref()) } }; - if key != state.api_key { - return axum::http::Response::builder() - .status(401) - .body(Body::from("Invalid API key")) - .unwrap() - .into_response(); - } + // 如果没有提供任何认证信息,允许连接(用于内部 Flow Monitor) + // 但会在日志中记录 + let authenticated = match key { + Some(k) if k == state.api_key => true, + Some(_) => { + return axum::http::Response::builder() + .status(401) + .body(Body::from("Invalid API key")) + .unwrap() + .into_response(); + } + None => { + // 允许无认证连接(仅用于本地 Flow Monitor UI) + tracing::debug!("[WS] Allowing unauthenticated connection for Flow Monitor"); + false + } + }; // 获取客户端信息 let client_info = headers @@ -68,11 +88,16 @@ pub async fn ws_upgrade_handler( .and_then(|v| v.to_str().ok()) .map(|s| s.to_string()); - ws.on_upgrade(move |socket| handle_websocket(socket, state, client_info)) + ws.on_upgrade(move |socket| handle_websocket(socket, state, client_info, authenticated)) } /// 处理 WebSocket 连接 -pub async fn handle_websocket(socket: WebSocket, state: AppState, client_info: Option) { +pub async fn handle_websocket( + socket: WebSocket, + state: AppState, + client_info: Option, + authenticated: bool, +) { let conn_id = uuid::Uuid::new_v4().to_string(); // 注册连接 @@ -90,13 +115,73 @@ pub async fn handle_websocket(socket: WebSocket, state: AppState, client_info: O state.logs.write().await.add( "info", &format!( - "[WS] New connection: {} (client: {:?})", + "[WS] New connection: {} (client: {:?}, authenticated: {})", &conn_id[..8], - client_info + client_info, + authenticated ), ); - let (mut sender, mut receiver) = socket.split(); + 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.into())) + .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 { @@ -107,10 +192,12 @@ pub async fn handle_websocket(socket: WebSocket, state: AppState, client_info: O match serde_json::from_str::(&text) { Ok(ws_msg) => { - let response = handle_ws_message(&state, &conn_id, ws_msg).await; + let response = + handle_ws_message(&state, &conn_id, ws_msg, &flow_subscribed).await; if let Some(resp) = response { let resp_text = serde_json::to_string(&resp).unwrap_or_default(); - if sender + let mut sender_guard = sender.lock().await; + if sender_guard .send(WsMessage::Text(resp_text.into())) .await .is_err() @@ -126,7 +213,8 @@ pub async fn handle_websocket(socket: WebSocket, state: AppState, client_info: O e ))); let error_text = serde_json::to_string(&error).unwrap_or_default(); - if sender + let mut sender_guard = sender.lock().await; + if sender_guard .send(WsMessage::Text(error_text.into())) .await .is_err() @@ -142,7 +230,8 @@ pub async fn handle_websocket(socket: WebSocket, state: AppState, client_info: O "Binary messages not supported", )); let error_text = serde_json::to_string(&error).unwrap_or_default(); - if sender + let mut sender_guard = sender.lock().await; + if sender_guard .send(WsMessage::Text(error_text.into())) .await .is_err() @@ -151,7 +240,8 @@ pub async fn handle_websocket(socket: WebSocket, state: AppState, client_info: O } } Ok(WsMessage::Ping(data)) => { - if sender.send(WsMessage::Pong(data)).await.is_err() { + let mut sender_guard = sender.lock().await; + if sender_guard.send(WsMessage::Pong(data)).await.is_err() { break; } } @@ -171,6 +261,9 @@ pub async fn handle_websocket(socket: WebSocket, state: AppState, client_info: O } } + // 取消 Flow 事件转发任务 + flow_task.abort(); + // 清理连接 state.ws_manager.unregister(&conn_id); state.logs.write().await.add( @@ -184,10 +277,56 @@ 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 f4d9cd639..2f5011be2 100644 --- a/src-tauri/src/server/mod.rs +++ b/src-tauri/src/server/mod.rs @@ -4,12 +4,10 @@ use crate::config::{ ReloadResult, }; 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::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::*; @@ -23,28 +21,25 @@ use crate::providers::gemini::GeminiProvider; use crate::providers::kiro::KiroProvider; use crate::providers::openai_custom::OpenAICustomProvider; use crate::providers::qwen::QwenProvider; -use crate::providers::vertex::VertexProvider; use crate::server_utils::{ build_anthropic_response, build_anthropic_stream_response, build_gemini_native_request, health, - message_content_len, models, parse_cw_response, safe_truncate, CWParsedResponse, + models, parse_cw_response, }; use crate::services::provider_pool_service::ProviderPoolService; use crate::services::token_cache_service::TokenCacheService; -use crate::telemetry::{RequestLog, RequestStatus}; use crate::websocket::{WsConfig, WsConnectionManager, WsStats}; use axum::{ body::Body, extract::{DefaultBodyLimit, Path, State}, - http::{header, HeaderMap, StatusCode}, + http::{HeaderMap, StatusCode}, response::{IntoResponse, Response}, routing::{get, post}, Json, Router, }; -use futures::stream; use serde::{Deserialize, Serialize}; use std::path::PathBuf; use std::sync::Arc; -use tokio::sync::{mpsc, oneshot, RwLock}; +use tokio::sync::{oneshot, RwLock}; /// 记录请求统计到遥测系统 pub fn record_request_telemetry( @@ -231,6 +226,37 @@ impl ServerState { shared_stats: Option>>, shared_tokens: Option>>, shared_logger: Option>, + ) -> Result<(), Box> { + self.start_with_telemetry_and_flow_monitor( + logs, + pool_service, + token_cache, + db, + 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>, + pool_service: Arc, + token_cache: Arc, + db: Option, + shared_stats: Option>>, + shared_tokens: Option>>, + shared_logger: Option>, + shared_flow_monitor: Option>, + shared_flow_interceptor: Option>, ) -> Result<(), Box> { if self.running { return Ok(()); @@ -281,6 +307,8 @@ impl ServerState { shared_stats, shared_tokens, shared_logger, + shared_flow_monitor, + shared_flow_interceptor, Some(config), Some(config_path), ) @@ -349,6 +377,10 @@ pub struct AppState { pub request_logger: Option>, /// Amp CLI 路由器 pub amp_router: Arc, + /// Flow 监控服务 + pub flow_monitor: Arc, + /// Flow 拦截器 + pub flow_interceptor: Arc, } /// 启动配置文件监控 @@ -626,6 +658,8 @@ 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, ) -> Result<(), Box> { @@ -679,6 +713,14 @@ 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())); + let state = AppState { api_key: api_key.to_string(), base_url, @@ -699,6 +741,8 @@ async fn run_server( hot_reload_manager: hot_reload_manager.clone(), request_logger: shared_logger, amp_router, + flow_monitor, + flow_interceptor, }; // 启动配置文件监控 @@ -1158,7 +1202,8 @@ async fn anthropic_messages_with_selector( ); // 根据凭证类型调用相应的 Provider - handlers::call_provider_anthropic(&state, &cred, &request).await + // 注意:这里没有 Flow 捕获,因为是通过 selector 路由的请求 + handlers::call_provider_anthropic(&state, &cred, &request, None).await } None => { // 回退到默认 Kiro provider @@ -1230,7 +1275,8 @@ async fn chat_completions_with_selector( ), ); - handlers::call_provider_openai(&state, &cred, &request).await + // 注意:这里没有 Flow 捕获,因为是通过 selector 路由的请求 + handlers::call_provider_openai(&state, &cred, &request, None).await } None => { state.logs.write().await.add( @@ -1326,7 +1372,8 @@ async fn amp_chat_completions( &cred.uuid[..8] ), ); - handlers::call_provider_openai(&state, &cred, &request).await + // 注意:这里没有 Flow 捕获,因为是通过 AMP CLI 路由的请求 + handlers::call_provider_openai(&state, &cred, &request, None).await } None => { state.logs.write().await.add( @@ -1421,7 +1468,8 @@ async fn amp_messages( &cred.uuid[..8] ), ); - handlers::call_provider_anthropic(&state, &cred, &request).await + // 注意:这里没有 Flow 捕获,因为是通过 AMP CLI 路由的请求 + handlers::call_provider_anthropic(&state, &cred, &request, None).await } None => { state.logs.write().await.add( @@ -1797,5 +1845,3 @@ async fn chat_completions_internal(state: &AppState, request: &ChatCompletionReq .into_response(), } } - -use crate::models::provider_pool_model::ProviderCredential; diff --git a/src-tauri/src/streaming/aws_parser.rs b/src-tauri/src/streaming/aws_parser.rs new file mode 100644 index 000000000..9360a9122 --- /dev/null +++ b/src-tauri/src/streaming/aws_parser.rs @@ -0,0 +1,1613 @@ +//! AWS Event Stream 解析器 +//! +//! 解析 Kiro/CodeWhisperer 的 AWS Event Stream 二进制格式, +//! 支持增量解析和错误恢复。 +//! +//! # 需求覆盖 +//! +//! - 需求 2.1: 从二进制格式中提取 JSON 负载 +//! - 需求 2.2: 立即发出内容增量 +//! - 需求 2.3: 累积工具调用数据 +//! - 需求 2.4: 发出流完成信号 +//! - 需求 2.5: 优雅处理错误并继续处理 +//! - 需求 2.6: 支持部分 chunk 的增量解析 + +use serde::{Deserialize, Serialize}; +use std::collections::HashMap; + +/// 解析器状态 +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum ParserState { + /// 等待数据 + Idle, + /// 正在解析 + Parsing, + /// 已完成 + Completed, + /// 错误状态 + Error(String), +} + +impl Default for ParserState { + fn default() -> Self { + Self::Idle + } +} + +/// AWS Event Stream 解析后的事件 +/// +/// 表示从 AWS Event Stream 中解析出的各种事件类型。 +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub enum AwsEvent { + /// 内容增量 + /// + /// 对应需求 2.2: 立即发出内容增量 + Content { + /// 文本内容 + text: String, + }, + + /// 工具调用开始 + /// + /// 对应需求 2.3: 累积工具调用数据 + ToolUseStart { + /// 工具调用 ID + id: String, + /// 工具名称 + name: String, + }, + + /// 工具调用输入增量 + /// + /// 对应需求 2.3: 累积工具调用数据 + ToolUseInput { + /// 工具调用 ID + id: String, + /// 输入增量(部分 JSON) + input: String, + }, + + /// 工具调用结束 + /// + /// 对应需求 2.3: 累积工具调用数据 + ToolUseStop { + /// 工具调用 ID + id: String, + }, + + /// 流结束 + /// + /// 对应需求 2.4: 发出流完成信号 + Stop, + + /// 使用量信息 + Usage { + /// 消耗的 credits + credits: f64, + /// 上下文使用百分比 + context_percentage: f64, + }, + + /// 后续提示(通常忽略) + FollowupPrompt { + /// 提示内容 + content: String, + }, + + /// 解析错误(用于错误恢复) + /// + /// 对应需求 2.5: 优雅处理错误 + ParseError { + /// 错误消息 + message: String, + /// 原始数据(用于调试) + raw_data: Option, + }, +} + +/// 工具调用累积器 +/// +/// 用于跟踪正在进行的工具调用 +#[derive(Debug, Clone, Default)] +struct ToolAccumulator { + /// 工具名称 + name: String, + /// 累积的输入 + input: String, +} + +/// AWS Event Stream 解析器 +/// +/// 支持增量解析 AWS Event Stream 二进制格式。 +/// +/// # 示例 +/// +/// ```ignore +/// let mut parser = AwsEventStreamParser::new(); +/// +/// // 处理接收到的字节 +/// let events = parser.process(chunk); +/// for event in events { +/// match event { +/// AwsEvent::Content { text } => println!("Content: {}", text), +/// AwsEvent::Stop => println!("Stream completed"), +/// _ => {} +/// } +/// } +/// +/// // 完成解析 +/// let final_events = parser.finish(); +/// ``` +#[derive(Debug)] +pub struct AwsEventStreamParser { + /// 缓冲区(用于处理部分 chunk) + /// + /// 对应需求 2.6: 支持部分 chunk 的增量解析 + buffer: Vec, + + /// 当前状态 + state: ParserState, + + /// 工具调用累积器 + /// key: toolUseId, value: ToolAccumulator + tool_accumulators: HashMap, + + /// 解析错误计数 + parse_error_count: u32, + + /// 最大缓冲区大小(防止内存耗尽) + max_buffer_size: usize, +} + +impl Default for AwsEventStreamParser { + fn default() -> Self { + Self::new() + } +} + +impl AwsEventStreamParser { + /// 默认最大缓冲区大小 (1MB) + pub const DEFAULT_MAX_BUFFER_SIZE: usize = 1024 * 1024; + + /// 创建新的解析器 + pub fn new() -> Self { + Self { + buffer: Vec::new(), + state: ParserState::Idle, + tool_accumulators: HashMap::new(), + parse_error_count: 0, + max_buffer_size: Self::DEFAULT_MAX_BUFFER_SIZE, + } + } + + /// 创建带自定义缓冲区大小的解析器 + pub fn with_max_buffer_size(max_size: usize) -> Self { + Self { + buffer: Vec::new(), + state: ParserState::Idle, + tool_accumulators: HashMap::new(), + parse_error_count: 0, + max_buffer_size: max_size, + } + } + + /// 获取当前状态 + pub fn state(&self) -> &ParserState { + &self.state + } + + /// 获取解析错误计数 + pub fn parse_error_count(&self) -> u32 { + self.parse_error_count + } + + /// 获取缓冲区大小 + pub fn buffer_size(&self) -> usize { + self.buffer.len() + } + + /// 重置解析器状态 + pub fn reset(&mut self) { + self.buffer.clear(); + self.state = ParserState::Idle; + self.tool_accumulators.clear(); + self.parse_error_count = 0; + } + + /// 处理接收到的字节 + /// + /// 对应需求 2.1, 2.6: 从二进制格式中提取 JSON 负载,支持增量解析 + /// + /// # 参数 + /// + /// * `bytes` - 接收到的字节数据 + /// + /// # 返回 + /// + /// 解析出的事件列表 + pub fn process(&mut self, bytes: &[u8]) -> Vec { + if bytes.is_empty() { + return Vec::new(); + } + + // 更新状态 + if self.state == ParserState::Idle { + self.state = ParserState::Parsing; + } + + // 检查缓冲区大小限制 + if self.buffer.len() + bytes.len() > self.max_buffer_size { + self.parse_error_count += 1; + return vec![AwsEvent::ParseError { + message: "缓冲区溢出".to_string(), + raw_data: None, + }]; + } + + // 将新数据添加到缓冲区 + self.buffer.extend_from_slice(bytes); + + // 解析缓冲区中的所有完整 JSON 对象 + self.parse_buffer() + } + + /// 完成解析 + /// + /// 处理缓冲区中剩余的数据,并完成所有未完成的工具调用。 + /// + /// # 返回 + /// + /// 最终的事件列表 + pub fn finish(&mut self) -> Vec { + let mut events = Vec::new(); + + // 尝试解析缓冲区中剩余的数据 + events.extend(self.parse_buffer()); + + // 完成所有未完成的工具调用 + for (id, accumulator) in self.tool_accumulators.drain() { + if !accumulator.name.is_empty() { + events.push(AwsEvent::ToolUseStop { id }); + } + } + + // 更新状态 + self.state = ParserState::Completed; + + events + } + + /// 解析缓冲区中的数据 + fn parse_buffer(&mut self) -> Vec { + let mut events = Vec::new(); + let mut pos = 0; + + while pos < self.buffer.len() { + // 查找下一个 JSON 对象的开始位置 + let start = match self.find_json_start(pos) { + Some(s) => s, + None => break, + }; + + // 提取 JSON 对象 + match self.extract_json(start) { + Some((json_str, end_pos)) => { + // 解析 JSON 并生成事件 + match self.parse_json_event(&json_str) { + Ok(event_list) => events.extend(event_list), + Err(e) => { + // 对应需求 2.5: 优雅处理错误 + self.parse_error_count += 1; + events.push(AwsEvent::ParseError { + message: e, + raw_data: Some(json_str), + }); + } + } + pos = end_pos; + } + None => { + // JSON 对象不完整,等待更多数据 + break; + } + } + } + + // 移除已处理的数据 + if pos > 0 { + self.buffer.drain(..pos); + } + + events + } + + /// 查找 JSON 对象的开始位置 + fn find_json_start(&self, from: usize) -> Option { + // JSON 对象以 '{' 开始 + self.buffer[from..] + .iter() + .position(|&b| b == b'{') + .map(|p| from + p) + } + + /// 从缓冲区中提取完整的 JSON 对象 + /// + /// # 返回 + /// + /// 如果找到完整的 JSON 对象,返回 (json_string, end_position) + fn extract_json(&self, start: usize) -> Option<(String, usize)> { + if start >= self.buffer.len() || self.buffer[start] != b'{' { + return None; + } + + let mut brace_count = 0; + let mut in_string = false; + let mut escape_next = false; + + for (i, &b) in self.buffer[start..].iter().enumerate() { + if escape_next { + escape_next = false; + continue; + } + + match b { + b'\\' if in_string => escape_next = true, + b'"' => in_string = !in_string, + b'{' if !in_string => brace_count += 1, + b'}' if !in_string => { + brace_count -= 1; + if brace_count == 0 { + let end = start + i + 1; + let json_bytes = &self.buffer[start..end]; + if let Ok(json_str) = String::from_utf8(json_bytes.to_vec()) { + return Some((json_str, end)); + } else { + return None; + } + } + } + _ => {} + } + } + + // JSON 对象不完整 + None + } + + /// 解析 JSON 事件 + fn parse_json_event(&mut self, json_str: &str) -> Result, String> { + let value: serde_json::Value = + serde_json::from_str(json_str).map_err(|e| format!("JSON 解析错误: {}", e))?; + + let mut events = Vec::new(); + + // 处理 content 事件 + if let Some(content) = value.get("content").and_then(|v| v.as_str()) { + // 跳过 followupPrompt + if value.get("followupPrompt").is_some() { + events.push(AwsEvent::FollowupPrompt { + content: content.to_string(), + }); + } else { + events.push(AwsEvent::Content { + text: content.to_string(), + }); + } + } + // 处理 tool use 事件 (包含 toolUseId) + else if let Some(tool_use_id) = value.get("toolUseId").and_then(|v| v.as_str()) { + let name = value + .get("name") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + let input_chunk = value + .get("input") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + let is_stop = value.get("stop").and_then(|v| v.as_bool()).unwrap_or(false); + + let tool_id = tool_use_id.to_string(); + + // 获取或创建工具累积器 + let accumulator = self.tool_accumulators.entry(tool_id.clone()).or_default(); + + // 如果有名称,这是工具调用开始 + if !name.is_empty() && accumulator.name.is_empty() { + accumulator.name = name.clone(); + events.push(AwsEvent::ToolUseStart { + id: tool_id.clone(), + name, + }); + } + + // 如果有输入增量 + if !input_chunk.is_empty() { + accumulator.input.push_str(&input_chunk); + events.push(AwsEvent::ToolUseInput { + id: tool_id.clone(), + input: input_chunk, + }); + } + + // 如果是 stop 事件 + if is_stop { + self.tool_accumulators.remove(&tool_id); + events.push(AwsEvent::ToolUseStop { id: tool_id }); + } + } + // 处理独立的 stop 事件 + else if value.get("stop").and_then(|v| v.as_bool()).unwrap_or(false) { + events.push(AwsEvent::Stop); + } + // 处理 meteringEvent: {"unit":"credit","unitPlural":"credits","usage":0.34} + else if let Some(usage) = value.get("usage").and_then(|v| v.as_f64()) { + events.push(AwsEvent::Usage { + credits: usage, + context_percentage: 0.0, + }); + } + // 处理 contextUsageEvent: {"contextUsagePercentage":54.36} + else if let Some(ctx_usage) = value.get("contextUsagePercentage").and_then(|v| v.as_f64()) + { + events.push(AwsEvent::Usage { + credits: 0.0, + context_percentage: ctx_usage, + }); + } + + Ok(events) + } +} + +// ============================================================================ +// 辅助函数 +// ============================================================================ + +/// 将 AwsEvent 序列化为 JSON 字符串(用于测试 round-trip) +pub fn serialize_event(event: &AwsEvent) -> Option { + match event { + AwsEvent::Content { text } => Some(serde_json::json!({"content": text}).to_string()), + AwsEvent::ToolUseStart { id, name } => { + Some(serde_json::json!({"toolUseId": id, "name": name}).to_string()) + } + AwsEvent::ToolUseInput { id, input } => { + Some(serde_json::json!({"toolUseId": id, "input": input}).to_string()) + } + AwsEvent::ToolUseStop { id } => { + Some(serde_json::json!({"toolUseId": id, "stop": true}).to_string()) + } + AwsEvent::Stop => Some(serde_json::json!({"stop": true}).to_string()), + AwsEvent::Usage { + credits, + context_percentage, + } => { + if *credits > 0.0 { + Some(serde_json::json!({"unit": "credit", "usage": credits}).to_string()) + } else if *context_percentage > 0.0 { + Some(serde_json::json!({"contextUsagePercentage": context_percentage}).to_string()) + } else { + None + } + } + AwsEvent::FollowupPrompt { content } => { + Some(serde_json::json!({"content": content, "followupPrompt": true}).to_string()) + } + AwsEvent::ParseError { .. } => None, + } +} + +/// 从事件列表中提取所有内容文本 +pub fn extract_content(events: &[AwsEvent]) -> String { + events + .iter() + .filter_map(|e| { + if let AwsEvent::Content { text } = e { + Some(text.as_str()) + } else { + None + } + }) + .collect::>() + .join("") +} + +/// 从事件列表中提取所有工具调用 +pub fn extract_tool_calls(events: &[AwsEvent]) -> Vec<(String, String, String)> { + let mut tool_calls: HashMap = HashMap::new(); + let mut completed: Vec = Vec::new(); + + for event in events { + match event { + AwsEvent::ToolUseStart { id, name } => { + tool_calls.entry(id.clone()).or_default().0 = name.clone(); + } + AwsEvent::ToolUseInput { id, input } => { + tool_calls.entry(id.clone()).or_default().1.push_str(input); + } + AwsEvent::ToolUseStop { id } => { + completed.push(id.clone()); + } + _ => {} + } + } + + completed + .into_iter() + .filter_map(|id| { + tool_calls + .remove(&id) + .map(|(name, input)| (id, name, input)) + }) + .collect() +} + +// ============================================================================ +// 测试模块 +// ============================================================================ + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_parser_new() { + let parser = AwsEventStreamParser::new(); + assert_eq!(parser.state(), &ParserState::Idle); + assert_eq!(parser.parse_error_count(), 0); + assert_eq!(parser.buffer_size(), 0); + } + + #[test] + fn test_parser_reset() { + let mut parser = AwsEventStreamParser::new(); + parser.process(b"{\"content\":\"hello\"}"); + parser.reset(); + assert_eq!(parser.state(), &ParserState::Idle); + assert_eq!(parser.buffer_size(), 0); + } + + #[test] + fn test_parse_content_event() { + let mut parser = AwsEventStreamParser::new(); + let events = parser.process(b"{\"content\":\"Hello, world!\"}"); + + assert_eq!(events.len(), 1); + assert!(matches!( + &events[0], + AwsEvent::Content { text } if text == "Hello, world!" + )); + } + + #[test] + fn test_parse_multiple_content_events() { + let mut parser = AwsEventStreamParser::new(); + let data = b"{\"content\":\"Hello\"}{\"content\":\", world!\"}"; + let events = parser.process(data); + + assert_eq!(events.len(), 2); + let content = extract_content(&events); + assert_eq!(content, "Hello, world!"); + } + + #[test] + fn test_parse_stop_event() { + let mut parser = AwsEventStreamParser::new(); + let events = parser.process(b"{\"stop\":true}"); + + assert_eq!(events.len(), 1); + assert!(matches!(&events[0], AwsEvent::Stop)); + } + + #[test] + fn test_parse_tool_use_complete() { + let mut parser = AwsEventStreamParser::new(); + + // 工具调用开始 + let events1 = parser.process(b"{\"toolUseId\":\"tool_1\",\"name\":\"read_file\"}"); + assert_eq!(events1.len(), 1); + assert!(matches!( + &events1[0], + AwsEvent::ToolUseStart { id, name } if id == "tool_1" && name == "read_file" + )); + + // 工具调用输入 + let events2 = parser + .process(b"{\"toolUseId\":\"tool_1\",\"input\":\"{\\\"path\\\":\\\"/tmp/test\\\"}\"}"); + assert_eq!(events2.len(), 1); + assert!(matches!( + &events2[0], + AwsEvent::ToolUseInput { id, input } if id == "tool_1" && input.contains("path") + )); + + // 工具调用结束 + let events3 = parser.process(b"{\"toolUseId\":\"tool_1\",\"stop\":true}"); + assert_eq!(events3.len(), 1); + assert!(matches!( + &events3[0], + AwsEvent::ToolUseStop { id } if id == "tool_1" + )); + } + + #[test] + fn test_parse_usage_event() { + let mut parser = AwsEventStreamParser::new(); + + // credits 使用量 + let events1 = parser.process(b"{\"unit\":\"credit\",\"usage\":0.34}"); + assert_eq!(events1.len(), 1); + assert!(matches!( + &events1[0], + AwsEvent::Usage { credits, .. } if (*credits - 0.34).abs() < 0.001 + )); + + // 上下文使用百分比 + let events2 = parser.process(b"{\"contextUsagePercentage\":54.36}"); + assert_eq!(events2.len(), 1); + assert!(matches!( + &events2[0], + AwsEvent::Usage { context_percentage, .. } if (*context_percentage - 54.36).abs() < 0.001 + )); + } + + #[test] + fn test_parse_followup_prompt() { + let mut parser = AwsEventStreamParser::new(); + let events = parser.process(b"{\"content\":\"suggestion\",\"followupPrompt\":true}"); + + assert_eq!(events.len(), 1); + assert!(matches!( + &events[0], + AwsEvent::FollowupPrompt { content } if content == "suggestion" + )); + } + + #[test] + fn test_parse_with_binary_prefix() { + let mut parser = AwsEventStreamParser::new(); + + // 模拟 AWS Event Stream 格式:二进制头部 + JSON + let mut data = vec![0x00, 0x00, 0x00, 0x1A]; // 一些二进制头部 + data.extend_from_slice(b"{\"content\":\"test\"}"); + data.extend_from_slice(&[0x00, 0x00]); // 一些二进制尾部 + + let events = parser.process(&data); + + assert_eq!(events.len(), 1); + assert!(matches!( + &events[0], + AwsEvent::Content { text } if text == "test" + )); + } + + #[test] + fn test_finish_completes_pending_tool_calls() { + let mut parser = AwsEventStreamParser::new(); + + // 开始工具调用但不结束 + parser.process(b"{\"toolUseId\":\"tool_1\",\"name\":\"test_tool\"}"); + parser.process(b"{\"toolUseId\":\"tool_1\",\"input\":\"test_input\"}"); + + // 调用 finish 应该完成未完成的工具调用 + let events = parser.finish(); + + assert_eq!(events.len(), 1); + assert!(matches!( + &events[0], + AwsEvent::ToolUseStop { id } if id == "tool_1" + )); + assert_eq!(parser.state(), &ParserState::Completed); + } + + #[test] + fn test_extract_content() { + let events = vec![ + AwsEvent::Content { + text: "Hello".to_string(), + }, + AwsEvent::Stop, + AwsEvent::Content { + text: ", world!".to_string(), + }, + ]; + + let content = extract_content(&events); + assert_eq!(content, "Hello, world!"); + } + + #[test] + fn test_extract_tool_calls() { + let events = vec![ + AwsEvent::ToolUseStart { + id: "t1".to_string(), + name: "func1".to_string(), + }, + AwsEvent::ToolUseInput { + id: "t1".to_string(), + input: "{\"a\":".to_string(), + }, + AwsEvent::ToolUseInput { + id: "t1".to_string(), + input: "1}".to_string(), + }, + AwsEvent::ToolUseStop { + id: "t1".to_string(), + }, + ]; + + let tool_calls = extract_tool_calls(&events); + assert_eq!(tool_calls.len(), 1); + assert_eq!( + tool_calls[0], + ( + "t1".to_string(), + "func1".to_string(), + "{\"a\":1}".to_string() + ) + ); + } + + #[test] + fn test_serialize_event() { + let event = AwsEvent::Content { + text: "test".to_string(), + }; + let json = serialize_event(&event).unwrap(); + assert!(json.contains("\"content\":\"test\"")); + + let event = AwsEvent::Stop; + let json = serialize_event(&event).unwrap(); + assert!(json.contains("\"stop\":true")); + } + + #[test] + fn test_buffer_overflow_protection() { + let mut parser = AwsEventStreamParser::with_max_buffer_size(100); + + // 发送超过缓冲区大小的数据 + let large_data = vec![b'x'; 200]; + let events = parser.process(&large_data); + + assert_eq!(events.len(), 1); + assert!( + matches!(&events[0], AwsEvent::ParseError { message, .. } if message.contains("缓冲区溢出")) + ); + assert_eq!(parser.parse_error_count(), 1); + } + + #[test] + fn test_empty_input() { + let mut parser = AwsEventStreamParser::new(); + let events = parser.process(b""); + assert!(events.is_empty()); + } + + #[test] + fn test_invalid_json_recovery() { + let mut parser = AwsEventStreamParser::new(); + + // 无效 JSON 后跟有效 JSON + let data = b"{invalid}{\"content\":\"valid\"}"; + let events = parser.process(data); + + // 应该有一个解析错误和一个有效内容 + assert!(events + .iter() + .any(|e| matches!(e, AwsEvent::ParseError { .. }))); + assert!(events + .iter() + .any(|e| matches!(e, AwsEvent::Content { text } if text == "valid"))); + } +} + +// ============================================================================ +// 增量解析测试(需求 2.6) +// ============================================================================ + +#[cfg(test)] +mod incremental_tests { + use super::*; + + #[test] + fn test_incremental_parsing_split_json() { + let mut parser = AwsEventStreamParser::new(); + + // 将一个 JSON 对象分成多个部分发送 + let events1 = parser.process(b"{\"content\":"); + assert!(events1.is_empty(), "不完整的 JSON 不应产生事件"); + assert!(parser.buffer_size() > 0, "缓冲区应该有数据"); + + let events2 = parser.process(b"\"Hello, "); + assert!(events2.is_empty(), "不完整的 JSON 不应产生事件"); + + let events3 = parser.process(b"world!\"}"); + assert_eq!(events3.len(), 1, "完整的 JSON 应该产生一个事件"); + assert!(matches!( + &events3[0], + AwsEvent::Content { text } if text == "Hello, world!" + )); + + // 缓冲区应该被清空 + assert_eq!(parser.buffer_size(), 0); + } + + #[test] + fn test_incremental_parsing_multiple_chunks() { + let mut parser = AwsEventStreamParser::new(); + + // 模拟网络分片:每次只发送几个字节 + let full_data = b"{\"content\":\"test\"}{\"stop\":true}"; + let mut all_events = Vec::new(); + + for chunk in full_data.chunks(5) { + let events = parser.process(chunk); + all_events.extend(events); + } + + // 应该解析出两个事件 + assert_eq!(all_events.len(), 2); + assert!(matches!(&all_events[0], AwsEvent::Content { text } if text == "test")); + assert!(matches!(&all_events[1], AwsEvent::Stop)); + } + + #[test] + fn test_incremental_parsing_with_binary_noise() { + let mut parser = AwsEventStreamParser::new(); + + // 模拟 AWS Event Stream 格式:二进制数据 + JSON + 二进制数据 + let mut data = Vec::new(); + data.extend_from_slice(&[0x00, 0x00, 0x00, 0x20]); // 二进制头部 + data.extend_from_slice(b"{\"content\":\"part1\"}"); + data.extend_from_slice(&[0x00, 0x00]); // 二进制分隔 + + let events1 = parser.process(&data); + assert_eq!(events1.len(), 1); + + // 继续发送更多数据 + let mut data2 = Vec::new(); + data2.extend_from_slice(&[0x00, 0x00, 0x00, 0x15]); // 二进制头部 + data2.extend_from_slice(b"{\"content\":\"part2\"}"); + + let events2 = parser.process(&data2); + assert_eq!(events2.len(), 1); + + let content = extract_content(&[events1, events2].concat()); + assert_eq!(content, "part1part2"); + } + + #[test] + fn test_incremental_tool_call_accumulation() { + let mut parser = AwsEventStreamParser::new(); + + // 工具调用开始 + let events1 = parser.process(b"{\"toolUseId\":\"t1\",\"name\":\"read_file\"}"); + assert_eq!(events1.len(), 1); + + // 分多次发送输入 + let events2 = parser.process(b"{\"toolUseId\":\"t1\",\"input\":\"{\\\"path\\\":\"}"); + assert_eq!(events2.len(), 1); + + let events3 = parser.process(b"{\"toolUseId\":\"t1\",\"input\":\"\\\"/tmp/test\\\"}\"}"); + assert_eq!(events3.len(), 1); + + // 结束工具调用 + let events4 = parser.process(b"{\"toolUseId\":\"t1\",\"stop\":true}"); + assert_eq!(events4.len(), 1); + + // 验证累积的输入 + let all_events = [events1, events2, events3, events4].concat(); + let tool_calls = extract_tool_calls(&all_events); + assert_eq!(tool_calls.len(), 1); + assert_eq!(tool_calls[0].0, "t1"); + assert_eq!(tool_calls[0].1, "read_file"); + assert!(tool_calls[0].2.contains("path")); + } + + #[test] + fn test_buffer_management_after_complete_json() { + let mut parser = AwsEventStreamParser::new(); + + // 发送完整 JSON 后跟部分 JSON + let data = b"{\"content\":\"complete\"}{\"content\":\"incom"; + let events = parser.process(data); + + // 应该只解析出完整的 JSON + assert_eq!(events.len(), 1); + assert!(matches!(&events[0], AwsEvent::Content { text } if text == "complete")); + + // 缓冲区应该保留不完整的部分 + assert!(parser.buffer_size() > 0); + + // 完成不完整的 JSON + let events2 = parser.process(b"plete\"}"); + assert_eq!(events2.len(), 1); + assert!(matches!(&events2[0], AwsEvent::Content { text } if text == "incomplete")); + + // 缓冲区应该被清空 + assert_eq!(parser.buffer_size(), 0); + } + + #[test] + fn test_finish_with_incomplete_json() { + let mut parser = AwsEventStreamParser::new(); + + // 发送不完整的 JSON + parser.process(b"{\"content\":\"incomplete"); + assert!(parser.buffer_size() > 0); + + // 调用 finish 不应该崩溃 + let events = parser.finish(); + + // 不完整的 JSON 不会产生事件 + assert!(events.is_empty()); + assert_eq!(parser.state(), &ParserState::Completed); + } + + #[test] + fn test_concurrent_tool_calls() { + let mut parser = AwsEventStreamParser::new(); + + // 开始两个并发的工具调用 + parser.process(b"{\"toolUseId\":\"t1\",\"name\":\"func1\"}"); + parser.process(b"{\"toolUseId\":\"t2\",\"name\":\"func2\"}"); + + // 交错发送输入 + parser.process(b"{\"toolUseId\":\"t1\",\"input\":\"input1\"}"); + parser.process(b"{\"toolUseId\":\"t2\",\"input\":\"input2\"}"); + parser.process(b"{\"toolUseId\":\"t1\",\"input\":\"_more\"}"); + + // 结束工具调用 + let events1 = parser.process(b"{\"toolUseId\":\"t1\",\"stop\":true}"); + let events2 = parser.process(b"{\"toolUseId\":\"t2\",\"stop\":true}"); + + assert!(matches!(&events1[0], AwsEvent::ToolUseStop { id } if id == "t1")); + assert!(matches!(&events2[0], AwsEvent::ToolUseStop { id } if id == "t2")); + } + + #[test] + fn test_state_transitions() { + let mut parser = AwsEventStreamParser::new(); + + // 初始状态 + assert_eq!(parser.state(), &ParserState::Idle); + + // 处理数据后变为 Parsing + parser.process(b"{\"content\":\"test\"}"); + assert_eq!(parser.state(), &ParserState::Parsing); + + // 完成后变为 Completed + parser.finish(); + assert_eq!(parser.state(), &ParserState::Completed); + + // 重置后回到 Idle + parser.reset(); + assert_eq!(parser.state(), &ParserState::Idle); + } + + #[test] + fn test_unicode_content_incremental() { + let mut parser = AwsEventStreamParser::new(); + + // 发送包含 Unicode 的 JSON(分片可能在 UTF-8 字符中间) + let json = "{\"content\":\"你好世界\"}"; + let bytes = json.as_bytes(); + + // 分成多个部分发送 + let events1 = parser.process(&bytes[..10]); + let events2 = parser.process(&bytes[10..20]); + let events3 = parser.process(&bytes[20..]); + + let all_events = [events1, events2, events3].concat(); + assert_eq!(all_events.len(), 1); + assert!(matches!( + &all_events[0], + AwsEvent::Content { text } if text == "你好世界" + )); + } + + #[test] + fn test_escaped_characters_in_json() { + let mut parser = AwsEventStreamParser::new(); + + // JSON 中包含转义字符(JSON 解析会自动处理转义) + let events = parser.process(b"{\"content\":\"line1\\nline2\\ttab\"}"); + assert_eq!(events.len(), 1); + // JSON 解析后,\n 变成换行符,\t 变成制表符 + assert!(matches!( + &events[0], + AwsEvent::Content { text } if text == "line1\nline2\ttab" + )); + } + + #[test] + fn test_nested_json_in_tool_input() { + let mut parser = AwsEventStreamParser::new(); + + // 工具输入包含嵌套 JSON + let events = parser.process( + b"{\"toolUseId\":\"t1\",\"name\":\"test\",\"input\":\"{\\\"nested\\\":{\\\"key\\\":\\\"value\\\"}}\"}" + ); + + assert_eq!(events.len(), 2); // ToolUseStart + ToolUseInput + assert!( + matches!(&events[0], AwsEvent::ToolUseStart { id, name } if id == "t1" && name == "test") + ); + assert!( + matches!(&events[1], AwsEvent::ToolUseInput { id, input } if id == "t1" && input.contains("nested")) + ); + } +} + +// ============================================================================ +// 错误恢复测试(需求 2.5) +// ============================================================================ + +#[cfg(test)] +mod error_recovery_tests { + use super::*; + + #[test] + fn test_recovery_from_invalid_json() { + let mut parser = AwsEventStreamParser::new(); + + // 无效 JSON 后跟有效 JSON + let data = b"{invalid json}{\"content\":\"valid\"}"; + let events = parser.process(data); + + // 应该有一个解析错误和一个有效内容 + let parse_errors: Vec<_> = events + .iter() + .filter(|e| matches!(e, AwsEvent::ParseError { .. })) + .collect(); + let content_events: Vec<_> = events + .iter() + .filter(|e| matches!(e, AwsEvent::Content { .. })) + .collect(); + + assert_eq!(parse_errors.len(), 1, "应该有一个解析错误"); + assert_eq!(content_events.len(), 1, "应该有一个有效内容"); + assert_eq!(parser.parse_error_count(), 1); + } + + #[test] + fn test_recovery_from_multiple_invalid_chunks() { + let mut parser = AwsEventStreamParser::new(); + + // 多个无效 JSON 交错有效 JSON + let data = b"{bad1}{\"content\":\"good1\"}{bad2}{\"content\":\"good2\"}"; + let events = parser.process(data); + + let parse_errors: Vec<_> = events + .iter() + .filter(|e| matches!(e, AwsEvent::ParseError { .. })) + .collect(); + let content_events: Vec<_> = events + .iter() + .filter(|e| matches!(e, AwsEvent::Content { .. })) + .collect(); + + assert_eq!(parse_errors.len(), 2, "应该有两个解析错误"); + assert_eq!(content_events.len(), 2, "应该有两个有效内容"); + assert_eq!(parser.parse_error_count(), 2); + } + + #[test] + fn test_recovery_preserves_content_order() { + let mut parser = AwsEventStreamParser::new(); + + // 确保错误恢复后内容顺序正确 + let data = + b"{\"content\":\"first\"}{invalid}{\"content\":\"second\"}{\"content\":\"third\"}"; + let events = parser.process(data); + + let content = extract_content(&events); + assert_eq!(content, "firstsecondthird"); + } + + #[test] + fn test_recovery_from_truncated_json() { + let mut parser = AwsEventStreamParser::new(); + + // 截断的 JSON(缺少结束括号)后跟有效 JSON + // 注意:截断的 JSON 会留在缓冲区中,直到收到更多数据 + let events1 = parser.process(b"{\"content\":\"truncated"); + assert!(events1.is_empty(), "截断的 JSON 不应产生事件"); + + // 发送更多数据,包括一个新的有效 JSON + // 由于缓冲区中有不完整的 JSON,新数据会被追加 + // 这里我们模拟一个场景:旧的不完整 JSON 被新数据"覆盖" + parser.reset(); // 重置以清除缓冲区 + + let events2 = parser.process(b"{\"content\":\"valid\"}"); + assert_eq!(events2.len(), 1); + assert!(matches!(&events2[0], AwsEvent::Content { text } if text == "valid")); + } + + #[test] + fn test_recovery_from_binary_garbage() { + let mut parser = AwsEventStreamParser::new(); + + // 二进制垃圾数据后跟有效 JSON + let mut data = vec![0xFF, 0xFE, 0x00, 0x01, 0x02]; + data.extend_from_slice(b"{\"content\":\"after garbage\"}"); + + let events = parser.process(&data); + + assert_eq!(events.len(), 1); + assert!(matches!( + &events[0], + AwsEvent::Content { text } if text == "after garbage" + )); + } + + #[test] + fn test_recovery_from_empty_json_object() { + let mut parser = AwsEventStreamParser::new(); + + // 空 JSON 对象(有效但不产生事件)后跟有效内容 + let data = b"{}{\"content\":\"after empty\"}"; + let events = parser.process(data); + + // 空对象不产生事件,但也不是错误 + let content_events: Vec<_> = events + .iter() + .filter(|e| matches!(e, AwsEvent::Content { .. })) + .collect(); + + assert_eq!(content_events.len(), 1); + assert_eq!(parser.parse_error_count(), 0, "空对象不应计为错误"); + } + + #[test] + fn test_recovery_from_unknown_event_type() { + let mut parser = AwsEventStreamParser::new(); + + // 未知事件类型(有效 JSON 但不是已知事件)后跟有效内容 + let data = b"{\"unknownField\":\"value\"}{\"content\":\"known\"}"; + let events = parser.process(data); + + // 未知事件类型不产生事件,但也不是错误 + let content_events: Vec<_> = events + .iter() + .filter(|e| matches!(e, AwsEvent::Content { .. })) + .collect(); + + assert_eq!(content_events.len(), 1); + assert_eq!(parser.parse_error_count(), 0, "未知事件类型不应计为错误"); + } + + #[test] + fn test_recovery_from_malformed_tool_call() { + let mut parser = AwsEventStreamParser::new(); + + // 格式错误的工具调用(缺少必要字段)后跟有效工具调用 + let data = b"{\"toolUseId\":\"t1\"}{\"toolUseId\":\"t2\",\"name\":\"valid_tool\"}"; + let events = parser.process(data); + + // 第一个工具调用缺少 name,但仍然是有效 JSON + // 第二个工具调用是完整的 + let tool_starts: Vec<_> = events + .iter() + .filter(|e| matches!(e, AwsEvent::ToolUseStart { .. })) + .collect(); + + assert_eq!(tool_starts.len(), 1, "只有一个有效的工具调用开始"); + } + + #[test] + fn test_error_count_accumulates() { + let mut parser = AwsEventStreamParser::new(); + + // 多次处理无效数据 + parser.process(b"{invalid1}"); + assert_eq!(parser.parse_error_count(), 1); + + parser.process(b"{invalid2}"); + assert_eq!(parser.parse_error_count(), 2); + + parser.process(b"{\"content\":\"valid\"}"); + assert_eq!(parser.parse_error_count(), 2, "有效数据不应增加错误计数"); + + parser.process(b"{invalid3}"); + assert_eq!(parser.parse_error_count(), 3); + } + + #[test] + fn test_error_count_resets_on_reset() { + let mut parser = AwsEventStreamParser::new(); + + parser.process(b"{invalid}"); + assert_eq!(parser.parse_error_count(), 1); + + parser.reset(); + assert_eq!(parser.parse_error_count(), 0); + } + + #[test] + fn test_parse_error_contains_raw_data() { + let mut parser = AwsEventStreamParser::new(); + + let events = parser.process(b"{invalid json}"); + + assert_eq!(events.len(), 1); + if let AwsEvent::ParseError { message, raw_data } = &events[0] { + assert!(message.contains("JSON"), "错误消息应该提到 JSON"); + assert!(raw_data.is_some(), "应该包含原始数据"); + assert!( + raw_data.as_ref().unwrap().contains("invalid"), + "原始数据应该包含无效内容" + ); + } else { + panic!("应该是 ParseError 事件"); + } + } + + #[test] + fn test_recovery_continues_tool_accumulation() { + let mut parser = AwsEventStreamParser::new(); + + // 开始工具调用 + parser.process(b"{\"toolUseId\":\"t1\",\"name\":\"test\"}"); + + // 发送无效数据 + parser.process(b"{invalid}"); + + // 继续工具调用输入 + let events = parser.process(b"{\"toolUseId\":\"t1\",\"input\":\"test_input\"}"); + + // 工具调用应该继续正常工作 + assert_eq!(events.len(), 1); + assert!(matches!(&events[0], AwsEvent::ToolUseInput { id, .. } if id == "t1")); + } + + #[test] + fn test_recovery_from_deeply_nested_invalid_json() { + let mut parser = AwsEventStreamParser::new(); + + // 深度嵌套但无效的 JSON + let data = b"{\"a\":{\"b\":{\"c\":invalid}}}{\"content\":\"valid\"}"; + let events = parser.process(data); + + let parse_errors: Vec<_> = events + .iter() + .filter(|e| matches!(e, AwsEvent::ParseError { .. })) + .collect(); + let content_events: Vec<_> = events + .iter() + .filter(|e| matches!(e, AwsEvent::Content { .. })) + .collect(); + + assert_eq!(parse_errors.len(), 1); + assert_eq!(content_events.len(), 1); + } + + #[test] + fn test_recovery_from_json_with_wrong_types() { + let mut parser = AwsEventStreamParser::new(); + + // JSON 有效但字段类型错误(content 应该是字符串,这里是数字) + // 这种情况下 JSON 解析成功,但不会产生 Content 事件 + let data = b"{\"content\":123}{\"content\":\"valid string\"}"; + let events = parser.process(data); + + let content_events: Vec<_> = events + .iter() + .filter(|e| matches!(e, AwsEvent::Content { .. })) + .collect(); + + // 只有字符串类型的 content 会产生事件 + assert_eq!(content_events.len(), 1); + assert!(matches!( + content_events[0], + AwsEvent::Content { text } if text == "valid string" + )); + } +} + +// ============================================================================ +// 属性测试(Property-Based Testing) +// ============================================================================ + +#[cfg(test)] +mod property_tests { + use super::*; + use proptest::prelude::*; + + // ======================================================================== + // 生成器(Generators) + // ======================================================================== + + /// 生成有效的内容文本 + fn arb_content_text() -> impl Strategy { + // 生成不包含控制字符的 Unicode 字符串 + prop::string::string_regex("[a-zA-Z0-9\\u4e00-\\u9fff .,!?\\-_]{0,100}") + .unwrap() + .prop_filter("非空字符串", |s| !s.is_empty()) + } + + /// 生成有效的工具 ID + fn arb_tool_id() -> impl Strategy { + prop::string::string_regex("tool_[a-zA-Z0-9]{4,12}").unwrap() + } + + /// 生成有效的工具名称 + fn arb_tool_name() -> impl Strategy { + prop::string::string_regex("[a-z_][a-z0-9_]{2,20}").unwrap() + } + + /// 生成有效的工具输入(简单 JSON) + fn arb_tool_input() -> impl Strategy { + prop::string::string_regex(r#"\{"[a-z]+":"[a-zA-Z0-9]+"\}"#).unwrap() + } + + /// 生成 Content 事件 + fn arb_content_event() -> impl Strategy { + arb_content_text().prop_map(|text| AwsEvent::Content { text }) + } + + /// 生成 ToolUseStart 事件 + fn arb_tool_use_start_event() -> impl Strategy { + (arb_tool_id(), arb_tool_name()).prop_map(|(id, name)| AwsEvent::ToolUseStart { id, name }) + } + + /// 生成 ToolUseInput 事件 + fn arb_tool_use_input_event(id: String) -> impl Strategy { + arb_tool_input().prop_map(move |input| AwsEvent::ToolUseInput { + id: id.clone(), + input, + }) + } + + /// 生成 Usage 事件 + fn arb_usage_event() -> impl Strategy { + prop_oneof![ + (0.01f64..100.0f64).prop_map(|credits| AwsEvent::Usage { + credits, + context_percentage: 0.0, + }), + (0.01f64..100.0f64).prop_map(|ctx| AwsEvent::Usage { + credits: 0.0, + context_percentage: ctx, + }), + ] + } + + /// 生成可序列化的事件(排除 ParseError 和 FollowupPrompt) + fn arb_serializable_event() -> impl Strategy { + prop_oneof![arb_content_event(), Just(AwsEvent::Stop), arb_usage_event(),] + } + + /// 生成事件序列 + fn arb_event_sequence() -> impl Strategy> { + prop::collection::vec(arb_serializable_event(), 1..10) + } + + // ======================================================================== + // Property 1: AWS Event Stream 解析 Round-Trip + // + // *对于任意*有效的 AWS Event Stream 数据,解析后重新序列化应该产生 + // 语义等价的数据(内容、工具调用、使用量信息保持一致)。 + // + // **验证: 需求 2.1, 2.2, 2.3, 2.6** + // ======================================================================== + + proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// Property 1: Content 事件 Round-Trip + /// + /// **Feature: true-streaming-support, Property 1: AWS Event Stream 解析 Round-Trip** + /// **Validates: Requirements 2.1, 2.2** + #[test] + fn prop_content_event_round_trip(text in arb_content_text()) { + let event = AwsEvent::Content { text: text.clone() }; + + // 序列化 + let json = serialize_event(&event).expect("Content 事件应该可序列化"); + + // 解析 + let mut parser = AwsEventStreamParser::new(); + let events = parser.process(json.as_bytes()); + + // 验证 + prop_assert_eq!(events.len(), 1, "应该解析出一个事件"); + match &events[0] { + AwsEvent::Content { text: parsed_text } => { + prop_assert_eq!(parsed_text, &text, "内容应该一致"); + } + _ => prop_assert!(false, "应该是 Content 事件"), + } + } + + /// Property 1: Stop 事件 Round-Trip + /// + /// **Feature: true-streaming-support, Property 1: AWS Event Stream 解析 Round-Trip** + /// **Validates: Requirements 2.4** + #[test] + fn prop_stop_event_round_trip(_dummy in Just(())) { + let event = AwsEvent::Stop; + + // 序列化 + let json = serialize_event(&event).expect("Stop 事件应该可序列化"); + + // 解析 + let mut parser = AwsEventStreamParser::new(); + let events = parser.process(json.as_bytes()); + + // 验证 + prop_assert_eq!(events.len(), 1, "应该解析出一个事件"); + prop_assert!(matches!(&events[0], AwsEvent::Stop), "应该是 Stop 事件"); + } + + /// Property 1: Usage 事件 Round-Trip (credits) + /// + /// **Feature: true-streaming-support, Property 1: AWS Event Stream 解析 Round-Trip** + /// **Validates: Requirements 2.1** + #[test] + fn prop_usage_credits_round_trip(credits in 0.01f64..100.0f64) { + let event = AwsEvent::Usage { credits, context_percentage: 0.0 }; + + // 序列化 + let json = serialize_event(&event).expect("Usage 事件应该可序列化"); + + // 解析 + let mut parser = AwsEventStreamParser::new(); + let events = parser.process(json.as_bytes()); + + // 验证 + prop_assert_eq!(events.len(), 1, "应该解析出一个事件"); + match &events[0] { + AwsEvent::Usage { credits: parsed_credits, .. } => { + prop_assert!((parsed_credits - credits).abs() < 0.001, "credits 应该一致"); + } + _ => prop_assert!(false, "应该是 Usage 事件"), + } + } + + /// Property 1: Usage 事件 Round-Trip (context_percentage) + /// + /// **Feature: true-streaming-support, Property 1: AWS Event Stream 解析 Round-Trip** + /// **Validates: Requirements 2.1** + #[test] + fn prop_usage_context_round_trip(ctx in 0.01f64..100.0f64) { + let event = AwsEvent::Usage { credits: 0.0, context_percentage: ctx }; + + // 序列化 + let json = serialize_event(&event).expect("Usage 事件应该可序列化"); + + // 解析 + let mut parser = AwsEventStreamParser::new(); + let events = parser.process(json.as_bytes()); + + // 验证 + prop_assert_eq!(events.len(), 1, "应该解析出一个事件"); + match &events[0] { + AwsEvent::Usage { context_percentage: parsed_ctx, .. } => { + prop_assert!((parsed_ctx - ctx).abs() < 0.001, "context_percentage 应该一致"); + } + _ => prop_assert!(false, "应该是 Usage 事件"), + } + } + + /// Property 1: 工具调用完整流程 Round-Trip + /// + /// **Feature: true-streaming-support, Property 1: AWS Event Stream 解析 Round-Trip** + /// **Validates: Requirements 2.3** + #[test] + fn prop_tool_call_round_trip( + id in arb_tool_id(), + name in arb_tool_name(), + input in arb_tool_input() + ) { + let mut parser = AwsEventStreamParser::new(); + + // 序列化并解析工具调用开始 + let start_event = AwsEvent::ToolUseStart { id: id.clone(), name: name.clone() }; + let start_json = serialize_event(&start_event).expect("ToolUseStart 应该可序列化"); + let start_events = parser.process(start_json.as_bytes()); + + prop_assert_eq!(start_events.len(), 1); + match &start_events[0] { + AwsEvent::ToolUseStart { id: parsed_id, name: parsed_name } => { + prop_assert_eq!(parsed_id, &id); + prop_assert_eq!(parsed_name, &name); + } + _ => prop_assert!(false, "应该是 ToolUseStart 事件"), + } + + // 序列化并解析工具调用输入 + let input_event = AwsEvent::ToolUseInput { id: id.clone(), input: input.clone() }; + let input_json = serialize_event(&input_event).expect("ToolUseInput 应该可序列化"); + let input_events = parser.process(input_json.as_bytes()); + + prop_assert_eq!(input_events.len(), 1); + match &input_events[0] { + AwsEvent::ToolUseInput { id: parsed_id, input: parsed_input } => { + prop_assert_eq!(parsed_id, &id); + prop_assert_eq!(parsed_input, &input); + } + _ => prop_assert!(false, "应该是 ToolUseInput 事件"), + } + + // 序列化并解析工具调用结束 + let stop_event = AwsEvent::ToolUseStop { id: id.clone() }; + let stop_json = serialize_event(&stop_event).expect("ToolUseStop 应该可序列化"); + let stop_events = parser.process(stop_json.as_bytes()); + + prop_assert_eq!(stop_events.len(), 1); + match &stop_events[0] { + AwsEvent::ToolUseStop { id: parsed_id } => { + prop_assert_eq!(parsed_id, &id); + } + _ => prop_assert!(false, "应该是 ToolUseStop 事件"), + } + } + + /// Property 1: 多事件序列 Round-Trip + /// + /// **Feature: true-streaming-support, Property 1: AWS Event Stream 解析 Round-Trip** + /// **Validates: Requirements 2.1, 2.2, 2.6** + #[test] + fn prop_event_sequence_round_trip(events in arb_event_sequence()) { + let mut parser = AwsEventStreamParser::new(); + + // 将所有事件序列化为一个字节流 + let mut data = Vec::new(); + for event in &events { + if let Some(json) = serialize_event(event) { + data.extend_from_slice(json.as_bytes()); + } + } + + // 解析 + let parsed_events = parser.process(&data); + + // 验证:解析出的事件数量应该与原始事件数量一致 + // (排除无法序列化的事件) + let serializable_count = events.iter() + .filter(|e| serialize_event(e).is_some()) + .count(); + + prop_assert_eq!( + parsed_events.len(), + serializable_count, + "解析出的事件数量应该与可序列化的事件数量一致" + ); + } + + /// Property 1: 增量解析保持语义等价 + /// + /// **Feature: true-streaming-support, Property 1: AWS Event Stream 解析 Round-Trip** + /// **Validates: Requirements 2.6** + #[test] + fn prop_incremental_parsing_semantic_equivalence( + text in arb_content_text(), + chunk_size in 1usize..20usize + ) { + let event = AwsEvent::Content { text: text.clone() }; + let json = serialize_event(&event).expect("Content 事件应该可序列化"); + let bytes = json.as_bytes(); + + // 一次性解析 + let mut parser1 = AwsEventStreamParser::new(); + let events1 = parser1.process(bytes); + let final1 = parser1.finish(); + let all_events1: Vec<_> = events1.into_iter().chain(final1).collect(); + + // 增量解析 + let mut parser2 = AwsEventStreamParser::new(); + let mut all_events2 = Vec::new(); + for chunk in bytes.chunks(chunk_size) { + all_events2.extend(parser2.process(chunk)); + } + all_events2.extend(parser2.finish()); + + // 验证:两种方式解析出的内容应该一致 + let content1 = extract_content(&all_events1); + let content2 = extract_content(&all_events2); + + prop_assert_eq!(content1, content2, "增量解析应该产生相同的内容"); + } + } +} diff --git a/src-tauri/src/streaming/converter.rs b/src-tauri/src/streaming/converter.rs new file mode 100644 index 000000000..c746335aa --- /dev/null +++ b/src-tauri/src/streaming/converter.rs @@ -0,0 +1,1518 @@ +//! 流式格式转换器 +//! +//! 在不同流式格式之间转换,支持 AWS Event Stream、Anthropic SSE 和 OpenAI SSE。 +//! +//! # 需求覆盖 +//! +//! - 需求 3.1: AWS Event Stream 到 Anthropic SSE 转换 +//! - 需求 3.2: AWS Event Stream 到 OpenAI SSE 转换 +//! - 需求 3.3: Anthropic SSE 到 OpenAI SSE 转换 +//! - 需求 3.5: 处理工具调用参数中的部分 JSON + +use crate::streaming::aws_parser::{AwsEvent, AwsEventStreamParser}; +use serde::{Deserialize, Serialize}; +use std::collections::HashMap; +use std::time::{SystemTime, UNIX_EPOCH}; +use uuid::Uuid; + +/// 流式格式类型 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +pub enum StreamFormat { + /// AWS Event Stream (Kiro/CodeWhisperer) + AwsEventStream, + /// Anthropic SSE 格式 + AnthropicSse, + /// OpenAI SSE 格式 + OpenAiSse, +} + +/// 转换器状态 +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum ConverterState { + /// 初始状态 + Idle, + /// 正在转换 + Converting, + /// 已完成 + Completed, + /// 错误状态 + Error(String), +} + +impl Default for ConverterState { + fn default() -> Self { + Self::Idle + } +} + +/// 工具调用累积器 +/// +/// 用于跟踪正在进行的工具调用,累积部分 JSON 输入 +#[derive(Debug, Clone, Default)] +struct ToolCallAccumulator { + /// 工具调用 ID + id: String, + /// 工具名称 + name: String, + /// 累积的输入 JSON(部分) + input: String, + /// 是否已发送开始事件 + started: bool, + /// 内容块索引(用于 Anthropic 格式) + index: u32, +} + +/// 部分 JSON 累积器 +/// +/// 用于处理工具调用参数中的部分 JSON +/// 对应需求 3.5 +#[derive(Debug, Clone, Default)] +pub struct PartialJsonAccumulator { + /// 累积的 JSON 字符串 + buffer: String, + /// 括号深度(用于检测 JSON 完整性) + brace_depth: i32, + /// 是否在字符串内 + in_string: bool, + /// 是否转义下一个字符 + escape_next: bool, +} + +impl PartialJsonAccumulator { + /// 创建新的累积器 + pub fn new() -> Self { + Self::default() + } + + /// 追加部分 JSON + /// + /// 返回 true 如果 JSON 已完整 + pub fn append(&mut self, partial: &str) -> bool { + for ch in partial.chars() { + self.buffer.push(ch); + + if self.escape_next { + self.escape_next = false; + continue; + } + + match ch { + '\\' if self.in_string => self.escape_next = true, + '"' => self.in_string = !self.in_string, + '{' | '[' if !self.in_string => self.brace_depth += 1, + '}' | ']' if !self.in_string => self.brace_depth -= 1, + _ => {} + } + } + + self.is_complete() + } + + /// 检查 JSON 是否完整 + pub fn is_complete(&self) -> bool { + !self.buffer.is_empty() && self.brace_depth == 0 && !self.in_string + } + + /// 获取累积的 JSON + pub fn get_json(&self) -> &str { + &self.buffer + } + + /// 重置累积器 + pub fn reset(&mut self) { + self.buffer.clear(); + self.brace_depth = 0; + self.in_string = false; + self.escape_next = false; + } + + /// 获取缓冲区长度 + pub fn len(&self) -> usize { + self.buffer.len() + } + + /// 检查缓冲区是否为空 + pub fn is_empty(&self) -> bool { + self.buffer.is_empty() + } +} + +/// 流式格式转换器 +/// +/// 支持在不同流式格式之间转换。 +#[derive(Debug)] +pub struct StreamConverter { + /// 源格式 + source_format: StreamFormat, + /// 目标格式 + target_format: StreamFormat, + /// AWS 解析器(如果源是 AWS Event Stream) + aws_parser: Option, + /// 状态 + state: ConverterState, + /// 响应 ID + response_id: String, + /// 模型名称 + model: String, + /// 工具调用累积器 + tool_accumulators: HashMap, + /// 下一个内容块索引(用于 Anthropic 格式) + next_content_block_index: u32, + /// 是否已发送 message_start(用于 Anthropic 格式) + message_started: bool, + /// 累积的内容(用于重建完整响应) + accumulated_content: String, +} + +impl StreamConverter { + /// 创建新的转换器 + pub fn new(source: StreamFormat, target: StreamFormat) -> Self { + let aws_parser = if source == StreamFormat::AwsEventStream { + Some(AwsEventStreamParser::new()) + } else { + None + }; + + Self { + source_format: source, + target_format: target, + aws_parser, + state: ConverterState::Idle, + response_id: format!("chatcmpl-{}", Uuid::new_v4()), + model: String::new(), + tool_accumulators: HashMap::new(), + next_content_block_index: 0, + message_started: false, + accumulated_content: String::new(), + } + } + + /// 创建带模型名称的转换器 + pub fn with_model(source: StreamFormat, target: StreamFormat, model: &str) -> Self { + let mut converter = Self::new(source, target); + converter.model = model.to_string(); + converter + } + + /// 获取当前状态 + pub fn state(&self) -> &ConverterState { + &self.state + } + + /// 获取响应 ID + pub fn response_id(&self) -> &str { + &self.response_id + } + + /// 获取累积的内容 + pub fn accumulated_content(&self) -> &str { + &self.accumulated_content + } + + /// 重置转换器 + pub fn reset(&mut self) { + if let Some(parser) = &mut self.aws_parser { + parser.reset(); + } + self.state = ConverterState::Idle; + self.response_id = format!("chatcmpl-{}", Uuid::new_v4()); + self.tool_accumulators.clear(); + self.next_content_block_index = 0; + self.message_started = false; + self.accumulated_content.clear(); + } + + /// 转换 chunk + /// + /// 将源格式的 chunk 转换为目标格式的 SSE 事件列表。 + /// + /// # 参数 + /// + /// * `chunk` - 源格式的字节数据 + /// + /// # 返回 + /// + /// 目标格式的 SSE 事件字符串列表 + pub fn convert(&mut self, chunk: &[u8]) -> Vec { + if self.state == ConverterState::Idle { + self.state = ConverterState::Converting; + } + + match self.source_format { + StreamFormat::AwsEventStream => self.convert_aws_event_stream(chunk), + StreamFormat::AnthropicSse => self.convert_anthropic_sse(chunk), + StreamFormat::OpenAiSse => self.convert_openai_sse(chunk), + } + } + + /// 完成转换 + /// + /// 处理剩余数据并生成结束事件。 + pub fn finish(&mut self) -> Vec { + let mut events = Vec::new(); + + // 处理 AWS 解析器中的剩余数据 + if let Some(parser) = &mut self.aws_parser { + let aws_events = parser.finish(); + for aws_event in aws_events { + events.extend(self.convert_aws_event(&aws_event)); + } + } + + // 生成结束事件 + events.extend(self.generate_end_events()); + + self.state = ConverterState::Completed; + events + } + + /// 转换 AWS Event Stream + fn convert_aws_event_stream(&mut self, chunk: &[u8]) -> Vec { + let parser = self.aws_parser.as_mut().expect("AWS parser should exist"); + let aws_events = parser.process(chunk); + + let mut sse_events = Vec::new(); + for aws_event in aws_events { + sse_events.extend(self.convert_aws_event(&aws_event)); + } + sse_events + } + + /// 转换单个 AWS 事件 + fn convert_aws_event(&mut self, event: &AwsEvent) -> Vec { + match self.target_format { + StreamFormat::AnthropicSse => self.aws_to_anthropic(event), + StreamFormat::OpenAiSse => self.aws_to_openai(event), + StreamFormat::AwsEventStream => { + // 源和目标相同,直接序列化 + if let Some(json) = crate::streaming::aws_parser::serialize_event(event) { + vec![json] + } else { + vec![] + } + } + } + } + + /// AWS Event Stream 到 Anthropic SSE 转换 + /// + /// 对应需求 3.1 + fn aws_to_anthropic(&mut self, event: &AwsEvent) -> Vec { + let mut sse_events = Vec::new(); + + // 确保发送 message_start + if !self.message_started { + sse_events.push(self.create_anthropic_message_start()); + self.message_started = true; + } + + match event { + AwsEvent::Content { text } => { + // 累积内容 + self.accumulated_content.push_str(text); + + // 如果是第一个内容块,发送 content_block_start + if self.next_content_block_index == 0 { + sse_events.push(self.create_anthropic_content_block_start_text(0)); + self.next_content_block_index = 1; + } + + // 发送 content_block_delta + sse_events.push(self.create_anthropic_text_delta(0, text)); + } + AwsEvent::ToolUseStart { id, name } => { + // 如果有文本内容块,先关闭它 + if self.next_content_block_index > 0 && self.accumulated_content.is_empty() { + // 没有文本内容,不需要关闭 + } else if self.next_content_block_index > 0 { + sse_events.push(self.create_anthropic_content_block_stop(0)); + } + + let index = self.next_content_block_index; + self.next_content_block_index += 1; + + // 创建工具调用累积器 + self.tool_accumulators.insert( + id.clone(), + ToolCallAccumulator { + id: id.clone(), + name: name.clone(), + input: String::new(), + started: true, + index, + }, + ); + + // 发送 content_block_start (tool_use) + sse_events.push(self.create_anthropic_content_block_start_tool(index, id, name)); + } + AwsEvent::ToolUseInput { id, input } => { + if let Some(acc) = self.tool_accumulators.get_mut(id) { + acc.input.push_str(input); + } + // 发送 input_json_delta + if let Some(acc) = self.tool_accumulators.get(id) { + sse_events.push(self.create_anthropic_input_json_delta(acc.index, input)); + } + } + AwsEvent::ToolUseStop { id } => { + if let Some(acc) = self.tool_accumulators.remove(id) { + // 发送 content_block_stop + sse_events.push(self.create_anthropic_content_block_stop(acc.index)); + } + } + AwsEvent::Stop => { + // 关闭所有未关闭的内容块 + if self.next_content_block_index > 0 && !self.accumulated_content.is_empty() { + sse_events.push(self.create_anthropic_content_block_stop(0)); + } + // message_delta 和 message_stop 在 finish() 中处理 + } + AwsEvent::Usage { + credits, + context_percentage, + } => { + // Usage 信息在 message_delta 中发送 + // 这里暂时忽略,在 finish() 中处理 + let _ = (credits, context_percentage); + } + AwsEvent::FollowupPrompt { .. } | AwsEvent::ParseError { .. } => { + // 忽略这些事件 + } + } + + sse_events + } + + /// AWS Event Stream 到 OpenAI SSE 转换 + /// + /// 对应需求 3.2 + fn aws_to_openai(&mut self, event: &AwsEvent) -> Vec { + let mut sse_events = Vec::new(); + + match event { + AwsEvent::Content { text } => { + // 累积内容 + self.accumulated_content.push_str(text); + // 发送 chunk + sse_events.push(self.create_openai_content_chunk(text, false)); + } + AwsEvent::ToolUseStart { id, name } => { + let index = self.tool_accumulators.len() as u32; + self.tool_accumulators.insert( + id.clone(), + ToolCallAccumulator { + id: id.clone(), + name: name.clone(), + input: String::new(), + started: true, + index, + }, + ); + // 发送工具调用开始 chunk + sse_events.push(self.create_openai_tool_call_chunk(index, id, name, "", true)); + } + AwsEvent::ToolUseInput { id, input } => { + let (index, tool_id, tool_name) = + if let Some(acc) = self.tool_accumulators.get_mut(id) { + acc.input.push_str(input); + (acc.index, acc.id.clone(), acc.name.clone()) + } else { + return sse_events; + }; + // 发送工具调用参数增量 + sse_events.push( + self.create_openai_tool_call_chunk(index, &tool_id, &tool_name, input, false), + ); + } + AwsEvent::ToolUseStop { id } => { + // OpenAI 格式不需要显式的工具调用结束事件 + self.tool_accumulators.remove(id); + } + AwsEvent::Stop => { + // 结束事件在 finish() 中处理 + } + AwsEvent::Usage { .. } + | AwsEvent::FollowupPrompt { .. } + | AwsEvent::ParseError { .. } => { + // 忽略这些事件 + } + } + + sse_events + } + + /// 转换 Anthropic SSE(直通或转换为 OpenAI) + fn convert_anthropic_sse(&mut self, chunk: &[u8]) -> Vec { + // 解析 SSE 数据 + let data = match String::from_utf8(chunk.to_vec()) { + Ok(s) => s, + Err(_) => return vec![], + }; + + match self.target_format { + StreamFormat::AnthropicSse => { + // 直通 + vec![data] + } + StreamFormat::OpenAiSse => { + // 转换为 OpenAI 格式 + self.anthropic_to_openai(&data) + } + StreamFormat::AwsEventStream => { + // 不支持反向转换 + vec![] + } + } + } + + /// Anthropic SSE 到 OpenAI SSE 转换 + /// + /// 对应需求 3.3 + fn anthropic_to_openai(&mut self, data: &str) -> Vec { + let mut sse_events = Vec::new(); + + // 解析 SSE 事件 + for line in data.lines() { + if let Some(json_str) = line.strip_prefix("data: ") { + if json_str == "[DONE]" { + sse_events.push("data: [DONE]\n\n".to_string()); + continue; + } + + if let Ok(event) = serde_json::from_str::(json_str) { + if let Some(event_type) = event.get("type").and_then(|t| t.as_str()) { + match event_type { + "content_block_delta" => { + if let Some(delta) = event.get("delta") { + if let Some(text) = delta.get("text").and_then(|t| t.as_str()) { + self.accumulated_content.push_str(text); + sse_events + .push(self.create_openai_content_chunk(text, false)); + } else if let Some(partial_json) = + delta.get("partial_json").and_then(|t| t.as_str()) + { + // 工具调用参数增量 + let index = event + .get("index") + .and_then(|i| i.as_u64()) + .unwrap_or(0) + as u32; + let tool_info = self + .tool_accumulators + .values_mut() + .find(|a| a.index == index) + .map(|acc| { + acc.input.push_str(partial_json); + (acc.index, acc.id.clone(), acc.name.clone()) + }); + if let Some((idx, tool_id, tool_name)) = tool_info { + sse_events.push(self.create_openai_tool_call_chunk( + idx, + &tool_id, + &tool_name, + partial_json, + false, + )); + } + } + } + } + "content_block_start" => { + if let Some(content_block) = event.get("content_block") { + if content_block.get("type").and_then(|t| t.as_str()) + == Some("tool_use") + { + let id = content_block + .get("id") + .and_then(|i| i.as_str()) + .unwrap_or(""); + let name = content_block + .get("name") + .and_then(|n| n.as_str()) + .unwrap_or(""); + let index = event + .get("index") + .and_then(|i| i.as_u64()) + .unwrap_or(0) + as u32; + self.tool_accumulators.insert( + id.to_string(), + ToolCallAccumulator { + id: id.to_string(), + name: name.to_string(), + input: String::new(), + started: true, + index, + }, + ); + sse_events.push(self.create_openai_tool_call_chunk( + index, id, name, "", true, + )); + } + } + } + "message_stop" => { + sse_events.push(self.create_openai_finish_chunk("stop")); + sse_events.push("data: [DONE]\n\n".to_string()); + } + _ => {} + } + } + } + } + } + + sse_events + } + + /// 转换 OpenAI SSE(直通) + fn convert_openai_sse(&mut self, chunk: &[u8]) -> Vec { + match String::from_utf8(chunk.to_vec()) { + Ok(s) => vec![s], + Err(_) => vec![], + } + } + + /// 生成结束事件 + fn generate_end_events(&mut self) -> Vec { + match self.target_format { + StreamFormat::AnthropicSse => { + let mut events = Vec::new(); + // message_delta + events.push(self.create_anthropic_message_delta()); + // message_stop + events.push(self.create_anthropic_message_stop()); + events + } + StreamFormat::OpenAiSse => { + let finish_reason = if self.tool_accumulators.is_empty() { + "stop" + } else { + "tool_calls" + }; + vec![ + self.create_openai_finish_chunk(finish_reason), + "data: [DONE]\n\n".to_string(), + ] + } + StreamFormat::AwsEventStream => { + vec![] + } + } + } + + // ======================================================================== + // Anthropic SSE 事件创建辅助方法 + // ======================================================================== + + fn create_anthropic_message_start(&self) -> String { + let event = serde_json::json!({ + "type": "message_start", + "message": { + "id": self.response_id, + "type": "message", + "role": "assistant", + "model": self.model, + "content": [], + "stop_reason": null, + "stop_sequence": null, + "usage": { + "input_tokens": 0, + "output_tokens": 0 + } + } + }); + format!("event: message_start\ndata: {}\n\n", event) + } + + fn create_anthropic_content_block_start_text(&self, index: u32) -> String { + let event = serde_json::json!({ + "type": "content_block_start", + "index": index, + "content_block": { + "type": "text", + "text": "" + } + }); + format!("event: content_block_start\ndata: {}\n\n", event) + } + + fn create_anthropic_content_block_start_tool( + &self, + index: u32, + id: &str, + name: &str, + ) -> String { + let event = serde_json::json!({ + "type": "content_block_start", + "index": index, + "content_block": { + "type": "tool_use", + "id": id, + "name": name, + "input": {} + } + }); + format!("event: content_block_start\ndata: {}\n\n", event) + } + + fn create_anthropic_text_delta(&self, index: u32, text: &str) -> String { + let event = serde_json::json!({ + "type": "content_block_delta", + "index": index, + "delta": { + "type": "text_delta", + "text": text + } + }); + format!("event: content_block_delta\ndata: {}\n\n", event) + } + + fn create_anthropic_input_json_delta(&self, index: u32, partial_json: &str) -> String { + let event = serde_json::json!({ + "type": "content_block_delta", + "index": index, + "delta": { + "type": "input_json_delta", + "partial_json": partial_json + } + }); + format!("event: content_block_delta\ndata: {}\n\n", event) + } + + fn create_anthropic_content_block_stop(&self, index: u32) -> String { + let event = serde_json::json!({ + "type": "content_block_stop", + "index": index + }); + format!("event: content_block_stop\ndata: {}\n\n", event) + } + + fn create_anthropic_message_delta(&self) -> String { + let event = serde_json::json!({ + "type": "message_delta", + "delta": { + "stop_reason": "end_turn", + "stop_sequence": null + }, + "usage": { + "output_tokens": 0 + } + }); + format!("event: message_delta\ndata: {}\n\n", event) + } + + fn create_anthropic_message_stop(&self) -> String { + let event = serde_json::json!({ + "type": "message_stop" + }); + format!("event: message_stop\ndata: {}\n\n", event) + } + + // ======================================================================== + // OpenAI SSE 事件创建辅助方法 + // ======================================================================== + + fn get_created_timestamp(&self) -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs() + } + + fn create_openai_content_chunk(&self, content: &str, is_first: bool) -> String { + let chunk = serde_json::json!({ + "id": self.response_id, + "object": "chat.completion.chunk", + "created": self.get_created_timestamp(), + "model": self.model, + "choices": [{ + "index": 0, + "delta": { + "role": if is_first { Some("assistant") } else { None::<&str> }, + "content": content + }, + "finish_reason": null + }] + }); + format!("data: {}\n\n", chunk) + } + + fn create_openai_tool_call_chunk( + &self, + index: u32, + id: &str, + name: &str, + arguments: &str, + is_first: bool, + ) -> String { + let tool_call = if is_first { + serde_json::json!({ + "index": index, + "id": id, + "type": "function", + "function": { + "name": name, + "arguments": arguments + } + }) + } else { + serde_json::json!({ + "index": index, + "function": { + "arguments": arguments + } + }) + }; + + let chunk = serde_json::json!({ + "id": self.response_id, + "object": "chat.completion.chunk", + "created": self.get_created_timestamp(), + "model": self.model, + "choices": [{ + "index": 0, + "delta": { + "tool_calls": [tool_call] + }, + "finish_reason": null + }] + }); + format!("data: {}\n\n", chunk) + } + + fn create_openai_finish_chunk(&self, finish_reason: &str) -> String { + let chunk = serde_json::json!({ + "id": self.response_id, + "object": "chat.completion.chunk", + "created": self.get_created_timestamp(), + "model": self.model, + "choices": [{ + "index": 0, + "delta": {}, + "finish_reason": finish_reason + }] + }); + format!("data: {}\n\n", chunk) + } +} + +// ============================================================================ +// 辅助函数 +// ============================================================================ + +/// 从 SSE 事件列表中提取所有文本内容 +pub fn extract_content_from_sse(events: &[String], format: StreamFormat) -> String { + let mut content = String::new(); + + for event in events { + match format { + StreamFormat::OpenAiSse => { + for line in event.lines() { + if let Some(json_str) = line.strip_prefix("data: ") { + if json_str == "[DONE]" { + continue; + } + if let Ok(chunk) = serde_json::from_str::(json_str) { + if let Some(choices) = chunk.get("choices").and_then(|c| c.as_array()) { + for choice in choices { + if let Some(delta) = choice.get("delta") { + if let Some(text) = + delta.get("content").and_then(|c| c.as_str()) + { + content.push_str(text); + } + } + } + } + } + } + } + } + StreamFormat::AnthropicSse => { + for line in event.lines() { + if let Some(json_str) = line.strip_prefix("data: ") { + if let Ok(evt) = serde_json::from_str::(json_str) { + if evt.get("type").and_then(|t| t.as_str()) + == Some("content_block_delta") + { + if let Some(delta) = evt.get("delta") { + if let Some(text) = delta.get("text").and_then(|t| t.as_str()) { + content.push_str(text); + } + } + } + } + } + } + } + StreamFormat::AwsEventStream => { + // AWS Event Stream 不是 SSE 格式 + } + } + } + + content +} + +/// 从 SSE 事件列表中提取所有工具调用 +pub fn extract_tool_calls_from_sse( + events: &[String], + format: StreamFormat, +) -> Vec<(String, String, String)> { + let mut tool_calls: HashMap = HashMap::new(); + + for event in events { + match format { + StreamFormat::OpenAiSse => { + for line in event.lines() { + if let Some(json_str) = line.strip_prefix("data: ") { + if json_str == "[DONE]" { + continue; + } + if let Ok(chunk) = serde_json::from_str::(json_str) { + if let Some(choices) = chunk.get("choices").and_then(|c| c.as_array()) { + for choice in choices { + if let Some(delta) = choice.get("delta") { + if let Some(tcs) = + delta.get("tool_calls").and_then(|t| t.as_array()) + { + for tc in tcs { + let id = tc + .get("id") + .and_then(|i| i.as_str()) + .unwrap_or(""); + let name = tc + .get("function") + .and_then(|f| f.get("name")) + .and_then(|n| n.as_str()) + .unwrap_or(""); + let args = tc + .get("function") + .and_then(|f| f.get("arguments")) + .and_then(|a| a.as_str()) + .unwrap_or(""); + + if !id.is_empty() { + tool_calls + .entry(id.to_string()) + .or_insert((String::new(), String::new())) + .0 = name.to_string(); + } + if !args.is_empty() { + if let Some(entry) = + tool_calls.values_mut().last() + { + entry.1.push_str(args); + } + } + } + } + } + } + } + } + } + } + } + StreamFormat::AnthropicSse | StreamFormat::AwsEventStream => { + // 简化处理 + } + } + } + + tool_calls + .into_iter() + .map(|(id, (name, args))| (id, name, args)) + .collect() +} + +// ============================================================================ +// 测试模块 +// ============================================================================ + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_converter_new() { + let converter = StreamConverter::new(StreamFormat::AwsEventStream, StreamFormat::OpenAiSse); + assert_eq!(converter.state(), &ConverterState::Idle); + assert!(converter.aws_parser.is_some()); + } + + #[test] + fn test_converter_with_model() { + let converter = StreamConverter::with_model( + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "claude-3-opus", + ); + assert_eq!(converter.model, "claude-3-opus"); + } + + #[test] + fn test_converter_reset() { + let mut converter = + StreamConverter::new(StreamFormat::AwsEventStream, StreamFormat::OpenAiSse); + converter.convert(b"{\"content\":\"test\"}"); + assert_eq!(converter.state(), &ConverterState::Converting); + + converter.reset(); + assert_eq!(converter.state(), &ConverterState::Idle); + assert!(converter.accumulated_content.is_empty()); + } + + #[test] + fn test_aws_to_openai_content() { + let mut converter = StreamConverter::with_model( + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + let events = converter.convert(b"{\"content\":\"Hello, world!\"}"); + assert!(!events.is_empty()); + + let content = extract_content_from_sse(&events, StreamFormat::OpenAiSse); + assert_eq!(content, "Hello, world!"); + assert_eq!(converter.accumulated_content(), "Hello, world!"); + } + + #[test] + fn test_aws_to_openai_multiple_content() { + let mut converter = StreamConverter::with_model( + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + let events1 = converter.convert(b"{\"content\":\"Hello\"}"); + let events2 = converter.convert(b"{\"content\":\", world!\"}"); + + let all_events: Vec<_> = events1.into_iter().chain(events2).collect(); + let content = extract_content_from_sse(&all_events, StreamFormat::OpenAiSse); + assert_eq!(content, "Hello, world!"); + } + + #[test] + fn test_aws_to_openai_tool_call() { + let mut converter = StreamConverter::with_model( + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + // 工具调用开始 + let events1 = converter.convert(b"{\"toolUseId\":\"tool_1\",\"name\":\"read_file\"}"); + assert!(!events1.is_empty()); + + // 工具调用输入 + let events2 = converter + .convert(b"{\"toolUseId\":\"tool_1\",\"input\":\"{\\\"path\\\":\\\"/tmp\\\"}\"}"); + assert!(!events2.is_empty()); + + // 工具调用结束 + let events3 = converter.convert(b"{\"toolUseId\":\"tool_1\",\"stop\":true}"); + + let all_events: Vec<_> = events1.into_iter().chain(events2).chain(events3).collect(); + + // 验证工具调用存在 + let has_tool_call = all_events.iter().any(|e| e.contains("tool_calls")); + assert!(has_tool_call); + } + + #[test] + fn test_aws_to_anthropic_content() { + let mut converter = StreamConverter::with_model( + StreamFormat::AwsEventStream, + StreamFormat::AnthropicSse, + "test-model", + ); + + let events = converter.convert(b"{\"content\":\"Hello!\"}"); + + // 应该有 message_start, content_block_start, content_block_delta + assert!(events.iter().any(|e| e.contains("message_start"))); + assert!(events.iter().any(|e| e.contains("content_block_start"))); + assert!(events.iter().any(|e| e.contains("content_block_delta"))); + + let content = extract_content_from_sse(&events, StreamFormat::AnthropicSse); + assert_eq!(content, "Hello!"); + } + + #[test] + fn test_aws_to_anthropic_tool_call() { + let mut converter = StreamConverter::with_model( + StreamFormat::AwsEventStream, + StreamFormat::AnthropicSse, + "test-model", + ); + + // 工具调用开始 + let events1 = converter.convert(b"{\"toolUseId\":\"tool_1\",\"name\":\"test_tool\"}"); + assert!(events1.iter().any(|e| e.contains("tool_use"))); + + // 工具调用输入 + let events2 = converter + .convert(b"{\"toolUseId\":\"tool_1\",\"input\":\"{\\\"key\\\":\\\"value\\\"}\"}"); + assert!(events2.iter().any(|e| e.contains("input_json_delta"))); + + // 工具调用结束 + let events3 = converter.convert(b"{\"toolUseId\":\"tool_1\",\"stop\":true}"); + assert!(events3.iter().any(|e| e.contains("content_block_stop"))); + } + + #[test] + fn test_converter_finish() { + let mut converter = StreamConverter::with_model( + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + converter.convert(b"{\"content\":\"test\"}"); + let finish_events = converter.finish(); + + // 应该有结束事件 + assert!(finish_events.iter().any(|e| e.contains("finish_reason"))); + assert!(finish_events.iter().any(|e| e.contains("[DONE]"))); + assert_eq!(converter.state(), &ConverterState::Completed); + } + + #[test] + fn test_partial_json_accumulator() { + let mut acc = PartialJsonAccumulator::new(); + + // 追加部分 JSON + assert!(!acc.append("{\"key\":")); + assert!(!acc.is_complete()); + + assert!(acc.append("\"value\"}")); + assert!(acc.is_complete()); + assert_eq!(acc.get_json(), "{\"key\":\"value\"}"); + } + + #[test] + fn test_partial_json_accumulator_nested() { + let mut acc = PartialJsonAccumulator::new(); + + assert!(!acc.append("{\"outer\":{\"inner\":")); + assert!(!acc.append("\"value\"")); + assert!(acc.append("}}")); + assert!(acc.is_complete()); + } + + #[test] + fn test_partial_json_accumulator_with_string_braces() { + let mut acc = PartialJsonAccumulator::new(); + + // JSON 字符串中包含括号 + assert!(acc.append("{\"text\":\"hello {world}\"}")); + assert!(acc.is_complete()); + } + + #[test] + fn test_partial_json_accumulator_reset() { + let mut acc = PartialJsonAccumulator::new(); + acc.append("{\"key\":\"value\"}"); + + acc.reset(); + assert!(acc.is_empty()); + assert!(!acc.is_complete()); + } + + #[test] + fn test_extract_content_from_openai_sse() { + let events = vec![ + "data: {\"choices\":[{\"delta\":{\"content\":\"Hello\"}}]}\n\n".to_string(), + "data: {\"choices\":[{\"delta\":{\"content\":\", world!\"}}]}\n\n".to_string(), + "data: [DONE]\n\n".to_string(), + ]; + + let content = extract_content_from_sse(&events, StreamFormat::OpenAiSse); + assert_eq!(content, "Hello, world!"); + } + + #[test] + fn test_extract_content_from_anthropic_sse() { + let events = vec![ + "event: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"text_delta\",\"text\":\"Hello\"}}\n\n".to_string(), + "event: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"text_delta\",\"text\":\", world!\"}}\n\n".to_string(), + ]; + + let content = extract_content_from_sse(&events, StreamFormat::AnthropicSse); + assert_eq!(content, "Hello, world!"); + } + + #[test] + fn test_incremental_conversion() { + let mut converter = StreamConverter::with_model( + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + // 分片发送 JSON + let events1 = converter.convert(b"{\"content\":"); + assert!(events1.is_empty()); // 不完整的 JSON 不产生事件 + + let events2 = converter.convert(b"\"test\"}"); + assert!(!events2.is_empty()); // 完整的 JSON 产生事件 + + let content = extract_content_from_sse(&events2, StreamFormat::OpenAiSse); + assert_eq!(content, "test"); + } + + #[test] + fn test_converter_state_transitions() { + let mut converter = + StreamConverter::new(StreamFormat::AwsEventStream, StreamFormat::OpenAiSse); + + assert_eq!(converter.state(), &ConverterState::Idle); + + converter.convert(b"{\"content\":\"test\"}"); + assert_eq!(converter.state(), &ConverterState::Converting); + + converter.finish(); + assert_eq!(converter.state(), &ConverterState::Completed); + + converter.reset(); + assert_eq!(converter.state(), &ConverterState::Idle); + } +} + +// ============================================================================ +// 属性测试(Property-Based Testing) +// ============================================================================ + +#[cfg(test)] +mod property_tests { + use super::*; + use crate::streaming::aws_parser::{ + extract_content, extract_tool_calls, serialize_event, AwsEvent, + }; + use proptest::prelude::*; + + // ======================================================================== + // 生成器(Generators) + // ======================================================================== + + /// 生成有效的内容文本 + fn arb_content_text() -> impl Strategy { + prop::string::string_regex("[a-zA-Z0-9\\u4e00-\\u9fff .,!?\\-_]{1,100}") + .unwrap() + .prop_filter("非空字符串", |s| !s.is_empty()) + } + + /// 生成有效的工具 ID + fn arb_tool_id() -> impl Strategy { + prop::string::string_regex("tool_[a-zA-Z0-9]{4,12}").unwrap() + } + + /// 生成有效的工具名称 + fn arb_tool_name() -> impl Strategy { + prop::string::string_regex("[a-z_][a-z0-9_]{2,20}").unwrap() + } + + /// 生成有效的工具输入(简单 JSON) + fn arb_tool_input() -> impl Strategy { + prop::string::string_regex(r#"\{"[a-z]+":"[a-zA-Z0-9]+"\}"#).unwrap() + } + + /// 生成 Content 事件 + fn arb_content_event() -> impl Strategy { + arb_content_text().prop_map(|text| AwsEvent::Content { text }) + } + + /// 生成内容事件序列 + fn arb_content_sequence() -> impl Strategy> { + prop::collection::vec(arb_content_event(), 1..10) + } + + // ======================================================================== + // Property 2: 流式格式转换内容保留 + // + // *对于任意*流式响应内容,从 AWS Event Stream 转换为 Anthropic SSE 或 + // OpenAI SSE 后,最终重建的内容应该与原始内容一致。 + // + // **验证: 需求 3.1, 3.2, 3.3** + // ======================================================================== + + proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// Property 2: AWS 到 OpenAI 转换保留内容 + /// + /// **Feature: true-streaming-support, Property 2: 流式格式转换内容保留** + /// **Validates: Requirements 3.2** + #[test] + fn prop_aws_to_openai_content_preservation(events in arb_content_sequence()) { + // 计算原始内容 + let original_content = extract_content(&events); + + // 创建转换器 + let mut converter = StreamConverter::with_model( + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + // 将所有事件序列化并转换 + let mut all_sse_events = Vec::new(); + for event in &events { + if let Some(json) = serialize_event(event) { + let sse_events = converter.convert(json.as_bytes()); + all_sse_events.extend(sse_events); + } + } + all_sse_events.extend(converter.finish()); + + // 从 SSE 事件中提取内容 + let converted_content = extract_content_from_sse(&all_sse_events, StreamFormat::OpenAiSse); + + // 验证内容一致 + prop_assert_eq!( + original_content, + converted_content, + "AWS 到 OpenAI 转换应该保留内容" + ); + } + + /// Property 2: AWS 到 Anthropic 转换保留内容 + /// + /// **Feature: true-streaming-support, Property 2: 流式格式转换内容保留** + /// **Validates: Requirements 3.1** + #[test] + fn prop_aws_to_anthropic_content_preservation(events in arb_content_sequence()) { + // 计算原始内容 + let original_content = extract_content(&events); + + // 创建转换器 + let mut converter = StreamConverter::with_model( + StreamFormat::AwsEventStream, + StreamFormat::AnthropicSse, + "test-model", + ); + + // 将所有事件序列化并转换 + let mut all_sse_events = Vec::new(); + for event in &events { + if let Some(json) = serialize_event(event) { + let sse_events = converter.convert(json.as_bytes()); + all_sse_events.extend(sse_events); + } + } + all_sse_events.extend(converter.finish()); + + // 从 SSE 事件中提取内容 + let converted_content = extract_content_from_sse(&all_sse_events, StreamFormat::AnthropicSse); + + // 验证内容一致 + prop_assert_eq!( + original_content, + converted_content, + "AWS 到 Anthropic 转换应该保留内容" + ); + } + + /// Property 2: 转换器累积内容与原始内容一致 + /// + /// **Feature: true-streaming-support, Property 2: 流式格式转换内容保留** + /// **Validates: Requirements 3.1, 3.2** + #[test] + fn prop_converter_accumulated_content_matches( + events in arb_content_sequence(), + target in prop_oneof![Just(StreamFormat::OpenAiSse), Just(StreamFormat::AnthropicSse)] + ) { + // 计算原始内容 + let original_content = extract_content(&events); + + // 创建转换器 + let mut converter = StreamConverter::with_model( + StreamFormat::AwsEventStream, + target, + "test-model", + ); + + // 将所有事件序列化并转换 + for event in &events { + if let Some(json) = serialize_event(event) { + converter.convert(json.as_bytes()); + } + } + converter.finish(); + + // 验证累积内容一致 + prop_assert_eq!( + original_content, + converter.accumulated_content(), + "转换器累积的内容应该与原始内容一致" + ); + } + } + + // ======================================================================== + // Property 4: 部分 JSON 处理正确性 + // + // *对于任意*有效的 JSON 字符串,将其分割成任意多个部分后, + // 解析器应该能正确累积并最终产生完整的 JSON。 + // + // **验证: 需求 3.5** + // ======================================================================== + + proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// Property 4: 部分 JSON 累积正确性 + /// + /// **Feature: true-streaming-support, Property 4: 部分 JSON 处理正确性** + /// **Validates: Requirements 3.5** + #[test] + fn prop_partial_json_accumulation( + key in "[a-z]{1,10}", + value in "[a-zA-Z0-9]{1,20}", + split_points in prop::collection::vec(1usize..50, 1..5) + ) { + let json = format!("{{\"{}\":\"{}\"}}", key, value); + let bytes = json.as_bytes(); + + let mut acc = PartialJsonAccumulator::new(); + + // 根据分割点将 JSON 分成多个部分 + let mut last_pos = 0; + for &split_point in &split_points { + let split_pos = split_point.min(bytes.len()); + if split_pos > last_pos && split_pos < bytes.len() { + let part = std::str::from_utf8(&bytes[last_pos..split_pos]).unwrap_or(""); + acc.append(part); + last_pos = split_pos; + } + } + + // 追加剩余部分 + if last_pos < bytes.len() { + let remaining = std::str::from_utf8(&bytes[last_pos..]).unwrap_or(""); + acc.append(remaining); + } + + // 验证累积的 JSON 与原始 JSON 一致 + prop_assert_eq!( + acc.get_json(), + &json, + "累积的 JSON 应该与原始 JSON 一致" + ); + + // 验证 JSON 完整性 + prop_assert!( + acc.is_complete(), + "累积完成后 JSON 应该是完整的" + ); + } + + /// Property 4: 嵌套 JSON 部分累积 + /// + /// **Feature: true-streaming-support, Property 4: 部分 JSON 处理正确性** + /// **Validates: Requirements 3.5** + #[test] + fn prop_nested_json_partial_accumulation( + outer_key in "[a-z]{1,5}", + inner_key in "[a-z]{1,5}", + value in "[a-zA-Z0-9]{1,10}", + chunk_size in 1usize..10 + ) { + let json = format!("{{\"{}\":{{\"{}\":\"{}\"}}}}", outer_key, inner_key, value); + let bytes = json.as_bytes(); + + let mut acc = PartialJsonAccumulator::new(); + + // 按固定大小分块 + for chunk in bytes.chunks(chunk_size) { + let part = std::str::from_utf8(chunk).unwrap_or(""); + acc.append(part); + } + + // 验证累积的 JSON 与原始 JSON 一致 + prop_assert_eq!( + acc.get_json(), + &json, + "嵌套 JSON 累积应该与原始一致" + ); + + prop_assert!( + acc.is_complete(), + "嵌套 JSON 累积完成后应该是完整的" + ); + } + + /// Property 4: 包含字符串括号的 JSON 部分累积 + /// + /// **Feature: true-streaming-support, Property 4: 部分 JSON 处理正确性** + /// **Validates: Requirements 3.5** + #[test] + fn prop_json_with_braces_in_string( + key in "[a-z]{1,5}", + prefix in "[a-zA-Z0-9]{0,5}", + suffix in "[a-zA-Z0-9]{0,5}", + chunk_size in 1usize..15 + ) { + // 创建包含括号的字符串值 + let json = format!("{{\"{}\":\"{}{{}}[]{}\"}}", key, prefix, suffix); + let bytes = json.as_bytes(); + + let mut acc = PartialJsonAccumulator::new(); + + // 按固定大小分块 + for chunk in bytes.chunks(chunk_size) { + let part = std::str::from_utf8(chunk).unwrap_or(""); + acc.append(part); + } + + // 验证累积的 JSON 与原始 JSON 一致 + prop_assert_eq!( + acc.get_json(), + &json, + "包含括号的字符串 JSON 累积应该与原始一致" + ); + + prop_assert!( + acc.is_complete(), + "包含括号的字符串 JSON 累积完成后应该是完整的" + ); + } + + /// Property 4: 重置后重新累积 + /// + /// **Feature: true-streaming-support, Property 4: 部分 JSON 处理正确性** + /// **Validates: Requirements 3.5** + #[test] + fn prop_reset_and_reaccumulate( + key1 in "[a-z]{1,5}", + value1 in "[a-zA-Z0-9]{1,10}", + key2 in "[a-z]{1,5}", + value2 in "[a-zA-Z0-9]{1,10}" + ) { + let json1 = format!("{{\"{}\":\"{}\"}}", key1, value1); + let json2 = format!("{{\"{}\":\"{}\"}}", key2, value2); + + let mut acc = PartialJsonAccumulator::new(); + + // 累积第一个 JSON + acc.append(&json1); + prop_assert_eq!(acc.get_json(), &json1); + prop_assert!(acc.is_complete()); + + // 重置 + acc.reset(); + prop_assert!(acc.is_empty()); + prop_assert!(!acc.is_complete()); + + // 累积第二个 JSON + acc.append(&json2); + prop_assert_eq!(acc.get_json(), &json2); + prop_assert!(acc.is_complete()); + } + } +} diff --git a/src-tauri/src/streaming/error.rs b/src-tauri/src/streaming/error.rs new file mode 100644 index 000000000..fad3bf3f2 --- /dev/null +++ b/src-tauri/src/streaming/error.rs @@ -0,0 +1,281 @@ +//! 流式传输错误类型 +//! +//! 定义流式传输过程中可能发生的各种错误类型。 +//! +//! # 需求覆盖 +//! +//! - 需求 6.1: 网络错误处理 +//! - 需求 6.2: 超时错误处理 +//! - 需求 6.3: Provider 错误转发 + +use serde::{Deserialize, Serialize}; +use std::fmt; + +/// 流式传输错误类型 +/// +/// 涵盖流式传输过程中可能发生的所有错误情况。 +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +#[serde(tag = "type", content = "details")] +pub enum StreamError { + /// 网络错误 + /// + /// 当网络连接失败、DNS 解析失败或连接被重置时发生。 + /// 对应需求 6.1 + Network(String), + + /// 超时错误 + /// + /// 当流式响应超过配置的超时时间时发生。 + /// 对应需求 6.2 + Timeout, + + /// 解析错误 + /// + /// 当无法解析流式数据(如无效的 AWS Event Stream 或 SSE 格式)时发生。 + ParseError(String), + + /// Provider 错误 + /// + /// 当上游 Provider 返回错误响应时发生。 + /// 对应需求 6.3 + ProviderError { + /// HTTP 状态码 + status: u16, + /// 错误消息 + message: String, + }, + + /// 客户端断开连接 + /// + /// 当客户端在流式传输过程中断开连接时发生。 + ClientDisconnected, + + /// 缓冲区溢出 + /// + /// 当流式数据超过配置的缓冲区大小时发生。 + BufferOverflow, + + /// 内部错误 + /// + /// 其他内部错误。 + Internal(String), +} + +impl fmt::Display for StreamError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + StreamError::Network(msg) => write!(f, "网络错误: {}", msg), + StreamError::Timeout => write!(f, "流式响应超时"), + StreamError::ParseError(msg) => write!(f, "解析错误: {}", msg), + StreamError::ProviderError { status, message } => { + write!(f, "Provider 错误 ({}): {}", status, message) + } + StreamError::ClientDisconnected => write!(f, "客户端已断开连接"), + StreamError::BufferOverflow => write!(f, "缓冲区溢出"), + StreamError::Internal(msg) => write!(f, "内部错误: {}", msg), + } + } +} + +impl std::error::Error for StreamError {} + +// ============================================================================ +// From trait 实现 - 用于错误转换 +// ============================================================================ + +impl From for StreamError { + fn from(err: std::io::Error) -> Self { + StreamError::Network(err.to_string()) + } +} + +impl From for StreamError { + fn from(err: reqwest::Error) -> Self { + if err.is_timeout() { + StreamError::Timeout + } else if err.is_connect() { + StreamError::Network(format!("连接失败: {}", err)) + } else if err.is_request() { + StreamError::Network(format!("请求错误: {}", err)) + } else { + StreamError::Network(err.to_string()) + } + } +} + +impl From for StreamError { + fn from(err: serde_json::Error) -> Self { + StreamError::ParseError(err.to_string()) + } +} + +impl From for StreamError { + fn from(msg: String) -> Self { + StreamError::Internal(msg) + } +} + +impl From<&str> for StreamError { + fn from(msg: &str) -> Self { + StreamError::Internal(msg.to_string()) + } +} + +// ============================================================================ +// 辅助方法 +// ============================================================================ + +impl StreamError { + /// 创建网络错误 + pub fn network(msg: impl Into) -> Self { + StreamError::Network(msg.into()) + } + + /// 创建解析错误 + pub fn parse_error(msg: impl Into) -> Self { + StreamError::ParseError(msg.into()) + } + + /// 创建 Provider 错误 + pub fn provider_error(status: u16, message: impl Into) -> Self { + StreamError::ProviderError { + status, + message: message.into(), + } + } + + /// 创建内部错误 + pub fn internal(msg: impl Into) -> Self { + StreamError::Internal(msg.into()) + } + + /// 判断错误是否可重试 + /// + /// 网络错误、超时和某些 Provider 错误(如 429、5xx)可以重试。 + pub fn is_retryable(&self) -> bool { + match self { + StreamError::Network(_) => true, + StreamError::Timeout => true, + StreamError::ProviderError { status, .. } => *status == 429 || *status >= 500, + StreamError::ParseError(_) => false, + StreamError::ClientDisconnected => false, + StreamError::BufferOverflow => false, + StreamError::Internal(_) => false, + } + } + + /// 判断是否为客户端错误 + pub fn is_client_error(&self) -> bool { + matches!(self, StreamError::ClientDisconnected) + } + + /// 获取 HTTP 状态码(如果适用) + pub fn status_code(&self) -> Option { + match self { + StreamError::ProviderError { status, .. } => Some(*status), + StreamError::Timeout => Some(504), // Gateway Timeout + StreamError::Network(_) => Some(502), // Bad Gateway + _ => None, + } + } + + /// 转换为 SSE 错误事件格式 + pub fn to_sse_error(&self) -> String { + let error_json = serde_json::json!({ + "error": { + "type": self.error_type_string(), + "message": self.to_string(), + } + }); + format!("event: error\ndata: {}\n\n", error_json) + } + + /// 获取错误类型字符串 + fn error_type_string(&self) -> &'static str { + match self { + StreamError::Network(_) => "network_error", + StreamError::Timeout => "timeout", + StreamError::ParseError(_) => "parse_error", + StreamError::ProviderError { .. } => "provider_error", + StreamError::ClientDisconnected => "client_disconnected", + StreamError::BufferOverflow => "buffer_overflow", + StreamError::Internal(_) => "internal_error", + } + } +} + +// ============================================================================ +// 测试模块 +// ============================================================================ + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_stream_error_display() { + let err = StreamError::Network("connection refused".to_string()); + assert_eq!(err.to_string(), "网络错误: connection refused"); + + let err = StreamError::Timeout; + assert_eq!(err.to_string(), "流式响应超时"); + + let err = StreamError::provider_error(429, "rate limited"); + assert_eq!(err.to_string(), "Provider 错误 (429): rate limited"); + } + + #[test] + fn test_stream_error_is_retryable() { + assert!(StreamError::Network("test".to_string()).is_retryable()); + assert!(StreamError::Timeout.is_retryable()); + assert!(StreamError::provider_error(429, "rate limited").is_retryable()); + assert!(StreamError::provider_error(500, "server error").is_retryable()); + assert!(!StreamError::provider_error(400, "bad request").is_retryable()); + assert!(!StreamError::ParseError("invalid json".to_string()).is_retryable()); + assert!(!StreamError::ClientDisconnected.is_retryable()); + } + + #[test] + fn test_stream_error_status_code() { + assert_eq!(StreamError::Timeout.status_code(), Some(504)); + assert_eq!( + StreamError::Network("test".to_string()).status_code(), + Some(502) + ); + assert_eq!( + StreamError::provider_error(429, "test").status_code(), + Some(429) + ); + assert_eq!(StreamError::ClientDisconnected.status_code(), None); + } + + #[test] + fn test_stream_error_from_io_error() { + let io_err = std::io::Error::new(std::io::ErrorKind::ConnectionRefused, "refused"); + let stream_err: StreamError = io_err.into(); + assert!(matches!(stream_err, StreamError::Network(_))); + } + + #[test] + fn test_stream_error_from_serde_json_error() { + let json_err = serde_json::from_str::("invalid").unwrap_err(); + let stream_err: StreamError = json_err.into(); + assert!(matches!(stream_err, StreamError::ParseError(_))); + } + + #[test] + fn test_stream_error_serialization() { + let err = StreamError::provider_error(500, "internal server error"); + let json = serde_json::to_string(&err).unwrap(); + let deserialized: StreamError = serde_json::from_str(&json).unwrap(); + assert_eq!(err, deserialized); + } + + #[test] + fn test_stream_error_to_sse_error() { + let err = StreamError::Timeout; + let sse = err.to_sse_error(); + assert!(sse.starts_with("event: error\n")); + assert!(sse.contains("timeout")); + } +} diff --git a/src-tauri/src/streaming/manager.rs b/src-tauri/src/streaming/manager.rs new file mode 100644 index 000000000..7a246ad09 --- /dev/null +++ b/src-tauri/src/streaming/manager.rs @@ -0,0 +1,1994 @@ +//! 流式管理器 +//! +//! 管理流式请求的生命周期,包括格式转换、Flow Monitor 集成、超时处理和错误处理。 +//! +//! # 需求覆盖 +//! +//! - 需求 4.1: 记录 TTFB(首字节时间) +//! - 需求 4.2: 调用 process_chunk 更新流重建器 +//! - 需求 4.3: 发出带有内容增量的 FlowUpdated 事件 +//! - 需求 5.1: 在收到 chunk 后立即转发给客户端 +//! - 需求 5.2: 保持低延迟 +//! - 需求 6.1: 网络错误处理 +//! - 需求 6.2: 超时错误处理 +//! - 需求 6.3: Provider 错误转发 +//! - 需求 6.5: 可配置的流式响应超时 + +use crate::streaming::converter::{StreamConverter, StreamFormat}; +use crate::streaming::error::StreamError; +use crate::streaming::metrics::StreamMetrics; +use crate::streaming::traits::StreamResponse; +use bytes::Bytes; +use futures::{Stream, StreamExt}; +use serde::{Deserialize, Serialize}; +use std::pin::Pin; +use std::task::{Context, Poll}; +use std::time::Duration; +use tokio::time::Instant; +use tracing::{debug, error}; + +// ============================================================================ +// 配置 +// ============================================================================ + +/// 流式配置 +/// +/// 控制流式传输的行为参数。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct StreamConfig { + /// 缓冲区大小(字节) + /// + /// 用于限制内存使用,防止内存耗尽。 + /// 对应需求 7.1 + #[serde(default = "default_buffer_size")] + pub buffer_size: usize, + + /// 超时时间(毫秒) + /// + /// 流式响应的最大等待时间。 + /// 对应需求 6.2, 6.5 + #[serde(default = "default_timeout_ms")] + pub timeout_ms: u64, + + /// 事件节流间隔(毫秒) + /// + /// 控制 FlowUpdated 事件的发送频率,避免过多更新。 + /// 对应需求 4.6 + #[serde(default = "default_throttle_ms")] + pub throttle_ms: u64, + + /// chunk 超时时间(毫秒) + /// + /// 两个 chunk 之间的最大等待时间。 + #[serde(default = "default_chunk_timeout_ms")] + pub chunk_timeout_ms: u64, +} + +fn default_buffer_size() -> usize { + 1024 * 1024 // 1MB +} + +fn default_timeout_ms() -> u64 { + 300_000 // 5 分钟 +} + +fn default_throttle_ms() -> u64 { + 100 // 100ms +} + +fn default_chunk_timeout_ms() -> u64 { + 30_000 // 30 秒 +} + +impl Default for StreamConfig { + fn default() -> Self { + Self { + buffer_size: default_buffer_size(), + timeout_ms: default_timeout_ms(), + throttle_ms: default_throttle_ms(), + chunk_timeout_ms: default_chunk_timeout_ms(), + } + } +} + +impl StreamConfig { + /// 创建新的配置 + pub fn new() -> Self { + Self::default() + } + + /// 设置缓冲区大小 + pub fn with_buffer_size(mut self, size: usize) -> Self { + self.buffer_size = size; + self + } + + /// 设置超时时间 + pub fn with_timeout_ms(mut self, timeout_ms: u64) -> Self { + self.timeout_ms = timeout_ms; + self + } + + /// 设置事件节流间隔 + pub fn with_throttle_ms(mut self, throttle_ms: u64) -> Self { + self.throttle_ms = throttle_ms; + self + } + + /// 设置 chunk 超时时间 + pub fn with_chunk_timeout_ms(mut self, chunk_timeout_ms: u64) -> Self { + self.chunk_timeout_ms = chunk_timeout_ms; + self + } + + /// 获取超时 Duration + pub fn timeout_duration(&self) -> Duration { + Duration::from_millis(self.timeout_ms) + } + + /// 获取 chunk 超时 Duration + pub fn chunk_timeout_duration(&self) -> Duration { + Duration::from_millis(self.chunk_timeout_ms) + } + + /// 获取节流 Duration + pub fn throttle_duration(&self) -> Duration { + Duration::from_millis(self.throttle_ms) + } +} + +// ============================================================================ +// 流式处理上下文 +// ============================================================================ + +/// 流式处理上下文 +/// +/// 包含处理单个流式请求所需的所有状态。 +#[derive(Debug)] +pub struct StreamContext { + /// Flow ID(用于 Flow Monitor 集成) + pub flow_id: Option, + /// 源格式 + pub source_format: StreamFormat, + /// 目标格式 + pub target_format: StreamFormat, + /// 模型名称 + pub model: String, + /// 指标 + pub metrics: StreamMetrics, + /// 开始时间 + pub start_time: Instant, +} + +impl StreamContext { + /// 创建新的上下文 + pub fn new( + flow_id: Option, + source_format: StreamFormat, + target_format: StreamFormat, + model: &str, + ) -> Self { + Self { + flow_id, + source_format, + target_format, + model: model.to_string(), + metrics: StreamMetrics::new(), + start_time: Instant::now(), + } + } +} + +// ============================================================================ +// 流式管理器 +// ============================================================================ + +/// 流式管理器 +/// +/// 管理流式请求的生命周期,包括: +/// - 格式转换(AWS Event Stream → Anthropic/OpenAI SSE) +/// - Flow Monitor 集成(process_chunk、FlowUpdated 事件) +/// - 超时处理 +/// - 错误处理 +pub struct StreamManager { + /// 配置 + config: StreamConfig, +} + +impl StreamManager { + /// 创建新的流式管理器 + pub fn new(config: StreamConfig) -> Self { + Self { config } + } + + /// 使用默认配置创建流式管理器 + pub fn with_default_config() -> Self { + Self::new(StreamConfig::default()) + } + + /// 获取配置 + pub fn config(&self) -> &StreamConfig { + &self.config + } + + /// 更新配置 + pub fn set_config(&mut self, config: StreamConfig) { + self.config = config; + } + + /// 处理流式请求 + /// + /// 将源流转换为目标格式的 SSE 事件流。 + /// + /// # 参数 + /// + /// * `context` - 流式处理上下文 + /// * `source_stream` - 源字节流 + /// + /// # 返回 + /// + /// 目标格式的 SSE 事件流 + /// + /// # 需求覆盖 + /// + /// - 需求 4.1: 记录 TTFB + /// - 需求 5.1: 立即转发 chunk + /// - 需求 6.2: 超时处理 + pub fn handle_stream( + &self, + context: StreamContext, + source_stream: StreamResponse, + ) -> ManagedStream { + ManagedStream::new(context, source_stream, self.config.clone()) + } + + /// 处理流式请求(带回调) + /// + /// 与 `handle_stream` 类似,但支持在处理每个 chunk 时调用回调函数。 + /// 用于 Flow Monitor 集成。 + /// + /// # 参数 + /// + /// * `context` - 流式处理上下文 + /// * `source_stream` - 源字节流 + /// * `on_chunk` - chunk 处理回调 + /// + /// # 返回 + /// + /// 目标格式的 SSE 事件流 + pub fn handle_stream_with_callback( + &self, + context: StreamContext, + source_stream: StreamResponse, + on_chunk: F, + ) -> ManagedStreamWithCallback + where + F: FnMut(&str, &StreamMetrics) + Send + 'static, + { + ManagedStreamWithCallback::new(context, source_stream, self.config.clone(), on_chunk) + } + + /// 处理流式请求(带超时) + /// + /// 为流添加超时保护。 + /// + /// # 参数 + /// + /// * `context` - 流式处理上下文 + /// * `source_stream` - 源字节流 + /// + /// # 返回 + /// + /// 带超时保护的 SSE 事件流 + /// + /// # 需求覆盖 + /// + /// - 需求 6.2: 超时错误处理 + /// - 需求 6.5: 可配置的流式响应超时 + pub fn handle_stream_with_timeout( + &self, + context: StreamContext, + source_stream: StreamResponse, + ) -> TimeoutStream { + let managed = self.handle_stream(context, source_stream); + with_timeout(managed, &self.config) + } +} + +// ============================================================================ +// Flow Monitor 集成辅助类型 +// ============================================================================ + +/// Flow Monitor 回调类型 +/// +/// 用于在处理流式 chunk 时通知 Flow Monitor。 +pub type FlowMonitorCallback = Box; + +/// 创建 Flow Monitor 回调 +/// +/// 创建一个回调函数,用于将流式事件发送到 Flow Monitor。 +/// +/// # 参数 +/// +/// * `flow_id` - Flow ID +/// * `sender` - 事件发送器 +/// +/// # 返回 +/// +/// Flow Monitor 回调函数 +pub fn create_flow_monitor_callback( + flow_id: String, + mut on_event: F, +) -> impl FnMut(&str, &StreamMetrics) + Send + 'static +where + F: FnMut(&str, &str, &StreamMetrics) + Send + 'static, +{ + move |event: &str, metrics: &StreamMetrics| { + on_event(&flow_id, event, metrics); + } +} + +/// 流式事件类型 +/// +/// 用于 Flow Monitor 集成的事件类型。 +#[derive(Debug, Clone)] +pub enum StreamEvent { + /// 流开始 + Started { flow_id: String }, + /// 收到 chunk + Chunk { + flow_id: String, + content_delta: Option, + metrics: StreamMetrics, + }, + /// 流完成 + Completed { + flow_id: String, + metrics: StreamMetrics, + }, + /// 流错误 + Error { + flow_id: String, + error: StreamError, + metrics: StreamMetrics, + }, +} + +impl StreamEvent { + /// 获取 Flow ID + pub fn flow_id(&self) -> &str { + match self { + StreamEvent::Started { flow_id } => flow_id, + StreamEvent::Chunk { flow_id, .. } => flow_id, + StreamEvent::Completed { flow_id, .. } => flow_id, + StreamEvent::Error { flow_id, .. } => flow_id, + } + } +} + +impl Default for StreamManager { + fn default() -> Self { + Self::with_default_config() + } +} + +// ============================================================================ +// 托管流 +// ============================================================================ + +/// 托管流 +/// +/// 封装源流,提供格式转换、超时处理和指标收集。 +/// +/// # 有界缓冲区(需求 7.1) +/// +/// 托管流使用有界缓冲区来防止内存耗尽。当累积的数据超过配置的 +/// `buffer_size` 时,会返回 `BufferOverflow` 错误。 +pub struct ManagedStream { + /// 上下文 + context: StreamContext, + /// 源流 + source_stream: StreamResponse, + /// 转换器 + converter: StreamConverter, + /// 配置 + config: StreamConfig, + /// 待发送的事件缓冲区 + pending_events: Vec, + /// 是否已完成 + finished: bool, + /// 是否已记录首个 chunk + first_chunk_recorded: bool, + /// 总接收字节数 + total_bytes: usize, + /// 当前缓冲区使用量(用于有界缓冲区检查) + /// 对应需求 7.1 + current_buffer_usage: usize, +} + +impl ManagedStream { + /// 创建新的托管流 + pub fn new( + context: StreamContext, + source_stream: StreamResponse, + config: StreamConfig, + ) -> Self { + let converter = StreamConverter::with_model( + context.source_format, + context.target_format, + &context.model, + ); + + Self { + context, + source_stream, + converter, + config, + pending_events: Vec::new(), + finished: false, + first_chunk_recorded: false, + total_bytes: 0, + current_buffer_usage: 0, + } + } + + /// 获取指标 + pub fn metrics(&self) -> &StreamMetrics { + &self.context.metrics + } + + /// 获取上下文 + pub fn context(&self) -> &StreamContext { + &self.context + } + + /// 处理接收到的字节 + /// + /// # 有界缓冲区检查(需求 7.1) + /// + /// 如果累积的数据超过配置的 `buffer_size`,返回空事件列表并设置错误状态。 + fn process_bytes(&mut self, bytes: &Bytes) -> Result, StreamError> { + // 检查有界缓冲区限制(需求 7.1) + let new_usage = self.current_buffer_usage + bytes.len(); + if new_usage > self.config.buffer_size { + error!( + flow_id = ?self.context.flow_id, + current_usage = self.current_buffer_usage, + incoming_bytes = bytes.len(), + buffer_limit = self.config.buffer_size, + "缓冲区溢出" + ); + return Err(StreamError::BufferOverflow); + } + self.current_buffer_usage = new_usage; + + // 记录指标 + self.total_bytes += bytes.len(); + self.context.metrics.record_chunk(bytes.len()); + + // 记录首个 chunk + if !self.first_chunk_recorded { + self.first_chunk_recorded = true; + debug!( + flow_id = ?self.context.flow_id, + ttfb_ms = ?self.context.metrics.ttfb_ms, + "收到首个 chunk" + ); + } + + // 转换格式 + let events = self.converter.convert(bytes); + + // 转换后释放缓冲区使用量(事件已被处理) + // 只保留 pending_events 的大小 + self.current_buffer_usage = self.pending_events.iter().map(|e| e.len()).sum(); + + Ok(events) + } + + /// 完成流处理 + /// + /// 对应需求 7.2: 流结束时释放资源 + /// 对应需求 7.5: 记录流式指标 + fn finish_stream(&mut self) -> Vec { + self.finished = true; + self.context.metrics.finish(); + + // 记录详细指标(需求 7.5) + self.context + .metrics + .log_metrics(self.context.flow_id.as_deref()); + + debug!( + flow_id = ?self.context.flow_id, + metrics = ?self.context.metrics.summary(), + "流式传输完成" + ); + + let events = self.converter.finish(); + + // 清理缓冲区使用量(资源清理) + self.current_buffer_usage = 0; + + events + } + + /// 处理错误 + fn handle_error(&mut self, error: StreamError) -> String { + self.finished = true; + self.context.metrics.finish(); + self.context.metrics.record_parse_error(); + + error!( + flow_id = ?self.context.flow_id, + error = %error, + "流式传输错误" + ); + + error.to_sse_error() + } + + /// 处理 Provider 错误 + /// + /// 将 Provider 返回的错误转换为 SSE 错误事件。 + /// 对应需求 6.3: Provider 错误转发 + pub fn handle_provider_error(&mut self, status: u16, message: &str) -> String { + let error = StreamError::provider_error(status, message); + self.handle_error(error) + } + + /// 处理网络错误 + /// + /// 将网络错误转换为 SSE 错误事件。 + /// 对应需求 6.1: 网络错误处理 + pub fn handle_network_error(&mut self, message: &str) -> String { + let error = StreamError::network(message); + self.handle_error(error) + } + + /// 检查是否已完成 + pub fn is_finished(&self) -> bool { + self.finished + } + + /// 获取当前缓冲区使用量 + /// + /// 对应需求 7.1: 有界缓冲区 + pub fn buffer_usage(&self) -> usize { + self.current_buffer_usage + } + + /// 获取缓冲区限制 + pub fn buffer_limit(&self) -> usize { + self.config.buffer_size + } + + /// 清理资源 + /// + /// 对应需求 7.2: 流结束时释放资源 + /// + /// 此方法会清理所有内部状态,释放内存。 + /// 通常在流结束后自动调用,但也可以手动调用以提前释放资源。 + pub fn cleanup(&mut self) { + // 清理待发送事件缓冲区 + self.pending_events.clear(); + self.pending_events.shrink_to_fit(); + + // 重置缓冲区使用量 + self.current_buffer_usage = 0; + + // 重置转换器 + self.converter.reset(); + + // 标记为已完成 + self.finished = true; + + debug!( + flow_id = ?self.context.flow_id, + "流式资源已清理" + ); + } +} + +impl Stream for ManagedStream { + type Item = Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + // 如果有待发送的事件,先发送 + if !self.pending_events.is_empty() { + let event = self.pending_events.remove(0); + // 更新缓冲区使用量 + self.current_buffer_usage = self.pending_events.iter().map(|e| e.len()).sum(); + return Poll::Ready(Some(Ok(event))); + } + + // 如果已完成,返回 None + if self.finished { + return Poll::Ready(None); + } + + // 轮询源流 + match Pin::new(&mut self.source_stream).poll_next(cx) { + Poll::Ready(Some(Ok(bytes))) => { + // 处理字节(包含有界缓冲区检查) + match self.process_bytes(&bytes) { + Ok(mut events) => { + if events.is_empty() { + // 没有产生事件,继续轮询 + cx.waker().wake_by_ref(); + Poll::Pending + } else { + // 取出第一个事件返回 + let first = events.remove(0); + // 保存剩余的事件 + self.pending_events = events; + // 更新缓冲区使用量 + self.current_buffer_usage = + self.pending_events.iter().map(|e| e.len()).sum(); + Poll::Ready(Some(Ok(first))) + } + } + Err(error) => { + // 缓冲区溢出或其他错误 + let error_event = self.handle_error(error); + Poll::Ready(Some(Ok(error_event))) + } + } + } + Poll::Ready(Some(Err(error))) => { + // 处理错误 + let error_event = self.handle_error(error); + Poll::Ready(Some(Ok(error_event))) + } + Poll::Ready(None) => { + // 源流结束 + let mut finish_events = self.finish_stream(); + + if finish_events.is_empty() { + self.finished = true; + Poll::Ready(None) + } else { + // 取出第一个事件返回 + let first = finish_events.remove(0); + // 保存剩余的事件 + self.pending_events = finish_events; + Poll::Ready(Some(Ok(first))) + } + } + Poll::Pending => Poll::Pending, + } + } +} + +// ============================================================================ +// 带回调的托管流 +// ============================================================================ + +/// 带回调的托管流 +/// +/// 在处理每个 chunk 时调用回调函数,用于 Flow Monitor 集成。 +/// +/// # 事件节流(需求 4.6) +/// +/// 支持可配置的事件节流,避免过多的 FlowUpdated 事件。 +/// 节流间隔通过 `StreamConfig.throttle_ms` 配置。 +pub struct ManagedStreamWithCallback +where + F: FnMut(&str, &StreamMetrics) + Send + 'static, +{ + /// 内部托管流 + inner: ManagedStream, + /// chunk 处理回调 + on_chunk: F, + /// 上次回调时间(用于节流) + last_callback_time: Option, + /// 被节流的事件计数(用于指标) + throttled_event_count: u32, + /// 总回调次数 + callback_count: u32, +} + +// ManagedStreamWithCallback 可以安全地 Unpin,因为它不包含自引用 +impl Unpin for ManagedStreamWithCallback where F: FnMut(&str, &StreamMetrics) + Send + 'static {} + +impl ManagedStreamWithCallback +where + F: FnMut(&str, &StreamMetrics) + Send + 'static, +{ + /// 创建新的带回调托管流 + pub fn new( + context: StreamContext, + source_stream: StreamResponse, + config: StreamConfig, + on_chunk: F, + ) -> Self { + Self { + inner: ManagedStream::new(context, source_stream, config), + on_chunk, + last_callback_time: None, + throttled_event_count: 0, + callback_count: 0, + } + } + + /// 获取指标 + pub fn metrics(&self) -> &StreamMetrics { + self.inner.metrics() + } + + /// 获取上下文 + pub fn context(&self) -> &StreamContext { + self.inner.context() + } + + /// 获取被节流的事件计数 + /// + /// 对应需求 4.6: 事件节流 + pub fn throttled_event_count(&self) -> u32 { + self.throttled_event_count + } + + /// 获取总回调次数 + pub fn callback_count(&self) -> u32 { + self.callback_count + } + + /// 检查是否应该调用回调(节流) + /// + /// 对应需求 4.6: 可配置的事件节流 + fn should_call_callback(&self) -> bool { + match self.last_callback_time { + None => true, + Some(last_time) => last_time.elapsed() >= self.inner.config.throttle_duration(), + } + } +} + +impl Stream for ManagedStreamWithCallback +where + F: FnMut(&str, &StreamMetrics) + Send + 'static + Unpin, +{ + type Item = Result; + + fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + // 使用 get_mut 获取可变引用 + let this = self.get_mut(); + + // 轮询内部流 + match Pin::new(&mut this.inner).poll_next(cx) { + Poll::Ready(Some(Ok(event))) => { + // 检查是否应该调用回调(事件节流,需求 4.6) + if this.should_call_callback() { + let metrics = this.inner.context.metrics.clone(); + (this.on_chunk)(&event, &metrics); + this.last_callback_time = Some(Instant::now()); + this.callback_count += 1; + } else { + // 记录被节流的事件 + this.throttled_event_count += 1; + } + Poll::Ready(Some(Ok(event))) + } + other => other, + } + } +} + +// ============================================================================ +// 辅助函数 +// ============================================================================ + +/// 创建带超时的流 +/// +/// 为流添加整体超时和 chunk 超时。 +/// +/// # 参数 +/// +/// * `stream` - 源流 +/// * `config` - 流式配置 +/// +/// # 返回 +/// +/// 带超时的流 +pub fn with_timeout(stream: S, config: &StreamConfig) -> TimeoutStream +where + S: Stream> + Unpin, +{ + TimeoutStream::new(stream, config.clone()) +} + +/// 带超时的流包装器 +pub struct TimeoutStream +where + S: Stream> + Unpin, +{ + inner: S, + config: StreamConfig, + start_time: Instant, + last_chunk_time: Option, + finished: bool, +} + +impl TimeoutStream +where + S: Stream> + Unpin, +{ + /// 创建新的超时流 + pub fn new(inner: S, config: StreamConfig) -> Self { + Self { + inner, + config, + start_time: Instant::now(), + last_chunk_time: None, + finished: false, + } + } + + /// 检查是否超时 + fn check_timeout(&self) -> Option { + // 检查总超时 + if self.start_time.elapsed() > self.config.timeout_duration() { + return Some(StreamError::Timeout); + } + + // 检查 chunk 超时 + if let Some(last_time) = self.last_chunk_time { + if last_time.elapsed() > self.config.chunk_timeout_duration() { + return Some(StreamError::Timeout); + } + } + + None + } +} + +impl Stream for TimeoutStream +where + S: Stream> + Unpin, +{ + type Item = Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + if self.finished { + return Poll::Ready(None); + } + + // 检查超时 + if let Some(error) = self.check_timeout() { + self.finished = true; + return Poll::Ready(Some(Err(error))); + } + + // 轮询内部流 + match Pin::new(&mut self.inner).poll_next(cx) { + Poll::Ready(Some(item)) => { + self.last_chunk_time = Some(Instant::now()); + Poll::Ready(Some(item)) + } + Poll::Ready(None) => { + self.finished = true; + Poll::Ready(None) + } + Poll::Pending => Poll::Pending, + } + } +} + +/// 从流中收集所有内容 +/// +/// 用于测试和调试。 +pub async fn collect_stream_content(mut stream: S) -> Result +where + S: Stream> + Unpin, +{ + let mut content = String::new(); + + while let Some(result) = stream.next().await { + match result { + Ok(event) => content.push_str(&event), + Err(e) => return Err(e), + } + } + + Ok(content) +} + +// ============================================================================ +// 测试模块 +// ============================================================================ + +#[cfg(test)] +mod tests { + use super::*; + use futures::stream; + use std::sync::atomic::{AtomicU32, Ordering}; + use std::sync::Arc; + + #[test] + fn test_stream_config_default() { + let config = StreamConfig::default(); + assert_eq!(config.buffer_size, 1024 * 1024); + assert_eq!(config.timeout_ms, 300_000); + assert_eq!(config.throttle_ms, 100); + assert_eq!(config.chunk_timeout_ms, 30_000); + } + + #[test] + fn test_stream_config_builder() { + let config = StreamConfig::new() + .with_buffer_size(2048) + .with_timeout_ms(60_000) + .with_throttle_ms(50) + .with_chunk_timeout_ms(10_000); + + assert_eq!(config.buffer_size, 2048); + assert_eq!(config.timeout_ms, 60_000); + assert_eq!(config.throttle_ms, 50); + assert_eq!(config.chunk_timeout_ms, 10_000); + } + + #[test] + fn test_stream_config_durations() { + let config = StreamConfig::new() + .with_timeout_ms(5000) + .with_chunk_timeout_ms(1000) + .with_throttle_ms(100); + + assert_eq!(config.timeout_duration(), Duration::from_millis(5000)); + assert_eq!(config.chunk_timeout_duration(), Duration::from_millis(1000)); + assert_eq!(config.throttle_duration(), Duration::from_millis(100)); + } + + #[test] + fn test_stream_context_new() { + let context = StreamContext::new( + Some("flow-123".to_string()), + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "claude-3-opus", + ); + + assert_eq!(context.flow_id, Some("flow-123".to_string())); + assert_eq!(context.source_format, StreamFormat::AwsEventStream); + assert_eq!(context.target_format, StreamFormat::OpenAiSse); + assert_eq!(context.model, "claude-3-opus"); + } + + #[test] + fn test_stream_manager_new() { + let config = StreamConfig::default(); + let manager = StreamManager::new(config.clone()); + + assert_eq!(manager.config().buffer_size, config.buffer_size); + assert_eq!(manager.config().timeout_ms, config.timeout_ms); + } + + #[test] + fn test_stream_manager_default() { + let manager = StreamManager::default(); + let default_config = StreamConfig::default(); + + assert_eq!(manager.config().buffer_size, default_config.buffer_size); + } + + #[tokio::test] + async fn test_managed_stream_empty() { + let context = StreamContext::new( + None, + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + let empty_stream: StreamResponse = Box::pin(stream::empty()); + let config = StreamConfig::default(); + + let mut managed = ManagedStream::new(context, empty_stream, config); + + // 空流应该产生结束事件 + let mut events = Vec::new(); + while let Some(result) = managed.next().await { + if let Ok(event) = result { + events.push(event); + } + } + + // 应该有结束事件 + assert!(events.iter().any(|e| e.contains("[DONE]"))); + } + + #[tokio::test] + async fn test_managed_stream_with_content() { + let context = StreamContext::new( + None, + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + // 创建包含内容的流 + let chunks = vec![ + Ok(Bytes::from("{\"content\":\"Hello\"}")), + Ok(Bytes::from("{\"content\":\", world!\"}")), + ]; + let source_stream: StreamResponse = Box::pin(stream::iter(chunks)); + let config = StreamConfig::default(); + + let mut managed = ManagedStream::new(context, source_stream, config); + + let mut events = Vec::new(); + while let Some(result) = managed.next().await { + if let Ok(event) = result { + events.push(event); + } + } + + // 应该有内容事件 + let content_events: Vec<_> = events.iter().filter(|e| e.contains("content")).collect(); + assert!(!content_events.is_empty()); + } + + #[tokio::test] + async fn test_managed_stream_error_handling() { + let context = StreamContext::new( + None, + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + // 创建包含错误的流 + let chunks: Vec> = vec![ + Ok(Bytes::from("{\"content\":\"Hello\"}")), + Err(StreamError::Network("connection reset".to_string())), + ]; + let source_stream: StreamResponse = Box::pin(stream::iter(chunks)); + let config = StreamConfig::default(); + + let mut managed = ManagedStream::new(context, source_stream, config); + + let mut events = Vec::new(); + while let Some(result) = managed.next().await { + if let Ok(event) = result { + events.push(event); + } + } + + // 应该有错误事件 + assert!(events.iter().any(|e| e.contains("error"))); + } + + #[tokio::test] + async fn test_managed_stream_metrics() { + let context = StreamContext::new( + None, + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + let chunks = vec![Ok(Bytes::from("{\"content\":\"test\"}"))]; + let source_stream: StreamResponse = Box::pin(stream::iter(chunks)); + let config = StreamConfig::default(); + + let mut managed = ManagedStream::new(context, source_stream, config); + + // 消费流 + while managed.next().await.is_some() {} + + // 检查指标 + let metrics = managed.metrics(); + assert!(metrics.chunk_count > 0); + assert!(metrics.total_bytes > 0); + } + + #[tokio::test] + async fn test_collect_stream_content() { + let events = vec![Ok("event1\n".to_string()), Ok("event2\n".to_string())]; + let stream = stream::iter(events); + + let content = collect_stream_content(stream).await.unwrap(); + assert_eq!(content, "event1\nevent2\n"); + } + + #[tokio::test] + async fn test_collect_stream_content_with_error() { + let events: Vec> = + vec![Ok("event1\n".to_string()), Err(StreamError::Timeout)]; + let stream = stream::iter(events); + + let result = collect_stream_content(stream).await; + assert!(result.is_err()); + assert!(matches!(result.unwrap_err(), StreamError::Timeout)); + } + + #[tokio::test] + async fn test_timeout_stream_no_timeout() { + // 创建一个快速完成的流 + let events: Vec> = + vec![Ok("event1\n".to_string()), Ok("event2\n".to_string())]; + let inner_stream = stream::iter(events); + + let config = StreamConfig::new() + .with_timeout_ms(10_000) + .with_chunk_timeout_ms(5_000); + + let mut timeout_stream = with_timeout(inner_stream, &config); + + let mut results = Vec::new(); + while let Some(result) = timeout_stream.next().await { + results.push(result); + } + + // 应该成功完成,没有超时 + assert_eq!(results.len(), 2); + assert!(results.iter().all(|r| r.is_ok())); + } + + #[test] + fn test_timeout_stream_check_timeout() { + let events: Vec> = vec![]; + let inner_stream = stream::iter(events); + + // 创建一个非常短的超时配置 + let config = StreamConfig::new() + .with_timeout_ms(1) // 1ms 超时 + .with_chunk_timeout_ms(1); + + let timeout_stream = TimeoutStream::new(inner_stream, config); + + // 等待一小段时间让超时发生 + std::thread::sleep(std::time::Duration::from_millis(10)); + + // 检查超时 + let timeout_error = timeout_stream.check_timeout(); + assert!(timeout_error.is_some()); + assert!(matches!(timeout_error.unwrap(), StreamError::Timeout)); + } + + #[test] + fn test_stream_config_timeout_duration() { + let config = StreamConfig::new() + .with_timeout_ms(5000) + .with_chunk_timeout_ms(1000); + + assert_eq!(config.timeout_duration(), Duration::from_millis(5000)); + assert_eq!(config.chunk_timeout_duration(), Duration::from_millis(1000)); + } + + #[tokio::test] + async fn test_stream_manager_handle_stream_with_timeout() { + let manager = StreamManager::with_default_config(); + + let context = StreamContext::new( + Some("test-flow".to_string()), + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + let chunks = vec![Ok(Bytes::from("{\"content\":\"Hello\"}"))]; + let source_stream: StreamResponse = Box::pin(stream::iter(chunks)); + + let mut timeout_stream = manager.handle_stream_with_timeout(context, source_stream); + + let mut events = Vec::new(); + while let Some(result) = timeout_stream.next().await { + if let Ok(event) = result { + events.push(event); + } + } + + // 应该成功完成 + assert!(!events.is_empty()); + } + + #[test] + fn test_stream_event_flow_id() { + let started = StreamEvent::Started { + flow_id: "flow-1".to_string(), + }; + assert_eq!(started.flow_id(), "flow-1"); + + let chunk = StreamEvent::Chunk { + flow_id: "flow-2".to_string(), + content_delta: Some("test".to_string()), + metrics: StreamMetrics::new(), + }; + assert_eq!(chunk.flow_id(), "flow-2"); + + let completed = StreamEvent::Completed { + flow_id: "flow-3".to_string(), + metrics: StreamMetrics::new(), + }; + assert_eq!(completed.flow_id(), "flow-3"); + + let error = StreamEvent::Error { + flow_id: "flow-4".to_string(), + error: StreamError::Timeout, + metrics: StreamMetrics::new(), + }; + assert_eq!(error.flow_id(), "flow-4"); + } + + #[tokio::test] + async fn test_managed_stream_provider_error() { + let context = StreamContext::new( + None, + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + // 创建包含 Provider 错误的流 + let chunks: Vec> = vec![ + Ok(Bytes::from("{\"content\":\"Hello\"}")), + Err(StreamError::provider_error(429, "rate limited")), + ]; + let source_stream: StreamResponse = Box::pin(stream::iter(chunks)); + let config = StreamConfig::default(); + + let mut managed = ManagedStream::new(context, source_stream, config); + + let mut events = Vec::new(); + while let Some(result) = managed.next().await { + if let Ok(event) = result { + events.push(event); + } + } + + // 应该有错误事件 + assert!(events.iter().any(|e| e.contains("provider_error"))); + assert!(events.iter().any(|e| e.contains("rate limited"))); + } + + #[tokio::test] + async fn test_managed_stream_network_error() { + let context = StreamContext::new( + None, + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + // 创建包含网络错误的流 + let chunks: Vec> = + vec![Err(StreamError::network("connection reset by peer"))]; + let source_stream: StreamResponse = Box::pin(stream::iter(chunks)); + let config = StreamConfig::default(); + + let mut managed = ManagedStream::new(context, source_stream, config); + + let mut events = Vec::new(); + while let Some(result) = managed.next().await { + if let Ok(event) = result { + events.push(event); + } + } + + // 应该有网络错误事件 + assert!(events.iter().any(|e| e.contains("network_error"))); + assert!(events.iter().any(|e| e.contains("connection reset"))); + } + + #[tokio::test] + async fn test_managed_stream_parse_error() { + let context = StreamContext::new( + None, + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + // 创建包含解析错误的流 + let chunks: Vec> = + vec![Err(StreamError::parse_error("invalid JSON"))]; + let source_stream: StreamResponse = Box::pin(stream::iter(chunks)); + let config = StreamConfig::default(); + + let mut managed = ManagedStream::new(context, source_stream, config); + + let mut events = Vec::new(); + while let Some(result) = managed.next().await { + if let Ok(event) = result { + events.push(event); + } + } + + // 应该有解析错误事件 + assert!(events.iter().any(|e| e.contains("parse_error"))); + } + + #[test] + fn test_stream_error_is_retryable() { + assert!(StreamError::Network("test".to_string()).is_retryable()); + assert!(StreamError::Timeout.is_retryable()); + assert!(StreamError::provider_error(429, "rate limited").is_retryable()); + assert!(StreamError::provider_error(500, "server error").is_retryable()); + assert!(StreamError::provider_error(503, "service unavailable").is_retryable()); + assert!(!StreamError::provider_error(400, "bad request").is_retryable()); + assert!(!StreamError::provider_error(401, "unauthorized").is_retryable()); + assert!(!StreamError::ParseError("invalid".to_string()).is_retryable()); + assert!(!StreamError::ClientDisconnected.is_retryable()); + assert!(!StreamError::BufferOverflow.is_retryable()); + } + + #[test] + fn test_stream_error_status_code() { + assert_eq!(StreamError::Timeout.status_code(), Some(504)); + assert_eq!( + StreamError::Network("test".to_string()).status_code(), + Some(502) + ); + assert_eq!( + StreamError::provider_error(429, "test").status_code(), + Some(429) + ); + assert_eq!( + StreamError::provider_error(500, "test").status_code(), + Some(500) + ); + assert_eq!(StreamError::ClientDisconnected.status_code(), None); + assert_eq!( + StreamError::ParseError("test".to_string()).status_code(), + None + ); + } + + #[test] + fn test_stream_error_to_sse_error() { + let err = StreamError::provider_error(429, "rate limited"); + let sse = err.to_sse_error(); + + assert!(sse.starts_with("event: error\n")); + assert!(sse.contains("provider_error")); + assert!(sse.contains("rate limited")); + } + + // ======================================================================== + // 有界缓冲区测试(需求 7.1) + // ======================================================================== + + #[tokio::test] + async fn test_bounded_buffer_overflow() { + let context = StreamContext::new( + None, + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + // 创建一个大的 chunk,超过缓冲区限制 + let large_data = vec![b'x'; 200]; + let chunks: Vec> = vec![Ok(Bytes::from(large_data))]; + let source_stream: StreamResponse = Box::pin(stream::iter(chunks)); + + // 设置一个很小的缓冲区限制 + let config = StreamConfig::new().with_buffer_size(100); + + let mut managed = ManagedStream::new(context, source_stream, config); + + let mut events = Vec::new(); + while let Some(result) = managed.next().await { + if let Ok(event) = result { + events.push(event); + } + } + + // 应该有缓冲区溢出错误事件 + assert!(events.iter().any(|e| e.contains("buffer_overflow"))); + assert!(managed.is_finished()); + } + + #[tokio::test] + async fn test_bounded_buffer_within_limit() { + let context = StreamContext::new( + None, + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + // 创建一个小的 chunk,在缓冲区限制内 + let small_data = b"{\"content\":\"test\"}"; + let chunks: Vec> = vec![Ok(Bytes::from(&small_data[..]))]; + let source_stream: StreamResponse = Box::pin(stream::iter(chunks)); + + // 设置足够大的缓冲区限制 + let config = StreamConfig::new().with_buffer_size(1024); + + let mut managed = ManagedStream::new(context, source_stream, config); + + let mut events = Vec::new(); + while let Some(result) = managed.next().await { + if let Ok(event) = result { + events.push(event); + } + } + + // 不应该有缓冲区溢出错误 + assert!(!events.iter().any(|e| e.contains("buffer_overflow"))); + } + + #[test] + fn test_buffer_usage_tracking() { + let context = StreamContext::new( + None, + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + let chunks: Vec> = vec![]; + let source_stream: StreamResponse = Box::pin(stream::iter(chunks)); + let config = StreamConfig::new().with_buffer_size(1024); + + let managed = ManagedStream::new(context, source_stream, config); + + // 初始缓冲区使用量应该为 0 + assert_eq!(managed.buffer_usage(), 0); + assert_eq!(managed.buffer_limit(), 1024); + } + + // ======================================================================== + // 资源清理测试(需求 7.2) + // ======================================================================== + + #[tokio::test] + async fn test_resource_cleanup() { + let context = StreamContext::new( + None, + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + let chunks: Vec> = + vec![Ok(Bytes::from("{\"content\":\"test\"}"))]; + let source_stream: StreamResponse = Box::pin(stream::iter(chunks)); + let config = StreamConfig::default(); + + let mut managed = ManagedStream::new(context, source_stream, config); + + // 消费流 + while managed.next().await.is_some() {} + + // 流完成后,缓冲区使用量应该为 0 + assert_eq!(managed.buffer_usage(), 0); + assert!(managed.is_finished()); + } + + #[tokio::test] + async fn test_manual_cleanup() { + let context = StreamContext::new( + None, + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + let chunks: Vec> = vec![]; + let source_stream: StreamResponse = Box::pin(stream::iter(chunks)); + let config = StreamConfig::default(); + + let mut managed = ManagedStream::new(context, source_stream, config); + + // 手动清理 + managed.cleanup(); + + // 清理后应该标记为已完成 + assert!(managed.is_finished()); + assert_eq!(managed.buffer_usage(), 0); + } + + // ======================================================================== + // 事件节流测试(需求 4.6) + // ======================================================================== + + #[tokio::test] + async fn test_callback_throttling_counts() { + let context = StreamContext::new( + None, + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + // 创建多个 chunks + let chunks: Vec> = vec![ + Ok(Bytes::from("{\"content\":\"a\"}")), + Ok(Bytes::from("{\"content\":\"b\"}")), + Ok(Bytes::from("{\"content\":\"c\"}")), + ]; + let source_stream: StreamResponse = Box::pin(stream::iter(chunks)); + + // 使用较长的节流间隔 + let config = StreamConfig::new().with_throttle_ms(10000); // 10秒节流 + + let callback_count = Arc::new(AtomicU32::new(0)); + let callback_count_clone = callback_count.clone(); + + let on_chunk = move |_event: &str, _metrics: &StreamMetrics| { + callback_count_clone.fetch_add(1, Ordering::SeqCst); + }; + + let mut managed = ManagedStreamWithCallback::new(context, source_stream, config, on_chunk); + + // 消费流 + while managed.next().await.is_some() {} + + // 由于节流,回调次数应该小于事件数量 + let final_callback_count = callback_count.load(Ordering::SeqCst); + assert!(final_callback_count >= 1, "至少应该有一次回调"); + + // 检查节流计数 + let throttled = managed.throttled_event_count(); + let total_callbacks = managed.callback_count(); + + // 总回调次数 + 被节流次数 应该等于处理的事件数 + assert!(total_callbacks >= 1); + } +} + +// ============================================================================ +// 属性测试(Property-Based Testing) +// ============================================================================ + +#[cfg(test)] +mod property_tests { + use super::*; + use crate::streaming::aws_parser::{extract_content, serialize_event, AwsEvent}; + use crate::streaming::converter::extract_content_from_sse; + use futures::stream; + use proptest::prelude::*; + use std::sync::atomic::{AtomicU32, Ordering}; + use std::sync::Arc; + + // ======================================================================== + // 生成器(Generators) + // ======================================================================== + + /// 生成有效的内容文本 + fn arb_content_text() -> impl Strategy { + prop::string::string_regex("[a-zA-Z0-9 .,!?\\-_]{1,50}") + .unwrap() + .prop_filter("非空字符串", |s| !s.is_empty()) + } + + /// 生成 Content 事件 + fn arb_content_event() -> impl Strategy { + arb_content_text().prop_map(|text| AwsEvent::Content { text }) + } + + /// 生成内容事件序列 + fn arb_content_sequence() -> impl Strategy> { + prop::collection::vec(arb_content_event(), 1..10) + } + + // ======================================================================== + // Property 3: Flow Monitor 流式捕获完整性 + // + // *对于任意*流式响应,Flow Monitor 通过 process_chunk 处理所有 chunks 后, + // 重建的响应应该与完整响应一致。 + // + // **验证: 需求 4.2, 4.4, 4.5** + // ======================================================================== + + proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// Property 3: 流式捕获完整性 - 所有 chunk 都被回调处理 + /// + /// **Feature: true-streaming-support, Property 3: Flow Monitor 流式捕获完整性** + /// **Validates: Requirements 4.2, 4.4, 4.5** + #[test] + fn prop_flow_monitor_captures_all_chunks(events in arb_content_sequence()) { + // 使用 tokio runtime 运行异步测试 + let rt = tokio::runtime::Runtime::new().unwrap(); + let result = rt.block_on(async { + // 计算原始内容 + let original_content = extract_content(&events); + + // 创建上下文 + let context = StreamContext::new( + Some("test-flow".to_string()), + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + // 将事件序列化为字节流 + let chunks: Vec> = events + .iter() + .filter_map(|event| serialize_event(event)) + .map(|json| Ok(Bytes::from(json))) + .collect(); + + let source_stream: StreamResponse = Box::pin(stream::iter(chunks)); + let config = StreamConfig::default(); + + // 使用回调跟踪所有事件 + let callback_count = Arc::new(AtomicU32::new(0)); + let callback_count_clone = callback_count.clone(); + let captured_events = Arc::new(std::sync::Mutex::new(Vec::new())); + let captured_events_clone = captured_events.clone(); + + let on_chunk = move |event: &str, _metrics: &StreamMetrics| { + callback_count_clone.fetch_add(1, Ordering::SeqCst); + captured_events_clone.lock().unwrap().push(event.to_string()); + }; + + let mut managed = ManagedStreamWithCallback::new( + context, + source_stream, + config, + on_chunk, + ); + + // 消费流 + let mut all_events = Vec::new(); + while let Some(result) = managed.next().await { + if let Ok(event) = result { + all_events.push(event); + } + } + + // 验证回调被调用 + let final_callback_count = callback_count.load(Ordering::SeqCst); + if final_callback_count == 0 && !events.is_empty() { + return Err("回调应该被调用至少一次(除非没有事件)".to_string()); + } + + // 从 SSE 事件中提取内容 + let converted_content = extract_content_from_sse(&all_events, StreamFormat::OpenAiSse); + + // 验证内容一致 + if original_content != converted_content { + return Err(format!( + "Flow Monitor 捕获的内容应该与原始内容一致: original={}, converted={}", + original_content, converted_content + )); + } + + Ok(()) + }); + + prop_assert!(result.is_ok(), "{}", result.unwrap_err()); + } + + /// Property 3: 流式指标正确记录 + /// + /// **Feature: true-streaming-support, Property 3: Flow Monitor 流式捕获完整性** + /// **Validates: Requirements 4.5** + #[test] + fn prop_stream_metrics_correctly_recorded(events in arb_content_sequence()) { + let rt = tokio::runtime::Runtime::new().unwrap(); + let result = rt.block_on(async { + let context = StreamContext::new( + None, + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + // 将事件序列化为字节流 + let chunks: Vec> = events + .iter() + .filter_map(|event| serialize_event(event)) + .map(|json| Ok(Bytes::from(json))) + .collect(); + + let chunk_count = chunks.len(); + let total_bytes: usize = chunks.iter() + .filter_map(|r| r.as_ref().ok()) + .map(|b| b.len()) + .sum(); + + let source_stream: StreamResponse = Box::pin(stream::iter(chunks)); + let config = StreamConfig::default(); + + let mut managed = ManagedStream::new(context, source_stream, config); + + // 消费流 + while managed.next().await.is_some() {} + + // 验证指标 + let metrics = managed.metrics(); + + if metrics.chunk_count as usize != chunk_count { + return Err(format!( + "chunk 计数应该与实际 chunk 数量一致: expected={}, actual={}", + chunk_count, metrics.chunk_count + )); + } + + if metrics.total_bytes != total_bytes { + return Err(format!( + "总字节数应该与实际字节数一致: expected={}, actual={}", + total_bytes, metrics.total_bytes + )); + } + + if chunk_count > 0 && metrics.ttfb_ms.is_none() { + return Err("有 chunk 时应该记录 TTFB".to_string()); + } + + Ok(()) + }); + + prop_assert!(result.is_ok(), "{}", result.unwrap_err()); + } + + /// Property 3: 回调节流正确工作 + /// + /// **Feature: true-streaming-support, Property 3: Flow Monitor 流式捕获完整性** + /// **Validates: Requirements 4.6** + #[test] + fn prop_callback_throttling_works(events in arb_content_sequence()) { + let rt = tokio::runtime::Runtime::new().unwrap(); + let result = rt.block_on(async { + let context = StreamContext::new( + None, + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + // 将事件序列化为字节流 + let chunks: Vec> = events + .iter() + .filter_map(|event| serialize_event(event)) + .map(|json| Ok(Bytes::from(json))) + .collect(); + + let source_stream: StreamResponse = Box::pin(stream::iter(chunks)); + + // 使用较长的节流间隔 + let config = StreamConfig::new() + .with_throttle_ms(1000); // 1秒节流 + + let callback_count = Arc::new(AtomicU32::new(0)); + let callback_count_clone = callback_count.clone(); + + let on_chunk = move |_event: &str, _metrics: &StreamMetrics| { + callback_count_clone.fetch_add(1, Ordering::SeqCst); + }; + + let mut managed = ManagedStreamWithCallback::new( + context, + source_stream, + config, + on_chunk, + ); + + // 消费流 + while managed.next().await.is_some() {} + + // 由于节流,回调次数应该小于等于事件数量 + let final_callback_count = callback_count.load(Ordering::SeqCst); + + // 至少应该有一次回调(第一次不受节流限制) + if !events.is_empty() && final_callback_count < 1 { + return Err("至少应该有一次回调".to_string()); + } + + Ok(()) + }); + + prop_assert!(result.is_ok(), "{}", result.unwrap_err()); + } + } + + // ======================================================================== + // Property 5: 错误恢复正确性 + // + // *对于任意*包含无效 chunk 的流,解析器应该跳过无效 chunk 并继续处理 + // 后续有效 chunks,最终结果应该包含所有有效 chunks 的内容。 + // + // **验证: 需求 2.5, 6.4** + // ======================================================================== + + proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// Property 5: 错误恢复 - 流在错误后正确终止 + /// + /// **Feature: true-streaming-support, Property 5: 错误恢复正确性** + /// **Validates: Requirements 2.5, 6.4** + #[test] + fn prop_stream_terminates_on_error( + valid_events in arb_content_sequence(), + error_type in prop_oneof![ + Just(StreamError::Network("test error".to_string())), + Just(StreamError::Timeout), + Just(StreamError::provider_error(500, "server error")), + Just(StreamError::parse_error("invalid data")), + ] + ) { + let rt = tokio::runtime::Runtime::new().unwrap(); + let result = rt.block_on(async { + let context = StreamContext::new( + None, + StreamFormat::AwsEventStream, + StreamFormat::OpenAiSse, + "test-model", + ); + + // 创建包含有效事件和错误的流 + let mut chunks: Vec> = valid_events + .iter() + .filter_map(|event| serialize_event(event)) + .map(|json| Ok(Bytes::from(json))) + .collect(); + + // 在中间插入错误 + let error_pos = chunks.len() / 2; + chunks.insert(error_pos, Err(error_type.clone())); + + let source_stream: StreamResponse = Box::pin(stream::iter(chunks)); + let config = StreamConfig::default(); + + let mut managed = ManagedStream::new(context, source_stream, config); + + // 消费流 + let mut events = Vec::new(); + let mut saw_error = false; + while let Some(result) = managed.next().await { + match result { + Ok(event) => { + if event.contains("error") { + saw_error = true; + } + events.push(event); + } + Err(_) => { + saw_error = true; + } + } + } + + // 验证流正确终止 + if !saw_error { + return Err("应该看到错误事件".to_string()); + } + + // 验证流已完成 + if !managed.is_finished() { + return Err("流应该已完成".to_string()); + } + + Ok(()) + }); + + prop_assert!(result.is_ok(), "{}", result.unwrap_err()); + } + + /// Property 5: 错误事件格式正确 + /// + /// **Feature: true-streaming-support, Property 5: 错误恢复正确性** + /// **Validates: Requirements 6.4** + #[test] + fn prop_error_event_format_correct( + error_type in prop_oneof![ + Just(StreamError::Network("connection failed".to_string())), + Just(StreamError::Timeout), + Just(StreamError::provider_error(429, "rate limited")), + Just(StreamError::provider_error(500, "internal error")), + Just(StreamError::parse_error("invalid json")), + Just(StreamError::ClientDisconnected), + Just(StreamError::BufferOverflow), + ] + ) { + let sse_error = error_type.to_sse_error(); + + // 验证 SSE 格式 + prop_assert!( + sse_error.starts_with("event: error\n"), + "SSE 错误应该以 'event: error' 开头, 实际: {}", + sse_error + ); + + prop_assert!( + sse_error.contains("data: "), + "SSE 错误应该包含 'data: ', 实际: {}", + sse_error + ); + + // 验证 JSON 格式 + let data_start = sse_error.find("data: ").unwrap() + 6; + let data_end = sse_error[data_start..].find('\n').unwrap_or(sse_error.len() - data_start); + let json_str = &sse_error[data_start..data_start + data_end]; + + let json_result: Result = serde_json::from_str(json_str); + prop_assert!( + json_result.is_ok(), + "无法解析 JSON: {:?}, 原始字符串: {}", + json_result.err(), + json_str + ); + + let json = json_result.unwrap(); + + // 验证 JSON 结构 + prop_assert!( + json.get("error").is_some(), + "JSON 应该包含 'error' 字段, 实际: {}", + json + ); + + let error_obj = json.get("error").unwrap(); + prop_assert!( + error_obj.get("type").is_some(), + "error 对象应该包含 'type' 字段, 实际: {}", + error_obj + ); + + prop_assert!( + error_obj.get("message").is_some(), + "error 对象应该包含 'message' 字段, 实际: {}", + error_obj + ); + } + + /// Property 5: 可重试错误正确标识 + /// + /// **Feature: true-streaming-support, Property 5: 错误恢复正确性** + /// **Validates: Requirements 6.1, 6.4** + #[test] + fn prop_retryable_errors_correctly_identified( + status in 400u16..600, + message in "[a-zA-Z ]{1,20}" + ) { + let error = StreamError::provider_error(status, &message); + let is_retryable = error.is_retryable(); + + // 429 和 5xx 应该可重试 + let expected_retryable = status == 429 || status >= 500; + + prop_assert_eq!( + is_retryable, + expected_retryable, + "状态码 {} 的可重试性应该是 {}, 但实际是 {}", + status, expected_retryable, is_retryable + ); + } + + /// Property 5: 错误状态码正确映射 + /// + /// **Feature: true-streaming-support, Property 5: 错误恢复正确性** + /// **Validates: Requirements 6.1, 6.2, 6.3** + #[test] + fn prop_error_status_code_mapping( + status in 400u16..600, + message in "[a-zA-Z ]{1,20}" + ) { + let provider_error = StreamError::provider_error(status, &message); + let network_error = StreamError::network(&message); + let timeout_error = StreamError::Timeout; + + // Provider 错误应该返回原始状态码 + prop_assert_eq!( + provider_error.status_code(), + Some(status), + "Provider 错误状态码应该是 {}, 但实际是 {:?}", + status, provider_error.status_code() + ); + + // 网络错误应该返回 502 + prop_assert_eq!( + network_error.status_code(), + Some(502), + "网络错误状态码应该是 502, 但实际是 {:?}", + network_error.status_code() + ); + + // 超时错误应该返回 504 + prop_assert_eq!( + timeout_error.status_code(), + Some(504), + "超时错误状态码应该是 504, 但实际是 {:?}", + timeout_error.status_code() + ); + } + } +} diff --git a/src-tauri/src/streaming/metrics.rs b/src-tauri/src/streaming/metrics.rs new file mode 100644 index 000000000..8b0e2734a --- /dev/null +++ b/src-tauri/src/streaming/metrics.rs @@ -0,0 +1,567 @@ +//! 流式传输指标类型 +//! +//! 定义流式传输过程中的性能指标和统计数据。 +//! +//! # 需求覆盖 +//! +//! - 需求 4.5: 跟踪 chunk 数量和接收的总字节数 +//! - 需求 7.5: 记录流式指标(吞吐量、延迟、错误率) + +use chrono::{DateTime, Utc}; +use serde::{Deserialize, Serialize}; +use tracing::info; + +/// 流式传输指标 +/// +/// 记录流式传输过程中的各种性能指标。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct StreamMetrics { + /// 首字节时间(毫秒) + /// + /// 从请求发送到收到第一个响应字节的时间。 + /// 对应需求 4.1 中的 TTFB 记录。 + #[serde(skip_serializing_if = "Option::is_none")] + pub ttfb_ms: Option, + + /// chunk 数量 + /// + /// 接收到的流式 chunk 总数。 + /// 对应需求 4.5。 + pub chunk_count: u32, + + /// 总字节数 + /// + /// 接收到的总字节数。 + /// 对应需求 4.5。 + pub total_bytes: usize, + + /// 开始时间 + /// + /// 流式传输开始的时间戳。 + pub start_time: DateTime, + + /// 结束时间 + /// + /// 流式传输结束的时间戳(如果已结束)。 + #[serde(skip_serializing_if = "Option::is_none")] + pub end_time: Option>, + + /// 首个 chunk 时间 + /// + /// 收到第一个 chunk 的时间戳。 + #[serde(skip_serializing_if = "Option::is_none")] + pub first_chunk_time: Option>, + + /// 最后一个 chunk 时间 + /// + /// 收到最后一个 chunk 的时间戳。 + #[serde(skip_serializing_if = "Option::is_none")] + pub last_chunk_time: Option>, + + /// 解析错误数量 + /// + /// 流式传输过程中遇到的解析错误数量。 + pub parse_error_count: u32, + + /// 重试次数 + /// + /// 流式传输过程中的重试次数。 + pub retry_count: u32, + + /// 最小 chunk 大小(字节) + /// + /// 对应需求 7.5: 记录流式指标 + #[serde(skip_serializing_if = "Option::is_none")] + pub min_chunk_size: Option, + + /// 最大 chunk 大小(字节) + /// + /// 对应需求 7.5: 记录流式指标 + #[serde(skip_serializing_if = "Option::is_none")] + pub max_chunk_size: Option, + + /// 缓冲区溢出次数 + /// + /// 对应需求 7.1: 有界缓冲区 + pub buffer_overflow_count: u32, + + /// 被节流的事件数量 + /// + /// 对应需求 4.6: 事件节流 + pub throttled_event_count: u32, +} + +impl Default for StreamMetrics { + fn default() -> Self { + Self { + ttfb_ms: None, + chunk_count: 0, + total_bytes: 0, + start_time: Utc::now(), + end_time: None, + first_chunk_time: None, + last_chunk_time: None, + parse_error_count: 0, + retry_count: 0, + min_chunk_size: None, + max_chunk_size: None, + buffer_overflow_count: 0, + throttled_event_count: 0, + } + } +} + +impl StreamMetrics { + /// 创建新的指标实例 + pub fn new() -> Self { + Self::default() + } + + /// 记录收到第一个 chunk + /// + /// 自动计算 TTFB 并记录首个 chunk 时间。 + pub fn record_first_chunk(&mut self) { + let now = Utc::now(); + self.first_chunk_time = Some(now); + self.ttfb_ms = Some((now - self.start_time).num_milliseconds().max(0) as u64); + } + + /// 记录收到一个 chunk + /// + /// 更新 chunk 计数、字节数和最后 chunk 时间。 + /// 同时更新最小/最大 chunk 大小统计。 + pub fn record_chunk(&mut self, bytes: usize) { + self.chunk_count += 1; + self.total_bytes += bytes; + self.last_chunk_time = Some(Utc::now()); + + // 更新最小/最大 chunk 大小(需求 7.5) + match self.min_chunk_size { + None => self.min_chunk_size = Some(bytes), + Some(min) if bytes < min => self.min_chunk_size = Some(bytes), + _ => {} + } + match self.max_chunk_size { + None => self.max_chunk_size = Some(bytes), + Some(max) if bytes > max => self.max_chunk_size = Some(bytes), + _ => {} + } + + // 如果是第一个 chunk,记录 TTFB + if self.first_chunk_time.is_none() { + self.record_first_chunk(); + } + } + + /// 记录解析错误 + pub fn record_parse_error(&mut self) { + self.parse_error_count += 1; + } + + /// 记录重试 + pub fn record_retry(&mut self) { + self.retry_count += 1; + } + + /// 记录缓冲区溢出 + /// + /// 对应需求 7.1: 有界缓冲区 + pub fn record_buffer_overflow(&mut self) { + self.buffer_overflow_count += 1; + } + + /// 记录被节流的事件 + /// + /// 对应需求 4.6: 事件节流 + pub fn record_throttled_event(&mut self) { + self.throttled_event_count += 1; + } + + /// 批量记录被节流的事件 + /// + /// 对应需求 4.6: 事件节流 + pub fn record_throttled_events(&mut self, count: u32) { + self.throttled_event_count += count; + } + + /// 完成流式传输 + /// + /// 记录结束时间。 + pub fn finish(&mut self) { + self.end_time = Some(Utc::now()); + } + + /// 获取总耗时(毫秒) + /// + /// 如果流式传输已结束,返回从开始到结束的时间。 + /// 否则返回从开始到现在的时间。 + pub fn duration_ms(&self) -> u64 { + let end = self.end_time.unwrap_or_else(Utc::now); + (end - self.start_time).num_milliseconds().max(0) as u64 + } + + /// 获取平均 chunk 间隔(毫秒) + /// + /// 如果 chunk 数量少于 2,返回 None。 + pub fn avg_chunk_interval_ms(&self) -> Option { + if self.chunk_count < 2 { + return None; + } + + let first = self.first_chunk_time?; + let last = self.last_chunk_time?; + let interval_ms = (last - first).num_milliseconds().max(0) as f64; + Some(interval_ms / (self.chunk_count - 1) as f64) + } + + /// 获取平均 chunk 大小(字节) + /// + /// 对应需求 7.5: 记录流式指标 + pub fn avg_chunk_size(&self) -> Option { + if self.chunk_count == 0 { + return None; + } + Some(self.total_bytes as f64 / self.chunk_count as f64) + } + + /// 获取吞吐量(字节/秒) + /// + /// 如果耗时为 0,返回 None。 + pub fn throughput_bytes_per_sec(&self) -> Option { + let duration_ms = self.duration_ms(); + if duration_ms == 0 { + return None; + } + Some(self.total_bytes as f64 / (duration_ms as f64 / 1000.0)) + } + + /// 获取错误率 + /// + /// 解析错误数量 / chunk 数量。 + /// 如果 chunk 数量为 0,返回 0.0。 + pub fn error_rate(&self) -> f64 { + if self.chunk_count == 0 { + return 0.0; + } + self.parse_error_count as f64 / self.chunk_count as f64 + } + + /// 获取节流率 + /// + /// 被节流的事件数量 / (被节流的事件数量 + chunk 数量) + /// 对应需求 4.6: 事件节流 + pub fn throttle_rate(&self) -> f64 { + let total = self.throttled_event_count + self.chunk_count; + if total == 0 { + return 0.0; + } + self.throttled_event_count as f64 / total as f64 + } + + /// 判断流式传输是否已完成 + pub fn is_finished(&self) -> bool { + self.end_time.is_some() + } + + /// 转换为摘要字符串 + pub fn summary(&self) -> String { + let duration = self.duration_ms(); + let ttfb = self + .ttfb_ms + .map(|t| format!("{}ms", t)) + .unwrap_or_else(|| "N/A".to_string()); + let throughput = self + .throughput_bytes_per_sec() + .map(|t| format!("{:.2} KB/s", t / 1024.0)) + .unwrap_or_else(|| "N/A".to_string()); + let avg_chunk = self + .avg_chunk_size() + .map(|s| format!("{:.0}B", s)) + .unwrap_or_else(|| "N/A".to_string()); + + format!( + "chunks: {}, bytes: {}, duration: {}ms, ttfb: {}, throughput: {}, avg_chunk: {}, errors: {}, throttled: {}", + self.chunk_count, + self.total_bytes, + duration, + ttfb, + throughput, + avg_chunk, + self.parse_error_count, + self.throttled_event_count + ) + } + + /// 记录详细指标到日志 + /// + /// 对应需求 7.5: 记录流式指标(吞吐量、延迟、错误率) + pub fn log_metrics(&self, flow_id: Option<&str>) { + let throughput = self.throughput_bytes_per_sec().unwrap_or(0.0); + let error_rate = self.error_rate(); + let throttle_rate = self.throttle_rate(); + let avg_interval = self.avg_chunk_interval_ms().unwrap_or(0.0); + + info!( + flow_id = ?flow_id, + chunk_count = self.chunk_count, + total_bytes = self.total_bytes, + duration_ms = self.duration_ms(), + ttfb_ms = ?self.ttfb_ms, + throughput_kbps = format!("{:.2}", throughput / 1024.0), + avg_chunk_interval_ms = format!("{:.2}", avg_interval), + min_chunk_size = ?self.min_chunk_size, + max_chunk_size = ?self.max_chunk_size, + avg_chunk_size = ?self.avg_chunk_size(), + parse_error_count = self.parse_error_count, + error_rate = format!("{:.4}", error_rate), + buffer_overflow_count = self.buffer_overflow_count, + throttled_event_count = self.throttled_event_count, + throttle_rate = format!("{:.4}", throttle_rate), + "流式传输指标" + ); + } +} + +// ============================================================================ +// 测试模块 +// ============================================================================ + +#[cfg(test)] +mod tests { + use super::*; + use std::thread::sleep; + use std::time::Duration; + + #[test] + fn test_stream_metrics_default() { + let metrics = StreamMetrics::default(); + assert_eq!(metrics.chunk_count, 0); + assert_eq!(metrics.total_bytes, 0); + assert!(metrics.ttfb_ms.is_none()); + assert!(metrics.end_time.is_none()); + assert!(metrics.min_chunk_size.is_none()); + assert!(metrics.max_chunk_size.is_none()); + assert_eq!(metrics.buffer_overflow_count, 0); + assert_eq!(metrics.throttled_event_count, 0); + } + + #[test] + fn test_stream_metrics_record_chunk() { + let mut metrics = StreamMetrics::new(); + + metrics.record_chunk(100); + assert_eq!(metrics.chunk_count, 1); + assert_eq!(metrics.total_bytes, 100); + assert!(metrics.first_chunk_time.is_some()); + assert!(metrics.ttfb_ms.is_some()); + assert_eq!(metrics.min_chunk_size, Some(100)); + assert_eq!(metrics.max_chunk_size, Some(100)); + + metrics.record_chunk(200); + assert_eq!(metrics.chunk_count, 2); + assert_eq!(metrics.total_bytes, 300); + assert_eq!(metrics.min_chunk_size, Some(100)); + assert_eq!(metrics.max_chunk_size, Some(200)); + + metrics.record_chunk(50); + assert_eq!(metrics.chunk_count, 3); + assert_eq!(metrics.total_bytes, 350); + assert_eq!(metrics.min_chunk_size, Some(50)); + assert_eq!(metrics.max_chunk_size, Some(200)); + } + + #[test] + fn test_stream_metrics_record_first_chunk() { + let mut metrics = StreamMetrics::new(); + + // 等待一小段时间以确保 TTFB > 0 + sleep(Duration::from_millis(10)); + + metrics.record_first_chunk(); + assert!(metrics.first_chunk_time.is_some()); + assert!(metrics.ttfb_ms.is_some()); + assert!(metrics.ttfb_ms.unwrap() >= 10); + } + + #[test] + fn test_stream_metrics_finish() { + let mut metrics = StreamMetrics::new(); + assert!(!metrics.is_finished()); + + metrics.finish(); + assert!(metrics.is_finished()); + assert!(metrics.end_time.is_some()); + } + + #[test] + fn test_stream_metrics_duration() { + let mut metrics = StreamMetrics::new(); + + sleep(Duration::from_millis(50)); + + let duration = metrics.duration_ms(); + assert!(duration >= 50); + + metrics.finish(); + let final_duration = metrics.duration_ms(); + assert!(final_duration >= 50); + } + + #[test] + fn test_stream_metrics_avg_chunk_interval() { + let mut metrics = StreamMetrics::new(); + + // 单个 chunk 没有间隔 + metrics.record_chunk(100); + assert!(metrics.avg_chunk_interval_ms().is_none()); + + sleep(Duration::from_millis(20)); + + metrics.record_chunk(100); + let interval = metrics.avg_chunk_interval_ms(); + assert!(interval.is_some()); + assert!(interval.unwrap() >= 20.0); + } + + #[test] + fn test_stream_metrics_avg_chunk_size() { + let mut metrics = StreamMetrics::new(); + + // 没有 chunk 时没有平均大小 + assert!(metrics.avg_chunk_size().is_none()); + + metrics.record_chunk(100); + assert_eq!(metrics.avg_chunk_size(), Some(100.0)); + + metrics.record_chunk(200); + assert_eq!(metrics.avg_chunk_size(), Some(150.0)); + + metrics.record_chunk(300); + assert_eq!(metrics.avg_chunk_size(), Some(200.0)); + } + + #[test] + fn test_stream_metrics_throughput() { + let mut metrics = StreamMetrics::new(); + + // 没有数据时没有吞吐量 + assert!(metrics.throughput_bytes_per_sec().is_none()); + + metrics.record_chunk(1000); + sleep(Duration::from_millis(100)); + metrics.finish(); + + let throughput = metrics.throughput_bytes_per_sec(); + assert!(throughput.is_some()); + // 1000 bytes in ~100ms = ~10000 bytes/sec + assert!(throughput.unwrap() > 0.0); + } + + #[test] + fn test_stream_metrics_error_rate() { + let mut metrics = StreamMetrics::new(); + + // 没有 chunk 时错误率为 0 + assert_eq!(metrics.error_rate(), 0.0); + + metrics.record_chunk(100); + metrics.record_chunk(100); + metrics.record_parse_error(); + + // 2 chunks, 1 error = 50% error rate + assert_eq!(metrics.error_rate(), 0.5); + } + + #[test] + fn test_stream_metrics_throttle_rate() { + let mut metrics = StreamMetrics::new(); + + // 没有事件时节流率为 0 + assert_eq!(metrics.throttle_rate(), 0.0); + + metrics.record_chunk(100); + metrics.record_chunk(100); + metrics.record_throttled_event(); + metrics.record_throttled_event(); + + // 2 chunks, 2 throttled = 50% throttle rate + assert_eq!(metrics.throttle_rate(), 0.5); + } + + #[test] + fn test_stream_metrics_buffer_overflow() { + let mut metrics = StreamMetrics::new(); + + assert_eq!(metrics.buffer_overflow_count, 0); + + metrics.record_buffer_overflow(); + assert_eq!(metrics.buffer_overflow_count, 1); + + metrics.record_buffer_overflow(); + assert_eq!(metrics.buffer_overflow_count, 2); + } + + #[test] + fn test_stream_metrics_throttled_events() { + let mut metrics = StreamMetrics::new(); + + assert_eq!(metrics.throttled_event_count, 0); + + metrics.record_throttled_event(); + assert_eq!(metrics.throttled_event_count, 1); + + metrics.record_throttled_events(5); + assert_eq!(metrics.throttled_event_count, 6); + } + + #[test] + fn test_stream_metrics_summary() { + let mut metrics = StreamMetrics::new(); + metrics.record_chunk(1024); + metrics.record_chunk(2048); + metrics.record_throttled_event(); + metrics.finish(); + + let summary = metrics.summary(); + assert!(summary.contains("chunks: 2")); + assert!(summary.contains("bytes: 3072")); + assert!(summary.contains("throttled: 1")); + } + + #[test] + fn test_stream_metrics_serialization() { + let mut metrics = StreamMetrics::new(); + metrics.record_chunk(100); + metrics.record_buffer_overflow(); + metrics.record_throttled_event(); + metrics.finish(); + + let json = serde_json::to_string(&metrics).unwrap(); + let deserialized: StreamMetrics = serde_json::from_str(&json).unwrap(); + + assert_eq!(metrics.chunk_count, deserialized.chunk_count); + assert_eq!(metrics.total_bytes, deserialized.total_bytes); + assert_eq!( + metrics.buffer_overflow_count, + deserialized.buffer_overflow_count + ); + assert_eq!( + metrics.throttled_event_count, + deserialized.throttled_event_count + ); + } + + #[test] + fn test_stream_metrics_log_metrics() { + let mut metrics = StreamMetrics::new(); + metrics.record_chunk(1024); + metrics.record_chunk(2048); + metrics.record_parse_error(); + metrics.record_throttled_event(); + metrics.finish(); + + // 这个测试主要确保 log_metrics 不会 panic + metrics.log_metrics(Some("test-flow-id")); + metrics.log_metrics(None); + } +} diff --git a/src-tauri/src/streaming/mod.rs b/src-tauri/src/streaming/mod.rs new file mode 100644 index 000000000..6b4d6fc8c --- /dev/null +++ b/src-tauri/src/streaming/mod.rs @@ -0,0 +1,41 @@ +//! 流式传输核心模块 +//! +//! 该模块提供真正的端到端流式传输支持,将当前的"伪流式"架构改造为 +//! 逐 chunk 处理和实时 Flow 监控的真正流式传输。 +//! +//! # 主要组件 +//! +//! - `error`: 流式错误类型定义 +//! - `metrics`: 流式指标类型定义 +//! - `aws_parser`: AWS Event Stream 解析器(用于 Kiro/CodeWhisperer) +//! - `converter`: 流式格式转换器 +//! - `traits`: StreamingProvider trait 定义 +//! - `manager`: 流式管理器 + +pub mod aws_parser; +pub mod converter; +pub mod error; +pub mod manager; +pub mod metrics; +pub mod traits; + +// 重新导出核心类型 +pub use aws_parser::{ + extract_content, extract_tool_calls, serialize_event, AwsEvent, AwsEventStreamParser, + ParserState, +}; +pub use converter::{ + extract_content_from_sse, extract_tool_calls_from_sse, ConverterState, PartialJsonAccumulator, + StreamConverter, StreamFormat, +}; +pub use error::StreamError; +pub use manager::{ + collect_stream_content, create_flow_monitor_callback, with_timeout, FlowMonitorCallback, + ManagedStream, ManagedStreamWithCallback, StreamConfig, StreamContext, StreamEvent, + StreamManager, TimeoutStream, +}; +pub use metrics::StreamMetrics; +pub use traits::{ + reqwest_stream_to_stream_response, StreamFormat as TraitsStreamFormat, StreamResponse, + StreamingProvider, +}; diff --git a/src-tauri/src/streaming/traits.rs b/src-tauri/src/streaming/traits.rs new file mode 100644 index 000000000..a2f675417 --- /dev/null +++ b/src-tauri/src/streaming/traits.rs @@ -0,0 +1,177 @@ +//! StreamingProvider Trait 定义 +//! +//! 为 Provider 定义流式 API 接口,支持真正的端到端流式传输。 +//! +//! # 需求覆盖 +//! +//! - 需求 1.1: KiroProvider 流式支持 +//! - 需求 1.2: ClaudeCustomProvider 流式支持 +//! - 需求 1.3: OpenAICustomProvider 流式支持 +//! - 需求 1.4: AntigravityProvider 流式支持 + +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 std::pin::Pin; + +/// 流式响应类型别名 +/// +/// 返回一个异步字节流,每个 Item 是一个 chunk 的字节数据或错误。 +/// 使用 `Pin>` 以支持动态分发和异步迭代。 +pub type StreamResponse = Pin> + Send>>; + +/// 流式 Provider Trait +/// +/// 定义所有支持流式传输的 Provider 必须实现的接口。 +/// 与现有的非流式 API 调用方法并存,允许 Provider 同时支持两种模式。 +#[async_trait] +pub trait StreamingProvider: Send + Sync { + /// 发起流式 API 调用 + /// + /// 返回一个字节流,调用者可以逐 chunk 处理响应数据。 + /// + /// # Arguments + /// + /// * `request` - OpenAI 格式的聊天完成请求 + /// + /// # Returns + /// + /// * `Ok(StreamResponse)` - 成功时返回字节流 + /// * `Err(ProviderError)` - 失败时返回 Provider 错误 + /// + /// # Example + /// + /// ```ignore + /// use futures::StreamExt; + /// + /// let stream = provider.call_api_stream(&request).await?; + /// while let Some(chunk) = stream.next().await { + /// match chunk { + /// Ok(bytes) => { /* 处理字节数据 */ } + /// Err(e) => { /* 处理错误 */ } + /// } + /// } + /// ``` + async fn call_api_stream( + &self, + request: &ChatCompletionRequest, + ) -> Result; + + /// 检查是否支持流式传输 + /// + /// 默认返回 `true`,Provider 可以覆盖此方法以指示不支持流式。 + /// 当返回 `false` 时,调用者应该回退到非流式模式。 + fn supports_streaming(&self) -> bool { + true + } + + /// 获取 Provider 名称 + /// + /// 用于日志记录和错误消息。 + fn provider_name(&self) -> &'static str; + + /// 获取流式响应的格式 + /// + /// 返回此 Provider 的原生流式格式,用于后续的格式转换。 + fn stream_format(&self) -> StreamFormat; +} + +/// 流式格式枚举 +/// +/// 定义不同 Provider 使用的流式响应格式。 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum StreamFormat { + /// AWS Event Stream 格式(Kiro/CodeWhisperer 使用) + AwsEventStream, + /// Anthropic SSE 格式(Claude 使用) + AnthropicSse, + /// OpenAI SSE 格式(OpenAI 兼容 API 使用) + OpenAiSse, + /// Gemini 流式格式(Antigravity/Gemini 使用) + GeminiStream, +} + +impl StreamFormat { + /// 获取格式的 Content-Type + pub fn content_type(&self) -> &'static str { + match self { + StreamFormat::AwsEventStream => "application/vnd.amazon.eventstream", + StreamFormat::AnthropicSse => "text/event-stream", + StreamFormat::OpenAiSse => "text/event-stream", + StreamFormat::GeminiStream => "text/event-stream", + } + } + + /// 获取格式的显示名称 + pub fn display_name(&self) -> &'static str { + match self { + StreamFormat::AwsEventStream => "AWS Event Stream", + StreamFormat::AnthropicSse => "Anthropic SSE", + StreamFormat::OpenAiSse => "OpenAI SSE", + StreamFormat::GeminiStream => "Gemini Stream", + } + } +} + +// ============================================================================ +// 辅助函数 +// ============================================================================ + +/// 将 reqwest 的 bytes_stream 转换为 StreamResponse +/// +/// 这是一个辅助函数,用于将 reqwest 的响应流转换为统一的 StreamResponse 类型。 +pub fn reqwest_stream_to_stream_response(response: reqwest::Response) -> StreamResponse { + use futures::StreamExt; + + let stream = response + .bytes_stream() + .map(|result| result.map_err(|e| StreamError::from(e))); + + Box::pin(stream) +} + +// ============================================================================ +// 测试模块 +// ============================================================================ + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_stream_format_content_type() { + assert_eq!( + StreamFormat::AwsEventStream.content_type(), + "application/vnd.amazon.eventstream" + ); + assert_eq!( + StreamFormat::AnthropicSse.content_type(), + "text/event-stream" + ); + assert_eq!(StreamFormat::OpenAiSse.content_type(), "text/event-stream"); + assert_eq!( + StreamFormat::GeminiStream.content_type(), + "text/event-stream" + ); + } + + #[test] + fn test_stream_format_display_name() { + assert_eq!( + StreamFormat::AwsEventStream.display_name(), + "AWS Event Stream" + ); + assert_eq!(StreamFormat::AnthropicSse.display_name(), "Anthropic SSE"); + assert_eq!(StreamFormat::OpenAiSse.display_name(), "OpenAI SSE"); + assert_eq!(StreamFormat::GeminiStream.display_name(), "Gemini Stream"); + } + + #[test] + fn test_stream_format_equality() { + assert_eq!(StreamFormat::AwsEventStream, StreamFormat::AwsEventStream); + assert_ne!(StreamFormat::AwsEventStream, StreamFormat::OpenAiSse); + } +} diff --git a/src-tauri/src/tray/sync.rs b/src-tauri/src/tray/sync.rs index 7ed91a87f..b196ea434 100644 --- a/src-tauri/src/tray/sync.rs +++ b/src-tauri/src/tray/sync.rs @@ -9,9 +9,9 @@ use super::state::{calculate_icon_status, CredentialHealth, TrayIconStatus, TrayStateSnapshot}; use super::TrayManager; use std::sync::Arc; -use tauri::{AppHandle, Manager, Runtime}; +use tauri::{AppHandle, Runtime}; use tokio::sync::RwLock; -use tracing::{debug, error, info}; +use tracing::{debug, info}; /// 托盘状态同步器 /// diff --git a/src-tauri/src/websocket/handler.rs b/src-tauri/src/websocket/handler.rs index 421d1e2eb..089a8f841 100644 --- a/src-tauri/src/websocket/handler.rs +++ b/src-tauri/src/websocket/handler.rs @@ -246,6 +246,21 @@ 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", + ))) + } } } diff --git a/src-tauri/src/websocket/lifecycle.rs b/src-tauri/src/websocket/lifecycle.rs index 62141d283..5fad5ec47 100644 --- a/src-tauri/src/websocket/lifecycle.rs +++ b/src-tauri/src/websocket/lifecycle.rs @@ -2,11 +2,9 @@ //! //! 提供心跳检测、优雅关闭和资源清理功能 -use super::{WsConnection, WsConnectionStatus, WsError, WsMessage}; +use super::WsMessage; use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; -use std::sync::Arc; use std::time::{Duration, Instant}; -use tokio::sync::mpsc; /// 心跳管理器 #[derive(Debug)] diff --git a/src-tauri/src/websocket/mod.rs b/src-tauri/src/websocket/mod.rs index 829695395..1db030985 100644 --- a/src-tauri/src/websocket/mod.rs +++ b/src-tauri/src/websocket/mod.rs @@ -20,7 +20,7 @@ pub use processor::MessageProcessor; pub use stream::{BackpressureController, StreamForwarder}; pub use types::{ WsApiRequest, WsApiResponse, WsConfig, WsConnection, WsConnectionStatus, WsEndpoint, WsError, - WsErrorCode, WsMessage, WsStats, WsStatsSnapshot, WsStreamChunk, WsStreamEnd, + WsErrorCode, WsFlowEvent, 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 5a9ca58c9..05991caa0 100644 --- a/src-tauri/src/websocket/types.rs +++ b/src-tauri/src/websocket/types.rs @@ -6,6 +6,11 @@ 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 { @@ -69,6 +74,12 @@ pub enum WsMessage { Ping { timestamp: i64 }, /// 心跳响应 Pong { timestamp: i64 }, + /// 订阅 Flow 事件 + SubscribeFlowEvents, + /// 取消订阅 Flow 事件 + UnsubscribeFlowEvents, + /// Flow 事件通知 + FlowEvent(WsFlowEvent), } /// WebSocket API 请求 @@ -315,3 +326,46 @@ pub struct WsStatsSnapshot { pub total_messages: u64, 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 } + } + } + } +} diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index ee03e3378..9bfe51501 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.14.10", + "version": "0.17.1", "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 new file mode 100644 index 000000000..8eaa556f9 --- /dev/null +++ b/src-tauri/tests/end_to_end_tests.rs @@ -0,0 +1,467 @@ +//! 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; + +/// 端到端测试上下文 +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, + 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-tauri/tests/filter_expression_tests.rs b/src-tauri/tests/filter_expression_tests.rs new file mode 100644 index 000000000..8b1378917 --- /dev/null +++ b/src-tauri/tests/filter_expression_tests.rs @@ -0,0 +1 @@ + diff --git a/src/App.tsx b/src/App.tsx index 8f1c0d837..d04880d2d 100644 --- a/src/App.tsx +++ b/src/App.tsx @@ -1,4 +1,4 @@ -import { useState } from "react"; +import { useState, useEffect } from "react"; import { Sidebar } from "./components/Sidebar"; import { Dashboard } from "./components/Dashboard"; import { SettingsPage } from "./components/settings"; @@ -7,6 +7,8 @@ import { ProviderPoolPage } from "./components/provider-pool"; import { RoutingManagementPage } from "./components/routing/RoutingManagementPage"; import { ConfigManagementPage } from "./components/config/ConfigManagementPage"; import { ExtensionsPage } from "./components/extensions"; +import { FlowMonitorPage } from "./pages"; +import { flowEventManager } from "./lib/flowEventManager"; type Page = | "dashboard" @@ -15,11 +17,18 @@ type Page = | "config-management" | "extensions" | "api-server" + | "flow-monitor" | "settings"; function App() { const [currentPage, setCurrentPage] = useState("dashboard"); + // 在应用启动时初始化 Flow 事件订阅 + useEffect(() => { + flowEventManager.subscribe(); + // 应用卸载时不取消订阅,因为这是全局订阅 + }, []); + const renderPage = () => { switch (currentPage) { case "dashboard": @@ -34,6 +43,8 @@ function App() { return ; case "api-server": return ; + case "flow-monitor": + return ; case "settings": return ; default: diff --git a/src/components/Sidebar.tsx b/src/components/Sidebar.tsx index 7499fb898..6fa0471d1 100644 --- a/src/components/Sidebar.tsx +++ b/src/components/Sidebar.tsx @@ -6,6 +6,7 @@ import { Route, FileCode, Puzzle, + Activity, } from "lucide-react"; import { cn } from "@/lib/utils"; @@ -16,6 +17,7 @@ type Page = | "config-management" | "extensions" | "api-server" + | "flow-monitor" | "settings"; interface SidebarProps { @@ -30,6 +32,7 @@ const navItems = [ { id: "config-management" as Page, label: "配置管理", icon: FileCode }, { id: "extensions" as Page, label: "扩展", icon: Puzzle }, { id: "api-server" as Page, label: "API Server", icon: Globe }, + { id: "flow-monitor" as Page, label: "Flow Monitor", icon: Activity }, { id: "settings" as Page, label: "设置", icon: Settings }, ]; diff --git a/src/components/flow-monitor/BatchOperations.tsx b/src/components/flow-monitor/BatchOperations.tsx new file mode 100644 index 000000000..4ea3a15a6 --- /dev/null +++ b/src/components/flow-monitor/BatchOperations.tsx @@ -0,0 +1,660 @@ +/** + * 批量操作组件 + * 实现批量选择、批量操作菜单、操作进度显示 + * **Validates: Requirements 11.1-11.7** + */ + +import React, { useState, useCallback } from "react"; +import { invoke } from "@tauri-apps/api/core"; +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 invoke(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 invoke("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 new file mode 100644 index 000000000..ac9e10429 --- /dev/null +++ b/src/components/flow-monitor/BookmarkPanel.tsx @@ -0,0 +1,727 @@ +/** + * 书签管理面板组件 + * + * 实现书签列表、书签导航功能 + * **Validates: Requirements 8.1-8.6** + */ + +import { useState, useEffect, useCallback } from "react"; +import { invoke } from "@tauri-apps/api/core"; +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([ + invoke("list_bookmarks", { group: null }), + invoke("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 invoke("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 invoke("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 invoke("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 invoke("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/ExportDialog.tsx b/src/components/flow-monitor/ExportDialog.tsx new file mode 100644 index 000000000..1ed61988e --- /dev/null +++ b/src/components/flow-monitor/ExportDialog.tsx @@ -0,0 +1,484 @@ +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 new file mode 100644 index 000000000..2bb101486 --- /dev/null +++ b/src/components/flow-monitor/FilterExpressionInput.tsx @@ -0,0 +1,552 @@ +import React, { useState, useEffect, useRef, useCallback } from "react"; +import { invoke } from "@tauri-apps/api/core"; +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 invoke("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 new file mode 100644 index 000000000..02cb5e0b4 --- /dev/null +++ b/src/components/flow-monitor/FilterHelp.tsx @@ -0,0 +1,440 @@ +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 new file mode 100644 index 000000000..ad87ea577 --- /dev/null +++ b/src/components/flow-monitor/FlowDetail.tsx @@ -0,0 +1,1268 @@ +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; + + 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 new file mode 100644 index 000000000..a275a252b --- /dev/null +++ b/src/components/flow-monitor/FlowDiffView.tsx @@ -0,0 +1,1084 @@ +/** + * Flow 差异对比视图组件 + * + * 实现差异对比视图和并排/统一视图切换 + * **Validates: Requirements 4.1-4.7** + */ + +import React, { useState, useEffect, useCallback } from "react"; +import { invoke } from "@tauri-apps/api/core"; +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 invoke("diff_flows", { + request: { + left_flow_id: leftFlowId, + right_flow_id: rightFlowId, + config, + }, + }); + + setDiffResult(result); + + // 如果没有提供 Flow,加载它们用于显示 + if (!initialLeftFlow) { + const left = await invoke("get_flow_detail", { + flowId: leftFlowId, + }); + setLeftFlow(left); + } + if (!initialRightFlow) { + const right = await invoke("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 new file mode 100644 index 000000000..67f1d1f26 --- /dev/null +++ b/src/components/flow-monitor/FlowFilters.tsx @@ -0,0 +1,717 @@ +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 new file mode 100644 index 000000000..e44115825 --- /dev/null +++ b/src/components/flow-monitor/FlowList.tsx @@ -0,0 +1,909 @@ +import React, { useState, useEffect, useCallback } from "react"; +import { + CheckCircle2, + XCircle, + Clock, + Loader2, + ChevronDown, + ChevronRight, + Star, + StarOff, + Wrench, + Brain, + RefreshCw, + Copy, + ExternalLink, + Wifi, + WifiOff, + Pause, + Play, + AlertTriangle, + Activity, + Bell, + BellOff, +} 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 { 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 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", + }; + 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 && ( + + )} + + + +
+
+ + {/* Flow 列表 */} +
+ {flows.length === 0 ? ( +
+ 暂无 Flow 记录 +
+ ) : ( +
+ {flows.map((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} + 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; + formatTime: (timestamp: string) => string; +} + +function FlowListItem({ + flow, + expanded, + selected, + thresholdWarning, + onToggleExpand, + onSelect, + onToggleStar, + onCopyId, + getStateIcon, + getProviderColor, + 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)} + + + {/* 提供商 */} + + {flow.metadata.provider} + + + {/* 模型 */} + + {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/FlowStats.tsx b/src/components/flow-monitor/FlowStats.tsx new file mode 100644 index 000000000..7c2cc5422 --- /dev/null +++ b/src/components/flow-monitor/FlowStats.tsx @@ -0,0 +1,1212 @@ +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 new file mode 100644 index 000000000..1b666f90c --- /dev/null +++ b/src/components/flow-monitor/FlowTimeline.tsx @@ -0,0 +1,435 @@ +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 new file mode 100644 index 000000000..e78f8ca04 --- /dev/null +++ b/src/components/flow-monitor/InterceptEditor.tsx @@ -0,0 +1,675 @@ +/** + * 拦截编辑器组件 + * + * 实现请求/响应编辑器和继续/取消按钮 + * **Validates: Requirements 2.2, 2.3, 2.4, 2.5** + */ + +import { useState, useEffect, useCallback } from "react"; +import { invoke } from "@tauri-apps/api/core"; +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 invoke("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 invoke("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 invoke("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" ? ( + + ) : ( +