diff --git a/eslint.config.js b/eslint.config.js index a5d2dcb99..4e6c69b1c 100644 --- a/eslint.config.js +++ b/eslint.config.js @@ -26,7 +26,48 @@ export default [ ...js.configs.recommended.rules, ...tseslint.configs.recommended.rules, ...reactHooks.configs.recommended.rules, - "react-refresh/only-export-components": ["warn", { allowConstantExport: true }], + "react-refresh/only-export-components": [ + "warn", + { + allowConstantExport: true, + allowExportNames: [ + // AddCustomProviderModal.tsx + "validateCustomProviderForm", + "isFormValid", + "hasRequiredFields", + // ApiKeyItem.tsx + "extractApiKeyDisplayInfo", + // ApiKeyList.tsx + "getApiKeyListStats", + // ApiKeyProviderSection.tsx + "verifyProviderSelectionSync", + "extractSelectionState", + // ConnectionTestButton.tsx + "getConnectionTestStatusInfo", + // DeleteProviderDialog.tsx + "canDeleteProvider", + "isSystemProvider", + // ProviderConfigForm.tsx + "getFieldsForProviderType", + "providerTypeRequiresField", + // ProviderGroup.tsx + "getGroupLabel", + "isProviderInGroup", + "getGroupOrder", + // ProviderList.tsx + "filterProviders", + "groupProviders", + "matchesSearchQuery", + // ProviderListItem.tsx + "extractListItemDisplayInfo", + "getApiKeyCount", + // ProviderSetting.tsx + "extractProviderSettingInfo", + // icons/providers/index.tsx + "iconComponents", + ], + }, + ], "@typescript-eslint/no-unused-vars": ["error", { argsIgnorePattern: "^_", varsIgnorePattern: "^_", caughtErrorsIgnorePattern: "^_" }], "@typescript-eslint/no-explicit-any": "off", }, diff --git a/jimeng-2025-12-25-7978-removebg-preview.png b/jimeng-2025-12-25-7978-removebg-preview.png deleted file mode 100644 index 9c0162490..000000000 Binary files a/jimeng-2025-12-25-7978-removebg-preview.png and /dev/null differ diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 6a06d7878..6838fedd8 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -3674,7 +3674,7 @@ dependencies = [ [[package]] name = "proxycast" -version = "0.26.0" +version = "0.27.0" dependencies = [ "anyhow", "arboard", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index e68ba5183..68e5260ac 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "proxycast" -version = "0.26.0" +version = "0.27.0" description = "AI API Proxy Desktop App" authors = ["you"] edition = "2021" diff --git a/src-tauri/src/agent/mod.rs b/src-tauri/src/agent/mod.rs index 16768426f..0259bfbf0 100644 --- a/src-tauri/src/agent/mod.rs +++ b/src-tauri/src/agent/mod.rs @@ -1,13 +1,24 @@ //! AI Agent 集成模块 //! -//! 提供基于 OpenAI 兼容 API 的 Agent 实现 +//! 使用策略模式支持多种 API 协议(OpenAI、Anthropic、Kiro、Gemini) //! 包含工具系统、流式处理和工具调用循环 +//! +//! ## 架构设计 +//! - protocols/ - 协议策略实现(策略模式) +//! - parsers/ - SSE 流解析器 +//! - native_agent - 核心 Agent 逻辑 +//! - tool_loop - 工具调用循环 +//! - tools/ - 工具实现 pub mod native_agent; +pub mod parsers; +pub mod protocols; pub mod tool_loop; pub mod tools; pub mod types; pub use native_agent::{NativeAgent, NativeAgentState}; +pub use parsers::{AnthropicSSEParser, OpenAISSEParser}; +pub use protocols::{create_protocol, AnthropicProtocol, OpenAIProtocol, Protocol}; pub use tool_loop::{ToolCallResult, ToolLoopConfig, ToolLoopEngine, ToolLoopError, ToolLoopState}; pub use types::*; diff --git a/src-tauri/src/agent/native_agent.rs b/src-tauri/src/agent/native_agent.rs index 4210e0f97..4c184857e 100644 --- a/src-tauri/src/agent/native_agent.rs +++ b/src-tauri/src/agent/native_agent.rs @@ -1,215 +1,38 @@ //! 原生 Rust Agent 实现 //! //! 支持连续对话(Conversation History)和工具调用(Tools) -//! 参考 goose 项目的 Agent 设计 +//! 使用策略模式支持多种 API 协议(OpenAI、Anthropic、Kiro、Gemini) +//! +//! ## 架构设计 +//! - protocols/ - 协议策略实现 +//! - parsers/ - SSE 流解析器 +//! - NativeAgent - 核心 Agent 逻辑 +//! - NativeAgentState - Tauri 状态管理 //! //! ## 流式处理 -//! - 实现 SSE 流解析,支持 text_delta 和 tool_calls 解析 //! - Requirements: 1.1, 1.3, 1.4 //! //! ## 工具调用循环 -//! - 实现工具调用检测、执行和结果收集 //! - Requirements: 7.1, 7.2, 7.3, 7.4, 7.5, 7.6 #![allow(dead_code)] +use crate::agent::protocols::{create_protocol, Protocol}; use crate::agent::tool_loop::{ToolCallResult, ToolLoopEngine, ToolLoopState}; +use crate::agent::tools::{create_default_registry, ToolRegistry}; use crate::agent::types::*; use crate::models::openai::{ ChatCompletionRequest, ChatCompletionResponse, ChatMessage, ContentPart as OpenAIContentPart, MessageContent as OpenAIMessageContent, }; -use futures::StreamExt; use parking_lot::RwLock; use reqwest::Client; -use serde_json::Value; use std::collections::HashMap; use std::sync::Arc; use std::time::Duration; use tokio::sync::mpsc; use tracing::{debug, error, info, warn}; -/// SSE 流解析器 -/// -/// 解析 Server-Sent Events 流,提取 text_delta 和 tool_calls -/// Requirements: 1.1, 1.3, 1.4 -#[derive(Debug, Default)] -struct SSEParser { - /// 累积的完整内容 - full_content: String, - /// 累积的工具调用 - tool_calls: Vec, - /// 当前正在构建的工具调用索引 - current_tool_indices: HashMap, -} - -/// 工具调用增量数据 -#[derive(Debug, Clone, Default)] -struct ToolCallDelta { - /// 工具调用索引 - index: usize, - /// 工具调用 ID - id: String, - /// 工具类型 - call_type: String, - /// 函数名 - function_name: String, - /// 函数参数(累积的 JSON 字符串) - function_arguments: String, -} - -impl SSEParser { - fn new() -> Self { - Self::default() - } - - /// 解析 SSE 数据行 - /// - /// 返回 (text_delta, is_done, usage) - fn parse_data(&mut self, data: &str) -> (Option, bool, Option) { - if data.trim() == "[DONE]" { - return (None, true, None); - } - - let json: Value = match serde_json::from_str(data) { - Ok(v) => v, - Err(e) => { - warn!("[SSEParser] 解析 JSON 失败: {} - data: {}", e, data); - return (None, false, None); - } - }; - - // 提取 usage 信息(如果存在) - let usage = json.get("usage").and_then(|u| { - let input = u.get("prompt_tokens").and_then(|v| v.as_u64()).unwrap_or(0) as u32; - let output = u - .get("completion_tokens") - .and_then(|v| v.as_u64()) - .unwrap_or(0) as u32; - if input > 0 || output > 0 { - Some(TokenUsage::new(input, output)) - } else { - None - } - }); - - // 检查是否有 choices - let choices = match json.get("choices").and_then(|c| c.as_array()) { - Some(c) => c, - None => return (None, false, usage), - }; - - if choices.is_empty() { - return (None, false, usage); - } - - let choice = &choices[0]; - let delta = match choice.get("delta") { - Some(d) => d, - None => return (None, false, usage), - }; - - // 检查 finish_reason - let finish_reason = choice - .get("finish_reason") - .and_then(|f| f.as_str()) - .unwrap_or(""); - let is_done = finish_reason == "stop" || finish_reason == "tool_calls"; - - // 提取文本内容 - let text_delta = delta - .get("content") - .and_then(|c| c.as_str()) - .filter(|s| !s.is_empty()) - .map(|s| { - self.full_content.push_str(s); - s.to_string() - }); - - // 提取工具调用 - if let Some(tool_calls) = delta.get("tool_calls").and_then(|tc| tc.as_array()) { - for tc in tool_calls { - self.parse_tool_call_delta(tc); - } - } - - (text_delta, is_done, usage) - } - - /// 解析工具调用增量 - fn parse_tool_call_delta(&mut self, tc: &Value) { - let index = tc.get("index").and_then(|i| i.as_u64()).unwrap_or(0) as usize; - - // 获取或创建工具调用 - let tool_call = self - .current_tool_indices - .entry(index) - .or_insert_with(|| ToolCallDelta { - index, - ..Default::default() - }); - - // 更新 ID - if let Some(id) = tc.get("id").and_then(|i| i.as_str()) { - tool_call.id = id.to_string(); - } - - // 更新类型 - if let Some(t) = tc.get("type").and_then(|t| t.as_str()) { - tool_call.call_type = t.to_string(); - } - - // 更新函数信息 - if let Some(function) = tc.get("function") { - if let Some(name) = function.get("name").and_then(|n| n.as_str()) { - tool_call.function_name = name.to_string(); - } - if let Some(args) = function.get("arguments").and_then(|a| a.as_str()) { - tool_call.function_arguments.push_str(args); - } - } - } - - /// 完成解析,返回最终的工具调用列表 - fn finalize_tool_calls(&mut self) -> Vec { - // 按索引排序并转换为 ToolCall - let mut indices: Vec<_> = self.current_tool_indices.keys().cloned().collect(); - indices.sort(); - - indices - .into_iter() - .filter_map(|idx| { - let delta = self.current_tool_indices.get(&idx)?; - if delta.id.is_empty() || delta.function_name.is_empty() { - return None; - } - Some(ToolCall { - id: delta.id.clone(), - call_type: if delta.call_type.is_empty() { - "function".to_string() - } else { - delta.call_type.clone() - }, - function: FunctionCall { - name: delta.function_name.clone(), - arguments: delta.function_arguments.clone(), - }, - }) - }) - .collect() - } - - /// 获取完整内容 - fn get_full_content(&self) -> String { - self.full_content.clone() - } - - /// 是否有工具调用 - fn has_tool_calls(&self) -> bool { - !self.current_tool_indices.is_empty() - } -} - /// 原生 Agent 实现 pub struct NativeAgent { client: Client, @@ -217,10 +40,18 @@ pub struct NativeAgent { api_key: String, sessions: Arc>>, config: AgentConfig, + /// Provider 类型,决定使用哪种协议 + provider_type: ProviderType, + /// 协议处理器 + protocol: Box, } impl NativeAgent { - pub fn new(base_url: String, api_key: String) -> Result { + pub fn new( + base_url: String, + api_key: String, + provider_type: ProviderType, + ) -> Result { let client = Client::builder() .timeout(Duration::from_secs(300)) .connect_timeout(Duration::from_secs(30)) @@ -228,12 +59,23 @@ impl NativeAgent { .build() .map_err(|e| format!("创建 HTTP 客户端失败: {}", e))?; + let protocol = create_protocol(provider_type); + + info!( + "[NativeAgent] 创建 Agent: base_url={}, provider={:?}, protocol_endpoint={}", + base_url, + provider_type, + protocol.endpoint() + ); + Ok(Self { client, base_url, api_key, sessions: Arc::new(RwLock::new(HashMap::new())), config: AgentConfig::default(), + provider_type, + protocol, }) } @@ -247,112 +89,7 @@ impl NativeAgent { self } - /// 将 AgentMessage 转换为 OpenAI ChatMessage - fn convert_to_chat_message(&self, msg: &AgentMessage) -> ChatMessage { - let content = match &msg.content { - MessageContent::Text(text) => Some(OpenAIMessageContent::Text(text.clone())), - MessageContent::Parts(parts) => { - let openai_parts: Vec = parts - .iter() - .map(|p| match p { - ContentPart::Text { text } => { - OpenAIContentPart::Text { text: text.clone() } - } - ContentPart::ImageUrl { image_url } => OpenAIContentPart::ImageUrl { - image_url: crate::models::openai::ImageUrl { - url: image_url.url.clone(), - detail: image_url.detail.clone(), - }, - }, - }) - .collect(); - Some(OpenAIMessageContent::Parts(openai_parts)) - } - }; - - ChatMessage { - role: msg.role.clone(), - content, - tool_calls: msg.tool_calls.as_ref().map(|calls| { - calls - .iter() - .map(|tc| crate::models::openai::ToolCall { - id: tc.id.clone(), - call_type: tc.call_type.clone(), - function: crate::models::openai::FunctionCall { - name: tc.function.name.clone(), - arguments: tc.function.arguments.clone(), - }, - }) - .collect() - }), - tool_call_id: msg.tool_call_id.clone(), - } - } - - /// 构建完整的消息列表(包含历史) - fn build_messages_with_history( - &self, - session: &AgentSession, - user_message: &str, - images: Option<&[ImageData]>, - ) -> Vec { - let mut messages = Vec::new(); - - // 1. 添加系统提示词 - let system_prompt = session - .system_prompt - .as_ref() - .or(self.config.system_prompt.as_ref()); - if let Some(prompt) = system_prompt { - messages.push(ChatMessage { - role: "system".to_string(), - content: Some(OpenAIMessageContent::Text(prompt.clone())), - tool_calls: None, - tool_call_id: None, - }); - } - - // 2. 添加历史消息 - for msg in &session.messages { - messages.push(self.convert_to_chat_message(msg)); - } - - // 3. 添加当前用户消息 - let user_msg = if let Some(imgs) = images { - let mut parts = vec![OpenAIContentPart::Text { - text: user_message.to_string(), - }]; - - for img in imgs { - parts.push(OpenAIContentPart::ImageUrl { - image_url: crate::models::openai::ImageUrl { - url: format!("data:{};base64,{}", img.media_type, img.data), - detail: None, - }, - }); - } - - ChatMessage { - role: "user".to_string(), - content: Some(OpenAIMessageContent::Parts(parts)), - tool_calls: None, - tool_call_id: None, - } - } else { - ChatMessage { - role: "user".to_string(), - content: Some(OpenAIMessageContent::Text(user_message.to_string())), - tool_calls: None, - tool_call_id: None, - } - }; - - messages.push(user_msg); - messages - } - - /// 发送聊天请求(支持连续对话) + /// 发送聊天请求(非流式,用于简单场景) pub async fn chat(&self, request: NativeChatRequest) -> Result { let model = request.model.unwrap_or_else(|| self.config.model.clone()); let session_id = request.session_id.clone(); @@ -363,42 +100,19 @@ impl NativeAgent { model, session_id, has_images ); - // 获取或创建会话 + // 获取会话 let session = if let Some(sid) = &session_id { self.sessions.read().get(sid).cloned() } else { None }; - let messages = if let Some(ref sess) = session { - // 使用会话历史构建消息 - self.build_messages_with_history(sess, &request.message, request.images.as_deref()) - } else { - // 无会话,单次对话 - self.build_single_messages(&request.message, request.images.as_deref()) - }; - - // 打印消息结构用于调试 - for (i, msg) in messages.iter().enumerate() { - let content_type = match &msg.content { - Some(OpenAIMessageContent::Text(_)) => "text", - Some(OpenAIMessageContent::Parts(parts)) => { - let has_image = parts - .iter() - .any(|p| matches!(p, OpenAIContentPart::ImageUrl { .. })); - if has_image { - "parts_with_image" - } else { - "parts_text_only" - } - } - None => "none", - }; - debug!( - "[NativeAgent] 消息[{}]: role={}, content_type={}", - i, msg.role, content_type - ); - } + // 构建消息 + let messages = self.build_openai_messages( + session.as_ref(), + &request.message, + request.images.as_deref(), + ); let chat_request = ChatCompletionRequest { model: model.clone(), @@ -407,7 +121,7 @@ impl NativeAgent { temperature: self.config.temperature, max_tokens: self.config.max_tokens, top_p: None, - tools: None, // TODO: 添加工具支持 + tools: None, tool_choice: None, reasoning_effort: None, }; @@ -480,28 +194,280 @@ impl NativeAgent { }) } - /// 构建单次对话消息(无历史) - fn build_single_messages( + /// 流式聊天(使用协议策略模式) + /// + /// Requirements: 1.1, 1.3, 1.4 + pub async fn chat_stream( &self, + request: NativeChatRequest, + tools: Option<&[crate::models::openai::Tool]>, + tx: mpsc::Sender, + ) -> Result { + let model = request + .model + .clone() + .unwrap_or_else(|| self.config.model.clone()); + let session_id = request.session_id.clone(); + + info!( + "[NativeAgent] 发送流式聊天请求: model={}, session={:?}, provider={:?}, tools_count={}", + model, + session_id, + self.provider_type, + tools.map(|t| t.len()).unwrap_or(0) + ); + + // 获取会话 + let session = if let Some(sid) = &session_id { + self.sessions.read().get(sid).cloned() + } else { + None + }; + + // 获取会话历史和配置 + let history: Vec = session + .as_ref() + .map(|s| s.messages.clone()) + .unwrap_or_default(); + + let config = if let Some(ref sess) = session { + let mut cfg = self.config.clone(); + if sess.system_prompt.is_some() { + cfg.system_prompt = sess.system_prompt.clone(); + } + cfg + } else { + self.config.clone() + }; + + // 使用协议策略发送请求 + let result = self + .protocol + .chat_stream( + &self.client, + &self.base_url, + &self.api_key, + &history, + &request.message, + request.images.as_deref(), + &model, + &config, + tools, + tx, + ) + .await?; + + // 更新会话历史 + if let Some(sid) = &session_id { + self.add_message_to_session( + sid, + "user", + MessageContent::Text(request.message.clone()), + request.images.as_deref(), + ); + self.add_assistant_message_to_session( + sid, + MessageContent::Text(result.content.clone()), + result.tool_calls.clone(), + ); + } + + Ok(result) + } + + /// 流式聊天(支持工具调用循环) + /// + /// Requirements: 7.1, 7.2, 7.3, 7.4, 7.5, 7.6 + pub async fn chat_stream_with_tools( + &self, + request: NativeChatRequest, + tx: mpsc::Sender, + tool_loop_engine: &ToolLoopEngine, + ) -> Result { + let session_id = request.session_id.clone(); + let mut state = ToolLoopState::new(); + + // 获取工具定义 + let tools = tool_loop_engine.registry().list_definitions_api(); + let tools_ref = if tools.is_empty() { + None + } else { + Some(tools.as_slice()) + }; + + // 首次请求 + let mut current_result = self + .chat_stream(request.clone(), tools_ref, tx.clone()) + .await?; + + // 工具调用循环 + // Requirements: 7.3 - THE Tool_Loop SHALL continue until the Agent produces a final response without tool_calls + while tool_loop_engine.should_continue(¤t_result, state.iteration) { + state.increment_iteration(); + + let tool_calls = current_result.tool_calls.as_ref().unwrap(); + state.add_tool_calls(tool_calls.len()); + + info!( + "[NativeAgent] 工具循环迭代 {}: 执行 {} 个工具调用", + state.iteration, + tool_calls.len() + ); + + // 执行所有工具调用 + let tool_results = tool_loop_engine + .execute_all_tool_calls(tool_calls, Some(&tx)) + .await; + + // 将工具结果添加到会话 + if let Some(sid) = &session_id { + for result in &tool_results { + self.add_tool_result_to_session(sid, result); + } + } + + // 继续对话 + let continue_request = NativeChatRequest { + session_id: session_id.clone(), + message: String::new(), + model: request.model.clone(), + images: None, + stream: true, + }; + + current_result = self + .chat_stream_continue(continue_request, tools_ref, tx.clone()) + .await?; + } + + // 检查是否因为达到最大迭代次数而停止 + if state.iteration >= tool_loop_engine.max_iterations() && current_result.has_tool_calls() { + warn!( + "[NativeAgent] 达到最大迭代次数 {},强制停止工具循环", + tool_loop_engine.max_iterations() + ); + let _ = tx + .send(StreamEvent::Error { + message: format!( + "达到最大工具调用迭代次数限制 ({})", + tool_loop_engine.max_iterations() + ), + }) + .await; + } + + state.mark_completed(current_result.content.clone()); + + info!( + "[NativeAgent] 工具循环完成: {} 次迭代, {} 个工具调用", + state.iteration, state.total_tool_calls + ); + + // 发送 FinalDone 事件,通知前端整个对话(包括工具循环)已完成 + let _ = tx + .send(StreamEvent::FinalDone { + usage: current_result.usage.clone(), + }) + .await; + + Ok(current_result) + } + + /// 继续流式对话(使用会话历史) + async fn chat_stream_continue( + &self, + request: NativeChatRequest, + tools: Option<&[crate::models::openai::Tool]>, + tx: mpsc::Sender, + ) -> Result { + let model = request.model.unwrap_or_else(|| self.config.model.clone()); + let session_id = request.session_id.as_ref().ok_or("需要 session_id")?; + + debug!( + "[NativeAgent] 继续流式对话: model={}, session={}, tools_count={}", + model, + session_id, + tools.map(|t| t.len()).unwrap_or(0) + ); + + // 获取会话 + let session = self + .sessions + .read() + .get(session_id) + .cloned() + .ok_or_else(|| format!("会话不存在: {}", session_id))?; + + // 获取配置 + let config = { + let mut cfg = self.config.clone(); + if session.system_prompt.is_some() { + cfg.system_prompt = session.system_prompt.clone(); + } + cfg + }; + + // 使用协议策略继续对话 + let result = self + .protocol + .chat_stream_continue( + &self.client, + &self.base_url, + &self.api_key, + &session.messages, + &model, + &config, + tools, + tx, + ) + .await?; + + // 更新会话历史 + self.add_assistant_message_to_session( + session_id, + MessageContent::Text(result.content.clone()), + result.tool_calls.clone(), + ); + + Ok(result) + } + + // ==================== 会话管理方法 ==================== + + /// 构建 OpenAI 格式消息(用于非流式请求) + fn build_openai_messages( + &self, + session: Option<&AgentSession>, user_message: &str, images: Option<&[ImageData]>, ) -> Vec { let mut messages = Vec::new(); - if let Some(system_prompt) = &self.config.system_prompt { + // 系统提示词 + let system_prompt = session + .and_then(|s| s.system_prompt.as_ref()) + .or(self.config.system_prompt.as_ref()); + if let Some(prompt) = system_prompt { messages.push(ChatMessage { role: "system".to_string(), - content: Some(OpenAIMessageContent::Text(system_prompt.clone())), + content: Some(OpenAIMessageContent::Text(prompt.clone())), tool_calls: None, tool_call_id: None, }); } + // 历史消息 + if let Some(sess) = session { + for msg in &sess.messages { + messages.push(self.convert_to_chat_message(msg)); + } + } + + // 用户消息 let user_msg = if let Some(imgs) = images { let mut parts = vec![OpenAIContentPart::Text { text: user_message.to_string(), }]; - for img in imgs { parts.push(OpenAIContentPart::ImageUrl { image_url: crate::models::openai::ImageUrl { @@ -510,7 +476,6 @@ impl NativeAgent { }, }); } - ChatMessage { role: "user".to_string(), content: Some(OpenAIMessageContent::Parts(parts)), @@ -530,6 +495,49 @@ impl NativeAgent { messages } + /// 将 AgentMessage 转换为 OpenAI ChatMessage + fn convert_to_chat_message(&self, msg: &AgentMessage) -> ChatMessage { + let content = match &msg.content { + MessageContent::Text(text) => Some(OpenAIMessageContent::Text(text.clone())), + MessageContent::Parts(parts) => { + let openai_parts: Vec = parts + .iter() + .map(|p| match p { + ContentPart::Text { text } => { + OpenAIContentPart::Text { text: text.clone() } + } + ContentPart::ImageUrl { image_url } => OpenAIContentPart::ImageUrl { + image_url: crate::models::openai::ImageUrl { + url: image_url.url.clone(), + detail: image_url.detail.clone(), + }, + }, + }) + .collect(); + Some(OpenAIMessageContent::Parts(openai_parts)) + } + }; + + ChatMessage { + role: msg.role.clone(), + content, + tool_calls: msg.tool_calls.as_ref().map(|calls| { + calls + .iter() + .map(|tc| crate::models::openai::ToolCall { + id: tc.id.clone(), + call_type: tc.call_type.clone(), + function: crate::models::openai::FunctionCall { + name: tc.function.name.clone(), + arguments: tc.function.arguments.clone(), + }, + }) + .collect() + }), + tool_call_id: msg.tool_call_id.clone(), + } + } + /// 添加消息到会话 fn add_message_to_session( &self, @@ -541,7 +549,6 @@ impl NativeAgent { let mut sessions = self.sessions.write(); if let Some(session) = sessions.get_mut(session_id) { let final_content = if let Some(imgs) = images { - // 如果有图片,转换为 Parts let mut parts = vec![ContentPart::Text { text: content.as_text(), }]; @@ -569,196 +576,6 @@ impl NativeAgent { } } - /// 流式聊天(支持连续对话) - /// - /// 实现 SSE 流解析,支持 text_delta 和 tool_calls 解析 - /// Requirements: 1.1, 1.3, 1.4 - pub async fn chat_stream( - &self, - request: NativeChatRequest, - tx: mpsc::Sender, - ) -> Result { - let model = request.model.unwrap_or_else(|| self.config.model.clone()); - let session_id = request.session_id.clone(); - - info!( - "[NativeAgent] 发送流式聊天请求: model={}, session={:?}", - model, session_id - ); - - // 获取会话 - let session = if let Some(sid) = &session_id { - self.sessions.read().get(sid).cloned() - } else { - None - }; - - let messages = if let Some(ref sess) = session { - self.build_messages_with_history(sess, &request.message, request.images.as_deref()) - } else { - self.build_single_messages(&request.message, request.images.as_deref()) - }; - - let chat_request = ChatCompletionRequest { - model: model.clone(), - messages, - stream: true, - temperature: self.config.temperature, - max_tokens: self.config.max_tokens, - top_p: None, - tools: None, - tool_choice: None, - reasoning_effort: None, - }; - - let url = format!("{}/v1/chat/completions", self.base_url); - - let response = self - .client - .post(&url) - .header("Authorization", format!("Bearer {}", self.api_key)) - .header("Content-Type", "application/json") - .json(&chat_request) - .send() - .await - .map_err(|e| format!("请求失败: {}", e))?; - - let status = response.status(); - if !status.is_success() { - let body = response.text().await.unwrap_or_default(); - error!("[NativeAgent] 流式请求失败: {} - {}", status, body); - let _ = tx - .send(StreamEvent::Error { - message: format!("API 错误 ({}): {}", status, body), - }) - .await; - return Err(format!("API 错误: {}", status)); - } - - let mut stream = response.bytes_stream(); - let mut buffer = String::new(); - let mut parser = SSEParser::new(); - let mut final_usage: Option = None; - - while let Some(chunk) = stream.next().await { - match chunk { - Ok(bytes) => { - let text = String::from_utf8_lossy(&bytes); - buffer.push_str(&text); - - // 处理完整的 SSE 事件(以 \n\n 分隔) - while let Some(pos) = buffer.find("\n\n") { - let event = buffer[..pos].to_string(); - buffer = buffer[pos + 2..].to_string(); - - for line in event.lines() { - if let Some(data) = line.strip_prefix("data: ") { - debug!("[NativeAgent] SSE data: {}", data); - let (text_delta, is_done, usage) = parser.parse_data(data); - debug!("[NativeAgent] 解析结果: text_delta={:?}, is_done={}, usage={:?}", - text_delta, is_done, usage); - - // 更新 usage - if usage.is_some() { - final_usage = usage; - } - - // 发送文本增量 - if let Some(text) = text_delta { - debug!("[NativeAgent] 发送 TextDelta: {}", text); - let _ = tx.send(StreamEvent::TextDelta { text }).await; - } - - // 检查是否完成 - if is_done { - // 获取最终结果 - let full_content = parser.get_full_content(); - let tool_calls = if parser.has_tool_calls() { - Some(parser.finalize_tool_calls()) - } else { - None - }; - - // 更新会话历史 - if let Some(sid) = &session_id { - self.add_message_to_session( - sid, - "user", - MessageContent::Text(request.message.clone()), - request.images.as_deref(), - ); - self.add_assistant_message_to_session( - sid, - MessageContent::Text(full_content.clone()), - tool_calls.clone(), - ); - } - - // 发送完成事件 - let _ = tx - .send(StreamEvent::Done { - usage: final_usage.clone(), - }) - .await; - - return Ok(StreamResult { - content: full_content, - tool_calls, - usage: final_usage, - }); - } - } - } - } - } - Err(e) => { - error!("[NativeAgent] 流读取错误: {}", e); - let _ = tx - .send(StreamEvent::Error { - message: format!("流读取错误: {}", e), - }) - .await; - return Err(format!("流读取错误: {}", e)); - } - } - } - - // 流正常结束但没有收到 [DONE] - let full_content = parser.get_full_content(); - let tool_calls = if parser.has_tool_calls() { - Some(parser.finalize_tool_calls()) - } else { - None - }; - - // 更新会话历史 - if let Some(sid) = &session_id { - self.add_message_to_session( - sid, - "user", - MessageContent::Text(request.message.clone()), - request.images.as_deref(), - ); - self.add_assistant_message_to_session( - sid, - MessageContent::Text(full_content.clone()), - tool_calls.clone(), - ); - } - - let _ = tx - .send(StreamEvent::Done { - usage: final_usage.clone(), - }) - .await; - - Ok(StreamResult { - content: full_content, - tool_calls, - usage: final_usage, - }) - } - /// 添加 assistant 消息到会话(支持工具调用) fn add_assistant_message_to_session( &self, @@ -780,8 +597,6 @@ impl NativeAgent { } /// 添加工具结果消息到会话 - /// - /// Requirements: 7.2 - THE Tool_Loop SHALL send tool results back to the Agent as tool role messages fn add_tool_result_to_session(&self, session_id: &str, tool_result: &ToolCallResult) { let mut sessions = self.sessions.write(); if let Some(session) = sessions.get_mut(session_id) { @@ -790,273 +605,7 @@ impl NativeAgent { } } - /// 流式聊天(支持工具调用循环) - /// - /// 实现完整的工具调用循环: - /// 1. 发送请求到 LLM - /// 2. 如果响应包含工具调用,执行工具 - /// 3. 将工具结果发送回 LLM - /// 4. 重复直到 LLM 产生最终响应或达到最大迭代次数 - /// - /// Requirements: 7.1, 7.2, 7.3, 7.4, 7.5, 7.6 - pub async fn chat_stream_with_tools( - &self, - request: NativeChatRequest, - tx: mpsc::Sender, - tool_loop_engine: &ToolLoopEngine, - ) -> Result { - let session_id = request.session_id.clone(); - let mut state = ToolLoopState::new(); - - // 首次请求 - let mut current_result = self.chat_stream(request.clone(), tx.clone()).await?; - - // 工具调用循环 - // Requirements: 7.3 - THE Tool_Loop SHALL continue until the Agent produces a final response without tool_calls - while tool_loop_engine.should_continue(¤t_result, state.iteration) { - state.increment_iteration(); - - let tool_calls = current_result.tool_calls.as_ref().unwrap(); - state.add_tool_calls(tool_calls.len()); - - info!( - "[NativeAgent] 工具循环迭代 {}: 执行 {} 个工具调用", - state.iteration, - tool_calls.len() - ); - - // 执行所有工具调用 - // Requirements: 7.1 - THE Tool_Loop SHALL execute each tool and collect results - // Requirements: 7.6 - WHILE the Tool_Loop is executing, THE Frontend SHALL display the current tool - let tool_results = tool_loop_engine - .execute_all_tool_calls(tool_calls, Some(&tx)) - .await; - - // 将工具结果添加到会话 - // Requirements: 7.2 - THE Tool_Loop SHALL send tool results back to the Agent as tool role messages - if let Some(sid) = &session_id { - for result in &tool_results { - self.add_tool_result_to_session(sid, result); - } - } - - // 构建继续对话的请求 - let continue_request = NativeChatRequest { - session_id: session_id.clone(), - message: String::new(), // 空消息,因为我们使用会话历史 - model: request.model.clone(), - images: None, - stream: true, - }; - - // 继续对话 - current_result = self - .chat_stream_continue(continue_request, tx.clone()) - .await?; - } - - // 检查是否因为达到最大迭代次数而停止 - // Requirements: 7.5 - THE Tool_Loop SHALL enforce a maximum iteration limit - if state.iteration >= tool_loop_engine.max_iterations() && current_result.has_tool_calls() { - warn!( - "[NativeAgent] 达到最大迭代次数 {},强制停止工具循环", - tool_loop_engine.max_iterations() - ); - let _ = tx - .send(StreamEvent::Error { - message: format!( - "达到最大工具调用迭代次数限制 ({})", - tool_loop_engine.max_iterations() - ), - }) - .await; - } - - state.mark_completed(current_result.content.clone()); - - info!( - "[NativeAgent] 工具循环完成: {} 次迭代, {} 个工具调用", - state.iteration, state.total_tool_calls - ); - - Ok(current_result) - } - - /// 继续流式对话(使用会话历史) - /// - /// 用于工具调用循环中继续对话 - async fn chat_stream_continue( - &self, - request: NativeChatRequest, - tx: mpsc::Sender, - ) -> Result { - let model = request.model.unwrap_or_else(|| self.config.model.clone()); - let session_id = request.session_id.as_ref().ok_or("需要 session_id")?; - - debug!( - "[NativeAgent] 继续流式对话: model={}, session={}", - model, session_id - ); - - // 获取会话 - let session = self - .sessions - .read() - .get(session_id) - .cloned() - .ok_or_else(|| format!("会话不存在: {}", session_id))?; - - // 构建消息(使用会话历史,不添加新的用户消息) - let messages = self.build_messages_from_session(&session); - - let chat_request = ChatCompletionRequest { - model: model.clone(), - messages, - stream: true, - temperature: self.config.temperature, - max_tokens: self.config.max_tokens, - top_p: None, - tools: None, // TODO: 添加工具定义 - tool_choice: None, - reasoning_effort: None, - }; - - let url = format!("{}/v1/chat/completions", self.base_url); - - let response = self - .client - .post(&url) - .header("Authorization", format!("Bearer {}", self.api_key)) - .header("Content-Type", "application/json") - .json(&chat_request) - .send() - .await - .map_err(|e| format!("请求失败: {}", e))?; - - let status = response.status(); - if !status.is_success() { - let body = response.text().await.unwrap_or_default(); - error!("[NativeAgent] 流式请求失败: {} - {}", status, body); - let _ = tx - .send(StreamEvent::Error { - message: format!("API 错误 ({}): {}", status, body), - }) - .await; - return Err(format!("API 错误: {}", status)); - } - - let mut stream = response.bytes_stream(); - let mut buffer = String::new(); - let mut parser = SSEParser::new(); - let mut final_usage: Option = None; - - while let Some(chunk) = stream.next().await { - match chunk { - Ok(bytes) => { - let text = String::from_utf8_lossy(&bytes); - buffer.push_str(&text); - - while let Some(pos) = buffer.find("\n\n") { - let event = buffer[..pos].to_string(); - buffer = buffer[pos + 2..].to_string(); - - for line in event.lines() { - if let Some(data) = line.strip_prefix("data: ") { - let (text_delta, is_done, usage) = parser.parse_data(data); - - if usage.is_some() { - final_usage = usage; - } - - if let Some(text) = text_delta { - let _ = tx.send(StreamEvent::TextDelta { text }).await; - } - - if is_done { - let full_content = parser.get_full_content(); - let tool_calls = if parser.has_tool_calls() { - Some(parser.finalize_tool_calls()) - } else { - None - }; - - // 更新会话历史 - self.add_assistant_message_to_session( - session_id, - MessageContent::Text(full_content.clone()), - tool_calls.clone(), - ); - - // 不发送 Done 事件,因为工具循环可能还会继续 - - return Ok(StreamResult { - content: full_content, - tool_calls, - usage: final_usage, - }); - } - } - } - } - } - Err(e) => { - error!("[NativeAgent] 流读取错误: {}", e); - let _ = tx - .send(StreamEvent::Error { - message: format!("流读取错误: {}", e), - }) - .await; - return Err(format!("流读取错误: {}", e)); - } - } - } - - // 流正常结束 - let full_content = parser.get_full_content(); - let tool_calls = if parser.has_tool_calls() { - Some(parser.finalize_tool_calls()) - } else { - None - }; - - self.add_assistant_message_to_session( - session_id, - MessageContent::Text(full_content.clone()), - tool_calls.clone(), - ); - - Ok(StreamResult { - content: full_content, - tool_calls, - usage: final_usage, - }) - } - - /// 从会话构建消息列表(不添加新的用户消息) - fn build_messages_from_session(&self, session: &AgentSession) -> Vec { - let mut messages = Vec::new(); - - // 添加系统提示词 - let system_prompt = session - .system_prompt - .as_ref() - .or(self.config.system_prompt.as_ref()); - if let Some(prompt) = system_prompt { - messages.push(ChatMessage { - role: "system".to_string(), - content: Some(OpenAIMessageContent::Text(prompt.clone())), - tool_calls: None, - tool_call_id: None, - }); - } - - // 添加所有历史消息 - for msg in &session.messages { - messages.push(self.convert_to_chat_message(msg)); - } - - messages - } + // ==================== 公开会话管理 API ==================== pub fn create_session(&self, model: Option, system_prompt: Option) -> String { let session_id = uuid::Uuid::new_v4().to_string(); @@ -1107,6 +656,8 @@ impl NativeAgent { } } +// ==================== Tauri 状态管理 ==================== + /// Tauri 状态:原生 Agent 管理器 #[derive(Clone, Default)] pub struct NativeAgentState { @@ -1120,8 +671,13 @@ impl NativeAgentState { } } - pub fn init(&self, base_url: String, api_key: String) -> Result<(), String> { - let agent = NativeAgent::new(base_url, api_key)?; + pub fn init( + &self, + base_url: String, + api_key: String, + provider_type: ProviderType, + ) -> Result<(), String> { + let agent = NativeAgent::new(base_url, api_key, provider_type)?; *self.agent.write() = Some(agent); Ok(()) } @@ -1134,32 +690,40 @@ impl NativeAgentState { *self.agent.write() = None; } + /// 获取工具注册表 + pub fn get_tool_registry(&self) -> Result, String> { + let base_dir = dirs::home_dir().ok_or_else(|| "无法获取用户 home 目录".to_string())?; + let registry = create_default_registry(base_dir); + Ok(Arc::new(registry)) + } + + /// 创建临时 Agent 用于异步操作 + fn create_temp_agent(&self) -> Result { + let guard = self.agent.read(); + let agent = guard.as_ref().ok_or_else(|| "Agent 未初始化".to_string())?; + + let client = Client::builder() + .timeout(Duration::from_secs(300)) + .connect_timeout(Duration::from_secs(30)) + .no_proxy() + .build() + .map_err(|e| format!("创建 HTTP 客户端失败: {}", e))?; + + let protocol = create_protocol(agent.provider_type); + + Ok(NativeAgent { + client, + base_url: agent.base_url.clone(), + api_key: agent.api_key.clone(), + sessions: agent.sessions.clone(), + config: agent.config.clone(), + provider_type: agent.provider_type, + protocol, + }) + } + pub async fn chat(&self, request: NativeChatRequest) -> Result { - let (base_url, api_key, config, sessions) = { - let guard = self.agent.read(); - let agent = guard.as_ref().ok_or_else(|| "Agent 未初始化".to_string())?; - ( - agent.base_url.clone(), - agent.api_key.clone(), - agent.config.clone(), - agent.sessions.clone(), - ) - }; - - // 创建临时 Agent,共享 sessions - let temp_agent = NativeAgent { - client: Client::builder() - .timeout(Duration::from_secs(300)) - .connect_timeout(Duration::from_secs(30)) - .no_proxy() - .build() - .map_err(|e| format!("创建 HTTP 客户端失败: {}", e))?, - base_url, - api_key, - sessions, - config, - }; - + let temp_agent = self.create_temp_agent()?; temp_agent.chat(request).await } @@ -1168,66 +732,17 @@ impl NativeAgentState { request: NativeChatRequest, tx: mpsc::Sender, ) -> Result { - let (base_url, api_key, config, sessions) = { - let guard = self.agent.read(); - let agent = guard.as_ref().ok_or_else(|| "Agent 未初始化".to_string())?; - ( - agent.base_url.clone(), - agent.api_key.clone(), - agent.config.clone(), - agent.sessions.clone(), - ) - }; - - let temp_agent = NativeAgent { - client: Client::builder() - .timeout(Duration::from_secs(300)) - .connect_timeout(Duration::from_secs(30)) - .no_proxy() - .build() - .map_err(|e| format!("创建 HTTP 客户端失败: {}", e))?, - base_url, - api_key, - sessions, - config, - }; - - temp_agent.chat_stream(request, tx).await + let temp_agent = self.create_temp_agent()?; + temp_agent.chat_stream(request, None, tx).await } - /// 流式聊天(支持工具调用循环) - /// - /// Requirements: 7.1, 7.2, 7.3, 7.4, 7.5, 7.6 pub async fn chat_stream_with_tools( &self, request: NativeChatRequest, tx: mpsc::Sender, tool_loop_engine: &ToolLoopEngine, ) -> Result { - let (base_url, api_key, config, sessions) = { - let guard = self.agent.read(); - let agent = guard.as_ref().ok_or_else(|| "Agent 未初始化".to_string())?; - ( - agent.base_url.clone(), - agent.api_key.clone(), - agent.config.clone(), - agent.sessions.clone(), - ) - }; - - let temp_agent = NativeAgent { - client: Client::builder() - .timeout(Duration::from_secs(300)) - .connect_timeout(Duration::from_secs(30)) - .no_proxy() - .build() - .map_err(|e| format!("创建 HTTP 客户端失败: {}", e))?, - base_url, - api_key, - sessions, - config, - }; - + let temp_agent = self.create_temp_agent()?; temp_agent .chat_stream_with_tools(request, tx, tool_loop_engine) .await @@ -1287,12 +802,12 @@ impl NativeAgentState { #[cfg(test)] mod tests { use super::*; + use crate::agent::parsers::OpenAISSEParser; #[test] fn test_sse_parser_text_delta() { - let mut parser = SSEParser::new(); + let mut parser = OpenAISSEParser::new(); - // 模拟 SSE 数据 let data1 = r#"{"choices":[{"delta":{"content":"Hello"}}]}"#; let data2 = r#"{"choices":[{"delta":{"content":" World"}}]}"#; let data3 = r#"{"choices":[{"delta":{},"finish_reason":"stop"}]}"#; @@ -1314,9 +829,8 @@ mod tests { #[test] fn test_sse_parser_tool_calls() { - let mut parser = SSEParser::new(); + let mut parser = OpenAISSEParser::new(); - // 模拟工具调用的 SSE 数据 let data1 = r#"{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"call_123","type":"function","function":{"name":"bash"}}]}}]}"#; let data2 = r#"{"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"{\"command\":"}}]}}]}"#; let data3 = r#"{"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"\"ls -la\"}"}}]}}]}"#; @@ -1336,255 +850,4 @@ mod tests { assert_eq!(tool_calls[0].function.name, "bash"); assert_eq!(tool_calls[0].function.arguments, r#"{"command":"ls -la"}"#); } - - #[test] - fn test_sse_parser_usage() { - let mut parser = SSEParser::new(); - - let data = r#"{"choices":[{"delta":{"content":"Hi"}}],"usage":{"prompt_tokens":10,"completion_tokens":5}}"#; - let (text, _, usage) = parser.parse_data(data); - - assert_eq!(text, Some("Hi".to_string())); - assert!(usage.is_some()); - let usage = usage.unwrap(); - assert_eq!(usage.input_tokens, 10); - assert_eq!(usage.output_tokens, 5); - } - - #[test] - fn test_sse_parser_done_signal() { - let mut parser = SSEParser::new(); - - let (_, done, _) = parser.parse_data("[DONE]"); - assert!(done); - } - - #[test] - fn test_sse_parser_invalid_json() { - let mut parser = SSEParser::new(); - - let (text, done, usage) = parser.parse_data("invalid json"); - assert!(text.is_none()); - assert!(!done); - assert!(usage.is_none()); - } - - #[test] - fn test_sse_parser_multiple_tool_calls() { - let mut parser = SSEParser::new(); - - // 两个工具调用 - let data1 = r#"{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"call_1","type":"function","function":{"name":"bash","arguments":"{}"}}]}}]}"#; - let data2 = r#"{"choices":[{"delta":{"tool_calls":[{"index":1,"id":"call_2","type":"function","function":{"name":"read_file","arguments":"{}"}}]}}]}"#; - - parser.parse_data(data1); - parser.parse_data(data2); - - let tool_calls = parser.finalize_tool_calls(); - assert_eq!(tool_calls.len(), 2); - assert_eq!(tool_calls[0].function.name, "bash"); - assert_eq!(tool_calls[1].function.name, "read_file"); - } -} - -#[cfg(test)] -mod proptests { - use super::*; - use proptest::prelude::*; - - /// 生成有效的文本内容(不包含特殊字符) - fn arb_text_content() -> impl Strategy { - "[a-zA-Z0-9 ,.!?]{1,100}".prop_map(|s| s) - } - - /// 生成文本片段列表 - fn arb_text_chunks() -> impl Strategy> { - prop::collection::vec(arb_text_content(), 1..10) - } - - /// 生成有效的工具名称 - fn arb_tool_name() -> impl Strategy { - prop_oneof![ - Just("bash".to_string()), - Just("read_file".to_string()), - Just("write_file".to_string()), - Just("edit_file".to_string()), - ] - } - - /// 生成有效的工具调用 ID - fn arb_tool_id() -> impl Strategy { - "call_[a-zA-Z0-9]{8}".prop_map(|s| s) - } - - proptest! { - #![proptest_config(ProptestConfig::with_cases(100))] - - /// **Feature: agent-tool-calling, Property 1: 流式事件完整性** - /// **Validates: Requirements 1.1, 1.3** - /// - /// *For any* Agent 响应流,流式处理器发送的所有 text_delta 事件的文本拼接后, - /// 应该等于最终的完整响应内容。 - #[test] - fn prop_streaming_text_completeness(chunks in arb_text_chunks()) { - let mut parser = SSEParser::new(); - let mut collected_deltas = String::new(); - - // 模拟流式处理 - for chunk in &chunks { - // 转义 JSON 特殊字符 - let escaped = chunk.replace('\\', "\\\\").replace('"', "\\\""); - let data = format!(r#"{{"choices":[{{"delta":{{"content":"{}"}}}}]}}"#, escaped); - let (text_delta, _, _) = parser.parse_data(&data); - - if let Some(text) = text_delta { - collected_deltas.push_str(&text); - } - } - - // 验证:收集的 text_delta 拼接后等于 parser 的完整内容 - prop_assert_eq!( - collected_deltas, - parser.get_full_content(), - "收集的 text_delta 应该等于完整内容" - ); - - // 验证:完整内容等于原始 chunks 拼接 - let expected = chunks.join(""); - prop_assert_eq!( - parser.get_full_content(), - expected, - "完整内容应该等于原始 chunks 拼接" - ); - } - - /// **Feature: agent-tool-calling, Property 1: 流式事件完整性 - 工具调用** - /// **Validates: Requirements 1.1, 1.3** - /// - /// *For any* 包含工具调用的响应流,工具调用信息应该被正确解析和累积。 - #[test] - fn prop_streaming_tool_calls_completeness( - tool_name in arb_tool_name(), - tool_id in arb_tool_id(), - arg_key in "[a-z]{3,10}", - arg_value in "[a-zA-Z0-9]{1,20}" - ) { - let mut parser = SSEParser::new(); - - // 使用 serde_json 构建正确的 JSON,避免手动转义问题 - let args_json = serde_json::json!({arg_key.clone(): arg_value.clone()}).to_string(); - - // 第一个 chunk: 工具 ID 和名称 - let data1 = serde_json::json!({ - "choices": [{ - "delta": { - "tool_calls": [{ - "index": 0, - "id": tool_id.clone(), - "type": "function", - "function": { - "name": tool_name.clone() - } - }] - } - }] - }).to_string(); - - // 第二个 chunk: 参数的前半部分 - let args_first_half = &args_json[..args_json.len()/2]; - let data2 = serde_json::json!({ - "choices": [{ - "delta": { - "tool_calls": [{ - "index": 0, - "function": { - "arguments": args_first_half - } - }] - } - }] - }).to_string(); - - // 第三个 chunk: 参数的后半部分 - let args_second_half = &args_json[args_json.len()/2..]; - let data3 = serde_json::json!({ - "choices": [{ - "delta": { - "tool_calls": [{ - "index": 0, - "function": { - "arguments": args_second_half - } - }] - } - }] - }).to_string(); - - parser.parse_data(&data1); - parser.parse_data(&data2); - parser.parse_data(&data3); - - prop_assert!(parser.has_tool_calls(), "应该检测到工具调用"); - - let tool_calls = parser.finalize_tool_calls(); - prop_assert_eq!(tool_calls.len(), 1, "应该有一个工具调用"); - prop_assert_eq!(&tool_calls[0].id, &tool_id, "工具调用 ID 应该匹配"); - prop_assert_eq!(&tool_calls[0].function.name, &tool_name, "工具名称应该匹配"); - - // 验证参数被正确累积 - prop_assert_eq!( - &tool_calls[0].function.arguments, - &args_json, - "工具参数应该被正确累积" - ); - } - - /// **Feature: agent-tool-calling, Property 1: 流式事件完整性 - Done 事件** - /// **Validates: Requirements 1.1, 1.3** - /// - /// *For any* 完成的响应流,应该正确识别 finish_reason。 - #[test] - fn prop_streaming_done_detection( - finish_reason in prop_oneof![Just("stop"), Just("tool_calls"), Just("length")] - ) { - let mut parser = SSEParser::new(); - - let data = format!( - r#"{{"choices":[{{"delta":{{}},"finish_reason":"{}"}}]}}"#, - finish_reason - ); - let (_, is_done, _) = parser.parse_data(&data); - - if finish_reason == "stop" || finish_reason == "tool_calls" { - prop_assert!(is_done, "finish_reason={} 应该标记为完成", finish_reason); - } else { - prop_assert!(!is_done, "finish_reason={} 不应该标记为完成", finish_reason); - } - } - - /// **Feature: agent-tool-calling, Property 1: 流式事件完整性 - Usage 统计** - /// **Validates: Requirements 1.3** - /// - /// *For any* 包含 usage 的响应,应该正确解析 token 使用量。 - #[test] - fn prop_streaming_usage_parsing( - input_tokens in 0u32..10000, - output_tokens in 0u32..10000 - ) { - let mut parser = SSEParser::new(); - - let data = format!( - r#"{{"choices":[{{"delta":{{"content":"test"}}}}],"usage":{{"prompt_tokens":{},"completion_tokens":{}}}}}"#, - input_tokens, output_tokens - ); - let (_, _, usage) = parser.parse_data(&data); - - if input_tokens > 0 || output_tokens > 0 { - prop_assert!(usage.is_some(), "应该解析出 usage"); - let usage = usage.unwrap(); - prop_assert_eq!(usage.input_tokens, input_tokens, "input_tokens 应该匹配"); - prop_assert_eq!(usage.output_tokens, output_tokens, "output_tokens 应该匹配"); - } - } - } } diff --git a/src-tauri/src/agent/parsers/anthropic_sse.rs b/src-tauri/src/agent/parsers/anthropic_sse.rs new file mode 100644 index 000000000..ef838c3db --- /dev/null +++ b/src-tauri/src/agent/parsers/anthropic_sse.rs @@ -0,0 +1,209 @@ +//! Anthropic SSE 流解析器 +//! +//! 解析 Anthropic Messages API 的 Server-Sent Events 流 + +use crate::agent::types::{FunctionCall, TokenUsage, ToolCall}; +use crate::models::anthropic::{AnthropicContentBlock, AnthropicDelta, AnthropicStreamEvent}; +use tracing::{debug, warn}; + +/// Anthropic 工具调用构建器 +#[derive(Debug, Clone, Default)] +struct AnthropicToolCallBuilder { + id: String, + name: String, + input_json: String, +} + +/// Anthropic SSE 流解析器 +/// +/// 解析 Anthropic Messages API 的 SSE 流 +#[derive(Debug, Default)] +pub struct AnthropicSSEParser { + /// 累积的完整内容 + full_content: String, + /// 累积的工具调用 + tool_calls: Vec, + /// 当前正在构建的工具调用 + current_tool: Option, + /// Usage 信息 + usage: Option, +} + +/// Anthropic SSE 解析结果 +#[derive(Debug, Clone)] +pub struct AnthropicParseResult { + /// 文本增量 + pub text_delta: Option, + /// 是否完成 + pub is_done: bool, + /// 工具调用开始(id, name) + pub tool_start: Option<(String, String)>, +} + +impl AnthropicSSEParser { + pub fn new() -> Self { + Self::default() + } + + /// 解析 SSE 数据行 + /// + /// 返回解析结果 + pub fn parse_data(&mut self, data: &str) -> AnthropicParseResult { + if data.trim().is_empty() { + return AnthropicParseResult { + text_delta: None, + is_done: false, + tool_start: None, + }; + } + + let event: AnthropicStreamEvent = match serde_json::from_str(data) { + Ok(e) => e, + Err(e) => { + warn!("[AnthropicSSEParser] 解析事件失败: {} - data: {}", e, data); + return AnthropicParseResult { + text_delta: None, + is_done: false, + tool_start: None, + }; + } + }; + + match event { + AnthropicStreamEvent::MessageStart { message } => { + debug!("[AnthropicSSEParser] 消息开始: id={}", message.id); + AnthropicParseResult { + text_delta: None, + is_done: false, + tool_start: None, + } + } + AnthropicStreamEvent::ContentBlockStart { + index, + content_block, + } => match content_block { + AnthropicContentBlock::ToolUse { id, name, .. } => { + debug!( + "[AnthropicSSEParser] 工具调用开始: id={}, name={}", + id, name + ); + self.current_tool = Some(AnthropicToolCallBuilder { + id: id.clone(), + name: name.clone(), + input_json: String::new(), + }); + AnthropicParseResult { + text_delta: None, + is_done: false, + tool_start: Some((id, name)), + } + } + AnthropicContentBlock::Text { .. } => { + debug!("[AnthropicSSEParser] 文本块开始: index={}", index); + AnthropicParseResult { + text_delta: None, + is_done: false, + tool_start: None, + } + } + _ => AnthropicParseResult { + text_delta: None, + is_done: false, + tool_start: None, + }, + }, + AnthropicStreamEvent::ContentBlockDelta { index: _, delta } => match delta { + AnthropicDelta::TextDelta { text } => { + self.full_content.push_str(&text); + AnthropicParseResult { + text_delta: Some(text), + is_done: false, + tool_start: None, + } + } + AnthropicDelta::InputJsonDelta { partial_json } => { + if let Some(ref mut tool) = self.current_tool { + tool.input_json.push_str(&partial_json); + } + AnthropicParseResult { + text_delta: None, + is_done: false, + tool_start: None, + } + } + AnthropicDelta::ThinkingDelta { thinking } => { + // 将思考内容添加到 full_content 中,用 标签包裹 + let thinking_text = format!("{}", thinking); + self.full_content.push_str(&thinking_text); + AnthropicParseResult { + text_delta: None, + is_done: false, + tool_start: None, + } + } + AnthropicDelta::SignatureDelta { .. } => { + // 忽略签名 delta + AnthropicParseResult { + text_delta: None, + is_done: false, + tool_start: None, + } + } + }, + AnthropicStreamEvent::ContentBlockStop { index: _ } => { + // 如果有正在构建的工具调用,完成它 + if let Some(tool) = self.current_tool.take() { + self.tool_calls.push(ToolCall { + id: tool.id, + call_type: "function".to_string(), + function: FunctionCall { + name: tool.name, + arguments: tool.input_json, + }, + }); + } + AnthropicParseResult { + text_delta: None, + is_done: false, + tool_start: None, + } + } + AnthropicStreamEvent::MessageDelta { delta: _, usage } => { + self.usage = Some(TokenUsage::new(usage.input_tokens, usage.output_tokens)); + AnthropicParseResult { + text_delta: None, + is_done: false, + tool_start: None, + } + } + AnthropicStreamEvent::MessageStop => { + debug!("[AnthropicSSEParser] 消息结束"); + AnthropicParseResult { + text_delta: None, + is_done: true, + tool_start: None, + } + } + } + } + + /// 完成解析,返回最终的工具调用列表 + pub fn finalize_tool_calls(&mut self) -> Vec { + std::mem::take(&mut self.tool_calls) + } + + /// 获取完整内容 + pub fn get_full_content(&self) -> String { + self.full_content.clone() + } + + /// 是否有工具调用 + pub fn has_tool_calls(&self) -> bool { + !self.tool_calls.is_empty() || self.current_tool.is_some() + } + + /// 获取 usage + pub fn get_usage(&self) -> Option { + self.usage.clone() + } +} diff --git a/src-tauri/src/agent/parsers/mod.rs b/src-tauri/src/agent/parsers/mod.rs new file mode 100644 index 000000000..cd898a324 --- /dev/null +++ b/src-tauri/src/agent/parsers/mod.rs @@ -0,0 +1,9 @@ +//! SSE 流解析器模块 +//! +//! 提供不同协议的 SSE 流解析器 + +mod anthropic_sse; +mod openai_sse; + +pub use anthropic_sse::{AnthropicParseResult, AnthropicSSEParser}; +pub use openai_sse::OpenAISSEParser; diff --git a/src-tauri/src/agent/parsers/openai_sse.rs b/src-tauri/src/agent/parsers/openai_sse.rs new file mode 100644 index 000000000..cd1891401 --- /dev/null +++ b/src-tauri/src/agent/parsers/openai_sse.rs @@ -0,0 +1,262 @@ +//! OpenAI SSE 流解析器 +//! +//! 解析 OpenAI 兼容 API 的 Server-Sent Events 流 +//! Requirements: 1.1, 1.3, 1.4 + +use crate::agent::types::{FunctionCall, TokenUsage, ToolCall}; +use serde_json::Value; +use std::collections::HashMap; +use tracing::warn; + +/// 工具调用增量数据 +#[derive(Debug, Clone, Default)] +struct ToolCallDelta { + /// 工具调用索引 + #[allow(dead_code)] + index: usize, + /// 工具调用 ID + id: String, + /// 工具类型 + call_type: String, + /// 函数名 + function_name: String, + /// 函数参数(累积的 JSON 字符串) + function_arguments: String, +} + +/// OpenAI SSE 流解析器 +/// +/// 解析 Server-Sent Events 流,提取 text_delta 和 tool_calls +#[derive(Debug, Default)] +pub struct OpenAISSEParser { + /// 累积的完整内容 + full_content: String, + /// 当前正在构建的工具调用索引 + current_tool_indices: HashMap, +} + +impl OpenAISSEParser { + pub fn new() -> Self { + Self::default() + } + + /// 解析 SSE 数据行 + /// + /// 返回 (text_delta, is_done, usage) + pub fn parse_data(&mut self, data: &str) -> (Option, bool, Option) { + if data.trim() == "[DONE]" { + return (None, true, None); + } + + let json: Value = match serde_json::from_str(data) { + Ok(v) => v, + Err(e) => { + warn!("[OpenAISSEParser] 解析 JSON 失败: {} - data: {}", e, data); + return (None, false, None); + } + }; + + // 提取 usage 信息(如果存在) + let usage = json.get("usage").and_then(|u| { + let input = u.get("prompt_tokens").and_then(|v| v.as_u64()).unwrap_or(0) as u32; + let output = u + .get("completion_tokens") + .and_then(|v| v.as_u64()) + .unwrap_or(0) as u32; + if input > 0 || output > 0 { + Some(TokenUsage::new(input, output)) + } else { + None + } + }); + + // 检查是否有 choices + let choices = match json.get("choices").and_then(|c| c.as_array()) { + Some(c) => c, + None => return (None, false, usage), + }; + + if choices.is_empty() { + return (None, false, usage); + } + + let choice = &choices[0]; + let delta = match choice.get("delta") { + Some(d) => d, + None => return (None, false, usage), + }; + + // 检查 finish_reason + let finish_reason = choice + .get("finish_reason") + .and_then(|f| f.as_str()) + .unwrap_or(""); + let is_done = finish_reason == "stop" || finish_reason == "tool_calls"; + + // 提取文本内容 + let text_delta = delta + .get("content") + .and_then(|c| c.as_str()) + .filter(|s| !s.is_empty()) + .map(|s| { + self.full_content.push_str(s); + s.to_string() + }); + + // 提取工具调用 + if let Some(tool_calls) = delta.get("tool_calls").and_then(|tc| tc.as_array()) { + for tc in tool_calls { + self.parse_tool_call_delta(tc); + } + } + + (text_delta, is_done, usage) + } + + /// 解析工具调用增量 + fn parse_tool_call_delta(&mut self, tc: &Value) { + let index = tc.get("index").and_then(|i| i.as_u64()).unwrap_or(0) as usize; + + // 获取或创建工具调用 + let tool_call = self + .current_tool_indices + .entry(index) + .or_insert_with(|| ToolCallDelta { + index, + ..Default::default() + }); + + // 更新 ID + if let Some(id) = tc.get("id").and_then(|i| i.as_str()) { + tool_call.id = id.to_string(); + } + + // 更新类型 + if let Some(t) = tc.get("type").and_then(|t| t.as_str()) { + tool_call.call_type = t.to_string(); + } + + // 更新函数信息 + if let Some(function) = tc.get("function") { + if let Some(name) = function.get("name").and_then(|n| n.as_str()) { + tool_call.function_name = name.to_string(); + } + if let Some(args) = function.get("arguments").and_then(|a| a.as_str()) { + tool_call.function_arguments.push_str(args); + } + } + } + + /// 完成解析,返回最终的工具调用列表 + pub fn finalize_tool_calls(&mut self) -> Vec { + // 按索引排序并转换为 ToolCall + let mut indices: Vec<_> = self.current_tool_indices.keys().cloned().collect(); + indices.sort(); + + indices + .into_iter() + .filter_map(|idx| { + let delta = self.current_tool_indices.get(&idx)?; + if delta.id.is_empty() || delta.function_name.is_empty() { + return None; + } + Some(ToolCall { + id: delta.id.clone(), + call_type: if delta.call_type.is_empty() { + "function".to_string() + } else { + delta.call_type.clone() + }, + function: FunctionCall { + name: delta.function_name.clone(), + arguments: delta.function_arguments.clone(), + }, + }) + }) + .collect() + } + + /// 获取完整内容 + pub fn get_full_content(&self) -> String { + self.full_content.clone() + } + + /// 是否有工具调用 + pub fn has_tool_calls(&self) -> bool { + !self.current_tool_indices.is_empty() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_text_delta() { + let mut parser = OpenAISSEParser::new(); + + let data1 = r#"{"choices":[{"delta":{"content":"Hello"}}]}"#; + let data2 = r#"{"choices":[{"delta":{"content":" World"}}]}"#; + let data3 = r#"{"choices":[{"delta":{},"finish_reason":"stop"}]}"#; + + let (text1, done1, _) = parser.parse_data(data1); + assert_eq!(text1, Some("Hello".to_string())); + assert!(!done1); + + let (text2, done2, _) = parser.parse_data(data2); + assert_eq!(text2, Some(" World".to_string())); + assert!(!done2); + + let (text3, done3, _) = parser.parse_data(data3); + assert!(text3.is_none()); + assert!(done3); + + assert_eq!(parser.get_full_content(), "Hello World"); + } + + #[test] + fn test_tool_calls() { + let mut parser = OpenAISSEParser::new(); + + let data1 = r#"{"choices":[{"delta":{"tool_calls":[{"index":0,"id":"call_123","type":"function","function":{"name":"bash"}}]}}]}"#; + let data2 = r#"{"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"{\"command\":"}}]}}]}"#; + let data3 = r#"{"choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"\"ls -la\"}"}}]}}]}"#; + let data4 = r#"{"choices":[{"delta":{},"finish_reason":"tool_calls"}]}"#; + + parser.parse_data(data1); + parser.parse_data(data2); + parser.parse_data(data3); + let (_, done, _) = parser.parse_data(data4); + + assert!(done); + assert!(parser.has_tool_calls()); + + let tool_calls = parser.finalize_tool_calls(); + assert_eq!(tool_calls.len(), 1); + assert_eq!(tool_calls[0].id, "call_123"); + assert_eq!(tool_calls[0].function.name, "bash"); + assert_eq!(tool_calls[0].function.arguments, r#"{"command":"ls -la"}"#); + } + + #[test] + fn test_usage() { + let mut parser = OpenAISSEParser::new(); + + let data = r#"{"choices":[{"delta":{"content":"Hi"}}],"usage":{"prompt_tokens":10,"completion_tokens":5}}"#; + let (text, _, usage) = parser.parse_data(data); + + assert_eq!(text, Some("Hi".to_string())); + assert!(usage.is_some()); + let usage = usage.unwrap(); + assert_eq!(usage.input_tokens, 10); + assert_eq!(usage.output_tokens, 5); + } + + #[test] + fn test_done_signal() { + let mut parser = OpenAISSEParser::new(); + + let (_, done, _) = parser.parse_data("[DONE]"); + assert!(done); + } +} diff --git a/src-tauri/src/agent/protocols/anthropic.rs b/src-tauri/src/agent/protocols/anthropic.rs new file mode 100644 index 000000000..3b28e79a8 --- /dev/null +++ b/src-tauri/src/agent/protocols/anthropic.rs @@ -0,0 +1,491 @@ +//! Anthropic 协议实现 +//! +//! 实现 Anthropic Messages API 协议 +//! 适用于 Claude、Claude OAuth 等 Anthropic 服务 + +use super::Protocol; +use crate::agent::parsers::AnthropicSSEParser; +use crate::agent::types::{ + AgentConfig, AgentMessage, ContentPart, ImageData, MessageContent, StreamEvent, StreamResult, +}; +use crate::models::anthropic::AnthropicMessage; +use crate::models::openai::Tool; +use async_trait::async_trait; +use futures::StreamExt; +use reqwest::Client; +use serde::Serialize; +use tokio::sync::mpsc; +use tracing::{debug, error, info}; + +/// Anthropic Messages API 请求 +#[derive(Debug, Serialize)] +struct AnthropicMessagesRequest { + model: String, + messages: Vec, + max_tokens: u32, + stream: bool, + #[serde(skip_serializing_if = "Option::is_none")] + system: Option, + #[serde(skip_serializing_if = "Option::is_none")] + temperature: Option, + #[serde(skip_serializing_if = "Option::is_none")] + tools: Option>, +} + +/// Anthropic 工具定义 +#[derive(Debug, Serialize)] +struct AnthropicTool { + name: String, + description: String, + input_schema: serde_json::Value, +} + +/// Anthropic 协议处理器 +pub struct AnthropicProtocol; + +impl AnthropicProtocol { + /// 将 OpenAI Tool 转换为 Anthropic Tool + fn convert_tools(tools: Option<&[Tool]>) -> Option> { + tools.map(|t| { + t.iter() + .filter_map(|tool| match tool { + Tool::Function { function } => Some(AnthropicTool { + name: function.name.clone(), + description: function.description.clone().unwrap_or_default(), + input_schema: function.parameters.clone().unwrap_or(serde_json::json!({ + "type": "object", + "properties": {} + })), + }), + // WebSearch 工具不支持转换为 Anthropic 格式,跳过 + Tool::WebSearch | Tool::WebSearch20250305 => None, + }) + .collect() + }) + } + + /// 将 AgentMessage 转换为 Anthropic Message + fn convert_to_anthropic_message(msg: &AgentMessage) -> AnthropicMessage { + let content = match &msg.content { + MessageContent::Text(text) => { + // 处理工具结果消息 + if msg.role == "tool" { + // Anthropic 使用 tool_result content block + if let Some(tool_call_id) = &msg.tool_call_id { + serde_json::json!([{ + "type": "tool_result", + "tool_use_id": tool_call_id, + "content": text + }]) + } else { + serde_json::json!(text) + } + } else if msg.role == "assistant" { + // 处理 assistant 消息 + let mut blocks = Vec::new(); + + if !text.is_empty() { + blocks.push(serde_json::json!({ + "type": "text", + "text": text + })); + } + + // 添加工具调用 + if let Some(tool_calls) = &msg.tool_calls { + for tc in tool_calls { + let input: serde_json::Value = + serde_json::from_str(&tc.function.arguments) + .unwrap_or(serde_json::json!({})); + blocks.push(serde_json::json!({ + "type": "tool_use", + "id": tc.id, + "name": tc.function.name, + "input": input + })); + } + } + + if blocks.is_empty() { + serde_json::json!("") + } else if blocks.len() == 1 && msg.tool_calls.is_none() { + serde_json::json!(text) + } else { + serde_json::json!(blocks) + } + } else { + serde_json::json!(text) + } + } + MessageContent::Parts(parts) => { + let blocks: Vec = parts + .iter() + .map(|p| match p { + ContentPart::Text { text } => serde_json::json!({ + "type": "text", + "text": text + }), + ContentPart::ImageUrl { image_url } => { + // 解析 data URL + if let Some(rest) = image_url.url.strip_prefix("data:") { + if let Some(comma_idx) = rest.find(',') { + let media_type = rest[..comma_idx] + .strip_suffix(";base64") + .unwrap_or(&rest[..comma_idx]); + let data = &rest[comma_idx + 1..]; + return serde_json::json!({ + "type": "image", + "source": { + "type": "base64", + "media_type": media_type, + "data": data + } + }); + } + } + // 普通 URL + serde_json::json!({ + "type": "image", + "source": { + "type": "url", + "url": image_url.url + } + }) + } + }) + .collect(); + serde_json::json!(blocks) + } + }; + + // Anthropic 没有 "tool" 角色,需要转换为 "user" + let role = if msg.role == "tool" { + "user".to_string() + } else { + msg.role.clone() + }; + + AnthropicMessage { role, content } + } + + /// 构建消息列表 + fn build_messages( + history: &[AgentMessage], + user_message: &str, + images: Option<&[ImageData]>, + config: &AgentConfig, + ) -> (Vec, Option) { + let mut messages = Vec::new(); + + // 系统提示词(Anthropic 使用单独的 system 字段) + let system_prompt = config.system_prompt.as_ref().map(|s| serde_json::json!(s)); + + // 添加历史消息(跳过 system 消息) + for msg in history { + if msg.role == "system" { + continue; + } + messages.push(Self::convert_to_anthropic_message(msg)); + } + + // 添加当前用户消息 + let user_content = if let Some(imgs) = images { + let mut parts = vec![serde_json::json!({ + "type": "text", + "text": user_message + })]; + + for img in imgs { + parts.push(serde_json::json!({ + "type": "image", + "source": { + "type": "base64", + "media_type": img.media_type, + "data": img.data + } + })); + } + serde_json::json!(parts) + } else { + serde_json::json!(user_message) + }; + + messages.push(AnthropicMessage { + role: "user".to_string(), + content: user_content, + }); + + (messages, system_prompt) + } + + /// 从历史构建消息(不添加新用户消息) + fn build_messages_from_history( + history: &[AgentMessage], + config: &AgentConfig, + ) -> (Vec, Option) { + let mut messages = Vec::new(); + + // 系统提示词 + let system_prompt = config.system_prompt.as_ref().map(|s| serde_json::json!(s)); + + // 添加所有历史消息(跳过 system) + for msg in history { + if msg.role == "system" { + continue; + } + messages.push(Self::convert_to_anthropic_message(msg)); + } + + (messages, system_prompt) + } + + /// 处理 SSE 流 + async fn process_stream( + response: reqwest::Response, + tx: mpsc::Sender, + send_done: bool, + ) -> Result { + let mut stream = response.bytes_stream(); + let mut buffer = String::new(); + let mut parser = AnthropicSSEParser::new(); + + while let Some(chunk) = stream.next().await { + match chunk { + Ok(bytes) => { + let text = String::from_utf8_lossy(&bytes); + buffer.push_str(&text); + + // 处理完整的 SSE 事件 + while let Some(pos) = buffer.find("\n\n") { + let event_block = buffer[..pos].to_string(); + buffer = buffer[pos + 2..].to_string(); + + // 提取 event 类型和 data + let mut event_type = String::new(); + let mut data = String::new(); + + for line in event_block.lines() { + if let Some(e) = line.strip_prefix("event: ") { + event_type = e.to_string(); + } else if let Some(d) = line.strip_prefix("data: ") { + data = d.to_string(); + } + } + + if data.is_empty() { + continue; + } + + debug!( + "[AnthropicProtocol] SSE event={}, data={}", + event_type, data + ); + let result = parser.parse_data(&data); + + // 发送工具开始事件 + if let Some((tool_id, tool_name)) = result.tool_start { + let _ = tx + .send(StreamEvent::ToolStart { + tool_name, + tool_id, + arguments: None, + }) + .await; + } + + // 发送文本增量 + if let Some(text) = result.text_delta { + let _ = tx.send(StreamEvent::TextDelta { text }).await; + } + + // 检查是否完成 + if result.is_done { + let full_content = parser.get_full_content(); + let tool_calls = if parser.has_tool_calls() { + Some(parser.finalize_tool_calls()) + } else { + None + }; + let usage = parser.get_usage(); + + if send_done { + let _ = tx + .send(StreamEvent::Done { + usage: usage.clone(), + }) + .await; + } + + return Ok(StreamResult { + content: full_content, + tool_calls, + usage, + }); + } + } + } + Err(e) => { + error!("[AnthropicProtocol] 流读取错误: {}", e); + let _ = tx + .send(StreamEvent::Error { + message: format!("流读取错误: {}", e), + }) + .await; + return Err(format!("流读取错误: {}", e)); + } + } + } + + // 流正常结束 + let full_content = parser.get_full_content(); + let tool_calls = if parser.has_tool_calls() { + Some(parser.finalize_tool_calls()) + } else { + None + }; + let usage = parser.get_usage(); + + if send_done { + let _ = tx + .send(StreamEvent::Done { + usage: usage.clone(), + }) + .await; + } + + Ok(StreamResult { + content: full_content, + tool_calls, + usage, + }) + } +} + +#[async_trait] +impl Protocol for AnthropicProtocol { + async fn chat_stream( + &self, + client: &Client, + base_url: &str, + api_key: &str, + messages: &[AgentMessage], + user_message: &str, + images: Option<&[ImageData]>, + model: &str, + config: &AgentConfig, + tools: Option<&[Tool]>, + tx: mpsc::Sender, + ) -> Result { + info!( + "[AnthropicProtocol] 发送流式请求: model={}, history_len={}, tools_count={}", + model, + messages.len(), + tools.map(|t| t.len()).unwrap_or(0) + ); + + let (anthropic_messages, system) = + Self::build_messages(messages, user_message, images, config); + + let anthropic_tools = Self::convert_tools(tools); + + let request = AnthropicMessagesRequest { + model: model.to_string(), + messages: anthropic_messages, + max_tokens: config.max_tokens.unwrap_or(4096), + stream: true, + system, + temperature: config.temperature, + tools: anthropic_tools, + }; + + let url = format!("{}{}", base_url, self.endpoint()); + + let response = client + .post(&url) + .header("Authorization", format!("Bearer {}", api_key)) + .header("Content-Type", "application/json") + .header("anthropic-version", "2023-06-01") + .json(&request) + .send() + .await + .map_err(|e| format!("请求失败: {}", e))?; + + let status = response.status(); + if !status.is_success() { + let body = response.text().await.unwrap_or_default(); + error!("[AnthropicProtocol] 请求失败: {} - {}", status, body); + let _ = tx + .send(StreamEvent::Error { + message: format!("API 错误 ({}): {}", status, body), + }) + .await; + return Err(format!("API 错误: {}", status)); + } + + Self::process_stream(response, tx, true).await + } + + async fn chat_stream_continue( + &self, + client: &Client, + base_url: &str, + api_key: &str, + messages: &[AgentMessage], + model: &str, + config: &AgentConfig, + tools: Option<&[Tool]>, + tx: mpsc::Sender, + ) -> Result { + debug!( + "[AnthropicProtocol] 继续流式对话: model={}, history_len={}, tools_count={}", + model, + messages.len(), + tools.map(|t| t.len()).unwrap_or(0) + ); + + let (anthropic_messages, system) = Self::build_messages_from_history(messages, config); + + let anthropic_tools = Self::convert_tools(tools); + + let request = AnthropicMessagesRequest { + model: model.to_string(), + messages: anthropic_messages, + max_tokens: config.max_tokens.unwrap_or(4096), + stream: true, + system, + temperature: config.temperature, + tools: anthropic_tools, + }; + + let url = format!("{}{}", base_url, self.endpoint()); + + let response = client + .post(&url) + .header("Authorization", format!("Bearer {}", api_key)) + .header("Content-Type", "application/json") + .header("anthropic-version", "2023-06-01") + .json(&request) + .send() + .await + .map_err(|e| format!("请求失败: {}", e))?; + + let status = response.status(); + if !status.is_success() { + let body = response.text().await.unwrap_or_default(); + error!("[AnthropicProtocol] 请求失败: {} - {}", status, body); + let _ = tx + .send(StreamEvent::Error { + message: format!("API 错误 ({}): {}", status, body), + }) + .await; + return Err(format!("API 错误: {}", status)); + } + + // 继续对话时不发送 Done 事件 + Self::process_stream(response, tx, false).await + } + + fn endpoint(&self) -> &'static str { + "/v1/messages" + } +} diff --git a/src-tauri/src/agent/protocols/mod.rs b/src-tauri/src/agent/protocols/mod.rs new file mode 100644 index 000000000..1674985bf --- /dev/null +++ b/src-tauri/src/agent/protocols/mod.rs @@ -0,0 +1,70 @@ +//! 协议策略模块 +//! +//! 使用策略模式处理不同 API 协议(OpenAI、Anthropic、Kiro、Gemini) + +mod anthropic; +mod openai; + +pub use anthropic::AnthropicProtocol; +pub use openai::OpenAIProtocol; + +use crate::agent::types::{ + AgentConfig, AgentMessage, ImageData, ProviderType, StreamEvent, StreamResult, +}; +use crate::models::openai::Tool; +use async_trait::async_trait; +use reqwest::Client; +use tokio::sync::mpsc; + +/// 协议处理器 trait +/// +/// 定义了所有协议必须实现的方法 +#[async_trait] +pub trait Protocol: Send + Sync { + /// 流式聊天 + /// + /// 发送消息并通过 channel 返回流式响应 + async fn chat_stream( + &self, + client: &Client, + base_url: &str, + api_key: &str, + messages: &[AgentMessage], + user_message: &str, + images: Option<&[ImageData]>, + model: &str, + config: &AgentConfig, + tools: Option<&[Tool]>, + tx: mpsc::Sender, + ) -> Result; + + /// 继续流式对话(工具调用后) + /// + /// 使用会话历史继续对话,不添加新的用户消息 + async fn chat_stream_continue( + &self, + client: &Client, + base_url: &str, + api_key: &str, + messages: &[AgentMessage], + model: &str, + config: &AgentConfig, + tools: Option<&[Tool]>, + tx: mpsc::Sender, + ) -> Result; + + /// 获取 API 端点 + fn endpoint(&self) -> &'static str; +} + +/// 根据 ProviderType 创建协议处理器 +pub fn create_protocol(provider_type: ProviderType) -> Box { + match provider_type { + // Claude 和 Kiro 都使用 Anthropic SSE 协议 + ProviderType::Claude | ProviderType::ClaudeOauth | ProviderType::Kiro => { + Box::new(AnthropicProtocol) + } + // 其他使用 OpenAI 兼容协议 + _ => Box::new(OpenAIProtocol), + } +} diff --git a/src-tauri/src/agent/protocols/openai.rs b/src-tauri/src/agent/protocols/openai.rs new file mode 100644 index 000000000..7ac82b4bf --- /dev/null +++ b/src-tauri/src/agent/protocols/openai.rs @@ -0,0 +1,380 @@ +//! OpenAI 协议实现 +//! +//! 实现 OpenAI Chat Completions API 协议 +//! 适用于 OpenAI、Qwen、Codex、Antigravity、IFlow、Kiro 等兼容服务 + +use super::Protocol; +use crate::agent::parsers::OpenAISSEParser; +use crate::agent::types::{ + AgentConfig, AgentMessage, ContentPart, ImageData, MessageContent, StreamEvent, StreamResult, +}; +use crate::models::openai::{ + ChatCompletionRequest, ChatMessage, ContentPart as OpenAIContentPart, + MessageContent as OpenAIMessageContent, Tool, +}; +use async_trait::async_trait; +use futures::StreamExt; +use reqwest::Client; +use tokio::sync::mpsc; +use tracing::{debug, error, info}; + +/// OpenAI 协议处理器 +pub struct OpenAIProtocol; + +impl OpenAIProtocol { + /// 将 AgentMessage 转换为 OpenAI ChatMessage + fn convert_to_chat_message(msg: &AgentMessage) -> ChatMessage { + let content = match &msg.content { + MessageContent::Text(text) => Some(OpenAIMessageContent::Text(text.clone())), + MessageContent::Parts(parts) => { + let openai_parts: Vec = parts + .iter() + .map(|p| match p { + ContentPart::Text { text } => { + OpenAIContentPart::Text { text: text.clone() } + } + ContentPart::ImageUrl { image_url } => OpenAIContentPart::ImageUrl { + image_url: crate::models::openai::ImageUrl { + url: image_url.url.clone(), + detail: image_url.detail.clone(), + }, + }, + }) + .collect(); + Some(OpenAIMessageContent::Parts(openai_parts)) + } + }; + + ChatMessage { + role: msg.role.clone(), + content, + tool_calls: msg.tool_calls.as_ref().map(|calls| { + calls + .iter() + .map(|tc| crate::models::openai::ToolCall { + id: tc.id.clone(), + call_type: tc.call_type.clone(), + function: crate::models::openai::FunctionCall { + name: tc.function.name.clone(), + arguments: tc.function.arguments.clone(), + }, + }) + .collect() + }), + tool_call_id: msg.tool_call_id.clone(), + } + } + + /// 构建消息列表 + fn build_messages( + history: &[AgentMessage], + user_message: &str, + images: Option<&[ImageData]>, + config: &AgentConfig, + ) -> Vec { + let mut messages = Vec::new(); + + // 添加系统提示词 + if let Some(prompt) = &config.system_prompt { + messages.push(ChatMessage { + role: "system".to_string(), + content: Some(OpenAIMessageContent::Text(prompt.clone())), + tool_calls: None, + tool_call_id: None, + }); + } + + // 添加历史消息 + for msg in history { + messages.push(Self::convert_to_chat_message(msg)); + } + + // 添加当前用户消息 + let user_msg = if let Some(imgs) = images { + let mut parts = vec![OpenAIContentPart::Text { + text: user_message.to_string(), + }]; + + for img in imgs { + parts.push(OpenAIContentPart::ImageUrl { + image_url: crate::models::openai::ImageUrl { + url: format!("data:{};base64,{}", img.media_type, img.data), + detail: None, + }, + }); + } + + ChatMessage { + role: "user".to_string(), + content: Some(OpenAIMessageContent::Parts(parts)), + tool_calls: None, + tool_call_id: None, + } + } else { + ChatMessage { + role: "user".to_string(), + content: Some(OpenAIMessageContent::Text(user_message.to_string())), + tool_calls: None, + tool_call_id: None, + } + }; + + messages.push(user_msg); + messages + } + + /// 从历史构建消息(不添加新用户消息) + fn build_messages_from_history( + history: &[AgentMessage], + config: &AgentConfig, + ) -> Vec { + let mut messages = Vec::new(); + + // 添加系统提示词 + if let Some(prompt) = &config.system_prompt { + messages.push(ChatMessage { + role: "system".to_string(), + content: Some(OpenAIMessageContent::Text(prompt.clone())), + tool_calls: None, + tool_call_id: None, + }); + } + + // 添加所有历史消息 + for msg in history { + messages.push(Self::convert_to_chat_message(msg)); + } + + messages + } + + /// 处理 SSE 流 + async fn process_stream( + response: reqwest::Response, + tx: mpsc::Sender, + send_done: bool, + ) -> Result { + let mut stream = response.bytes_stream(); + let mut buffer = String::new(); + let mut parser = OpenAISSEParser::new(); + let mut final_usage = None; + + while let Some(chunk) = stream.next().await { + match chunk { + Ok(bytes) => { + let text = String::from_utf8_lossy(&bytes); + buffer.push_str(&text); + + // 处理完整的 SSE 事件(以 \n\n 分隔) + while let Some(pos) = buffer.find("\n\n") { + let event = buffer[..pos].to_string(); + buffer = buffer[pos + 2..].to_string(); + + for line in event.lines() { + if let Some(data) = line.strip_prefix("data: ") { + debug!("[OpenAIProtocol] SSE data: {}", data); + let (text_delta, is_done, usage) = parser.parse_data(data); + + if usage.is_some() { + final_usage = usage; + } + + if let Some(text) = text_delta { + let _ = tx.send(StreamEvent::TextDelta { text }).await; + } + + if is_done { + let full_content = parser.get_full_content(); + let tool_calls = if parser.has_tool_calls() { + Some(parser.finalize_tool_calls()) + } else { + None + }; + + if send_done { + let _ = tx + .send(StreamEvent::Done { + usage: final_usage.clone(), + }) + .await; + } + + return Ok(StreamResult { + content: full_content, + tool_calls, + usage: final_usage, + }); + } + } + } + } + } + Err(e) => { + error!("[OpenAIProtocol] 流读取错误: {}", e); + let _ = tx + .send(StreamEvent::Error { + message: format!("流读取错误: {}", e), + }) + .await; + return Err(format!("流读取错误: {}", e)); + } + } + } + + // 流正常结束但没有收到 [DONE] + let full_content = parser.get_full_content(); + let tool_calls = if parser.has_tool_calls() { + Some(parser.finalize_tool_calls()) + } else { + None + }; + + if send_done { + let _ = tx + .send(StreamEvent::Done { + usage: final_usage.clone(), + }) + .await; + } + + Ok(StreamResult { + content: full_content, + tool_calls, + usage: final_usage, + }) + } +} + +#[async_trait] +impl Protocol for OpenAIProtocol { + async fn chat_stream( + &self, + client: &Client, + base_url: &str, + api_key: &str, + messages: &[AgentMessage], + user_message: &str, + images: Option<&[ImageData]>, + model: &str, + config: &AgentConfig, + tools: Option<&[Tool]>, + tx: mpsc::Sender, + ) -> Result { + info!( + "[OpenAIProtocol] 发送流式请求: model={}, history_len={}, tools_count={}", + model, + messages.len(), + tools.map(|t| t.len()).unwrap_or(0) + ); + + let chat_messages = Self::build_messages(messages, user_message, images, config); + + let request = ChatCompletionRequest { + model: model.to_string(), + messages: chat_messages, + stream: true, + temperature: config.temperature, + max_tokens: config.max_tokens, + top_p: None, + tools: tools.map(|t| t.to_vec()), + tool_choice: if tools.is_some() { + Some(serde_json::json!("auto")) + } else { + None + }, + reasoning_effort: None, + }; + + let url = format!("{}{}", base_url, self.endpoint()); + + let response = client + .post(&url) + .header("Authorization", format!("Bearer {}", api_key)) + .header("Content-Type", "application/json") + .json(&request) + .send() + .await + .map_err(|e| format!("请求失败: {}", e))?; + + let status = response.status(); + if !status.is_success() { + let body = response.text().await.unwrap_or_default(); + error!("[OpenAIProtocol] 请求失败: {} - {}", status, body); + let _ = tx + .send(StreamEvent::Error { + message: format!("API 错误 ({}): {}", status, body), + }) + .await; + return Err(format!("API 错误: {}", status)); + } + + Self::process_stream(response, tx, true).await + } + + async fn chat_stream_continue( + &self, + client: &Client, + base_url: &str, + api_key: &str, + messages: &[AgentMessage], + model: &str, + config: &AgentConfig, + tools: Option<&[Tool]>, + tx: mpsc::Sender, + ) -> Result { + debug!( + "[OpenAIProtocol] 继续流式对话: model={}, history_len={}, tools_count={}", + model, + messages.len(), + tools.map(|t| t.len()).unwrap_or(0) + ); + + let chat_messages = Self::build_messages_from_history(messages, config); + + let request = ChatCompletionRequest { + model: model.to_string(), + messages: chat_messages, + stream: true, + temperature: config.temperature, + max_tokens: config.max_tokens, + top_p: None, + tools: tools.map(|t| t.to_vec()), + tool_choice: if tools.is_some() { + Some(serde_json::json!("auto")) + } else { + None + }, + reasoning_effort: None, + }; + + let url = format!("{}{}", base_url, self.endpoint()); + + let response = client + .post(&url) + .header("Authorization", format!("Bearer {}", api_key)) + .header("Content-Type", "application/json") + .json(&request) + .send() + .await + .map_err(|e| format!("请求失败: {}", e))?; + + let status = response.status(); + if !status.is_success() { + let body = response.text().await.unwrap_or_default(); + error!("[OpenAIProtocol] 请求失败: {} - {}", status, body); + let _ = tx + .send(StreamEvent::Error { + message: format!("API 错误 ({}): {}", status, body), + }) + .await; + return Err(format!("API 错误: {}", status)); + } + + // 继续对话时不发送 Done 事件(工具循环可能还会继续) + Self::process_stream(response, tx, false).await + } + + fn endpoint(&self) -> &'static str { + "/v1/chat/completions" + } +} diff --git a/src-tauri/src/agent/tool_loop.rs b/src-tauri/src/agent/tool_loop.rs index 597f70e03..51dca0072 100644 --- a/src-tauri/src/agent/tool_loop.rs +++ b/src-tauri/src/agent/tool_loop.rs @@ -150,6 +150,11 @@ impl ToolLoopEngine { self.config.max_iterations } + /// 获取工具注册表引用 + pub fn registry(&self) -> &ToolRegistry { + &self.registry + } + /// 检查响应是否包含工具调用 /// /// Requirements: 7.1 - WHEN the Agent response contains tool_calls @@ -227,6 +232,7 @@ impl ToolLoopEngine { .send(StreamEvent::ToolStart { tool_name: tool_call.function.name.clone(), tool_id: tool_call.id.clone(), + arguments: Some(tool_call.function.arguments.clone()), }) .await; } @@ -891,7 +897,7 @@ mod proptests { // 验证:收到 ToolStart 事件 let event1 = rx.recv().await; prop_assert!(event1.is_some(), "应该收到 ToolStart 事件"); - if let Some(StreamEvent::ToolStart { tool_name, tool_id: event_tool_id }) = event1 { + if let Some(StreamEvent::ToolStart { tool_name, tool_id: event_tool_id, .. }) = event1 { prop_assert_eq!(tool_name, "echo", "工具名称应该为 'echo'"); prop_assert_eq!(event_tool_id, tool_id.clone(), "工具 ID 应该匹配"); } else { diff --git a/src-tauri/src/agent/tools/mod.rs b/src-tauri/src/agent/tools/mod.rs index d3f82fa2b..2b0dba6ab 100644 --- a/src-tauri/src/agent/tools/mod.rs +++ b/src-tauri/src/agent/tools/mod.rs @@ -30,3 +30,44 @@ pub use registry::{Tool, ToolRegistry}; pub use security::{SecurityError, SecurityManager}; pub use types::*; pub use write_file::{WriteFileResult, WriteFileTool}; + +use std::path::Path; +use std::sync::Arc; +use tracing::info; + +/// 创建包含所有默认工具的注册表 +/// +/// # Arguments +/// * `base_dir` - 基础目录,所有文件操作必须在此目录内 +/// +/// # Returns +/// 包含 bash, read_file, write_file, edit_file 工具的注册表 +pub fn create_default_registry(base_dir: impl AsRef) -> ToolRegistry { + let security = Arc::new(SecurityManager::new(base_dir.as_ref())); + let registry = ToolRegistry::new(); + + // 注册核心工具 + if let Err(e) = registry.register(BashTool::new(Arc::clone(&security))) { + tracing::error!("注册 BashTool 失败: {}", e); + } + + if let Err(e) = registry.register(ReadFileTool::new(Arc::clone(&security))) { + tracing::error!("注册 ReadFileTool 失败: {}", e); + } + + if let Err(e) = registry.register(WriteFileTool::new(Arc::clone(&security))) { + tracing::error!("注册 WriteFileTool 失败: {}", e); + } + + if let Err(e) = registry.register(EditFileTool::new(Arc::clone(&security))) { + tracing::error!("注册 EditFileTool 失败: {}", e); + } + + info!( + "[Tools] 已创建默认工具注册表,共 {} 个工具: {:?}", + registry.len(), + registry.list_names() + ); + + registry +} diff --git a/src-tauri/src/agent/tools/prompt.rs b/src-tauri/src/agent/tools/prompt.rs index 91fa21b97..ad2c6e90f 100644 --- a/src-tauri/src/agent/tools/prompt.rs +++ b/src-tauri/src/agent/tools/prompt.rs @@ -49,11 +49,18 @@ impl ToolPromptGenerator { self } - /// 生成包含工具定义的 System Prompt + /// 生成包含工具使用指导的 System Prompt /// - /// Requirements: 2.3 - THE System_Prompt SHALL include all available tool definitions - /// in a format the LLM can understand - pub fn generate_system_prompt(&self, tools: &[ToolDefinition]) -> String { + /// 注意:工具定义已通过 API 的 tools 字段发送,不需要在 system prompt 中重复 + /// 此方法只返回工具使用指导 + pub fn generate_system_prompt(&self, _tools: &[ToolDefinition]) -> String { + // 只返回使用指导,工具定义由 API 原生处理 + TOOL_USAGE_INSTRUCTIONS.to_string() + } + + /// 生成包含工具定义的完整 System Prompt(旧版本,保留兼容性) + #[allow(dead_code)] + pub fn generate_full_system_prompt(&self, tools: &[ToolDefinition]) -> String { match self.format { PromptFormat::Xml => self.generate_xml_prompt(tools), PromptFormat::Json => self.generate_json_prompt(tools), @@ -177,25 +184,38 @@ impl ToolPromptGenerator { } } -/// 工具使用说明模板 -const TOOL_USAGE_INSTRUCTIONS: &str = r#"You have access to a set of tools that you can use to help accomplish tasks. When you need to use a tool, respond with a tool call in the following format: +/// 工具使用说明模板(适合桌面软件) +const TOOL_USAGE_INSTRUCTIONS: &str = r#"你是一个友好的 AI 助手。 - -tool_name - -{ - "param1": "value1", - "param2": "value2" -} - - +# 核心原则 -Important guidelines for tool usage: -1. Only use tools when necessary to accomplish the task -2. Provide all required parameters for each tool call -3. Wait for tool results before making additional tool calls that depend on them -4. If a tool call fails, analyze the error and try an alternative approach -5. Always explain your reasoning before and after using tools"#; +1. **自然交流**:对于问候、闲聊、问答,直接用文字回复,不要调用任何工具 +2. **显式授权**:只有当用户**明确提供**文件路径或目录时,才能操作 +3. **不要主动探索**:不要自作主张读取目录或文件来"了解环境" + +# 可用工具 + +- **read_file**:读取用户指定的文件或目录 +- **write_file**:创建/覆盖用户指定的文件 +- **edit_file**:修改用户指定的文件 +- **bash**:执行用户要求的命令 + +# 重要限制 + +⚠️ **禁止行为**: +- 用户说"你好"时,不要读取任何文件 +- 用户没有给路径时,不要自己猜测或使用 "." +- 不要为了"打招呼"或"了解用户"而调用工具 + +✅ **正确做法**: +- 用户说"你好" → 直接回复问候 +- 用户说"看看 /path/to/file" → 调用 read_file +- 用户说"列出目录内容" → 询问用户要查看哪个目录 + +# 输出格式 +- 使用 Markdown 格式 +- 简洁明了 +- 使用中文回复"#; /// XML 特殊字符转义 fn escape_xml(s: &str) -> String { @@ -305,13 +325,25 @@ mod tests { } #[test] - fn test_generate_xml_prompt() { - let generator = ToolPromptGenerator::new().with_format(PromptFormat::Xml); + fn test_generate_system_prompt() { + let generator = ToolPromptGenerator::new(); let tools = create_test_tools(); let prompt = generator.generate_system_prompt(&tools); + // 验证包含工具使用说明(新版本只返回指导,不包含工具定义) + assert!(prompt.contains("你是一个友好的 AI 助手")); + assert!(prompt.contains("可用工具")); + assert!(prompt.contains("read_file")); // 在说明中提到 + } + + #[test] + fn test_generate_full_xml_prompt() { + let generator = ToolPromptGenerator::new().with_format(PromptFormat::Xml); + let tools = create_test_tools(); + let prompt = generator.generate_full_system_prompt(&tools); + // 验证包含工具使用说明 - assert!(prompt.contains("You have access to a set of tools")); + assert!(prompt.contains("可用工具")); // 验证包含 tools 标签 assert!(prompt.contains("")); assert!(prompt.contains("")); @@ -321,13 +353,13 @@ mod tests { } #[test] - fn test_generate_json_prompt() { + fn test_generate_full_json_prompt() { let generator = ToolPromptGenerator::new().with_format(PromptFormat::Json); let tools = create_test_tools(); - let prompt = generator.generate_system_prompt(&tools); + let prompt = generator.generate_full_system_prompt(&tools); // 验证包含工具使用说明 - assert!(prompt.contains("You have access to a set of tools")); + assert!(prompt.contains("可用工具")); // 验证包含 JSON 代码块 assert!(prompt.contains("```json")); // 验证包含所有工具 @@ -339,11 +371,9 @@ mod tests { fn test_generate_tools_prompt_convenience_function() { let tools = create_test_tools(); - let xml_prompt = generate_tools_prompt(&tools, PromptFormat::Xml); - assert!(xml_prompt.contains("")); - - let json_prompt = generate_tools_prompt(&tools, PromptFormat::Json); - assert!(json_prompt.contains("```json")); + // generate_tools_prompt 使用 generate_system_prompt,只返回指导 + let prompt = generate_tools_prompt(&tools, PromptFormat::Xml); + assert!(prompt.contains("可用工具")); } #[test] @@ -361,9 +391,8 @@ mod tests { let prompt = generator.generate_system_prompt(&[]); // 即使没有工具,也应该包含使用说明 - assert!(prompt.contains("You have access to a set of tools")); - assert!(prompt.contains("")); - assert!(prompt.contains("")); + assert!(prompt.contains("你是一个友好的 AI 助手")); + assert!(prompt.contains("可用工具")); } #[test] @@ -399,7 +428,7 @@ mod tests { } #[test] - fn test_prompt_contains_all_tool_names_and_descriptions() { + fn test_full_prompt_contains_all_tool_names_and_descriptions() { let tools = vec![ ToolDefinition::new("tool_a", "Description for tool A"), ToolDefinition::new("tool_b", "Description for tool B"), @@ -407,7 +436,8 @@ mod tests { ]; let generator = ToolPromptGenerator::new(); - let prompt = generator.generate_system_prompt(&tools); + // 使用 generate_full_system_prompt 来包含工具定义 + let prompt = generator.generate_full_system_prompt(&tools); // 验证所有工具名称都在 prompt 中 for tool in &tools { @@ -494,14 +524,14 @@ mod proptests { proptest! { #![proptest_config(ProptestConfig::with_cases(100))] - /// **Feature: agent-tool-calling, Property 3: System Prompt 工具包含** + /// **Feature: agent-tool-calling, Property 3: Full System Prompt 工具包含** /// **Validates: Requirements 2.3** /// - /// *For any* 已注册的工具集合,生成的 System Prompt 应该包含所有工具的 name 和 description。 + /// *For any* 已注册的工具集合,生成的完整 System Prompt 应该包含所有工具的 name 和 description。 #[test] - fn prop_system_prompt_contains_all_tool_names(tools in arb_tool_definitions()) { + fn prop_full_system_prompt_contains_all_tool_names(tools in arb_tool_definitions()) { let generator = ToolPromptGenerator::new().with_format(PromptFormat::Xml); - let prompt = generator.generate_system_prompt(&tools); + let prompt = generator.generate_full_system_prompt(&tools); // 验证所有工具名称都在 prompt 中 for tool in &tools { @@ -514,14 +544,14 @@ mod proptests { } } - /// **Feature: agent-tool-calling, Property 3: System Prompt 工具包含 - 描述** + /// **Feature: agent-tool-calling, Property 3: Full System Prompt 工具包含 - 描述** /// **Validates: Requirements 2.3** /// - /// *For any* 已注册的工具集合,生成的 System Prompt 应该包含所有工具的 description。 + /// *For any* 已注册的工具集合,生成的完整 System Prompt 应该包含所有工具的 description。 #[test] - fn prop_system_prompt_contains_all_tool_descriptions(tools in arb_tool_definitions()) { + fn prop_full_system_prompt_contains_all_tool_descriptions(tools in arb_tool_definitions()) { let generator = ToolPromptGenerator::new().with_format(PromptFormat::Xml); - let prompt = generator.generate_system_prompt(&tools); + let prompt = generator.generate_full_system_prompt(&tools); // 验证所有工具描述都在 prompt 中 for tool in &tools { @@ -534,14 +564,14 @@ mod proptests { } } - /// **Feature: agent-tool-calling, Property 3: System Prompt 工具包含 - JSON 格式** + /// **Feature: agent-tool-calling, Property 3: Full System Prompt 工具包含 - JSON 格式** /// **Validates: Requirements 2.3** /// - /// *For any* 已注册的工具集合,JSON 格式的 System Prompt 也应该包含所有工具的 name 和 description。 + /// *For any* 已注册的工具集合,JSON 格式的完整 System Prompt 也应该包含所有工具的 name 和 description。 #[test] - fn prop_system_prompt_json_contains_all_tools(tools in arb_tool_definitions()) { + fn prop_full_system_prompt_json_contains_all_tools(tools in arb_tool_definitions()) { let generator = ToolPromptGenerator::new().with_format(PromptFormat::Json); - let prompt = generator.generate_system_prompt(&tools); + let prompt = generator.generate_full_system_prompt(&tools); // 验证所有工具名称和描述都在 prompt 中 for tool in &tools { @@ -560,18 +590,18 @@ mod proptests { } } - /// **Feature: agent-tool-calling, Property 3: System Prompt 工具包含 - 格式一致性** + /// **Feature: agent-tool-calling, Property 3: Full System Prompt 工具包含 - 格式一致性** /// **Validates: Requirements 2.3** /// /// *For any* 已注册的工具集合,无论使用 XML 还是 JSON 格式, - /// 生成的 System Prompt 都应该包含相同的工具信息。 + /// 生成的完整 System Prompt 都应该包含相同的工具信息。 #[test] - fn prop_system_prompt_format_consistency(tools in arb_tool_definitions()) { + fn prop_full_system_prompt_format_consistency(tools in arb_tool_definitions()) { let xml_generator = ToolPromptGenerator::new().with_format(PromptFormat::Xml); let json_generator = ToolPromptGenerator::new().with_format(PromptFormat::Json); - let xml_prompt = xml_generator.generate_system_prompt(&tools); - let json_prompt = json_generator.generate_system_prompt(&tools); + let xml_prompt = xml_generator.generate_full_system_prompt(&tools); + let json_prompt = json_generator.generate_full_system_prompt(&tools); // 两种格式都应该包含所有工具名称和描述 for tool in &tools { @@ -588,15 +618,15 @@ mod proptests { } } - /// **Feature: agent-tool-calling, Property 3: System Prompt 工具包含 - 工具数量** + /// **Feature: agent-tool-calling, Property 3: Full System Prompt 工具包含 - 工具数量** /// **Validates: Requirements 2.3** /// - /// *For any* 已注册的工具集合,生成的 System Prompt 中工具名称出现的次数 + /// *For any* 已注册的工具集合,生成的完整 System Prompt 中工具名称出现的次数 /// 应该至少等于工具数量(每个工具至少出现一次)。 #[test] - fn prop_system_prompt_tool_count(tools in arb_tool_definitions()) { + fn prop_full_system_prompt_tool_count(tools in arb_tool_definitions()) { let generator = ToolPromptGenerator::new().with_format(PromptFormat::Xml); - let prompt = generator.generate_system_prompt(&tools); + let prompt = generator.generate_full_system_prompt(&tools); // 统计每个工具名称在 prompt 中出现的次数 for tool in &tools { diff --git a/src-tauri/src/agent/tools/registry.rs b/src-tauri/src/agent/tools/registry.rs index 6ced93a55..fa1f68ee2 100644 --- a/src-tauri/src/agent/tools/registry.rs +++ b/src-tauri/src/agent/tools/registry.rs @@ -160,6 +160,15 @@ impl ToolRegistry { self.tools.read().values().map(|t| t.definition()).collect() } + /// 获取所有工具定义(OpenAI API 格式) + pub fn list_definitions_api(&self) -> Vec { + self.tools + .read() + .values() + .map(|t| t.definition().to_api_format()) + .collect() + } + /// 获取所有工具名称 pub fn list_names(&self) -> Vec { self.tools.read().keys().cloned().collect() diff --git a/src-tauri/src/agent/tools/types.rs b/src-tauri/src/agent/tools/types.rs index e7085ac6e..c9381c5fe 100644 --- a/src-tauri/src/agent/tools/types.rs +++ b/src-tauri/src/agent/tools/types.rs @@ -48,6 +48,17 @@ impl ToolDefinition { self.parameters.validate()?; Ok(()) } + + /// 转换为 OpenAI API 格式的工具定义 + pub fn to_api_format(&self) -> crate::models::openai::Tool { + crate::models::openai::Tool::Function { + function: crate::models::openai::FunctionDef { + name: self.name.clone(), + description: Some(self.description.clone()), + parameters: Some(serde_json::to_value(&self.parameters).unwrap_or_default()), + }, + } + } } /// JSON Schema 参数定义 diff --git a/src-tauri/src/agent/types.rs b/src-tauri/src/agent/types.rs index 8512832fd..bdf21a727 100644 --- a/src-tauri/src/agent/types.rs +++ b/src-tauri/src/agent/types.rs @@ -5,6 +5,74 @@ use serde::{Deserialize, Serialize}; +/// Provider 类型枚举 +/// +/// 决定使用哪种 API 协议 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)] +#[serde(rename_all = "snake_case")] +pub enum ProviderType { + /// Claude (Anthropic 协议) + Claude, + /// Claude OAuth (Anthropic 协议) + ClaudeOauth, + /// Kiro/CodeWhisperer (AWS Event Stream 协议) + Kiro, + /// Gemini (Gemini 协议) + Gemini, + /// OpenAI 及其兼容服务 (默认) + #[default] + OpenAI, + /// 通义千问 (OpenAI 兼容) + Qwen, + /// Codex (OpenAI 兼容) + Codex, + /// Antigravity (OpenAI 兼容) + Antigravity, + /// iFlow (OpenAI 兼容) + IFlow, +} + +impl ProviderType { + /// 从字符串解析 provider 类型 + pub fn from_str(s: &str) -> Self { + match s.to_lowercase().as_str() { + "claude" => Self::Claude, + "claude_oauth" => Self::ClaudeOauth, + "kiro" => Self::Kiro, + "gemini" => Self::Gemini, + "openai" => Self::OpenAI, + "qwen" => Self::Qwen, + "codex" => Self::Codex, + "antigravity" => Self::Antigravity, + "iflow" => Self::IFlow, + _ => Self::OpenAI, // 默认使用 OpenAI 协议 + } + } + + /// 获取 API 端点路径 + pub fn endpoint(&self) -> &'static str { + match self { + Self::Claude | Self::ClaudeOauth => "/v1/messages", + Self::Kiro => "/v1/chat/completions", // Kiro 使用 OpenAI 兼容格式,但后端会转换 + Self::Gemini => "/v1/gemini/chat/completions", + _ => "/v1/chat/completions", + } + } + + /// 是否使用 Anthropic 协议 + pub fn is_anthropic(&self) -> bool { + matches!(self, Self::Claude | Self::ClaudeOauth) + } + + /// 是否使用 OpenAI 兼容协议 + pub fn is_openai_compatible(&self) -> bool { + matches!( + self, + Self::OpenAI | Self::Qwen | Self::Codex | Self::Antigravity | Self::IFlow | Self::Kiro + ) + } +} + /// Agent 会话状态 #[derive(Debug, Clone, Serialize, Deserialize)] pub struct AgentSession { @@ -239,6 +307,9 @@ pub enum StreamEvent { tool_name: String, /// 工具调用 ID tool_id: String, + /// 工具参数(JSON 字符串) + #[serde(skip_serializing_if = "Option::is_none")] + arguments: Option, }, /// 工具调用结束 @@ -251,11 +322,16 @@ pub enum StreamEvent { result: ToolExecutionResult, }, - /// 完成 + /// 完成(单次 API 响应完成,工具循环可能继续) /// Requirements: 1.3 - THE Streaming_Handler SHALL emit a done event with token usage statistics #[serde(rename = "done")] Done { usage: Option }, + /// 最终完成(整个对话完成,包括所有工具调用循环) + /// 前端收到此事件后才能取消监听 + #[serde(rename = "final_done")] + FinalDone { usage: Option }, + /// 错误 /// Requirements: 1.4 - IF a streaming error occurs, THEN THE Streaming_Handler SHALL emit an error event #[serde(rename = "error")] diff --git a/src-tauri/src/backends/mod.rs b/src-tauri/src/backends/mod.rs new file mode 100644 index 000000000..8c1690953 --- /dev/null +++ b/src-tauri/src/backends/mod.rs @@ -0,0 +1,43 @@ +//! 后端调用层 +//! +//! 提供与各种 AI 后端服务的 HTTP 通信能力。 +//! 后端层只负责 HTTP 请求/响应,不包含任何协议转换逻辑。 +//! +//! # 架构设计 +//! +//! ```text +//! backends/ +//! ├── traits.rs # Backend trait 定义 +//! ├── kiro.rs # Kiro/CodeWhisperer 后端 (待迁移) +//! ├── codex.rs # Codex 后端 (待迁移) +//! └── claude.rs # Claude API 后端 (待迁移) +//! ``` +//! +//! # 职责说明 +//! +//! - **只做 HTTP 调用**: 构建 HTTP 请求,发送,接收响应 +//! - **不做协议转换**: 协议转换在 translator 层完成 +//! - **处理认证**: 管理 access token,刷新过期凭证 +//! - **处理重试**: 可选的重试逻辑 +//! +//! # 使用示例 +//! +//! ```ignore +//! use proxycast::backends::KiroBackend; +//! use proxycast::backends::traits::Backend; +//! +//! let backend = KiroBackend::new(credentials); +//! let response = backend.call_stream(&cw_request).await?; +//! ``` + +pub mod traits; + +// 重新导出核心类型 +pub use traits::{ + AuthenticatedBackend, Backend, BackendError, BackendErrorKind, BackendResult, ByteStream, +}; + +// TODO: 后续阶段迁移以下后端 +// pub mod kiro; +// pub mod codex; +// pub mod claude; diff --git a/src-tauri/src/backends/traits.rs b/src-tauri/src/backends/traits.rs new file mode 100644 index 000000000..9ef73bf87 --- /dev/null +++ b/src-tauri/src/backends/traits.rs @@ -0,0 +1,194 @@ +//! 后端调用层 Trait 定义 +//! +//! 定义后端 HTTP 调用的核心接口。 +//! 后端层只负责 HTTP 请求/响应,不包含任何协议转换逻辑。 + +use async_trait::async_trait; +use bytes::Bytes; +use futures::Stream; +use std::error::Error; +use std::pin::Pin; + +/// 字节流类型 +pub type ByteStream = + Pin>> + Send>>; + +/// 后端调用结果 +pub type BackendResult = Result; + +/// 后端错误类型 +#[derive(Debug, Clone)] +pub struct BackendError { + /// 错误类型 + pub kind: BackendErrorKind, + /// 错误消息 + pub message: String, + /// HTTP 状态码(如果有) + pub status_code: Option, +} + +impl std::fmt::Display for BackendError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + if let Some(code) = self.status_code { + write!(f, "{} ({}): {}", self.kind, code, self.message) + } else { + write!(f, "{}: {}", self.kind, self.message) + } + } +} + +impl std::error::Error for BackendError {} + +/// 后端错误类型枚举 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum BackendErrorKind { + /// 认证错误 + AuthenticationError, + /// 网络错误 + NetworkError, + /// 请求超时 + Timeout, + /// 服务端错误 + ServerError, + /// 请求格式错误 + BadRequest, + /// 速率限制 + RateLimited, + /// 其他错误 + Other, +} + +impl std::fmt::Display for BackendErrorKind { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::AuthenticationError => write!(f, "AuthenticationError"), + Self::NetworkError => write!(f, "NetworkError"), + Self::Timeout => write!(f, "Timeout"), + Self::ServerError => write!(f, "ServerError"), + Self::BadRequest => write!(f, "BadRequest"), + Self::RateLimited => write!(f, "RateLimited"), + Self::Other => write!(f, "Other"), + } + } +} + +impl BackendError { + /// 创建新的后端错误 + pub fn new(kind: BackendErrorKind, message: impl Into) -> Self { + Self { + kind, + message: message.into(), + status_code: None, + } + } + + /// 带 HTTP 状态码创建错误 + pub fn with_status(kind: BackendErrorKind, message: impl Into, status: u16) -> Self { + Self { + kind, + message: message.into(), + status_code: Some(status), + } + } + + /// 从 HTTP 状态码推断错误类型 + pub fn from_status(status: u16, message: impl Into) -> Self { + let kind = match status { + 401 | 403 => BackendErrorKind::AuthenticationError, + 400 => BackendErrorKind::BadRequest, + 429 => BackendErrorKind::RateLimited, + 500..=599 => BackendErrorKind::ServerError, + _ => BackendErrorKind::Other, + }; + Self::with_status(kind, message, status) + } + + /// 是否可重试 + pub fn is_retryable(&self) -> bool { + matches!( + self.kind, + BackendErrorKind::NetworkError + | BackendErrorKind::Timeout + | BackendErrorKind::ServerError + | BackendErrorKind::RateLimited + ) + } +} + +/// 后端 Trait +/// +/// 定义后端 HTTP 调用的接口。 +#[async_trait] +pub trait Backend: Send + Sync { + /// 后端请求类型 + type Request: Send; + + /// 非流式调用 + /// + /// # 返回 + /// + /// 响应的原始字节 + async fn call(&self, request: &Self::Request) -> BackendResult; + + /// 流式调用 + /// + /// # 返回 + /// + /// 字节流 + async fn call_stream(&self, request: &Self::Request) -> BackendResult; + + /// 获取后端名称 + fn name(&self) -> &str; + + /// 检查后端是否可用 + async fn is_available(&self) -> bool { + true + } +} + +/// 带认证的后端 Trait +#[async_trait] +pub trait AuthenticatedBackend: Backend { + /// 刷新认证凭证 + async fn refresh_credentials(&mut self) -> BackendResult<()>; + + /// 检查凭证是否有效 + fn credentials_valid(&self) -> bool; +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_backend_error_display() { + let err = BackendError::new(BackendErrorKind::NetworkError, "connection refused"); + assert_eq!(format!("{}", err), "NetworkError: connection refused"); + + let err = BackendError::with_status(BackendErrorKind::ServerError, "internal error", 500); + assert_eq!(format!("{}", err), "ServerError (500): internal error"); + } + + #[test] + fn test_backend_error_from_status() { + let err = BackendError::from_status(401, "unauthorized"); + assert_eq!(err.kind, BackendErrorKind::AuthenticationError); + assert_eq!(err.status_code, Some(401)); + + let err = BackendError::from_status(429, "too many requests"); + assert_eq!(err.kind, BackendErrorKind::RateLimited); + + let err = BackendError::from_status(503, "service unavailable"); + assert_eq!(err.kind, BackendErrorKind::ServerError); + } + + #[test] + fn test_backend_error_retryable() { + assert!(BackendError::new(BackendErrorKind::NetworkError, "").is_retryable()); + assert!(BackendError::new(BackendErrorKind::Timeout, "").is_retryable()); + assert!(BackendError::new(BackendErrorKind::ServerError, "").is_retryable()); + assert!(BackendError::new(BackendErrorKind::RateLimited, "").is_retryable()); + assert!(!BackendError::new(BackendErrorKind::AuthenticationError, "").is_retryable()); + assert!(!BackendError::new(BackendErrorKind::BadRequest, "").is_retryable()); + } +} diff --git a/src-tauri/src/commands/agent_cmd.rs b/src-tauri/src/commands/agent_cmd.rs index f3eacb70b..fb3f96c0a 100644 --- a/src-tauri/src/commands/agent_cmd.rs +++ b/src-tauri/src/commands/agent_cmd.rs @@ -2,7 +2,7 @@ //! //! 提供原生 Agent 的 Tauri 命令(兼容旧 API) -use crate::agent::{ImageData, NativeAgentState, NativeChatRequest}; +use crate::agent::{ImageData, NativeAgentState, NativeChatRequest, ProviderType}; use crate::AppState; use serde::{Deserialize, Serialize}; use tauri::State; @@ -34,12 +34,13 @@ pub async fn agent_start_process( ) -> Result { tracing::info!("[Agent] 初始化原生 Agent"); - let (port, api_key, running) = { + let (port, api_key, running, default_provider) = { let state = app_state.read().await; ( state.config.server.port, state.running_api_key.clone(), state.running, + state.config.routing.default_provider.clone(), ) }; @@ -49,8 +50,9 @@ pub async fn agent_start_process( let api_key = api_key.ok_or_else(|| "ProxyCast API Server 未配置 API Key".to_string())?; let base_url = format!("http://127.0.0.1:{}", port); + let provider_type = ProviderType::from_str(&default_provider); - agent_state.init(base_url.clone(), api_key)?; + agent_state.init(base_url.clone(), api_key, provider_type)?; Ok(AgentProcessStatus { running: true, @@ -118,12 +120,13 @@ pub async fn agent_create_session( // 如果未初始化,自动初始化 if !agent_state.is_initialized() { - let (port, api_key, running) = { + let (port, api_key, running, default_provider) = { let state = app_state.read().await; ( state.config.server.port, state.running_api_key.clone(), state.running, + state.config.routing.default_provider.clone(), ) }; @@ -133,7 +136,8 @@ pub async fn agent_create_session( let api_key = api_key.ok_or_else(|| "未配置 API Key".to_string())?; let base_url = format!("http://127.0.0.1:{}", port); - agent_state.init(base_url, api_key)?; + let provider_type = ProviderType::from_str(&default_provider); + agent_state.init(base_url, api_key, provider_type)?; } // 构建包含 Skills 的 System Prompt @@ -222,12 +226,13 @@ pub async fn agent_send_message( // 如果未初始化,自动初始化 if !agent_state.is_initialized() { - let (port, api_key, running) = { + let (port, api_key, running, default_provider) = { let state = app_state.read().await; ( state.config.server.port, state.running_api_key.clone(), state.running, + state.config.routing.default_provider.clone(), ) }; @@ -237,7 +242,8 @@ pub async fn agent_send_message( let api_key = api_key.ok_or_else(|| "未配置 API Key".to_string())?; let base_url = format!("http://127.0.0.1:{}", port); - agent_state.init(base_url, api_key)?; + let provider_type = ProviderType::from_str(&default_provider); + agent_state.init(base_url, api_key, provider_type)?; } // 根据启用的模式构建最终消息 diff --git a/src-tauri/src/commands/api_key_provider_cmd.rs b/src-tauri/src/commands/api_key_provider_cmd.rs new file mode 100644 index 000000000..739fb152f --- /dev/null +++ b/src-tauri/src/commands/api_key_provider_cmd.rs @@ -0,0 +1,398 @@ +//! API Key Provider Tauri 命令 +//! +//! 提供 API Key Provider 管理的前端调用接口。 +//! +//! **Feature: provider-ui-refactor** +//! **Validates: Requirements 9.1** + +use crate::database::dao::api_key_provider::{ + ApiKeyEntry, ApiKeyProvider, ApiProviderType, ProviderGroup, ProviderWithKeys, +}; +use crate::database::DbConnection; +use crate::services::api_key_provider_service::{ApiKeyProviderService, ImportResult}; +use serde::{Deserialize, Serialize}; +use std::sync::Arc; +use tauri::State; + +/// API Key Provider 服务状态封装 +pub struct ApiKeyProviderServiceState(pub Arc); + +// ============================================================================ +// 请求/响应类型 +// ============================================================================ + +/// 添加自定义 Provider 请求 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct AddCustomProviderRequest { + pub name: String, + #[serde(rename = "type")] + pub provider_type: String, + pub api_host: String, + pub api_version: Option, + pub project: Option, + pub location: Option, + pub region: Option, +} + +/// 更新 Provider 请求 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct UpdateProviderRequest { + pub name: Option, + pub api_host: Option, + pub enabled: Option, + pub sort_order: Option, + pub api_version: Option, + pub project: Option, + pub location: Option, + pub region: Option, +} + +/// 添加 API Key 请求 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct AddApiKeyRequest { + pub provider_id: String, + pub api_key: String, + pub alias: Option, +} + +/// Provider 显示数据(用于前端) +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ProviderDisplay { + pub id: String, + pub name: String, + #[serde(rename = "type")] + pub provider_type: String, + pub api_host: String, + pub is_system: bool, + pub group: String, + pub enabled: bool, + pub sort_order: i32, + pub api_version: Option, + pub project: Option, + pub location: Option, + pub region: Option, + pub api_key_count: usize, + pub created_at: String, + pub updated_at: String, +} + +/// API Key 显示数据(用于前端,掩码显示) +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ApiKeyDisplay { + pub id: String, + pub provider_id: String, + /// 掩码后的 API Key + pub api_key_masked: String, + pub alias: Option, + pub enabled: bool, + pub usage_count: i64, + pub error_count: i64, + pub last_used_at: Option, + pub created_at: String, +} + +/// Provider 完整显示数据(包含 API Keys) +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ProviderWithKeysDisplay { + #[serde(flatten)] + pub provider: ProviderDisplay, + pub api_keys: Vec, +} + +// ============================================================================ +// 辅助函数 +// ============================================================================ + +/// 将 API Key 转换为掩码显示 +fn mask_api_key(key: &str) -> String { + let chars: Vec = key.chars().collect(); + if chars.len() <= 12 { + "****".to_string() + } else { + let prefix: String = chars[..6].iter().collect(); + let suffix: String = chars[chars.len() - 4..].iter().collect(); + format!("{}****{}", prefix, suffix) + } +} + +/// 将 ApiKeyProvider 转换为 ProviderDisplay +fn provider_to_display(provider: &ApiKeyProvider, api_key_count: usize) -> ProviderDisplay { + ProviderDisplay { + id: provider.id.clone(), + name: provider.name.clone(), + provider_type: provider.provider_type.to_string(), + api_host: provider.api_host.clone(), + is_system: provider.is_system, + group: provider.group.to_string(), + enabled: provider.enabled, + sort_order: provider.sort_order, + api_version: provider.api_version.clone(), + project: provider.project.clone(), + location: provider.location.clone(), + region: provider.region.clone(), + api_key_count, + created_at: provider.created_at.to_rfc3339(), + updated_at: provider.updated_at.to_rfc3339(), + } +} + +/// 将 ApiKeyEntry 转换为 ApiKeyDisplay(需要解密后掩码) +fn api_key_to_display(key: &ApiKeyEntry, service: &ApiKeyProviderService) -> ApiKeyDisplay { + // 解密后掩码显示 + let masked = match service.decrypt_api_key(&key.api_key_encrypted) { + Ok(decrypted) => mask_api_key(&decrypted), + Err(_) => "****".to_string(), + }; + + ApiKeyDisplay { + id: key.id.clone(), + provider_id: key.provider_id.clone(), + api_key_masked: masked, + alias: key.alias.clone(), + enabled: key.enabled, + usage_count: key.usage_count, + error_count: key.error_count, + last_used_at: key.last_used_at.map(|t| t.to_rfc3339()), + created_at: key.created_at.to_rfc3339(), + } +} + +/// 将 ProviderWithKeys 转换为 ProviderWithKeysDisplay +fn provider_with_keys_to_display( + pwk: &ProviderWithKeys, + service: &ApiKeyProviderService, +) -> ProviderWithKeysDisplay { + let api_keys: Vec = pwk + .api_keys + .iter() + .map(|k| api_key_to_display(k, service)) + .collect(); + + ProviderWithKeysDisplay { + provider: provider_to_display(&pwk.provider, pwk.api_keys.len()), + api_keys, + } +} + +// ============================================================================ +// Tauri 命令 +// ============================================================================ + +/// 获取所有 API Key Provider(包含 API Keys) +#[tauri::command] +pub fn get_api_key_providers( + db: State<'_, DbConnection>, + service: State<'_, ApiKeyProviderServiceState>, +) -> Result, String> { + let providers = service.0.get_all_providers(&db)?; + Ok(providers + .iter() + .map(|p| provider_with_keys_to_display(p, &service.0)) + .collect()) +} + +/// 获取单个 API Key Provider(包含 API Keys) +#[tauri::command] +pub fn get_api_key_provider( + db: State<'_, DbConnection>, + service: State<'_, ApiKeyProviderServiceState>, + id: String, +) -> Result, String> { + let provider = service.0.get_provider(&db, &id)?; + Ok(provider.map(|p| provider_with_keys_to_display(&p, &service.0))) +} + +/// 添加自定义 Provider +#[tauri::command] +pub fn add_custom_api_key_provider( + db: State<'_, DbConnection>, + service: State<'_, ApiKeyProviderServiceState>, + request: AddCustomProviderRequest, +) -> Result { + let provider_type: ApiProviderType = request + .provider_type + .parse() + .map_err(|e: String| format!("无效的 Provider 类型: {}", e))?; + + let provider = service.0.add_custom_provider( + &db, + request.name, + provider_type, + request.api_host, + request.api_version, + request.project, + request.location, + request.region, + )?; + + Ok(provider_to_display(&provider, 0)) +} + +/// 更新 Provider 配置 +#[tauri::command] +pub fn update_api_key_provider( + db: State<'_, DbConnection>, + service: State<'_, ApiKeyProviderServiceState>, + id: String, + request: UpdateProviderRequest, +) -> Result { + let provider = service.0.update_provider( + &db, + &id, + request.name, + request.api_host, + request.enabled, + request.sort_order, + request.api_version, + request.project, + request.location, + request.region, + )?; + + // 获取 API Key 数量 + let full_provider = service.0.get_provider(&db, &id)?; + let api_key_count = full_provider.map(|p| p.api_keys.len()).unwrap_or(0); + + Ok(provider_to_display(&provider, api_key_count)) +} + +/// 删除自定义 Provider +#[tauri::command] +pub fn delete_custom_api_key_provider( + db: State<'_, DbConnection>, + service: State<'_, ApiKeyProviderServiceState>, + id: String, +) -> Result { + service.0.delete_custom_provider(&db, &id) +} + +/// 添加 API Key +#[tauri::command] +pub fn add_api_key( + db: State<'_, DbConnection>, + service: State<'_, ApiKeyProviderServiceState>, + request: AddApiKeyRequest, +) -> Result { + let key = service + .0 + .add_api_key(&db, &request.provider_id, &request.api_key, request.alias)?; + + Ok(api_key_to_display(&key, &service.0)) +} + +/// 删除 API Key +#[tauri::command] +pub fn delete_api_key( + db: State<'_, DbConnection>, + service: State<'_, ApiKeyProviderServiceState>, + key_id: String, +) -> Result { + service.0.delete_api_key(&db, &key_id) +} + +/// 切换 API Key 启用状态 +#[tauri::command] +pub fn toggle_api_key( + db: State<'_, DbConnection>, + service: State<'_, ApiKeyProviderServiceState>, + key_id: String, + enabled: bool, +) -> Result { + let key = service.0.toggle_api_key(&db, &key_id, enabled)?; + Ok(api_key_to_display(&key, &service.0)) +} + +/// 更新 API Key 别名 +#[tauri::command] +pub fn update_api_key_alias( + db: State<'_, DbConnection>, + service: State<'_, ApiKeyProviderServiceState>, + key_id: String, + alias: Option, +) -> Result { + let key = service.0.update_api_key_alias(&db, &key_id, alias)?; + Ok(api_key_to_display(&key, &service.0)) +} + +/// 获取下一个可用的 API Key(用于 API 调用) +#[tauri::command] +pub fn get_next_api_key( + db: State<'_, DbConnection>, + service: State<'_, ApiKeyProviderServiceState>, + provider_id: String, +) -> Result, String> { + service.0.get_next_api_key(&db, &provider_id) +} + +/// 记录 API Key 使用 +#[tauri::command] +pub fn record_api_key_usage( + db: State<'_, DbConnection>, + service: State<'_, ApiKeyProviderServiceState>, + key_id: String, +) -> Result<(), String> { + service.0.record_usage(&db, &key_id) +} + +/// 记录 API Key 错误 +#[tauri::command] +pub fn record_api_key_error( + db: State<'_, DbConnection>, + service: State<'_, ApiKeyProviderServiceState>, + key_id: String, +) -> Result<(), String> { + service.0.record_error(&db, &key_id) +} + +/// 获取 UI 状态 +#[tauri::command] +pub fn get_provider_ui_state( + db: State<'_, DbConnection>, + service: State<'_, ApiKeyProviderServiceState>, + key: String, +) -> Result, String> { + service.0.get_ui_state(&db, &key) +} + +/// 设置 UI 状态 +#[tauri::command] +pub fn set_provider_ui_state( + db: State<'_, DbConnection>, + service: State<'_, ApiKeyProviderServiceState>, + key: String, + value: String, +) -> Result<(), String> { + service.0.set_ui_state(&db, &key, &value) +} + +/// 批量更新 Provider 排序顺序 +/// **Validates: Requirements 8.4** +#[tauri::command] +pub fn update_provider_sort_orders( + db: State<'_, DbConnection>, + service: State<'_, ApiKeyProviderServiceState>, + sort_orders: Vec<(String, i32)>, +) -> Result<(), String> { + service.0.update_provider_sort_orders(&db, sort_orders) +} + +/// 导出 Provider 配置 +#[tauri::command] +pub fn export_api_key_providers( + db: State<'_, DbConnection>, + service: State<'_, ApiKeyProviderServiceState>, + include_keys: bool, +) -> Result { + let config = service.0.export_config(&db, include_keys)?; + serde_json::to_string_pretty(&config).map_err(|e| format!("序列化失败: {}", e)) +} + +/// 导入 Provider 配置 +#[tauri::command] +pub fn import_api_key_providers( + db: State<'_, DbConnection>, + service: State<'_, ApiKeyProviderServiceState>, + config_json: String, +) -> Result { + service.0.import_config(&db, &config_json) +} diff --git a/src-tauri/src/commands/mod.rs b/src-tauri/src/commands/mod.rs index 975d9c587..21a20e87d 100644 --- a/src-tauri/src/commands/mod.rs +++ b/src-tauri/src/commands/mod.rs @@ -1,4 +1,5 @@ pub mod agent_cmd; +pub mod api_key_provider_cmd; pub mod auto_fix_cmd; pub mod browser_interceptor_cmd; pub mod config_cmd; diff --git a/src-tauri/src/commands/native_agent_cmd.rs b/src-tauri/src/commands/native_agent_cmd.rs index 3076fe137..2beb2826c 100644 --- a/src-tauri/src/commands/native_agent_cmd.rs +++ b/src-tauri/src/commands/native_agent_cmd.rs @@ -3,8 +3,8 @@ //! 提供原生 Rust Agent 的 Tauri 命令,替代 aster sidecar 方案 use crate::agent::{ - AgentSession, ImageData, NativeAgent, NativeAgentState, NativeChatRequest, NativeChatResponse, - StreamEvent, + AgentSession, ImageData, NativeAgentState, NativeChatRequest, NativeChatResponse, ProviderType, + StreamEvent, ToolLoopEngine, }; use crate::AppState; use serde::{Deserialize, Serialize}; @@ -24,12 +24,13 @@ pub async fn native_agent_init( ) -> Result { tracing::info!("[NativeAgent] 初始化 Agent"); - let (port, api_key, running) = { + let (port, api_key, running, default_provider) = { let state = app_state.read().await; ( state.config.server.port, state.running_api_key.clone(), state.running, + state.config.routing.default_provider.clone(), ) }; @@ -40,8 +41,15 @@ pub async fn native_agent_init( let api_key = api_key.ok_or_else(|| "ProxyCast API Server 未配置 API Key".to_string())?; let base_url = format!("http://127.0.0.1:{}", port); + let provider_type = ProviderType::from_str(&default_provider); - agent_state.init(base_url.clone(), api_key)?; + tracing::info!( + "[NativeAgent] 初始化 Agent: base_url={}, provider={:?}", + base_url, + provider_type + ); + + agent_state.init(base_url.clone(), api_key, provider_type)?; tracing::info!("[NativeAgent] Agent 初始化成功: {}", base_url); @@ -90,12 +98,13 @@ pub async fn native_agent_chat( // 如果 Agent 未初始化,自动初始化 if !agent_state.is_initialized() { - let (port, api_key, running) = { + let (port, api_key, running, default_provider) = { let state = app_state.read().await; ( state.config.server.port, state.running_api_key.clone(), state.running, + state.config.routing.default_provider.clone(), ) }; @@ -105,7 +114,8 @@ pub async fn native_agent_chat( let api_key = api_key.ok_or_else(|| "未配置 API Key".to_string())?; let base_url = format!("http://127.0.0.1:{}", port); - agent_state.init(base_url, api_key)?; + let provider_type = ProviderType::from_str(&default_provider); + agent_state.init(base_url, api_key, provider_type)?; } let request = NativeChatRequest { @@ -133,25 +143,28 @@ pub async fn native_agent_chat_stream( agent_state: State<'_, NativeAgentState>, app_state: State<'_, AppState>, message: String, + event_name: String, + session_id: Option, model: Option, images: Option>, - event_name: String, ) -> Result<(), String> { tracing::info!( - "[NativeAgent] 发送流式消息: message_len={}, model={:?}, event={}", + "[NativeAgent] 发送流式消息: message_len={}, model={:?}, event={}, session={:?}", message.len(), model, - event_name + event_name, + session_id ); // 如果 Agent 未初始化,自动初始化 if !agent_state.is_initialized() { - let (port, api_key, running) = { + let (port, api_key, running, default_provider) = { let state = app_state.read().await; ( state.config.server.port, state.running_api_key.clone(), state.running, + state.config.routing.default_provider.clone(), ) }; @@ -161,22 +174,15 @@ pub async fn native_agent_chat_stream( let api_key = api_key.ok_or_else(|| "未配置 API Key".to_string())?; let base_url = format!("http://127.0.0.1:{}", port); - agent_state.init(base_url, api_key)?; + let provider_type = ProviderType::from_str(&default_provider); + agent_state.init(base_url, api_key, provider_type)?; } - // 获取配置用于创建独立的 Agent - let (base_url, api_key) = { - let state = app_state.read().await; - let base_url = format!("http://127.0.0.1:{}", state.config.server.port); - let api_key = state - .running_api_key - .clone() - .ok_or_else(|| "未配置 API Key".to_string())?; - (base_url, api_key) - }; + // 获取工具注册表(用于创建 ToolLoopEngine) + let tool_registry = agent_state.get_tool_registry()?; let request = NativeChatRequest { - session_id: None, + session_id, // 使用前端传递的 session_id 以保持上下文 message, model, images: images.map(|imgs| { @@ -190,27 +196,40 @@ pub async fn native_agent_chat_stream( stream: true, }; + // 克隆 agent_state 用于后台任务(共享 sessions) + let agent_state_clone = agent_state.inner().clone(); + // 在后台任务中处理流式响应 let event_name_clone = event_name.clone(); + eprintln!( + "[native_agent_chat_stream] 启动后台任务, event_name={}", + event_name_clone + ); tauri::async_runtime::spawn(async move { - let agent = match NativeAgent::new(base_url, api_key) { - Ok(a) => a, - Err(e) => { - let _ = app_handle.emit( - &event_name_clone, - StreamEvent::Error { - message: e.to_string(), - }, - ); - return; - } - }; + eprintln!("[native_agent_chat_stream] 后台任务开始执行"); + + // 创建工具循环引擎(使用共享的 tool_registry) + let tool_loop_engine = ToolLoopEngine::new(tool_registry); + eprintln!("[native_agent_chat_stream] 工具循环引擎创建成功"); let (tx, mut rx) = mpsc::channel::(100); - let stream_task = tokio::spawn(async move { agent.chat_stream(request, tx).await }); + // 使用 agent_state 的方法(共享 sessions) + eprintln!( + "[native_agent_chat_stream] 开始 chat_stream_with_tools, request.session_id={:?}", + request.session_id + ); + let stream_task = tokio::spawn(async move { + agent_state_clone + .chat_stream_with_tools(request, tx, &tool_loop_engine) + .await + }); + eprintln!("[native_agent_chat_stream] 开始接收流式事件..."); + // 注意:不要在收到 Done 事件后立即 break,因为工具循环可能还在执行 + // 继续接收直到 channel 关闭(stream_task 完成) while let Some(event) = rx.recv().await { + eprintln!("[native_agent_chat_stream] 收到事件: {:?}", event); tracing::debug!( "[NativeAgent] 收到流式事件: {:?}, 发送到: {}", event, @@ -218,17 +237,30 @@ pub async fn native_agent_chat_stream( ); if let Err(e) = app_handle.emit(&event_name_clone, &event) { tracing::error!("[NativeAgent] 发送事件失败: {}", e); + eprintln!("[native_agent_chat_stream] 发送事件失败: {}", e); break; } tracing::debug!("[NativeAgent] 事件发送成功"); - if matches!(event, StreamEvent::Done { .. } | StreamEvent::Error { .. }) { - tracing::info!("[NativeAgent] 流式响应完成"); + // 只在 Error 时 break,Done 不 break 因为工具循环可能还会发送更多事件 + if matches!(event, StreamEvent::Error { .. }) { + tracing::info!("[NativeAgent] 流式响应错误,停止接收"); + eprintln!("[native_agent_chat_stream] 流式响应错误"); break; } } + eprintln!("[native_agent_chat_stream] channel 关闭,事件接收完成"); - let _ = stream_task.await; + eprintln!("[native_agent_chat_stream] 等待 stream_task 完成..."); + match stream_task.await { + Ok(result) => { + eprintln!("[native_agent_chat_stream] stream_task 完成: {:?}", result); + } + Err(e) => { + eprintln!("[native_agent_chat_stream] stream_task 错误: {}", e); + } + } + eprintln!("[native_agent_chat_stream] 后台任务结束"); }); Ok(()) diff --git a/src-tauri/src/config/export.rs b/src-tauri/src/config/export.rs index c4f5c5902..7ebf77ad2 100644 --- a/src-tauri/src/config/export.rs +++ b/src-tauri/src/config/export.rs @@ -501,6 +501,7 @@ mod base64 { } pub use self::base64::decode as base64_decode; +#[allow(unused_imports)] pub use self::base64::encode as base64_encode; #[cfg(test)] diff --git a/src-tauri/src/config/mod.rs b/src-tauri/src/config/mod.rs index 860580481..718506c45 100644 --- a/src-tauri/src/config/mod.rs +++ b/src-tauri/src/config/mod.rs @@ -3,6 +3,8 @@ //! 提供 YAML 配置文件支持、热重载和配置导入导出功能 //! 同时保持与旧版 JSON 配置的向后兼容性 +#![allow(unused_imports)] + mod export; mod hot_reload; mod import; diff --git a/src-tauri/src/database/dao/api_key_provider.rs b/src-tauri/src/database/dao/api_key_provider.rs new file mode 100644 index 000000000..21f0ec129 --- /dev/null +++ b/src-tauri/src/database/dao/api_key_provider.rs @@ -0,0 +1,609 @@ +//! API Key Provider 数据访问对象 +//! +//! 提供 API Key Provider 的 CRUD 操作。 +//! +//! **Feature: provider-ui-refactor** +//! **Validates: Requirements 9.1** + +use chrono::{DateTime, Utc}; +use rusqlite::{params, Connection}; +use serde::{Deserialize, Serialize}; + +// ============================================================================ +// 数据模型 +// ============================================================================ + +/// Provider API 类型 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "kebab-case")] +pub enum ApiProviderType { + Openai, + OpenaiResponse, + Anthropic, + Gemini, + AzureOpenai, + Vertexai, + AwsBedrock, + Ollama, + NewApi, + Gateway, +} + +impl std::fmt::Display for ApiProviderType { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + ApiProviderType::Openai => write!(f, "openai"), + ApiProviderType::OpenaiResponse => write!(f, "openai-response"), + ApiProviderType::Anthropic => write!(f, "anthropic"), + ApiProviderType::Gemini => write!(f, "gemini"), + ApiProviderType::AzureOpenai => write!(f, "azure-openai"), + ApiProviderType::Vertexai => write!(f, "vertexai"), + ApiProviderType::AwsBedrock => write!(f, "aws-bedrock"), + ApiProviderType::Ollama => write!(f, "ollama"), + ApiProviderType::NewApi => write!(f, "new-api"), + ApiProviderType::Gateway => write!(f, "gateway"), + } + } +} + +impl std::str::FromStr for ApiProviderType { + type Err = String; + + fn from_str(s: &str) -> Result { + match s.to_lowercase().as_str() { + "openai" => Ok(ApiProviderType::Openai), + "openai-response" => Ok(ApiProviderType::OpenaiResponse), + "anthropic" => Ok(ApiProviderType::Anthropic), + "gemini" => Ok(ApiProviderType::Gemini), + "azure-openai" => Ok(ApiProviderType::AzureOpenai), + "vertexai" => Ok(ApiProviderType::Vertexai), + "aws-bedrock" => Ok(ApiProviderType::AwsBedrock), + "ollama" => Ok(ApiProviderType::Ollama), + "new-api" => Ok(ApiProviderType::NewApi), + "gateway" => Ok(ApiProviderType::Gateway), + _ => Err(format!("Invalid provider type: {}", s)), + } + } +} + +/// Provider 分组类型 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum ProviderGroup { + Mainstream, + Chinese, + Cloud, + Aggregator, + Local, + Specialized, + Custom, +} + +impl std::fmt::Display for ProviderGroup { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + ProviderGroup::Mainstream => write!(f, "mainstream"), + ProviderGroup::Chinese => write!(f, "chinese"), + ProviderGroup::Cloud => write!(f, "cloud"), + ProviderGroup::Aggregator => write!(f, "aggregator"), + ProviderGroup::Local => write!(f, "local"), + ProviderGroup::Specialized => write!(f, "specialized"), + ProviderGroup::Custom => write!(f, "custom"), + } + } +} + +impl std::str::FromStr for ProviderGroup { + type Err = String; + + fn from_str(s: &str) -> Result { + match s.to_lowercase().as_str() { + "mainstream" => Ok(ProviderGroup::Mainstream), + "chinese" => Ok(ProviderGroup::Chinese), + "cloud" => Ok(ProviderGroup::Cloud), + "aggregator" => Ok(ProviderGroup::Aggregator), + "local" => Ok(ProviderGroup::Local), + "specialized" => Ok(ProviderGroup::Specialized), + "custom" => Ok(ProviderGroup::Custom), + _ => Err(format!("Invalid provider group: {}", s)), + } + } +} + +/// API Key Provider 配置 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ApiKeyProvider { + pub id: String, + pub name: String, + #[serde(rename = "type")] + pub provider_type: ApiProviderType, + pub api_host: String, + pub is_system: bool, + pub group: ProviderGroup, + pub enabled: bool, + pub sort_order: i32, + pub api_version: Option, + pub project: Option, + pub location: Option, + pub region: Option, + pub created_at: DateTime, + pub updated_at: DateTime, +} + +/// API Key 条目 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ApiKeyEntry { + pub id: String, + pub provider_id: String, + /// 加密后的 API Key + pub api_key_encrypted: String, + pub alias: Option, + pub enabled: bool, + pub usage_count: i64, + pub error_count: i64, + pub last_used_at: Option>, + pub created_at: DateTime, +} + +/// Provider 完整数据(包含 API Keys) +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ProviderWithKeys { + #[serde(flatten)] + pub provider: ApiKeyProvider, + pub api_keys: Vec, +} + +// ============================================================================ +// DAO 实现 +// ============================================================================ + +pub struct ApiKeyProviderDao; + +impl ApiKeyProviderDao { + // ==================== Provider 操作 ==================== + + /// 获取所有 Provider + pub fn get_all_providers(conn: &Connection) -> Result, rusqlite::Error> { + let mut stmt = conn.prepare( + "SELECT id, name, type, api_host, is_system, group_name, enabled, sort_order, + api_version, project, location, region, created_at, updated_at + FROM api_key_providers + ORDER BY sort_order ASC, created_at ASC", + )?; + + let rows = stmt.query_map([], Self::row_to_provider)?; + let mut providers = Vec::new(); + for provider in rows.flatten() { + providers.push(provider); + } + Ok(providers) + } + + /// 根据 ID 获取 Provider + pub fn get_provider_by_id( + conn: &Connection, + id: &str, + ) -> Result, rusqlite::Error> { + let mut stmt = conn.prepare( + "SELECT id, name, type, api_host, is_system, group_name, enabled, sort_order, + api_version, project, location, region, created_at, updated_at + FROM api_key_providers + WHERE id = ?1", + )?; + + let mut rows = stmt.query([id])?; + if let Some(row) = rows.next()? { + Ok(Some(Self::row_to_provider(row)?)) + } else { + Ok(None) + } + } + + /// 根据分组获取 Provider + pub fn get_providers_by_group( + conn: &Connection, + group: ProviderGroup, + ) -> Result, rusqlite::Error> { + let mut stmt = conn.prepare( + "SELECT id, name, type, api_host, is_system, group_name, enabled, sort_order, + api_version, project, location, region, created_at, updated_at + FROM api_key_providers + WHERE group_name = ?1 + ORDER BY sort_order ASC, created_at ASC", + )?; + + let rows = stmt.query_map([group.to_string()], Self::row_to_provider)?; + let mut providers = Vec::new(); + for provider in rows.flatten() { + providers.push(provider); + } + Ok(providers) + } + + /// 插入新 Provider + pub fn insert_provider( + conn: &Connection, + provider: &ApiKeyProvider, + ) -> Result<(), rusqlite::Error> { + conn.execute( + "INSERT INTO api_key_providers + (id, name, type, api_host, is_system, group_name, enabled, sort_order, + api_version, project, location, region, created_at, updated_at) + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14)", + params![ + provider.id, + provider.name, + provider.provider_type.to_string(), + provider.api_host, + provider.is_system, + provider.group.to_string(), + provider.enabled, + provider.sort_order, + provider.api_version, + provider.project, + provider.location, + provider.region, + provider.created_at.to_rfc3339(), + provider.updated_at.to_rfc3339(), + ], + )?; + Ok(()) + } + + /// 更新 Provider + pub fn update_provider( + conn: &Connection, + provider: &ApiKeyProvider, + ) -> Result<(), rusqlite::Error> { + conn.execute( + "UPDATE api_key_providers SET + name = ?2, type = ?3, api_host = ?4, is_system = ?5, group_name = ?6, + enabled = ?7, sort_order = ?8, api_version = ?9, project = ?10, + location = ?11, region = ?12, updated_at = ?13 + WHERE id = ?1", + params![ + provider.id, + provider.name, + provider.provider_type.to_string(), + provider.api_host, + provider.is_system, + provider.group.to_string(), + provider.enabled, + provider.sort_order, + provider.api_version, + provider.project, + provider.location, + provider.region, + provider.updated_at.to_rfc3339(), + ], + )?; + Ok(()) + } + + /// 删除 Provider(仅限自定义 Provider) + pub fn delete_provider(conn: &Connection, id: &str) -> Result { + // 先检查是否为系统 Provider + let is_system: bool = conn.query_row( + "SELECT is_system FROM api_key_providers WHERE id = ?1", + [id], + |row| row.get(0), + )?; + + if is_system { + return Ok(false); // 不允许删除系统 Provider + } + + let affected = conn.execute("DELETE FROM api_key_providers WHERE id = ?1", [id])?; + Ok(affected > 0) + } + + /// 从数据库行转换为 ApiKeyProvider + fn row_to_provider(row: &rusqlite::Row) -> Result { + let id: String = row.get(0)?; + let name: String = row.get(1)?; + let type_str: String = row.get(2)?; + let api_host: String = row.get(3)?; + let is_system: bool = row.get(4)?; + let group_str: String = row.get(5)?; + let enabled: bool = row.get(6)?; + let sort_order: i32 = row.get(7)?; + let api_version: Option = row.get(8)?; + let project: Option = row.get(9)?; + let location: Option = row.get(10)?; + let region: Option = row.get(11)?; + let created_at_str: String = row.get(12)?; + let updated_at_str: String = row.get(13)?; + + let provider_type: ApiProviderType = type_str.parse().unwrap_or(ApiProviderType::Openai); + let group: ProviderGroup = group_str.parse().unwrap_or(ProviderGroup::Custom); + + let created_at = DateTime::parse_from_rfc3339(&created_at_str) + .map(|dt| dt.with_timezone(&Utc)) + .unwrap_or_else(|_| Utc::now()); + let updated_at = DateTime::parse_from_rfc3339(&updated_at_str) + .map(|dt| dt.with_timezone(&Utc)) + .unwrap_or_else(|_| Utc::now()); + + Ok(ApiKeyProvider { + id, + name, + provider_type, + api_host, + is_system, + group, + enabled, + sort_order, + api_version, + project, + location, + region, + created_at, + updated_at, + }) + } + + // ==================== API Key 操作 ==================== + + /// 获取 Provider 的所有 API Keys + pub fn get_api_keys_by_provider( + conn: &Connection, + provider_id: &str, + ) -> Result, rusqlite::Error> { + let mut stmt = conn.prepare( + "SELECT id, provider_id, api_key_encrypted, alias, enabled, + usage_count, error_count, last_used_at, created_at + FROM api_keys + WHERE provider_id = ?1 + ORDER BY created_at ASC", + )?; + + let rows = stmt.query_map([provider_id], Self::row_to_api_key)?; + let mut keys = Vec::new(); + for key in rows.flatten() { + keys.push(key); + } + Ok(keys) + } + + /// 获取所有启用的 API Keys + pub fn get_enabled_api_keys_by_provider( + conn: &Connection, + provider_id: &str, + ) -> Result, rusqlite::Error> { + let mut stmt = conn.prepare( + "SELECT id, provider_id, api_key_encrypted, alias, enabled, + usage_count, error_count, last_used_at, created_at + FROM api_keys + WHERE provider_id = ?1 AND enabled = 1 + ORDER BY created_at ASC", + )?; + + let rows = stmt.query_map([provider_id], Self::row_to_api_key)?; + let mut keys = Vec::new(); + for key in rows.flatten() { + keys.push(key); + } + Ok(keys) + } + + /// 根据 ID 获取 API Key + pub fn get_api_key_by_id( + conn: &Connection, + id: &str, + ) -> Result, rusqlite::Error> { + let mut stmt = conn.prepare( + "SELECT id, provider_id, api_key_encrypted, alias, enabled, + usage_count, error_count, last_used_at, created_at + FROM api_keys + WHERE id = ?1", + )?; + + let mut rows = stmt.query([id])?; + if let Some(row) = rows.next()? { + Ok(Some(Self::row_to_api_key(row)?)) + } else { + Ok(None) + } + } + + /// 插入新 API Key + pub fn insert_api_key(conn: &Connection, key: &ApiKeyEntry) -> Result<(), rusqlite::Error> { + conn.execute( + "INSERT INTO api_keys + (id, provider_id, api_key_encrypted, alias, enabled, + usage_count, error_count, last_used_at, created_at) + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)", + params![ + key.id, + key.provider_id, + key.api_key_encrypted, + key.alias, + key.enabled, + key.usage_count, + key.error_count, + key.last_used_at.map(|t| t.to_rfc3339()), + key.created_at.to_rfc3339(), + ], + )?; + Ok(()) + } + + /// 更新 API Key + pub fn update_api_key(conn: &Connection, key: &ApiKeyEntry) -> Result<(), rusqlite::Error> { + conn.execute( + "UPDATE api_keys SET + alias = ?2, enabled = ?3, usage_count = ?4, error_count = ?5, last_used_at = ?6 + WHERE id = ?1", + params![ + key.id, + key.alias, + key.enabled, + key.usage_count, + key.error_count, + key.last_used_at.map(|t| t.to_rfc3339()), + ], + )?; + Ok(()) + } + + /// 删除 API Key + pub fn delete_api_key(conn: &Connection, id: &str) -> Result { + let affected = conn.execute("DELETE FROM api_keys WHERE id = ?1", [id])?; + Ok(affected > 0) + } + + /// 更新 API Key 使用统计 + pub fn update_api_key_usage( + conn: &Connection, + id: &str, + usage_count: i64, + last_used_at: DateTime, + ) -> Result<(), rusqlite::Error> { + conn.execute( + "UPDATE api_keys SET usage_count = ?2, last_used_at = ?3 WHERE id = ?1", + params![id, usage_count, last_used_at.to_rfc3339()], + )?; + Ok(()) + } + + /// 增加 API Key 错误计数 + pub fn increment_api_key_error(conn: &Connection, id: &str) -> Result<(), rusqlite::Error> { + conn.execute( + "UPDATE api_keys SET error_count = error_count + 1 WHERE id = ?1", + [id], + )?; + Ok(()) + } + + /// 从数据库行转换为 ApiKeyEntry + fn row_to_api_key(row: &rusqlite::Row) -> Result { + let id: String = row.get(0)?; + let provider_id: String = row.get(1)?; + let api_key_encrypted: String = row.get(2)?; + let alias: Option = row.get(3)?; + let enabled: bool = row.get(4)?; + let usage_count: i64 = row.get(5)?; + let error_count: i64 = row.get(6)?; + let last_used_at_str: Option = row.get(7)?; + let created_at_str: String = row.get(8)?; + + let last_used_at = last_used_at_str.and_then(|s| { + DateTime::parse_from_rfc3339(&s) + .ok() + .map(|dt| dt.with_timezone(&Utc)) + }); + let created_at = DateTime::parse_from_rfc3339(&created_at_str) + .map(|dt| dt.with_timezone(&Utc)) + .unwrap_or_else(|_| Utc::now()); + + Ok(ApiKeyEntry { + id, + provider_id, + api_key_encrypted, + alias, + enabled, + usage_count, + error_count, + last_used_at, + created_at, + }) + } + + // ==================== UI 状态操作 ==================== + + /// 获取 UI 状态 + pub fn get_ui_state(conn: &Connection, key: &str) -> Result, rusqlite::Error> { + let mut stmt = conn.prepare("SELECT value FROM provider_ui_state WHERE key = ?1")?; + let mut rows = stmt.query([key])?; + if let Some(row) = rows.next()? { + Ok(Some(row.get(0)?)) + } else { + Ok(None) + } + } + + /// 设置 UI 状态 + pub fn set_ui_state(conn: &Connection, key: &str, value: &str) -> Result<(), rusqlite::Error> { + conn.execute( + "INSERT OR REPLACE INTO provider_ui_state (key, value) VALUES (?1, ?2)", + params![key, value], + )?; + Ok(()) + } + + /// 删除 UI 状态 + pub fn delete_ui_state(conn: &Connection, key: &str) -> Result { + let affected = conn.execute("DELETE FROM provider_ui_state WHERE key = ?1", [key])?; + Ok(affected > 0) + } + + // ==================== 复合查询 ==================== + + /// 获取所有 Provider 及其 API Keys + pub fn get_all_providers_with_keys( + conn: &Connection, + ) -> Result, rusqlite::Error> { + let providers = Self::get_all_providers(conn)?; + let mut result = Vec::new(); + + for provider in providers { + let api_keys = Self::get_api_keys_by_provider(conn, &provider.id)?; + result.push(ProviderWithKeys { provider, api_keys }); + } + + Ok(result) + } + + /// 获取启用的 Provider 及其启用的 API Keys + pub fn get_enabled_providers_with_keys( + conn: &Connection, + ) -> Result, rusqlite::Error> { + let mut stmt = conn.prepare( + "SELECT id, name, type, api_host, is_system, group_name, enabled, sort_order, + api_version, project, location, region, created_at, updated_at + FROM api_key_providers + WHERE enabled = 1 + ORDER BY sort_order ASC, created_at ASC", + )?; + + let rows = stmt.query_map([], Self::row_to_provider)?; + let mut result = Vec::new(); + + for provider in rows.flatten() { + let api_keys = Self::get_enabled_api_keys_by_provider(conn, &provider.id)?; + if !api_keys.is_empty() { + result.push(ProviderWithKeys { provider, api_keys }); + } + } + + Ok(result) + } + + /// 统计 Provider 的 API Key 数量 + pub fn count_api_keys_by_provider( + conn: &Connection, + provider_id: &str, + ) -> Result { + conn.query_row( + "SELECT COUNT(*) FROM api_keys WHERE provider_id = ?1", + [provider_id], + |row| row.get(0), + ) + } + + /// 批量更新 Provider 排序顺序 + /// **Validates: Requirements 8.4** + pub fn update_provider_sort_orders( + conn: &Connection, + sort_orders: &[(String, i32)], + ) -> Result<(), rusqlite::Error> { + let now = chrono::Utc::now().to_rfc3339(); + for (id, sort_order) in sort_orders { + conn.execute( + "UPDATE api_key_providers SET sort_order = ?2, updated_at = ?3 WHERE id = ?1", + params![id, sort_order, now], + )?; + } + Ok(()) + } +} diff --git a/src-tauri/src/database/dao/mod.rs b/src-tauri/src/database/dao/mod.rs index 13aef57fa..db90bc9af 100644 --- a/src-tauri/src/database/dao/mod.rs +++ b/src-tauri/src/database/dao/mod.rs @@ -1,3 +1,4 @@ +pub mod api_key_provider; pub mod installed_plugins; pub mod mcp; pub mod prompts; diff --git a/src-tauri/src/database/migration.rs b/src-tauri/src/database/migration.rs index d0b748bdf..91d0c997b 100644 --- a/src-tauri/src/database/migration.rs +++ b/src-tauri/src/database/migration.rs @@ -1,4 +1,4 @@ -use rusqlite::Connection; +use rusqlite::{params, Connection}; /// 从旧的 JSON 配置迁移数据到 SQLite #[allow(dead_code)] @@ -44,3 +44,226 @@ pub fn migrate_from_json(conn: &Connection) -> Result<(), String> { Ok(()) } + +/// 将 api_keys 表中的数据迁移到 provider_pool_credentials 表 +/// +/// 迁移逻辑: +/// 1. 读取 api_keys 表中的所有 API Key +/// 2. 根据 provider_id 查找对应的 api_key_providers 配置 +/// 3. 将 API Key 转换为 CredentialData::OpenAIKey 或 CredentialData::ClaudeKey +/// 4. 插入到 provider_pool_credentials 表 +pub fn migrate_api_keys_to_pool(conn: &Connection) -> Result { + // 检查是否已经迁移过 + let migrated: bool = conn + .query_row( + "SELECT value FROM settings WHERE key = 'migrated_api_keys_to_pool'", + [], + |row| row.get::<_, String>(0), + ) + .map(|v| v == "true") + .unwrap_or(false); + + if migrated { + tracing::debug!("[迁移] API Keys 已迁移过,跳过"); + return Ok(0); + } + + tracing::info!("[迁移] 开始将 api_keys 迁移到 provider_pool_credentials"); + + // 查询所有 API Keys 及其对应的 Provider 信息 + let mut stmt = conn + .prepare( + "SELECT k.id, k.provider_id, k.api_key_encrypted, k.alias, k.enabled, + k.usage_count, k.error_count, k.last_used_at, k.created_at, + p.type, p.api_host, p.name as provider_name + FROM api_keys k + JOIN api_key_providers p ON k.provider_id = p.id + ORDER BY k.created_at ASC", + ) + .map_err(|e| format!("准备查询语句失败: {}", e))?; + + let rows = stmt + .query_map([], |row| { + Ok(ApiKeyMigrationRow { + id: row.get(0)?, + provider_id: row.get(1)?, + api_key_encrypted: row.get(2)?, + alias: row.get(3)?, + enabled: row.get(4)?, + usage_count: row.get::<_, i64>(5)? as u64, + error_count: row.get::<_, i64>(6)? as u32, + last_used_at: row.get(7)?, + created_at: row.get(8)?, + provider_type: row.get(9)?, + api_host: row.get(10)?, + provider_name: row.get(11)?, + }) + }) + .map_err(|e| format!("查询 API Keys 失败: {}", e))?; + + let mut migrated_count = 0; + let now = chrono::Utc::now().timestamp(); + + for row_result in rows { + let row = row_result.map_err(|e| format!("读取行数据失败: {}", e))?; + + // 检查是否已存在相同的凭证(通过 api_key_encrypted 判断) + let exists: bool = conn + .query_row( + "SELECT COUNT(*) > 0 FROM provider_pool_credentials + WHERE credential_data LIKE ?1", + params![format!("%{}%", row.api_key_encrypted)], + |r| r.get(0), + ) + .unwrap_or(false); + + if exists { + tracing::debug!( + "[迁移] 跳过已存在的 API Key: {} (provider: {})", + row.alias.as_deref().unwrap_or(&row.id), + row.provider_id + ); + continue; + } + + // 根据 provider_type 确定 pool_provider_type 和 credential_data + let (pool_provider_type, credential_data) = match row.provider_type.to_lowercase().as_str() + { + "anthropic" => { + let cred = serde_json::json!({ + "type": "claude_key", + "api_key": row.api_key_encrypted, + "base_url": if row.api_host.is_empty() { None } else { Some(&row.api_host) } + }); + ("claude", cred) + } + "openai" | "openai-response" => { + let cred = serde_json::json!({ + "type": "openai_key", + "api_key": row.api_key_encrypted, + "base_url": if row.api_host.is_empty() { None } else { Some(&row.api_host) } + }); + ("openai", cred) + } + "gemini" => { + let cred = serde_json::json!({ + "type": "gemini_api_key", + "api_key": row.api_key_encrypted, + "base_url": if row.api_host.is_empty() { None } else { Some(&row.api_host) }, + "excluded_models": [] + }); + ("gemini_api_key", cred) + } + "vertex" | "vertexai" => { + let cred = serde_json::json!({ + "type": "vertex_key", + "api_key": row.api_key_encrypted, + "base_url": if row.api_host.is_empty() { None } else { Some(&row.api_host) }, + "model_aliases": {} + }); + ("vertex", cred) + } + // 其他类型默认作为 OpenAI 兼容处理 + _ => { + let cred = serde_json::json!({ + "type": "openai_key", + "api_key": row.api_key_encrypted, + "base_url": if row.api_host.is_empty() { None } else { Some(&row.api_host) } + }); + ("openai", cred) + } + }; + + // 生成名称:优先使用 alias,否则使用 provider_name + let name = row + .alias + .clone() + .or_else(|| Some(format!("{} (迁移)", row.provider_name))); + + // 解析时间 + let created_at_ts = row + .created_at + .as_ref() + .and_then(|s| chrono::DateTime::parse_from_rfc3339(s).ok()) + .map(|dt| dt.timestamp()) + .unwrap_or(now); + + let last_used_ts = row + .last_used_at + .as_ref() + .and_then(|s| chrono::DateTime::parse_from_rfc3339(s).ok()) + .map(|dt| dt.timestamp()); + + // 插入到 provider_pool_credentials + let uuid = uuid::Uuid::new_v4().to_string(); + let credential_json = credential_data.to_string(); + + conn.execute( + "INSERT INTO provider_pool_credentials + (uuid, provider_type, credential_data, name, is_healthy, is_disabled, + check_health, check_model_name, not_supported_models, usage_count, error_count, + last_used, last_error_time, last_error_message, last_health_check_time, + last_health_check_model, created_at, updated_at, source, proxy_url) + VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9, ?10, ?11, ?12, ?13, ?14, ?15, ?16, ?17, ?18, ?19, ?20)", + params![ + uuid, + pool_provider_type, + credential_json, + name, + true, // is_healthy + !row.enabled, // is_disabled (反转 enabled) + true, // check_health + Option::::None, // check_model_name + "[]", // not_supported_models + row.usage_count as i64, + row.error_count as i32, + last_used_ts, + Option::::None, // last_error_time + Option::::None, // last_error_message + Option::::None, // last_health_check_time + Option::::None, // last_health_check_model + created_at_ts, + now, + "imported", // source: 标记为导入 + Option::::None, // proxy_url + ], + ) + .map_err(|e| format!("插入凭证失败: {}", e))?; + + tracing::info!( + "[迁移] 已迁移 API Key: {} -> {} (provider_type: {})", + row.alias.as_deref().unwrap_or(&row.id), + uuid, + pool_provider_type + ); + + migrated_count += 1; + } + + // 标记迁移完成 + conn.execute( + "INSERT OR REPLACE INTO settings (key, value) VALUES ('migrated_api_keys_to_pool', 'true')", + [], + ) + .map_err(|e| format!("标记迁移完成失败: {}", e))?; + + tracing::info!("[迁移] API Keys 迁移完成,共迁移 {} 条记录", migrated_count); + + Ok(migrated_count) +} + +/// API Key 迁移行数据 +struct ApiKeyMigrationRow { + id: String, + provider_id: String, + api_key_encrypted: String, + alias: Option, + enabled: bool, + usage_count: u64, + error_count: u32, + last_used_at: Option, + created_at: Option, + provider_type: String, + api_host: String, + provider_name: String, +} diff --git a/src-tauri/src/database/mod.rs b/src-tauri/src/database/mod.rs index b13e4ea0a..c5ff95165 100644 --- a/src-tauri/src/database/mod.rs +++ b/src-tauri/src/database/mod.rs @@ -1,6 +1,7 @@ pub mod dao; pub mod migration; pub mod schema; +pub mod system_providers; use rusqlite::Connection; use std::path::PathBuf; @@ -26,5 +27,17 @@ pub fn init_database() -> Result { schema::create_tables(&conn).map_err(|e| e.to_string())?; migration::migrate_from_json(&conn)?; + // 执行 API Keys 到 Provider Pool 的迁移 + match migration::migrate_api_keys_to_pool(&conn) { + Ok(count) => { + if count > 0 { + tracing::info!("[数据库] 已将 {} 条 API Key 迁移到凭证池", count); + } + } + Err(e) => { + tracing::warn!("[数据库] API Key 迁移失败(非致命): {}", e); + } + } + Ok(Arc::new(Mutex::new(conn))) } diff --git a/src-tauri/src/database/schema.rs b/src-tauri/src/database/schema.rs index 999c4931e..c23719516 100644 --- a/src-tauri/src/database/schema.rs +++ b/src-tauri/src/database/schema.rs @@ -1,6 +1,68 @@ use rusqlite::Connection; pub fn create_tables(conn: &Connection) -> Result<(), rusqlite::Error> { + // API Key Provider 配置表 + // _Requirements: 9.1_ + conn.execute( + "CREATE TABLE IF NOT EXISTS api_key_providers ( + id TEXT PRIMARY KEY, + name TEXT NOT NULL, + type TEXT NOT NULL, + api_host TEXT NOT NULL, + is_system INTEGER NOT NULL DEFAULT 0, + group_name TEXT NOT NULL, + enabled INTEGER NOT NULL DEFAULT 0, + sort_order INTEGER NOT NULL DEFAULT 0, + api_version TEXT, + project TEXT, + location TEXT, + region TEXT, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL + )", + [], + )?; + + // 创建 api_key_providers 索引 + conn.execute( + "CREATE INDEX IF NOT EXISTS idx_api_key_providers_group ON api_key_providers(group_name)", + [], + )?; + + // API Key 条目表 + // _Requirements: 9.1, 9.2_ + conn.execute( + "CREATE TABLE IF NOT EXISTS api_keys ( + id TEXT PRIMARY KEY, + provider_id TEXT NOT NULL, + api_key_encrypted TEXT NOT NULL, + alias TEXT, + enabled INTEGER NOT NULL DEFAULT 1, + usage_count INTEGER NOT NULL DEFAULT 0, + error_count INTEGER NOT NULL DEFAULT 0, + last_used_at TEXT, + created_at TEXT NOT NULL, + FOREIGN KEY (provider_id) REFERENCES api_key_providers(id) ON DELETE CASCADE + )", + [], + )?; + + // 创建 api_keys 索引 + conn.execute( + "CREATE INDEX IF NOT EXISTS idx_api_keys_provider ON api_keys(provider_id)", + [], + )?; + + // Provider UI 状态表 + // _Requirements: 8.4_ + conn.execute( + "CREATE TABLE IF NOT EXISTS provider_ui_state ( + key TEXT PRIMARY KEY, + value TEXT NOT NULL + )", + [], + )?; + // Providers 表 conn.execute( "CREATE TABLE IF NOT EXISTS providers ( diff --git a/src-tauri/src/database/system_providers.rs b/src-tauri/src/database/system_providers.rs new file mode 100644 index 000000000..3606a82d5 --- /dev/null +++ b/src-tauri/src/database/system_providers.rs @@ -0,0 +1,632 @@ +//! 系统预设 Provider 配置 +//! +//! 定义所有 60+ 系统预设 Provider 的配置,与前端 `SYSTEM_PROVIDERS` 保持一致。 +//! +//! **Feature: provider-ui-refactor** +//! **Validates: Requirements 3.1-3.6** + +use crate::database::dao::api_key_provider::{ApiKeyProvider, ApiProviderType, ProviderGroup}; +use chrono::Utc; + +/// 系统 Provider 配置定义 +pub struct SystemProviderDef { + pub id: &'static str, + pub name: &'static str, + pub provider_type: ApiProviderType, + pub api_host: &'static str, + pub group: ProviderGroup, + pub sort_order: i32, + pub api_version: Option<&'static str>, +} + +/// 获取所有系统 Provider 配置 +pub fn get_system_providers() -> Vec { + vec![ + // ========================================================================= + // 主流 AI (10个) - Requirements 3.1 + // ========================================================================= + SystemProviderDef { + id: "openai", + name: "OpenAI", + provider_type: ApiProviderType::OpenaiResponse, + api_host: "https://api.openai.com", + group: ProviderGroup::Mainstream, + sort_order: 1, + api_version: None, + }, + SystemProviderDef { + id: "anthropic", + name: "Anthropic", + provider_type: ApiProviderType::Anthropic, + api_host: "https://api.anthropic.com", + group: ProviderGroup::Mainstream, + sort_order: 2, + api_version: None, + }, + SystemProviderDef { + id: "gemini", + name: "Gemini", + provider_type: ApiProviderType::Gemini, + api_host: "https://generativelanguage.googleapis.com", + group: ProviderGroup::Mainstream, + sort_order: 3, + api_version: None, + }, + SystemProviderDef { + id: "deepseek", + name: "DeepSeek", + provider_type: ApiProviderType::Openai, + api_host: "https://api.deepseek.com", + group: ProviderGroup::Mainstream, + sort_order: 4, + api_version: None, + }, + SystemProviderDef { + id: "moonshot", + name: "Moonshot", + provider_type: ApiProviderType::Openai, + api_host: "https://api.moonshot.cn", + group: ProviderGroup::Mainstream, + sort_order: 5, + api_version: None, + }, + SystemProviderDef { + id: "groq", + name: "Groq", + provider_type: ApiProviderType::Openai, + api_host: "https://api.groq.com/openai", + group: ProviderGroup::Mainstream, + sort_order: 6, + api_version: None, + }, + SystemProviderDef { + id: "grok", + name: "Grok (xAI)", + provider_type: ApiProviderType::Openai, + api_host: "https://api.x.ai", + group: ProviderGroup::Mainstream, + sort_order: 7, + api_version: None, + }, + SystemProviderDef { + id: "mistral", + name: "Mistral", + provider_type: ApiProviderType::Openai, + api_host: "https://api.mistral.ai", + group: ProviderGroup::Mainstream, + sort_order: 8, + api_version: None, + }, + SystemProviderDef { + id: "perplexity", + name: "Perplexity", + provider_type: ApiProviderType::Openai, + api_host: "https://api.perplexity.ai/", + group: ProviderGroup::Mainstream, + sort_order: 9, + api_version: None, + }, + SystemProviderDef { + id: "cohere", + name: "Cohere", + provider_type: ApiProviderType::Openai, + api_host: "https://api.cohere.ai", + group: ProviderGroup::Mainstream, + sort_order: 10, + api_version: None, + }, + // ========================================================================= + // 国内 AI (15个) - Requirements 3.2 + // ========================================================================= + SystemProviderDef { + id: "zhipu", + name: "智谱 (ZhiPu)", + provider_type: ApiProviderType::Openai, + api_host: "https://open.bigmodel.cn/api/paas/v4/", + group: ProviderGroup::Chinese, + sort_order: 11, + api_version: None, + }, + SystemProviderDef { + id: "baichuan", + name: "百川 (Baichuan)", + provider_type: ApiProviderType::Openai, + api_host: "https://api.baichuan-ai.com", + group: ProviderGroup::Chinese, + sort_order: 12, + api_version: None, + }, + SystemProviderDef { + id: "dashscope", + name: "百炼/通义千问 (Dashscope)", + provider_type: ApiProviderType::Openai, + api_host: "https://dashscope.aliyuncs.com/compatible-mode/v1/", + group: ProviderGroup::Chinese, + sort_order: 13, + api_version: None, + }, + SystemProviderDef { + id: "stepfun", + name: "阶跃星辰 (StepFun)", + provider_type: ApiProviderType::Openai, + api_host: "https://api.stepfun.com", + group: ProviderGroup::Chinese, + sort_order: 14, + api_version: None, + }, + SystemProviderDef { + id: "doubao", + name: "豆包 (Doubao)", + provider_type: ApiProviderType::Openai, + api_host: "https://ark.cn-beijing.volces.com/api/v3/", + group: ProviderGroup::Chinese, + sort_order: 15, + api_version: None, + }, + SystemProviderDef { + id: "minimax", + name: "MiniMax", + provider_type: ApiProviderType::Openai, + api_host: "https://api.minimaxi.com/v1", + group: ProviderGroup::Chinese, + sort_order: 16, + api_version: None, + }, + SystemProviderDef { + id: "yi", + name: "零一万物 (Yi)", + provider_type: ApiProviderType::Openai, + api_host: "https://api.lingyiwanwu.com", + group: ProviderGroup::Chinese, + sort_order: 17, + api_version: None, + }, + SystemProviderDef { + id: "hunyuan", + name: "腾讯混元 (Hunyuan)", + provider_type: ApiProviderType::Openai, + api_host: "https://api.hunyuan.cloud.tencent.com", + group: ProviderGroup::Chinese, + sort_order: 18, + api_version: None, + }, + SystemProviderDef { + id: "tencent-cloud-ti", + name: "腾讯云 TI", + provider_type: ApiProviderType::Openai, + api_host: "https://api.lkeap.cloud.tencent.com", + group: ProviderGroup::Chinese, + sort_order: 19, + api_version: None, + }, + SystemProviderDef { + id: "baidu-cloud", + name: "百度云 (Baidu Cloud)", + provider_type: ApiProviderType::Openai, + api_host: "https://qianfan.baidubce.com/v2/", + group: ProviderGroup::Chinese, + sort_order: 20, + api_version: None, + }, + SystemProviderDef { + id: "infini", + name: "无问芯穹 (Infini)", + provider_type: ApiProviderType::Openai, + api_host: "https://cloud.infini-ai.com/maas", + group: ProviderGroup::Chinese, + sort_order: 21, + api_version: None, + }, + SystemProviderDef { + id: "modelscope", + name: "魔搭 (ModelScope)", + provider_type: ApiProviderType::Openai, + api_host: "https://api-inference.modelscope.cn/v1/", + group: ProviderGroup::Chinese, + sort_order: 22, + api_version: None, + }, + SystemProviderDef { + id: "xirang", + name: "息壤 (Xirang)", + provider_type: ApiProviderType::Openai, + api_host: "https://wishub-x1.ctyun.cn", + group: ProviderGroup::Chinese, + sort_order: 23, + api_version: None, + }, + SystemProviderDef { + id: "mimo", + name: "小米 MiMo", + provider_type: ApiProviderType::Openai, + api_host: "https://api.xiaomimimo.com", + group: ProviderGroup::Chinese, + sort_order: 24, + api_version: None, + }, + SystemProviderDef { + id: "zhinao", + name: "360 智脑 (Zhinao)", + provider_type: ApiProviderType::Openai, + api_host: "https://api.360.cn", + group: ProviderGroup::Chinese, + sort_order: 25, + api_version: None, + }, + // ========================================================================= + // 云服务 (5个) - Requirements 3.3 + // ========================================================================= + SystemProviderDef { + id: "azure-openai", + name: "Azure OpenAI", + provider_type: ApiProviderType::AzureOpenai, + api_host: "", + group: ProviderGroup::Cloud, + sort_order: 26, + api_version: Some("2024-02-15-preview"), + }, + SystemProviderDef { + id: "vertexai", + name: "VertexAI", + provider_type: ApiProviderType::Vertexai, + api_host: "", + group: ProviderGroup::Cloud, + sort_order: 27, + api_version: None, + }, + SystemProviderDef { + id: "aws-bedrock", + name: "AWS Bedrock", + provider_type: ApiProviderType::AwsBedrock, + api_host: "", + group: ProviderGroup::Cloud, + sort_order: 28, + api_version: None, + }, + SystemProviderDef { + id: "github", + name: "Github Models", + provider_type: ApiProviderType::Openai, + api_host: "https://models.github.ai/inference", + group: ProviderGroup::Cloud, + sort_order: 29, + api_version: None, + }, + SystemProviderDef { + id: "copilot", + name: "Github Copilot", + provider_type: ApiProviderType::Openai, + api_host: "https://api.githubcopilot.com/", + group: ProviderGroup::Cloud, + sort_order: 30, + api_version: None, + }, + // ========================================================================= + // API 聚合/中转服务 (25个) - Requirements 3.4 + // ========================================================================= + SystemProviderDef { + id: "silicon", + name: "Silicon Flow", + provider_type: ApiProviderType::Openai, + api_host: "https://api.siliconflow.cn", + group: ProviderGroup::Aggregator, + sort_order: 31, + api_version: None, + }, + SystemProviderDef { + id: "openrouter", + name: "OpenRouter", + provider_type: ApiProviderType::Openai, + api_host: "https://openrouter.ai/api/v1/", + group: ProviderGroup::Aggregator, + sort_order: 32, + api_version: None, + }, + SystemProviderDef { + id: "aihubmix", + name: "AiHubMix", + provider_type: ApiProviderType::Openai, + api_host: "https://aihubmix.com", + group: ProviderGroup::Aggregator, + sort_order: 33, + api_version: None, + }, + SystemProviderDef { + id: "302ai", + name: "302.AI", + provider_type: ApiProviderType::Openai, + api_host: "https://api.302.ai", + group: ProviderGroup::Aggregator, + sort_order: 34, + api_version: None, + }, + SystemProviderDef { + id: "together", + name: "Together", + provider_type: ApiProviderType::Openai, + api_host: "https://api.together.xyz", + group: ProviderGroup::Aggregator, + sort_order: 35, + api_version: None, + }, + SystemProviderDef { + id: "fireworks", + name: "Fireworks", + provider_type: ApiProviderType::Openai, + api_host: "https://api.fireworks.ai/inference", + group: ProviderGroup::Aggregator, + sort_order: 36, + api_version: None, + }, + SystemProviderDef { + id: "nvidia", + name: "NVIDIA", + provider_type: ApiProviderType::Openai, + api_host: "https://integrate.api.nvidia.com", + group: ProviderGroup::Aggregator, + sort_order: 37, + api_version: None, + }, + SystemProviderDef { + id: "hyperbolic", + name: "Hyperbolic", + provider_type: ApiProviderType::Openai, + api_host: "https://api.hyperbolic.xyz", + group: ProviderGroup::Aggregator, + sort_order: 38, + api_version: None, + }, + SystemProviderDef { + id: "cerebras", + name: "Cerebras", + provider_type: ApiProviderType::Openai, + api_host: "https://api.cerebras.ai/v1", + group: ProviderGroup::Aggregator, + sort_order: 39, + api_version: None, + }, + SystemProviderDef { + id: "ppio", + name: "PPIO", + provider_type: ApiProviderType::Openai, + api_host: "https://api.ppinfra.com/v3/openai/", + group: ProviderGroup::Aggregator, + sort_order: 40, + api_version: None, + }, + SystemProviderDef { + id: "qiniu", + name: "七牛 (Qiniu)", + provider_type: ApiProviderType::Openai, + api_host: "https://api.qnaigc.com", + group: ProviderGroup::Aggregator, + sort_order: 41, + api_version: None, + }, + SystemProviderDef { + id: "tokenflux", + name: "TokenFlux", + provider_type: ApiProviderType::Openai, + api_host: "https://api.tokenflux.ai/openai/v1", + group: ProviderGroup::Aggregator, + sort_order: 42, + api_version: None, + }, + SystemProviderDef { + id: "cephalon", + name: "Cephalon", + provider_type: ApiProviderType::Openai, + api_host: "https://cephalon.cloud/user-center/v1/model", + group: ProviderGroup::Aggregator, + sort_order: 43, + api_version: None, + }, + SystemProviderDef { + id: "lanyun", + name: "蓝云 (Lanyun)", + provider_type: ApiProviderType::Openai, + api_host: "https://maas-api.lanyun.net", + group: ProviderGroup::Aggregator, + sort_order: 44, + api_version: None, + }, + SystemProviderDef { + id: "ph8", + name: "PH8", + provider_type: ApiProviderType::Openai, + api_host: "https://ph8.co", + group: ProviderGroup::Aggregator, + sort_order: 45, + api_version: None, + }, + SystemProviderDef { + id: "sophnet", + name: "SophNet", + provider_type: ApiProviderType::Openai, + api_host: "https://www.sophnet.com/api/open-apis/v1", + group: ProviderGroup::Aggregator, + sort_order: 46, + api_version: None, + }, + SystemProviderDef { + id: "ocoolai", + name: "ocoolAI", + provider_type: ApiProviderType::Openai, + api_host: "https://api.ocoolai.com", + group: ProviderGroup::Aggregator, + sort_order: 47, + api_version: None, + }, + SystemProviderDef { + id: "dmxapi", + name: "DMXAPI", + provider_type: ApiProviderType::Openai, + api_host: "https://www.dmxapi.cn", + group: ProviderGroup::Aggregator, + sort_order: 48, + api_version: None, + }, + SystemProviderDef { + id: "aionly", + name: "AIOnly", + provider_type: ApiProviderType::Openai, + api_host: "https://api.aiionly.com", + group: ProviderGroup::Aggregator, + sort_order: 49, + api_version: None, + }, + SystemProviderDef { + id: "burncloud", + name: "BurnCloud", + provider_type: ApiProviderType::Openai, + api_host: "https://ai.burncloud.com", + group: ProviderGroup::Aggregator, + sort_order: 50, + api_version: None, + }, + SystemProviderDef { + id: "alayanew", + name: "AlayaNew", + provider_type: ApiProviderType::Openai, + api_host: "https://deepseek.alayanew.com", + group: ProviderGroup::Aggregator, + sort_order: 51, + api_version: None, + }, + SystemProviderDef { + id: "longcat", + name: "LongCat", + provider_type: ApiProviderType::Openai, + api_host: "https://api.longcat.chat/openai", + group: ProviderGroup::Aggregator, + sort_order: 52, + api_version: None, + }, + SystemProviderDef { + id: "poe", + name: "Poe", + provider_type: ApiProviderType::Openai, + api_host: "https://api.poe.com/v1/", + group: ProviderGroup::Aggregator, + sort_order: 53, + api_version: None, + }, + SystemProviderDef { + id: "huggingface", + name: "Hugging Face", + provider_type: ApiProviderType::OpenaiResponse, + api_host: "https://router.huggingface.co/v1/", + group: ProviderGroup::Aggregator, + sort_order: 54, + api_version: None, + }, + SystemProviderDef { + id: "vercel-gateway", + name: "Vercel AI Gateway", + provider_type: ApiProviderType::Gateway, + api_host: "https://ai-gateway.vercel.sh/v1/ai", + group: ProviderGroup::Aggregator, + sort_order: 55, + api_version: None, + }, + // ========================================================================= + // 本地/自托管服务 (5个) - Requirements 3.5 + // ========================================================================= + SystemProviderDef { + id: "ollama", + name: "Ollama", + provider_type: ApiProviderType::Ollama, + api_host: "http://localhost:11434", + group: ProviderGroup::Local, + sort_order: 56, + api_version: None, + }, + SystemProviderDef { + id: "lmstudio", + name: "LM Studio", + provider_type: ApiProviderType::Openai, + api_host: "http://localhost:1234", + group: ProviderGroup::Local, + sort_order: 57, + api_version: None, + }, + SystemProviderDef { + id: "new-api", + name: "New API", + provider_type: ApiProviderType::NewApi, + api_host: "http://localhost:3000", + group: ProviderGroup::Local, + sort_order: 58, + api_version: None, + }, + SystemProviderDef { + id: "gpustack", + name: "GPUStack", + provider_type: ApiProviderType::Openai, + api_host: "", + group: ProviderGroup::Local, + sort_order: 59, + api_version: None, + }, + SystemProviderDef { + id: "ovms", + name: "OpenVINO Model Server", + provider_type: ApiProviderType::Openai, + api_host: "http://localhost:8000/v3/", + group: ProviderGroup::Local, + sort_order: 60, + api_version: None, + }, + // ========================================================================= + // 专用服务 (3个) - Requirements 3.6 + // ========================================================================= + SystemProviderDef { + id: "jina", + name: "Jina (Embedding/Rerank)", + provider_type: ApiProviderType::Openai, + api_host: "https://api.jina.ai", + group: ProviderGroup::Specialized, + sort_order: 61, + api_version: None, + }, + SystemProviderDef { + id: "voyageai", + name: "VoyageAI (Embedding)", + provider_type: ApiProviderType::Openai, + api_host: "https://api.voyageai.com", + group: ProviderGroup::Specialized, + sort_order: 62, + api_version: None, + }, + SystemProviderDef { + id: "cherryin", + name: "CherryIN", + provider_type: ApiProviderType::Openai, + api_host: "https://open.cherryin.net", + group: ProviderGroup::Specialized, + sort_order: 63, + api_version: None, + }, + ] +} + +/// 将 SystemProviderDef 转换为 ApiKeyProvider +pub fn to_api_key_provider(def: &SystemProviderDef) -> ApiKeyProvider { + let now = Utc::now(); + ApiKeyProvider { + id: def.id.to_string(), + name: def.name.to_string(), + provider_type: def.provider_type, + api_host: def.api_host.to_string(), + is_system: true, + group: def.group, + enabled: false, + sort_order: def.sort_order, + api_version: def.api_version.map(|s| s.to_string()), + project: None, + location: None, + region: None, + created_at: now, + updated_at: now, + } +} diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index 025342df8..b898cc8fa 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -1,10 +1,11 @@ pub mod agent; +pub mod backends; pub mod browser_interceptor; mod commands; mod config; mod converter; pub mod credential; -mod database; +pub mod database; pub mod flow_monitor; pub mod injection; mod logger; @@ -18,9 +19,11 @@ pub mod resilience; pub mod router; mod server; mod server_utils; -mod services; +pub mod services; +pub mod stream; pub mod streaming; pub mod telemetry; +pub mod translator; pub mod tray; pub mod websocket; @@ -30,6 +33,7 @@ use tauri::{Manager, Runtime}; use tokio::sync::RwLock; use agent::NativeAgentState; +use commands::api_key_provider_cmd::ApiKeyProviderServiceState; use commands::browser_interceptor_cmd::BrowserInterceptorState; use commands::flow_monitor_cmd::{ BatchOperationsState, BookmarkManagerState, EnhancedStatsServiceState, FlowInterceptorState, @@ -48,6 +52,7 @@ use flow_monitor::{ FlowMonitor, FlowMonitorConfig, FlowQueryService, FlowReplayer, InterceptConfig, QuickFilterManager, SessionManager, }; +use services::api_key_provider_service::ApiKeyProviderService; use services::provider_pool_service::ProviderPoolService; use services::skill_service::SkillService; use services::token_cache_service::TokenCacheService; @@ -1528,6 +1533,11 @@ pub fn run() { let provider_pool_service = ProviderPoolService::new(); let provider_pool_service_state = ProviderPoolServiceState(Arc::new(provider_pool_service)); + // Initialize ApiKeyProviderService + let api_key_provider_service = ApiKeyProviderService::new(); + let api_key_provider_service_state = + ApiKeyProviderServiceState(Arc::new(api_key_provider_service)); + // Initialize CredentialSyncService (optional - only if config manager is available) // For now, we initialize it as None since ConfigManager requires async setup // This can be enhanced later to properly initialize with ConfigManager @@ -1775,6 +1785,7 @@ pub fn run() { .manage(db) .manage(skill_service_state) .manage(provider_pool_service_state) + .manage(api_key_provider_service_state) .manage(credential_sync_service_state) .manage(token_cache_service_state) .manage(machine_id_service_state) @@ -2170,6 +2181,24 @@ pub fn run() { commands::provider_pool_cmd::install_playwright, commands::provider_pool_cmd::start_kiro_playwright_login, commands::provider_pool_cmd::cancel_kiro_playwright_login, + // API Key Provider commands + commands::api_key_provider_cmd::get_api_key_providers, + commands::api_key_provider_cmd::get_api_key_provider, + commands::api_key_provider_cmd::add_custom_api_key_provider, + commands::api_key_provider_cmd::update_api_key_provider, + commands::api_key_provider_cmd::delete_custom_api_key_provider, + commands::api_key_provider_cmd::add_api_key, + commands::api_key_provider_cmd::delete_api_key, + commands::api_key_provider_cmd::toggle_api_key, + commands::api_key_provider_cmd::update_api_key_alias, + commands::api_key_provider_cmd::get_next_api_key, + commands::api_key_provider_cmd::record_api_key_usage, + commands::api_key_provider_cmd::record_api_key_error, + commands::api_key_provider_cmd::get_provider_ui_state, + commands::api_key_provider_cmd::set_provider_ui_state, + commands::api_key_provider_cmd::update_provider_sort_orders, + commands::api_key_provider_cmd::export_api_key_providers, + commands::api_key_provider_cmd::import_api_key_providers, // Route commands commands::route_cmd::get_available_routes, commands::route_cmd::get_route_curl_examples, diff --git a/src-tauri/src/providers/claude_custom.rs b/src-tauri/src/providers/claude_custom.rs index ed56b4691..50167bcd1 100644 --- a/src-tauri/src/providers/claude_custom.rs +++ b/src-tauri/src/providers/claude_custom.rs @@ -409,12 +409,36 @@ impl StreamingProvider for ClaudeCustomProvider { // 转换 OpenAI 请求为 Anthropic 格式 let mut anthropic_messages = Vec::new(); let mut system_content = None; + // 收集 tool 角色消息的 tool_result,稍后合并到 user 消息中 + let mut pending_tool_results: Vec = Vec::new(); for msg in &request.messages { let role = &msg.role; + // 处理 tool 角色消息(工具调用结果) + if role == "tool" { + // 转换为 Anthropic tool_result content block + let tool_call_id = msg.tool_call_id.clone().unwrap_or_default(); + let content = msg.get_content_text(); + pending_tool_results.push(serde_json::json!({ + "type": "tool_result", + "tool_use_id": tool_call_id, + "content": content + })); + continue; + } + + // 如果有待处理的 tool_results 且当前不是 assistant 消息,先添加一个 user 消息 + if !pending_tool_results.is_empty() && role != "assistant" { + anthropic_messages.push(serde_json::json!({ + "role": "user", + "content": pending_tool_results.clone() + })); + pending_tool_results.clear(); + } + // 提取消息内容,转换为 Anthropic 格式的 content 数组 - let content_blocks: Vec = match &msg.content { + let mut content_blocks: Vec = match &msg.content { Some(MessageContent::Text(text)) => { if text.is_empty() { vec![] @@ -443,6 +467,23 @@ impl StreamingProvider for ClaudeCustomProvider { None => vec![], }; + // 处理 assistant 消息中的 tool_calls + if role == "assistant" { + if let Some(ref tool_calls) = msg.tool_calls { + for tc in tool_calls { + // 解析 arguments JSON + let input: serde_json::Value = serde_json::from_str(&tc.function.arguments) + .unwrap_or(serde_json::json!({})); + content_blocks.push(serde_json::json!({ + "type": "tool_use", + "id": tc.id, + "name": tc.function.name, + "input": input + })); + } + } + } + if role == "system" { // system 消息只提取文本 let text = content_blocks @@ -464,6 +505,14 @@ impl StreamingProvider for ClaudeCustomProvider { } } + // 处理末尾的 tool_results + if !pending_tool_results.is_empty() { + anthropic_messages.push(serde_json::json!({ + "role": "user", + "content": pending_tool_results + })); + } + let mut anthropic_body = serde_json::json!({ "model": request.model, "max_tokens": request.max_tokens.unwrap_or(4096), @@ -475,6 +524,76 @@ impl StreamingProvider for ClaudeCustomProvider { anthropic_body["system"] = serde_json::json!(sys); } + // 转换 tools: OpenAI 格式 -> Anthropic 格式 + if let Some(ref tools) = request.tools { + let anthropic_tools: Vec = tools + .iter() + .filter_map(|tool| { + match tool { + crate::models::openai::Tool::Function { function } => { + Some(serde_json::json!({ + "name": function.name, + "description": function.description.clone().unwrap_or_default(), + "input_schema": function.parameters.clone().unwrap_or_else(|| serde_json::json!({"type": "object", "properties": {}})) + })) + } + // WebSearch 等其他工具类型暂不处理 + _ => None, + } + }) + .collect(); + + if !anthropic_tools.is_empty() { + anthropic_body["tools"] = serde_json::json!(anthropic_tools); + tracing::info!( + "[CLAUDE_STREAM] 添加 {} 个工具到请求", + anthropic_tools.len() + ); + } + } + + // 转换 tool_choice: OpenAI 格式 -> Anthropic 格式 + if let Some(ref tool_choice) = request.tool_choice { + let anthropic_tool_choice = match tool_choice { + serde_json::Value::String(s) => { + match s.as_str() { + "none" => Some(serde_json::json!({"type": "none"})), + "auto" => Some(serde_json::json!({"type": "auto"})), + "required" | "any" => Some(serde_json::json!({"type": "any"})), + _ => None, // 未知值,不设置 + } + } + serde_json::Value::Object(obj) => { + // 处理 {"type": "function", "function": {"name": "xxx"}} 格式 + if let Some(func) = obj.get("function") { + if let Some(name) = func.get("name").and_then(|n| n.as_str()) { + Some(serde_json::json!({"type": "tool", "name": name})) + } else { + None + } + } else if let Some(t) = obj.get("type").and_then(|t| t.as_str()) { + match t { + "any" | "tool" => Some(serde_json::json!({"type": "any"})), + "auto" => Some(serde_json::json!({"type": "auto"})), + "none" => Some(serde_json::json!({"type": "none"})), + _ => None, + } + } else { + None + } + } + _ => None, + }; + + if let Some(tc) = anthropic_tool_choice { + anthropic_body["tool_choice"] = tc; + tracing::info!( + "[CLAUDE_STREAM] 设置 tool_choice: {:?}", + anthropic_body["tool_choice"] + ); + } + } + let url = self.build_url("messages"); tracing::info!( diff --git a/src-tauri/src/providers/kiro.rs b/src-tauri/src/providers/kiro.rs index 0c74e2b9b..fc498ffc6 100644 --- a/src-tauri/src/providers/kiro.rs +++ b/src-tauri/src/providers/kiro.rs @@ -2,9 +2,12 @@ #![allow(dead_code)] -use crate::converter::openai_to_cw::convert_openai_to_codewhisperer; +// 使用新的 translator 模块替代旧的 converter +use crate::models::anthropic::AnthropicMessagesRequest; use crate::models::openai::*; use crate::providers::traits::{CredentialProvider, ProviderResult}; +use crate::translator::kiro::anthropic::request::convert_anthropic_to_codewhisperer; +use crate::translator::kiro::openai::request::convert_openai_to_codewhisperer; use async_trait::async_trait; use reqwest::Client; use serde::{Deserialize, Serialize}; @@ -1226,3 +1229,90 @@ impl StreamingProvider for KiroProvider { StreamFormat::AwsEventStream } } + +// ============================================================================ +// Anthropic 格式直接支持 +// ============================================================================ + +impl KiroProvider { + /// 直接处理 Anthropic 格式的流式请求 + /// + /// 绕过 OpenAI 中间格式,直接从 Anthropic → CodeWhisperer + /// 这样可以保留 Anthropic 特有的字段(如 tool_choice) + pub async fn call_api_stream_anthropic( + &self, + request: &AnthropicMessagesRequest, + ) -> 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 + }; + + // 直接转换 Anthropic → CodeWhisperer(不经过 OpenAI) + let cw_request = convert_anthropic_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::info!( + "[KIRO_STREAM_ANTHROPIC] 直接 Anthropic→CodeWhisperer 流式请求: 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}" + ), + ) + .json(&cw_request) + .send() + .await + .map_err(|e| { + tracing::error!("[KIRO_STREAM_ANTHROPIC] 请求发送失败: {}", e); + ProviderError::from_reqwest_error(&e) + })?; + + tracing::info!("[KIRO_STREAM_ANTHROPIC] 收到响应: status={}", resp.status()); + + // 检查响应状态 + let status = resp.status(); + if !status.is_success() { + let body = resp.text().await.unwrap_or_default(); + tracing::error!("[KIRO_STREAM_ANTHROPIC] 请求失败: {} - {}", status, body); + return Err(ProviderError::from_http_status(status.as_u16(), &body)); + } + + tracing::info!("[KIRO_STREAM_ANTHROPIC] 流式响应开始: status={}", status); + + // 将 reqwest 响应转换为 StreamResponse + Ok(reqwest_stream_to_stream_response(resp)) + } +} diff --git a/src-tauri/src/server/handlers/provider_calls.rs b/src-tauri/src/server/handlers/provider_calls.rs index b8028e32d..e929c9299 100644 --- a/src-tauri/src/server/handlers/provider_calls.rs +++ b/src-tauri/src/server/handlers/provider_calls.rs @@ -66,10 +66,11 @@ use crate::server_utils::{ build_anthropic_response, build_anthropic_stream_response, parse_cw_response, safe_truncate, CWParsedResponse, }; +use crate::stream::{PipelineConfig, StreamPipeline}; use crate::streaming::traits::StreamingProvider; use crate::streaming::{ - AnthropicSseGenerator, AwsEvent, AwsEventStreamParser, StreamConfig, StreamContext, - StreamError, StreamFormat as StreamingFormat, StreamManager, StreamResponse, + StreamConfig, StreamContext, StreamError, StreamFormat as StreamingFormat, StreamManager, + StreamResponse, }; /// 根据凭证调用 Provider (Anthropic 格式) @@ -765,6 +766,27 @@ pub async fn call_provider_openai( _flow_id: Option<&str>, ) -> Response { let _start_time = std::time::Instant::now(); + + // 调试:打印凭证类型 + let cred_type = match &credential.credential { + CredentialData::KiroOAuth { .. } => "KiroOAuth", + CredentialData::ClaudeKey { .. } => "ClaudeKey", + CredentialData::OpenAIKey { .. } => "OpenAIKey", + CredentialData::GeminiOAuth { .. } => "GeminiOAuth", + CredentialData::GeminiApiKey { .. } => "GeminiApiKey", + CredentialData::VertexKey { .. } => "VertexKey", + CredentialData::QwenOAuth { .. } => "QwenOAuth", + CredentialData::AntigravityOAuth { .. } => "AntigravityOAuth", + _ => "Other", + }; + tracing::info!( + "[CALL_PROVIDER_OPENAI] 凭证类型={}, 凭证名称={:?}, provider_type={}, uuid={}", + cred_type, + credential.name, + credential.provider_type, + &credential.uuid[..8] + ); + match &credential.credential { CredentialData::KiroOAuth { creds_file_path } => { // 优先使用 token cache,避免每次都刷新 token @@ -842,17 +864,15 @@ pub async fn call_provider_openai( tracing::info!("[OPENAI_STREAM] 开始转换流式响应"); - // 创建 StreamConverter 将 AWS Event Stream 转换为 OpenAI SSE - let converter = std::sync::Arc::new(tokio::sync::Mutex::new( - crate::streaming::converter::StreamConverter::with_model( - crate::streaming::converter::StreamFormat::AwsEventStream, - crate::streaming::converter::StreamFormat::OpenAiSse, - &request.model, - ), + // 使用新的统一流处理管道 (Kiro → OpenAI) + let config = PipelineConfig::kiro_to_openai(request.model.clone()); + let pipeline = std::sync::Arc::new(tokio::sync::Mutex::new( + StreamPipeline::new(config), )); // 创建转换流 - let converter_for_stream = converter.clone(); + let pipeline_for_stream = pipeline.clone(); + let pipeline_for_finalize = pipeline.clone(); let final_stream = async_stream::stream! { use futures::StreamExt; @@ -866,10 +886,10 @@ pub async fn call_provider_openai( bytes.len() ); - // 转换 chunk + // 使用 Pipeline 处理 chunk let sse_events = { - let mut converter_guard = converter_for_stream.lock().await; - converter_guard.convert(&bytes) + let mut pipeline_guard = pipeline_for_stream.lock().await; + pipeline_guard.process_chunk(&bytes) }; tracing::debug!( @@ -879,7 +899,7 @@ pub async fn call_provider_openai( // yield 每个 SSE 事件 for sse_str in sse_events { - yield Ok::(sse_str); + yield Ok::(sse_str); } } Err(e) => { @@ -892,16 +912,16 @@ pub async fn call_provider_openai( tracing::info!("[OPENAI_STREAM] 流结束,生成 finalize 事件"); - // 流结束,生成结束事件 + // 流结束,使用 Pipeline 生成结束事件 let final_events = { - let mut converter_guard = converter_for_stream.lock().await; - converter_guard.finish() + let mut pipeline_guard = pipeline_for_finalize.lock().await; + pipeline_guard.finish() }; tracing::info!("[OPENAI_STREAM] finalize 生成 {} 个事件", final_events.len()); for sse_str in final_events { - yield Ok::(sse_str); + yield Ok::(sse_str); } }; @@ -1954,13 +1974,10 @@ pub async fn handle_kiro_stream( // 使用缓存的 token 覆盖文件中的 token(缓存的 token 更新) kiro.credentials.access_token = Some(token); - // 转换请求格式 - let openai_request = convert_anthropic_to_openai(request); + tracing::info!("[KIRO_STREAM] 准备调用 call_api_stream_anthropic (直接转换)"); - tracing::info!("[KIRO_STREAM] 准备调用 call_api_stream"); - - // 调用流式 API(需求 4.1, 4.2, 4.3: 401/403 错误重试逻辑) - let stream_response = match kiro.call_api_stream(&openai_request).await { + // 调用流式 API - 直接使用 Anthropic 格式(需求 4.1, 4.2, 4.3: 401/403 错误重试逻辑) + let stream_response = match kiro.call_api_stream_anthropic(request).await { Ok(stream) => { tracing::info!("[KIRO_STREAM] call_api_stream 成功返回流"); stream @@ -2009,7 +2026,7 @@ pub async fn handle_kiro_stream( // 使用新 token 重试(需求 4.2) kiro.credentials.access_token = Some(new_token); - match kiro.call_api_stream(&openai_request).await { + match kiro.call_api_stream_anthropic(request).await { Ok(stream) => stream, Err(retry_err) => { let _ = state.pool_service.mark_unhealthy( @@ -2060,24 +2077,20 @@ pub async fn handle_kiro_stream( flow_id ); - // 创建 AWS Event Stream 解析器和 Anthropic SSE 生成器 - let parser = std::sync::Arc::new(tokio::sync::Mutex::new(AwsEventStreamParser::new())); - let generator = std::sync::Arc::new(tokio::sync::Mutex::new(AnthropicSseGenerator::new( - &request.model, - ))); + // 使用新的统一流处理管道 (Kiro → Anthropic) + let config = PipelineConfig::kiro_to_anthropic(request.model.clone()); + let pipeline = std::sync::Arc::new(tokio::sync::Mutex::new(StreamPipeline::new(config))); // 获取 flow_id 的克隆用于回调 let flow_id_owned = flow_id.map(|s| s.to_string()); let flow_monitor = state.flow_monitor.clone(); - // 创建转换流 - 使用 map 而不是 then,避免异步闭包的复杂性 - let parser_clone = parser.clone(); - let generator_clone = generator.clone(); + // 创建转换流 + let pipeline_clone = pipeline.clone(); let flow_id_for_stream = flow_id_owned.clone(); let flow_monitor_for_stream = flow_monitor.clone(); - // 使用 async_stream 直接处理整个流 - let generator_for_finalize = generator.clone(); + let pipeline_for_finalize = pipeline.clone(); let flow_id_for_finalize = flow_id_owned.clone(); let flow_monitor_for_finalize = flow_monitor.clone(); @@ -2089,112 +2102,60 @@ pub async fn handle_kiro_stream( while let Some(chunk_result) = stream_response.next().await { match chunk_result { Ok(bytes) => { - // 调试日志:记录接收到的字节数和原始数据预览 - let bytes_preview = if bytes.len() > 200 { - format!("{}...", String::from_utf8_lossy(&bytes[..200])) - } else { - String::from_utf8_lossy(&bytes).to_string() - }; - tracing::info!( - "[KIRO_STREAM] 收到 {} 字节数据, 预览: {}", - bytes.len(), - bytes_preview.replace('\n', "\\n") + // 调试日志:记录接收到的字节数 + tracing::debug!( + "[KIRO_STREAM] 收到 {} 字节数据", + bytes.len() ); - // 解析 AWS Event Stream - let events = { - let mut parser_guard = parser_clone.lock().await; - parser_guard.process(&bytes) + // 使用 Pipeline 处理字节块 + let sse_strings = { + let mut pipeline_guard = pipeline_clone.lock().await; + pipeline_guard.process_chunk(&bytes) }; - // 调试日志:记录解析出的事件数量 - tracing::info!( - "[KIRO_STREAM] 解析出 {} 个事件", - events.len() + // 调试日志:记录生成的 SSE 事件数量 + tracing::debug!( + "[KIRO_STREAM] 生成 {} 个 SSE 事件", + sse_strings.len() ); - // 转换为 Anthropic SSE 事件 - for event in events { - // 需求 5.2: 当 AWS Event Stream 解析失败时记录警告,跳过无效数据继续处理 - if let AwsEvent::ParseError { message, raw_data } = &event { - tracing::warn!( - "[KIRO_STREAM] AWS Event Stream 解析错误: {}, 原始数据: {:?}", - message, - raw_data.as_ref().map(|s| if s.len() > 100 { &s[..100] } else { s }) - ); - // 跳过无效数据,继续处理后续 chunks - continue; - } + for sse_str in sse_strings { + // 调用 FlowMonitor.process_chunk()(需求 3.2) + if let Some(ref fid) = flow_id_for_stream { + // 解析 SSE 事件类型和数据 + let lines: Vec<&str> = sse_str.lines().collect(); + let mut event_type: Option<&str> = None; + let mut data: Option<&str> = None; - // 调试日志:记录事件类型 - tracing::info!( - "[KIRO_STREAM] 处理事件: {:?}", - match &event { - AwsEvent::Content { text } => format!("Content({}字符): {}", text.len(), if text.len() > 50 { &text[..50] } else { text }), - AwsEvent::ToolUseStart { id, name } => format!("ToolUseStart({}, {})", id, name), - AwsEvent::ToolUseInput { id, input } => format!("ToolUseInput({}, {}字符)", id, input.len()), - AwsEvent::ToolUseStop { id } => format!("ToolUseStop({})", id), - AwsEvent::Stop => "Stop".to_string(), - AwsEvent::Usage { credits, context_percentage } => format!("Usage({}, {})", credits, context_percentage), - AwsEvent::FollowupPrompt { content } => format!("FollowupPrompt({}字符)", content.len()), - AwsEvent::ParseError { message, .. } => format!("ParseError({})", message), - } - ); - - let sse_strings = { - let mut generator_guard = generator_clone.lock().await; - generator_guard.process_event(event) - }; - - // 调试日志:记录生成的 SSE 事件数量和内容预览 - tracing::info!( - "[KIRO_STREAM] 生成 {} 个 SSE 事件", - sse_strings.len() - ); - - for sse_str in sse_strings { - let preview = if sse_str.len() > 200 { &sse_str[..200] } else { &sse_str }; - tracing::info!( - "[KIRO_STREAM] SSE 事件: {}", - preview.replace('\n', "\\n") - ); - - // 调用 FlowMonitor.process_chunk()(需求 3.2) - if let Some(ref fid) = flow_id_for_stream { - // 解析 SSE 事件类型和数据 - let lines: Vec<&str> = sse_str.lines().collect(); - let mut event_type: Option<&str> = None; - let mut data: Option<&str> = None; - - for line in &lines { - if line.starts_with("event: ") { - event_type = Some(&line[7..]); - } else if line.starts_with("data: ") { - data = Some(&line[6..]); - } - } - - if let Some(d) = data { - let flow_monitor_clone = flow_monitor_for_stream.clone(); - let fid_clone = fid.clone(); - let event_type_owned = event_type.map(|s| s.to_string()); - let data_owned = d.to_string(); - - tokio::spawn(async move { - flow_monitor_clone - .process_chunk( - &fid_clone, - event_type_owned.as_deref(), - &data_owned, - ) - .await; - }); + for line in &lines { + if line.starts_with("event: ") { + event_type = Some(&line[7..]); + } else if line.starts_with("data: ") { + data = Some(&line[6..]); } } - // 立即 yield SSE 事件 - yield Ok::(sse_str); + if let Some(d) = data { + let flow_monitor_clone = flow_monitor_for_stream.clone(); + let fid_clone = fid.clone(); + let event_type_owned = event_type.map(|s| s.to_string()); + let data_owned = d.to_string(); + + tokio::spawn(async move { + flow_monitor_clone + .process_chunk( + &fid_clone, + event_type_owned.as_deref(), + &data_owned, + ) + .await; + }); + } } + + // 立即 yield SSE 事件 + yield Ok::(sse_str); } } Err(e) => { @@ -2232,19 +2193,15 @@ pub async fn handle_kiro_stream( tracing::info!("[KIRO_STREAM] 流结束,生成 finalize 事件"); - // 流结束,生成 finalize 事件 + // 流结束,使用 Pipeline 生成 finalize 事件 let final_events = { - let mut generator_guard = generator_for_finalize.lock().await; - generator_guard.finalize() + let mut pipeline_guard = pipeline_for_finalize.lock().await; + pipeline_guard.finish() }; - tracing::info!("[KIRO_STREAM] finalize 生成 {} 个事件", final_events.len()); + tracing::debug!("[KIRO_STREAM] finalize 生成 {} 个事件", final_events.len()); for sse_str in final_events { - tracing::info!( - "[KIRO_STREAM] finalize 事件: {}", - if sse_str.len() > 200 { &sse_str[..200] } else { &sse_str }.replace('\n', "\\n") - ); // 调用 FlowMonitor.process_chunk() if let Some(ref fid) = flow_id_for_finalize { let lines: Vec<&str> = sse_str.lines().collect(); diff --git a/src-tauri/src/services/api_key_provider_service.rs b/src-tauri/src/services/api_key_provider_service.rs new file mode 100644 index 000000000..a07ba6755 --- /dev/null +++ b/src-tauri/src/services/api_key_provider_service.rs @@ -0,0 +1,650 @@ +//! API Key Provider 管理服务 +//! +//! 提供 API Key Provider 的 CRUD 操作、加密存储和轮询负载均衡功能。 +//! +//! **Feature: provider-ui-refactor** +//! **Validates: Requirements 7.3, 9.1, 9.2, 9.3** + +use crate::database::dao::api_key_provider::{ + ApiKeyEntry, ApiKeyProvider, ApiKeyProviderDao, ApiProviderType, ProviderGroup, + ProviderWithKeys, +}; +use crate::database::system_providers::{get_system_providers, to_api_key_provider}; +use crate::database::DbConnection; +use base64::{engine::general_purpose::STANDARD as BASE64, Engine}; +use chrono::Utc; +use sha2::{Digest, Sha256}; +use std::collections::HashMap; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::RwLock; + +// ============================================================================ +// 加密服务 +// ============================================================================ + +/// 简单的 API Key 加密服务 +/// 使用 XOR 加密 + Base64 编码 +/// 注意:这是一个简单的混淆方案,不是强加密 +struct EncryptionService { + /// 加密密钥(从机器 ID 派生) + key: Vec, +} + +impl EncryptionService { + /// 创建新的加密服务 + fn new() -> Self { + // 使用机器特定信息生成密钥 + let machine_id = Self::get_machine_id(); + let mut hasher = Sha256::new(); + hasher.update(machine_id.as_bytes()); + hasher.update(b"proxycast-api-key-encryption-salt"); + let key = hasher.finalize().to_vec(); + + Self { key } + } + + /// 获取机器 ID + fn get_machine_id() -> String { + // 尝试获取机器 ID,失败则使用默认值 + if let Ok(id) = std::fs::read_to_string("/etc/machine-id") { + return id.trim().to_string(); + } + if let Ok(id) = std::fs::read_to_string("/var/lib/dbus/machine-id") { + return id.trim().to_string(); + } + // macOS: 使用 IOPlatformUUID + #[cfg(target_os = "macos")] + { + if let Ok(output) = std::process::Command::new("ioreg") + .args(["-rd1", "-c", "IOPlatformExpertDevice"]) + .output() + { + let stdout = String::from_utf8_lossy(&output.stdout); + for line in stdout.lines() { + if line.contains("IOPlatformUUID") { + if let Some(uuid) = line.split('"').nth(3) { + return uuid.to_string(); + } + } + } + } + } + // 默认值 + "proxycast-default-machine-id".to_string() + } + + /// 加密 API Key + fn encrypt(&self, plaintext: &str) -> String { + let encrypted: Vec = plaintext + .as_bytes() + .iter() + .enumerate() + .map(|(i, b)| b ^ self.key[i % self.key.len()]) + .collect(); + BASE64.encode(encrypted) + } + + /// 解密 API Key + fn decrypt(&self, ciphertext: &str) -> Result { + let encrypted = BASE64 + .decode(ciphertext) + .map_err(|e| format!("Base64 解码失败: {}", e))?; + let decrypted: Vec = encrypted + .iter() + .enumerate() + .map(|(i, b)| b ^ self.key[i % self.key.len()]) + .collect(); + String::from_utf8(decrypted).map_err(|e| format!("UTF-8 解码失败: {}", e)) + } + + /// 检查是否为加密后的值(非明文) + fn is_encrypted(&self, value: &str) -> bool { + // 加密后的值是 Base64 编码的,通常不包含常见的 API Key 前缀 + !value.starts_with("sk-") + && !value.starts_with("pk-") + && !value.starts_with("api-") + && BASE64.decode(value).is_ok() + } +} + +// ============================================================================ +// API Key Provider 服务 +// ============================================================================ + +/// API Key Provider 管理服务 +pub struct ApiKeyProviderService { + /// 加密服务 + encryption: EncryptionService, + /// 轮询索引(按 provider_id 分组) + round_robin_index: RwLock>, +} + +impl Default for ApiKeyProviderService { + fn default() -> Self { + Self::new() + } +} + +impl ApiKeyProviderService { + /// 创建新的服务实例 + pub fn new() -> Self { + Self { + encryption: EncryptionService::new(), + round_robin_index: RwLock::new(HashMap::new()), + } + } + + // ==================== Provider 操作 ==================== + + /// 初始化系统 Provider + /// 检查数据库中是否存在系统 Provider,如果不存在则插入 + /// **Validates: Requirements 9.3** + pub fn initialize_system_providers(&self, db: &DbConnection) -> Result { + let conn = db.lock().map_err(|e| e.to_string())?; + let system_providers = get_system_providers(); + let mut inserted_count = 0; + + for def in &system_providers { + // 检查是否已存在 + let existing = + ApiKeyProviderDao::get_provider_by_id(&conn, def.id).map_err(|e| e.to_string())?; + + if existing.is_none() { + // 插入新的系统 Provider + let provider = to_api_key_provider(def); + ApiKeyProviderDao::insert_provider(&conn, &provider).map_err(|e| e.to_string())?; + inserted_count += 1; + } + } + + if inserted_count > 0 { + tracing::info!("初始化了 {} 个系统 Provider", inserted_count); + } + + Ok(inserted_count) + } + + /// 获取所有 Provider(包含 API Keys) + /// 首次调用时会自动初始化系统 Provider + pub fn get_all_providers(&self, db: &DbConnection) -> Result, String> { + // 首先确保系统 Provider 已初始化 + self.initialize_system_providers(db)?; + + let conn = db.lock().map_err(|e| e.to_string())?; + let mut providers = + ApiKeyProviderDao::get_all_providers_with_keys(&conn).map_err(|e| e.to_string())?; + + // 解密 API Keys(用于前端显示掩码) + for provider in &mut providers { + for _key in &mut provider.api_keys { + // 保持加密状态,前端会显示掩码 + } + } + + Ok(providers) + } + + /// 获取单个 Provider(包含 API Keys) + pub fn get_provider( + &self, + db: &DbConnection, + id: &str, + ) -> Result, String> { + let conn = db.lock().map_err(|e| e.to_string())?; + let provider = + ApiKeyProviderDao::get_provider_by_id(&conn, id).map_err(|e| e.to_string())?; + + match provider { + Some(p) => { + let api_keys = ApiKeyProviderDao::get_api_keys_by_provider(&conn, id) + .map_err(|e| e.to_string())?; + Ok(Some(ProviderWithKeys { + provider: p, + api_keys, + })) + } + None => Ok(None), + } + } + + /// 添加自定义 Provider + pub fn add_custom_provider( + &self, + db: &DbConnection, + name: String, + provider_type: ApiProviderType, + api_host: String, + api_version: Option, + project: Option, + location: Option, + region: Option, + ) -> Result { + let now = Utc::now(); + let id = format!("custom-{}", uuid::Uuid::new_v4()); + + let provider = ApiKeyProvider { + id: id.clone(), + name, + provider_type, + api_host, + is_system: false, + group: ProviderGroup::Custom, + enabled: true, + sort_order: 9999, // 自定义 Provider 排在最后 + api_version, + project, + location, + region, + created_at: now, + updated_at: now, + }; + + let conn = db.lock().map_err(|e| e.to_string())?; + ApiKeyProviderDao::insert_provider(&conn, &provider).map_err(|e| e.to_string())?; + + Ok(provider) + } + + /// 更新 Provider 配置 + pub fn update_provider( + &self, + db: &DbConnection, + id: &str, + name: Option, + api_host: Option, + enabled: Option, + sort_order: Option, + api_version: Option, + project: Option, + location: Option, + region: Option, + ) -> Result { + let conn = db.lock().map_err(|e| e.to_string())?; + let mut provider = ApiKeyProviderDao::get_provider_by_id(&conn, id) + .map_err(|e| e.to_string())? + .ok_or_else(|| format!("Provider not found: {}", id))?; + + // 更新字段 + if let Some(n) = name { + provider.name = n; + } + if let Some(h) = api_host { + provider.api_host = h; + } + if let Some(e) = enabled { + provider.enabled = e; + } + if let Some(s) = sort_order { + provider.sort_order = s; + } + if let Some(v) = api_version { + provider.api_version = if v.is_empty() { None } else { Some(v) }; + } + if let Some(p) = project { + provider.project = if p.is_empty() { None } else { Some(p) }; + } + if let Some(l) = location { + provider.location = if l.is_empty() { None } else { Some(l) }; + } + if let Some(r) = region { + provider.region = if r.is_empty() { None } else { Some(r) }; + } + provider.updated_at = Utc::now(); + + ApiKeyProviderDao::update_provider(&conn, &provider).map_err(|e| e.to_string())?; + + Ok(provider) + } + + /// 删除自定义 Provider + /// 系统 Provider 不允许删除 + pub fn delete_custom_provider(&self, db: &DbConnection, id: &str) -> Result { + let conn = db.lock().map_err(|e| e.to_string())?; + + // 检查是否为系统 Provider + let provider = ApiKeyProviderDao::get_provider_by_id(&conn, id) + .map_err(|e| e.to_string())? + .ok_or_else(|| format!("Provider not found: {}", id))?; + + if provider.is_system { + return Err("不允许删除系统 Provider".to_string()); + } + + ApiKeyProviderDao::delete_provider(&conn, id).map_err(|e| e.to_string()) + } + + // ==================== API Key 操作 ==================== + + /// 添加 API Key + pub fn add_api_key( + &self, + db: &DbConnection, + provider_id: &str, + api_key: &str, + alias: Option, + ) -> Result { + // 验证 Provider 存在 + let conn = db.lock().map_err(|e| e.to_string())?; + let _ = ApiKeyProviderDao::get_provider_by_id(&conn, provider_id) + .map_err(|e| e.to_string())? + .ok_or_else(|| format!("Provider not found: {}", provider_id))?; + + // 加密 API Key + let encrypted_key = self.encryption.encrypt(api_key); + + let now = Utc::now(); + let key = ApiKeyEntry { + id: uuid::Uuid::new_v4().to_string(), + provider_id: provider_id.to_string(), + api_key_encrypted: encrypted_key, + alias, + enabled: true, + usage_count: 0, + error_count: 0, + last_used_at: None, + created_at: now, + }; + + ApiKeyProviderDao::insert_api_key(&conn, &key).map_err(|e| e.to_string())?; + + Ok(key) + } + + /// 删除 API Key + pub fn delete_api_key(&self, db: &DbConnection, key_id: &str) -> Result { + let conn = db.lock().map_err(|e| e.to_string())?; + ApiKeyProviderDao::delete_api_key(&conn, key_id).map_err(|e| e.to_string()) + } + + /// 切换 API Key 启用状态 + pub fn toggle_api_key( + &self, + db: &DbConnection, + key_id: &str, + enabled: bool, + ) -> Result { + let conn = db.lock().map_err(|e| e.to_string())?; + let mut key = ApiKeyProviderDao::get_api_key_by_id(&conn, key_id) + .map_err(|e| e.to_string())? + .ok_or_else(|| format!("API Key not found: {}", key_id))?; + + key.enabled = enabled; + ApiKeyProviderDao::update_api_key(&conn, &key).map_err(|e| e.to_string())?; + + Ok(key) + } + + /// 更新 API Key 别名 + pub fn update_api_key_alias( + &self, + db: &DbConnection, + key_id: &str, + alias: Option, + ) -> Result { + let conn = db.lock().map_err(|e| e.to_string())?; + let mut key = ApiKeyProviderDao::get_api_key_by_id(&conn, key_id) + .map_err(|e| e.to_string())? + .ok_or_else(|| format!("API Key not found: {}", key_id))?; + + key.alias = alias; + ApiKeyProviderDao::update_api_key(&conn, &key).map_err(|e| e.to_string())?; + + Ok(key) + } + + // ==================== 轮询负载均衡 ==================== + + /// 获取下一个可用的 API Key(轮询负载均衡) + /// **Validates: Requirements 7.3** + pub fn get_next_api_key( + &self, + db: &DbConnection, + provider_id: &str, + ) -> Result, String> { + let conn = db.lock().map_err(|e| e.to_string())?; + + // 获取所有启用的 API Keys + let keys = ApiKeyProviderDao::get_enabled_api_keys_by_provider(&conn, provider_id) + .map_err(|e| e.to_string())?; + + if keys.is_empty() { + return Ok(None); + } + + // 获取或创建轮询索引 + let index = { + let mut indices = self.round_robin_index.write().map_err(|e| e.to_string())?; + indices + .entry(provider_id.to_string()) + .or_insert_with(|| AtomicUsize::new(0)) + .fetch_add(1, Ordering::SeqCst) + }; + + // 选择 API Key + let selected_key = &keys[index % keys.len()]; + + // 解密并返回 + let decrypted = self.encryption.decrypt(&selected_key.api_key_encrypted)?; + Ok(Some(decrypted)) + } + + /// 获取下一个可用的 API Key 条目(包含 ID,用于记录使用) + pub fn get_next_api_key_entry( + &self, + db: &DbConnection, + provider_id: &str, + ) -> Result, String> { + let conn = db.lock().map_err(|e| e.to_string())?; + + // 获取所有启用的 API Keys + let keys = ApiKeyProviderDao::get_enabled_api_keys_by_provider(&conn, provider_id) + .map_err(|e| e.to_string())?; + + if keys.is_empty() { + return Ok(None); + } + + // 获取或创建轮询索引 + let index = { + let mut indices = self.round_robin_index.write().map_err(|e| e.to_string())?; + indices + .entry(provider_id.to_string()) + .or_insert_with(|| AtomicUsize::new(0)) + .fetch_add(1, Ordering::SeqCst) + }; + + // 选择 API Key + let selected_key = &keys[index % keys.len()]; + + // 解密并返回 + let decrypted = self.encryption.decrypt(&selected_key.api_key_encrypted)?; + Ok(Some((selected_key.id.clone(), decrypted))) + } + + /// 记录 API Key 使用 + pub fn record_usage(&self, db: &DbConnection, key_id: &str) -> Result<(), String> { + let conn = db.lock().map_err(|e| e.to_string())?; + let key = ApiKeyProviderDao::get_api_key_by_id(&conn, key_id) + .map_err(|e| e.to_string())? + .ok_or_else(|| format!("API Key not found: {}", key_id))?; + + ApiKeyProviderDao::update_api_key_usage(&conn, key_id, key.usage_count + 1, Utc::now()) + .map_err(|e| e.to_string()) + } + + /// 记录 API Key 错误 + pub fn record_error(&self, db: &DbConnection, key_id: &str) -> Result<(), String> { + let conn = db.lock().map_err(|e| e.to_string())?; + ApiKeyProviderDao::increment_api_key_error(&conn, key_id).map_err(|e| e.to_string()) + } + + // ==================== 加密相关 ==================== + + /// 检查 API Key 是否已加密 + pub fn is_encrypted(&self, value: &str) -> bool { + self.encryption.is_encrypted(value) + } + + /// 解密 API Key(用于 API 调用) + pub fn decrypt_api_key(&self, encrypted: &str) -> Result { + self.encryption.decrypt(encrypted) + } + + /// 加密 API Key(用于存储) + pub fn encrypt_api_key(&self, plaintext: &str) -> String { + self.encryption.encrypt(plaintext) + } + + // ==================== UI 状态 ==================== + + /// 获取 UI 状态 + pub fn get_ui_state(&self, db: &DbConnection, key: &str) -> Result, String> { + let conn = db.lock().map_err(|e| e.to_string())?; + ApiKeyProviderDao::get_ui_state(&conn, key).map_err(|e| e.to_string()) + } + + /// 设置 UI 状态 + pub fn set_ui_state(&self, db: &DbConnection, key: &str, value: &str) -> Result<(), String> { + let conn = db.lock().map_err(|e| e.to_string())?; + ApiKeyProviderDao::set_ui_state(&conn, key, value).map_err(|e| e.to_string()) + } + + /// 批量更新 Provider 排序顺序 + /// **Validates: Requirements 8.4** + pub fn update_provider_sort_orders( + &self, + db: &DbConnection, + sort_orders: Vec<(String, i32)>, + ) -> Result<(), String> { + let conn = db.lock().map_err(|e| e.to_string())?; + ApiKeyProviderDao::update_provider_sort_orders(&conn, &sort_orders) + .map_err(|e| e.to_string()) + } + + // ==================== 导入导出 ==================== + + /// 导出配置 + pub fn export_config( + &self, + db: &DbConnection, + include_keys: bool, + ) -> Result { + let conn = db.lock().map_err(|e| e.to_string())?; + let providers = + ApiKeyProviderDao::get_all_providers_with_keys(&conn).map_err(|e| e.to_string())?; + + let export_data = if include_keys { + // 包含 API Keys(但不包含实际的 key 值) + let providers_json: Vec = providers + .iter() + .map(|p| { + let keys: Vec = p + .api_keys + .iter() + .map(|k| { + serde_json::json!({ + "id": k.id, + "alias": k.alias, + "enabled": k.enabled, + }) + }) + .collect(); + serde_json::json!({ + "provider": p.provider, + "api_keys": keys, + }) + }) + .collect(); + serde_json::json!({ + "version": "1.0", + "exported_at": Utc::now().to_rfc3339(), + "providers": providers_json, + }) + } else { + // 不包含 API Keys + let providers_json: Vec = providers + .iter() + .map(|p| serde_json::json!(p.provider)) + .collect(); + serde_json::json!({ + "version": "1.0", + "exported_at": Utc::now().to_rfc3339(), + "providers": providers_json, + }) + }; + + Ok(export_data) + } + + /// 导入配置 + pub fn import_config( + &self, + db: &DbConnection, + config_json: &str, + ) -> Result { + let config: serde_json::Value = + serde_json::from_str(config_json).map_err(|e| format!("JSON 解析失败: {}", e))?; + + let providers = config["providers"] + .as_array() + .ok_or_else(|| "配置格式错误: 缺少 providers 数组".to_string())?; + + let conn = db.lock().map_err(|e| e.to_string())?; + let mut imported_providers = 0; + let mut skipped_providers = 0; + let mut errors = Vec::new(); + + for provider_json in providers { + let provider_data = if provider_json.get("provider").is_some() { + &provider_json["provider"] + } else { + provider_json + }; + + let id = provider_data["id"] + .as_str() + .ok_or_else(|| "Provider 缺少 id".to_string())?; + + // 检查是否已存在 + if ApiKeyProviderDao::get_provider_by_id(&conn, id) + .map_err(|e| e.to_string())? + .is_some() + { + skipped_providers += 1; + continue; + } + + // 解析 Provider + let provider: ApiKeyProvider = serde_json::from_value(provider_data.clone()) + .map_err(|e| format!("Provider 解析失败: {}", e))?; + + // 插入 Provider + if let Err(e) = ApiKeyProviderDao::insert_provider(&conn, &provider) { + errors.push(format!("导入 Provider {} 失败: {}", id, e)); + continue; + } + + imported_providers += 1; + } + + Ok(ImportResult { + success: errors.is_empty(), + imported_providers, + imported_api_keys: 0, // API Keys 不在导入中包含实际值 + skipped_providers, + errors, + }) + } +} + +/// 导入结果 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ImportResult { + pub success: bool, + pub imported_providers: usize, + pub imported_api_keys: usize, + pub skipped_providers: usize, + pub errors: Vec, +} + +use serde::{Deserialize, Serialize}; diff --git a/src-tauri/src/services/mod.rs b/src-tauri/src/services/mod.rs index b34c9bbed..e6f5d5229 100644 --- a/src-tauri/src/services/mod.rs +++ b/src-tauri/src/services/mod.rs @@ -1,3 +1,4 @@ +pub mod api_key_provider_service; pub mod backup_service; pub mod kiro_event_service; pub mod live_sync; diff --git a/src-tauri/src/stream/events.rs b/src-tauri/src/stream/events.rs new file mode 100644 index 000000000..504ab32fb --- /dev/null +++ b/src-tauri/src/stream/events.rs @@ -0,0 +1,317 @@ +//! 统一流事件类型 +//! +//! 定义流式传输的中间表示 (Intermediate Representation), +//! 用于解耦解析器 (parsers) 和生成器 (generators)。 +//! +//! # 设计原则 +//! +//! - Parsers 输出 `StreamEvent` +//! - Generators 消费 `StreamEvent` 生成目标格式 +//! - 不同后端的解析器都输出相同的 `StreamEvent` 类型 +//! - 不同前端的生成器都消费相同的 `StreamEvent` 类型 + +use serde::{Deserialize, Serialize}; + +/// 统一流事件类型 +/// +/// 作为不同协议之间的中间表示,解耦: +/// - 后端流格式解析 (AWS Event Stream, OpenAI SSE, Anthropic SSE) +/// - 前端流格式生成 (OpenAI SSE, Anthropic SSE, Gemini SSE) +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +pub enum StreamEvent { + /// 消息开始 + /// + /// 表示一个新的消息/响应开始 + MessageStart { + /// 消息 ID + id: String, + /// 模型名称 + model: String, + }, + + /// 内容块开始 + /// + /// 表示一个新的内容块开始(文本或工具调用) + ContentBlockStart { + /// 内容块索引 + index: u32, + /// 内容块类型 + block_type: ContentBlockType, + }, + + /// 文本内容增量 + /// + /// 对应文本内容的增量输出 + TextDelta { + /// 文本内容 + text: String, + }, + + /// 工具调用开始 + /// + /// 表示一个新的工具调用开始 + ToolUseStart { + /// 工具调用 ID + id: String, + /// 工具名称 + name: String, + }, + + /// 工具调用参数增量 + /// + /// 工具调用参数的增量输出(部分 JSON) + ToolUseInputDelta { + /// 工具调用 ID + id: String, + /// 参数增量(部分 JSON 字符串) + partial_json: String, + }, + + /// 工具调用结束 + /// + /// 表示工具调用参数传输完成 + ToolUseStop { + /// 工具调用 ID + id: String, + }, + + /// 内容块结束 + /// + /// 表示一个内容块传输完成 + ContentBlockStop { + /// 内容块索引 + index: u32, + }, + + /// 消息结束 + /// + /// 表示整个消息/响应结束 + MessageStop { + /// 停止原因 + stop_reason: StopReason, + }, + + /// 使用量信息 + /// + /// Token 使用统计 + Usage { + /// 输入 token 数 + input_tokens: u32, + /// 输出 token 数 + output_tokens: u32, + /// 缓存读取 token 数(可选) + cache_read_input_tokens: Option, + /// 缓存创建 token 数(可选) + cache_creation_input_tokens: Option, + }, + + /// 后端特定使用量(如 CodeWhisperer credits) + /// + /// 用于传递后端特定的使用量信息 + BackendUsage { + /// 消耗的 credits + credits: f64, + /// 上下文使用百分比 + context_percentage: f64, + }, + + /// 错误事件 + /// + /// 流处理过程中的错误 + Error { + /// 错误类型 + error_type: String, + /// 错误消息 + message: String, + }, + + /// Ping/心跳事件 + /// + /// 保持连接活跃 + Ping, +} + +/// 内容块类型 +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub enum ContentBlockType { + /// 文本内容 + Text, + /// 工具调用 + ToolUse { + /// 工具调用 ID + id: String, + /// 工具名称 + name: String, + }, +} + +/// 停止原因 +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub enum StopReason { + /// 正常结束 + EndTurn, + /// 达到最大 token 数 + MaxTokens, + /// 需要工具调用 + ToolUse, + /// 用户停止 + StopSequence, + /// 其他原因 + Other(String), +} + +impl Default for StopReason { + fn default() -> Self { + Self::EndTurn + } +} + +impl StopReason { + /// 从字符串解析停止原因 + pub fn from_str(s: &str) -> Self { + match s.to_lowercase().as_str() { + "end_turn" | "stop" => Self::EndTurn, + "max_tokens" | "length" => Self::MaxTokens, + "tool_use" | "tool_calls" => Self::ToolUse, + "stop_sequence" => Self::StopSequence, + _ => Self::Other(s.to_string()), + } + } + + /// 转换为 OpenAI 格式的字符串 + pub fn to_openai_str(&self) -> &str { + match self { + Self::EndTurn => "stop", + Self::MaxTokens => "length", + Self::ToolUse => "tool_calls", + Self::StopSequence => "stop", + Self::Other(_) => "stop", + } + } + + /// 转换为 Anthropic 格式的字符串 + pub fn to_anthropic_str(&self) -> &str { + match self { + Self::EndTurn => "end_turn", + Self::MaxTokens => "max_tokens", + Self::ToolUse => "tool_use", + Self::StopSequence => "stop_sequence", + Self::Other(s) => s, + } + } +} + +/// 流事件上下文 +/// +/// 用于在流处理过程中跟踪状态 +#[derive(Debug, Clone, Default)] +pub struct StreamContext { + /// 消息 ID + pub message_id: Option, + /// 模型名称 + pub model: Option, + /// 当前内容块索引 + pub current_block_index: u32, + /// 活跃的工具调用 ID 列表 + pub active_tool_calls: Vec, + /// 累计输入 tokens + pub input_tokens: u32, + /// 累计输出 tokens + pub output_tokens: u32, +} + +impl StreamContext { + /// 创建新的上下文 + pub fn new() -> Self { + Self::default() + } + + /// 使用消息 ID 和模型创建上下文 + pub fn with_message(id: String, model: String) -> Self { + Self { + message_id: Some(id), + model: Some(model), + ..Default::default() + } + } + + /// 获取下一个内容块索引 + pub fn next_block_index(&mut self) -> u32 { + let index = self.current_block_index; + self.current_block_index += 1; + index + } + + /// 添加活跃的工具调用 + pub fn add_tool_call(&mut self, id: String) { + if !self.active_tool_calls.contains(&id) { + self.active_tool_calls.push(id); + } + } + + /// 移除工具调用 + pub fn remove_tool_call(&mut self, id: &str) { + self.active_tool_calls.retain(|x| x != id); + } + + /// 检查是否有活跃的工具调用 + pub fn has_active_tool_calls(&self) -> bool { + !self.active_tool_calls.is_empty() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_stop_reason_from_str() { + assert_eq!(StopReason::from_str("end_turn"), StopReason::EndTurn); + assert_eq!(StopReason::from_str("stop"), StopReason::EndTurn); + assert_eq!(StopReason::from_str("max_tokens"), StopReason::MaxTokens); + assert_eq!(StopReason::from_str("length"), StopReason::MaxTokens); + assert_eq!(StopReason::from_str("tool_use"), StopReason::ToolUse); + assert_eq!(StopReason::from_str("tool_calls"), StopReason::ToolUse); + } + + #[test] + fn test_stop_reason_to_openai() { + assert_eq!(StopReason::EndTurn.to_openai_str(), "stop"); + assert_eq!(StopReason::MaxTokens.to_openai_str(), "length"); + assert_eq!(StopReason::ToolUse.to_openai_str(), "tool_calls"); + } + + #[test] + fn test_stop_reason_to_anthropic() { + assert_eq!(StopReason::EndTurn.to_anthropic_str(), "end_turn"); + assert_eq!(StopReason::MaxTokens.to_anthropic_str(), "max_tokens"); + assert_eq!(StopReason::ToolUse.to_anthropic_str(), "tool_use"); + } + + #[test] + fn test_stream_context_block_index() { + let mut ctx = StreamContext::new(); + assert_eq!(ctx.next_block_index(), 0); + assert_eq!(ctx.next_block_index(), 1); + assert_eq!(ctx.next_block_index(), 2); + } + + #[test] + fn test_stream_context_tool_calls() { + let mut ctx = StreamContext::new(); + assert!(!ctx.has_active_tool_calls()); + + ctx.add_tool_call("tool_1".to_string()); + assert!(ctx.has_active_tool_calls()); + + ctx.add_tool_call("tool_2".to_string()); + assert_eq!(ctx.active_tool_calls.len(), 2); + + ctx.remove_tool_call("tool_1"); + assert_eq!(ctx.active_tool_calls.len(), 1); + assert!(ctx.has_active_tool_calls()); + + ctx.remove_tool_call("tool_2"); + assert!(!ctx.has_active_tool_calls()); + } +} diff --git a/src-tauri/src/stream/generators/anthropic_sse.rs b/src-tauri/src/stream/generators/anthropic_sse.rs new file mode 100644 index 000000000..0269ea231 --- /dev/null +++ b/src-tauri/src/stream/generators/anthropic_sse.rs @@ -0,0 +1,465 @@ +//! Anthropic SSE 生成器 +//! +//! 将 `StreamEvent` 转换为 Anthropic Messages API SSE 格式。 +//! +//! # 格式说明 +//! +//! Anthropic SSE 格式: +//! ```text +//! event: message_start +//! data: {"type":"message_start","message":{...}} +//! +//! 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"}} +//! +//! event: message_stop +//! data: {"type":"message_stop"} +//! ``` + +use crate::stream::events::{ContentBlockType, StopReason, StreamEvent}; +use std::collections::HashMap; +use uuid::Uuid; + +/// 工具调用状态 +#[derive(Debug, Clone, Default)] +struct ToolCallState { + /// 工具调用 ID + id: String, + /// 工具名称 + name: String, + /// 累积的输入 JSON + input: String, + /// 内容块索引 + index: u32, +} + +/// Anthropic SSE 生成器 +#[derive(Debug)] +pub struct AnthropicSseGenerator { + /// 消息 ID + message_id: String, + /// 模型名称 + model: String, + /// 是否已发送 message_start 事件 + message_started: bool, + /// 工具调用状态映射 + tool_calls: HashMap, + /// 输入 token 数量 + input_tokens: u32, + /// 输出 token 数量 + output_tokens: u32, + /// 缓存读取 token 数 + cache_read_input_tokens: u32, + /// 缓存创建 token 数 + cache_creation_input_tokens: u32, + /// 累积的停止原因 + stop_reason: Option, +} + +impl Default for AnthropicSseGenerator { + fn default() -> Self { + Self::new("unknown".to_string()) + } +} + +impl AnthropicSseGenerator { + /// 创建新的生成器 + pub fn new(model: String) -> Self { + Self { + message_id: format!("msg_{}", Uuid::new_v4().simple()), + model, + message_started: false, + tool_calls: HashMap::new(), + input_tokens: 0, + output_tokens: 0, + cache_read_input_tokens: 0, + cache_creation_input_tokens: 0, + stop_reason: None, + } + } + + /// 使用指定的消息 ID 创建生成器 + pub fn with_id(id: String, model: String) -> Self { + Self { + message_id: id, + model, + message_started: false, + tool_calls: HashMap::new(), + input_tokens: 0, + output_tokens: 0, + cache_read_input_tokens: 0, + cache_creation_input_tokens: 0, + stop_reason: None, + } + } + + /// 将 StreamEvent 转换为 Anthropic SSE 字符串列表 + /// + /// # 返回 + /// + /// SSE 事件字符串列表,每个字符串都是完整的 SSE 事件(包含 `event:` 和 `data:` 行) + pub fn generate(&mut self, event: &StreamEvent) -> Vec { + let mut sse_events = Vec::new(); + + // 确保发送 message_start + if !self.message_started { + match event { + StreamEvent::MessageStart { id, model } => { + self.message_id = id.clone(); + self.model = model.clone(); + } + _ => {} + } + sse_events.push(self.create_message_start()); + self.message_started = true; + } + + match event { + StreamEvent::MessageStart { .. } => { + // 已经在上面处理了 + } + + StreamEvent::ContentBlockStart { index, block_type } => { + match block_type { + ContentBlockType::Text => { + sse_events.push(self.create_content_block_start_text(*index)); + } + ContentBlockType::ToolUse { id, name } => { + // 记录工具调用状态 + self.tool_calls.insert( + id.clone(), + ToolCallState { + id: id.clone(), + name: name.clone(), + input: String::new(), + index: *index, + }, + ); + sse_events.push(self.create_content_block_start_tool(*index, id, name)); + } + } + } + + StreamEvent::TextDelta { text } => { + // 假设文本块索引为 0(大多数情况下) + sse_events.push(self.create_text_delta(0, text)); + } + + StreamEvent::ToolUseStart { id, name } => { + // 如果还没有对应的 ContentBlockStart,创建工具调用状态 + if !self.tool_calls.contains_key(id) { + let index = self.tool_calls.len() as u32; + self.tool_calls.insert( + id.clone(), + ToolCallState { + id: id.clone(), + name: name.clone(), + input: String::new(), + index, + }, + ); + } + // ToolUseStart 已经在 ContentBlockStart 中处理 + } + + StreamEvent::ToolUseInputDelta { id, partial_json } => { + if let Some(state) = self.tool_calls.get_mut(id) { + state.input.push_str(partial_json); + let index = state.index; + sse_events.push(self.create_input_json_delta(index, partial_json)); + } + } + + StreamEvent::ToolUseStop { id } => { + // 工具调用结束,但保留状态直到 ContentBlockStop + let _ = id; + } + + StreamEvent::ContentBlockStop { index } => { + sse_events.push(self.create_content_block_stop(*index)); + } + + StreamEvent::MessageStop { stop_reason } => { + self.stop_reason = Some(stop_reason.clone()); + // 移除所有工具调用状态 + self.tool_calls.clear(); + sse_events.push(self.create_message_delta(stop_reason)); + sse_events.push(self.create_message_stop()); + } + + StreamEvent::Usage { + input_tokens, + output_tokens, + cache_read_input_tokens, + cache_creation_input_tokens, + } => { + self.input_tokens = *input_tokens; + self.output_tokens = *output_tokens; + if let Some(cache_read) = cache_read_input_tokens { + self.cache_read_input_tokens = *cache_read; + } + if let Some(cache_creation) = cache_creation_input_tokens { + self.cache_creation_input_tokens = *cache_creation; + } + // Anthropic 在流中发送 ping 事件而不是 usage 事件 + } + + StreamEvent::BackendUsage { .. } => { + // 后端特定的使用量信息,不转换 + } + + StreamEvent::Error { + error_type, + message, + } => { + sse_events.push(self.create_error(error_type, message)); + } + + StreamEvent::Ping => { + sse_events.push(self.create_ping()); + } + } + + sse_events + } + + /// 获取消息 ID + pub fn message_id(&self) -> &str { + &self.message_id + } + + /// 获取模型名称 + pub fn model(&self) -> &str { + &self.model + } + + // ======================================================================== + // SSE 事件创建方法 + // ======================================================================== + + fn create_message_start(&self) -> String { + let event = serde_json::json!({ + "type": "message_start", + "message": { + "id": self.message_id, + "type": "message", + "role": "assistant", + "model": self.model, + "content": [], + "stop_reason": serde_json::Value::Null, + "stop_sequence": serde_json::Value::Null, + "usage": { + "input_tokens": self.input_tokens, + "output_tokens": self.output_tokens, + "cache_read_input_tokens": self.cache_read_input_tokens, + "cache_creation_input_tokens": self.cache_creation_input_tokens + } + } + }); + format!("event: message_start\ndata: {}\n\n", event) + } + + fn create_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_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_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_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_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_message_delta(&self, stop_reason: &StopReason) -> String { + let event = serde_json::json!({ + "type": "message_delta", + "delta": { + "stop_reason": stop_reason.to_anthropic_str(), + "stop_sequence": serde_json::Value::Null + }, + "usage": { + "output_tokens": self.output_tokens + } + }); + format!("event: message_delta\ndata: {}\n\n", event) + } + + fn create_message_stop(&self) -> String { + let event = serde_json::json!({ + "type": "message_stop" + }); + format!("event: message_stop\ndata: {}\n\n", event) + } + + fn create_ping(&self) -> String { + let event = serde_json::json!({ + "type": "ping" + }); + format!("event: ping\ndata: {}\n\n", event) + } + + fn create_error(&self, error_type: &str, message: &str) -> String { + let event = serde_json::json!({ + "type": "error", + "error": { + "type": error_type, + "message": message + } + }); + format!("event: error\ndata: {}\n\n", event) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_generate_message_start() { + let mut generator = AnthropicSseGenerator::new("claude-3-sonnet".to_string()); + let event = StreamEvent::MessageStart { + id: "msg_123".to_string(), + model: "claude-3-sonnet".to_string(), + }; + + let sse = generator.generate(&event); + assert_eq!(sse.len(), 1); + assert!(sse[0].starts_with("event: message_start\ndata: ")); + assert!(sse[0].contains("\"id\":\"msg_123\"")); + } + + #[test] + fn test_generate_text_content() { + let mut generator = AnthropicSseGenerator::new("claude-3-sonnet".to_string()); + + // 先发送 message_start + let _ = generator.generate(&StreamEvent::MessageStart { + id: "msg_123".to_string(), + model: "claude-3-sonnet".to_string(), + }); + + // 内容块开始 + let sse = generator.generate(&StreamEvent::ContentBlockStart { + index: 0, + block_type: ContentBlockType::Text, + }); + assert!(sse[0].contains("content_block_start")); + assert!(sse[0].contains("\"type\":\"text\"")); + + // 文本增量 + let sse = generator.generate(&StreamEvent::TextDelta { + text: "Hello".to_string(), + }); + assert!(sse[0].contains("content_block_delta")); + assert!(sse[0].contains("text_delta")); + assert!(sse[0].contains("Hello")); + } + + #[test] + fn test_generate_tool_use() { + let mut generator = AnthropicSseGenerator::new("claude-3-sonnet".to_string()); + + // 发送 message_start + let _ = generator.generate(&StreamEvent::MessageStart { + id: "msg_123".to_string(), + model: "claude-3-sonnet".to_string(), + }); + + // 工具调用内容块开始 + let sse = generator.generate(&StreamEvent::ContentBlockStart { + index: 1, + block_type: ContentBlockType::ToolUse { + id: "tool_abc".to_string(), + name: "read_file".to_string(), + }, + }); + assert!(sse[0].contains("content_block_start")); + assert!(sse[0].contains("\"type\":\"tool_use\"")); + assert!(sse[0].contains("\"id\":\"tool_abc\"")); + assert!(sse[0].contains("\"name\":\"read_file\"")); + + // 工具参数增量 + let sse = generator.generate(&StreamEvent::ToolUseInputDelta { + id: "tool_abc".to_string(), + partial_json: "{\"path\":".to_string(), + }); + assert!(sse[0].contains("content_block_delta")); + assert!(sse[0].contains("input_json_delta")); + } + + #[test] + fn test_generate_message_stop() { + let mut generator = AnthropicSseGenerator::new("claude-3-sonnet".to_string()); + + // 发送 message_start + let _ = generator.generate(&StreamEvent::MessageStart { + id: "msg_123".to_string(), + model: "claude-3-sonnet".to_string(), + }); + + // 消息结束 + let sse = generator.generate(&StreamEvent::MessageStop { + stop_reason: StopReason::EndTurn, + }); + assert_eq!(sse.len(), 2); + assert!(sse[0].contains("message_delta")); + assert!(sse[0].contains("end_turn")); + assert!(sse[1].contains("message_stop")); + } +} diff --git a/src-tauri/src/stream/generators/mod.rs b/src-tauri/src/stream/generators/mod.rs new file mode 100644 index 000000000..c92108e12 --- /dev/null +++ b/src-tauri/src/stream/generators/mod.rs @@ -0,0 +1,15 @@ +//! SSE 流生成器 +//! +//! 将统一的 `StreamEvent` 转换为不同前端协议的 SSE 格式。 +//! +//! # 支持的格式 +//! +//! - OpenAI SSE (data: {...}) +//! - Anthropic SSE (event: xxx\ndata: {...}) +//! - Gemini SSE (待实现) + +pub mod anthropic_sse; +pub mod openai_sse; + +pub use anthropic_sse::AnthropicSseGenerator; +pub use openai_sse::OpenAiSseGenerator; diff --git a/src-tauri/src/stream/generators/openai_sse.rs b/src-tauri/src/stream/generators/openai_sse.rs new file mode 100644 index 000000000..ffc59d0c7 --- /dev/null +++ b/src-tauri/src/stream/generators/openai_sse.rs @@ -0,0 +1,402 @@ +//! OpenAI SSE 生成器 +//! +//! 将 `StreamEvent` 转换为 OpenAI Chat Completions SSE 格式。 +//! +//! # 格式说明 +//! +//! OpenAI SSE 格式: +//! ```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] +//! ``` + +use crate::stream::events::{ContentBlockType, StopReason, StreamEvent}; +use serde::Serialize; +use std::collections::HashMap; +use std::time::{SystemTime, UNIX_EPOCH}; + +/// OpenAI SSE 生成器 +#[derive(Debug)] +pub struct OpenAiSseGenerator { + /// 响应 ID + response_id: String, + /// 模型名称 + model: String, + /// 创建时间戳 + created: u64, + /// 工具调用状态 (tool_call_id -> (index, name, accumulated_args)) + tool_calls: HashMap, + /// 下一个工具调用索引 + next_tool_index: usize, +} + +#[derive(Debug, Clone)] +struct ToolCallState { + /// 工具调用在 tool_calls 数组中的索引 + index: usize, + /// 工具名称 + name: String, + /// 累积的参数 + arguments: String, +} + +impl Default for OpenAiSseGenerator { + fn default() -> Self { + Self::new("unknown".to_string()) + } +} + +impl OpenAiSseGenerator { + /// 创建新的生成器 + pub fn new(model: String) -> Self { + let created = SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|d| d.as_secs()) + .unwrap_or(0); + + Self { + response_id: format!("chatcmpl-{}", uuid::Uuid::new_v4().simple()), + model, + created, + tool_calls: HashMap::new(), + next_tool_index: 0, + } + } + + /// 使用指定的响应 ID 创建生成器 + pub fn with_id(id: String, model: String) -> Self { + let created = SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|d| d.as_secs()) + .unwrap_or(0); + + Self { + response_id: id, + model, + created, + tool_calls: HashMap::new(), + next_tool_index: 0, + } + } + + /// 将 StreamEvent 转换为 OpenAI SSE 字符串 + /// + /// # 返回 + /// + /// - `Some(String)` - 生成的 SSE 字符串(包含 `data: ` 前缀和换行) + /// - `None` - 该事件不需要生成 SSE 输出 + pub fn generate(&mut self, event: &StreamEvent) -> Option { + match event { + StreamEvent::MessageStart { id, model } => { + self.response_id = id.clone(); + self.model = model.clone(); + // OpenAI 格式不需要单独的 message_start 事件 + None + } + + StreamEvent::ContentBlockStart { block_type, .. } => { + // OpenAI 格式不需要单独的 content_block_start 事件 + // 但我们需要跟踪工具调用 + if let ContentBlockType::ToolUse { id, name } = block_type { + let index = self.next_tool_index; + self.next_tool_index += 1; + self.tool_calls.insert( + id.clone(), + ToolCallState { + index, + name: name.clone(), + arguments: String::new(), + }, + ); + } + None + } + + StreamEvent::TextDelta { text } => { + let chunk = OpenAiStreamChunk { + id: &self.response_id, + object: "chat.completion.chunk", + created: self.created, + model: &self.model, + choices: vec![OpenAiChoice { + index: 0, + delta: OpenAiDelta { + role: None, + content: Some(text.as_str()), + tool_calls: None, + }, + finish_reason: None, + }], + }; + Some(format!("data: {}\n\n", serde_json::to_string(&chunk).ok()?)) + } + + StreamEvent::ToolUseStart { id, name } => { + // 确保工具调用状态存在 + let index = if let Some(state) = self.tool_calls.get(id) { + state.index + } else { + let index = self.next_tool_index; + self.next_tool_index += 1; + self.tool_calls.insert( + id.clone(), + ToolCallState { + index, + name: name.clone(), + arguments: String::new(), + }, + ); + index + }; + + let chunk = OpenAiStreamChunk { + id: &self.response_id, + object: "chat.completion.chunk", + created: self.created, + model: &self.model, + choices: vec![OpenAiChoice { + index: 0, + delta: OpenAiDelta { + role: None, + content: None, + tool_calls: Some(vec![OpenAiToolCallDelta { + index, + id: Some(id.as_str()), + r#type: Some("function"), + function: Some(OpenAiFunctionDelta { + name: Some(name.as_str()), + arguments: None, + }), + }]), + }, + finish_reason: None, + }], + }; + Some(format!("data: {}\n\n", serde_json::to_string(&chunk).ok()?)) + } + + StreamEvent::ToolUseInputDelta { id, partial_json } => { + let index = self.tool_calls.get(id)?.index; + + // 累积参数 + if let Some(state) = self.tool_calls.get_mut(id) { + state.arguments.push_str(partial_json); + } + + let chunk = OpenAiStreamChunk { + id: &self.response_id, + object: "chat.completion.chunk", + created: self.created, + model: &self.model, + choices: vec![OpenAiChoice { + index: 0, + delta: OpenAiDelta { + role: None, + content: None, + tool_calls: Some(vec![OpenAiToolCallDelta { + index, + id: None, + r#type: None, + function: Some(OpenAiFunctionDelta { + name: None, + arguments: Some(partial_json.as_str()), + }), + }]), + }, + finish_reason: None, + }], + }; + Some(format!("data: {}\n\n", serde_json::to_string(&chunk).ok()?)) + } + + StreamEvent::ToolUseStop { id } => { + // OpenAI 格式不需要单独的工具调用结束事件 + self.tool_calls.remove(id); + None + } + + StreamEvent::ContentBlockStop { .. } => { + // OpenAI 格式不需要单独的 content_block_stop 事件 + None + } + + StreamEvent::MessageStop { stop_reason } => { + let finish_reason = stop_reason.to_openai_str(); + + let chunk = OpenAiStreamChunk { + id: &self.response_id, + object: "chat.completion.chunk", + created: self.created, + model: &self.model, + choices: vec![OpenAiChoice { + index: 0, + delta: OpenAiDelta { + role: None, + content: None, + tool_calls: None, + }, + finish_reason: Some(finish_reason), + }], + }; + + let chunk_str = format!("data: {}\n\n", serde_json::to_string(&chunk).ok()?); + Some(format!("{}data: [DONE]\n\n", chunk_str)) + } + + StreamEvent::Usage { + input_tokens, + output_tokens, + .. + } => { + // OpenAI 在流式响应中通常不发送 usage + // 但某些实现可能需要,这里可以选择性地生成 + let _ = (input_tokens, output_tokens); + None + } + + StreamEvent::BackendUsage { .. } => { + // 后端特定的使用量信息,不转换为 OpenAI 格式 + None + } + + StreamEvent::Error { + error_type, + message, + } => { + // 生成错误响应 + let error_obj = serde_json::json!({ + "error": { + "type": error_type, + "message": message, + } + }); + Some(format!("data: {}\n\n", error_obj)) + } + + StreamEvent::Ping => { + // 心跳事件,生成空的 SSE 注释 + Some(": ping\n\n".to_string()) + } + } + } + + /// 生成 [DONE] 事件 + pub fn generate_done(&self) -> String { + "data: [DONE]\n\n".to_string() + } + + /// 获取响应 ID + pub fn response_id(&self) -> &str { + &self.response_id + } +} + +// ============================================================================ +// OpenAI SSE 数据结构 +// ============================================================================ + +#[derive(Debug, Serialize)] +struct OpenAiStreamChunk<'a> { + id: &'a str, + object: &'a str, + created: u64, + model: &'a str, + choices: Vec>, +} + +#[derive(Debug, Serialize)] +struct OpenAiChoice<'a> { + index: usize, + delta: OpenAiDelta<'a>, + #[serde(skip_serializing_if = "Option::is_none")] + finish_reason: Option<&'a str>, +} + +#[derive(Debug, Serialize)] +struct OpenAiDelta<'a> { + #[serde(skip_serializing_if = "Option::is_none")] + role: Option<&'a str>, + #[serde(skip_serializing_if = "Option::is_none")] + content: Option<&'a str>, + #[serde(skip_serializing_if = "Option::is_none")] + tool_calls: Option>>, +} + +#[derive(Debug, Serialize)] +struct OpenAiToolCallDelta<'a> { + index: usize, + #[serde(skip_serializing_if = "Option::is_none")] + id: Option<&'a str>, + #[serde(skip_serializing_if = "Option::is_none")] + r#type: Option<&'a str>, + #[serde(skip_serializing_if = "Option::is_none")] + function: Option>, +} + +#[derive(Debug, Serialize)] +struct OpenAiFunctionDelta<'a> { + #[serde(skip_serializing_if = "Option::is_none")] + name: Option<&'a str>, + #[serde(skip_serializing_if = "Option::is_none")] + arguments: Option<&'a str>, +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_generate_text_delta() { + let mut generator = OpenAiSseGenerator::new("gpt-4".to_string()); + let event = StreamEvent::TextDelta { + text: "Hello".to_string(), + }; + + let sse = generator.generate(&event); + assert!(sse.is_some()); + let sse = sse.unwrap(); + assert!(sse.starts_with("data: ")); + assert!(sse.contains("\"content\":\"Hello\"")); + } + + #[test] + fn test_generate_tool_call() { + let mut generator = OpenAiSseGenerator::new("gpt-4".to_string()); + + // 工具调用开始 + let event = StreamEvent::ToolUseStart { + id: "call_123".to_string(), + name: "read_file".to_string(), + }; + let sse = generator.generate(&event); + assert!(sse.is_some()); + let sse = sse.unwrap(); + assert!(sse.contains("\"tool_calls\"")); + assert!(sse.contains("\"name\":\"read_file\"")); + + // 工具参数增量 + let event = StreamEvent::ToolUseInputDelta { + id: "call_123".to_string(), + partial_json: "{\"path\":".to_string(), + }; + let sse = generator.generate(&event); + assert!(sse.is_some()); + let sse = sse.unwrap(); + assert!(sse.contains("\"arguments\":\"{\\\"path\\\":\"")); + } + + #[test] + fn test_generate_message_stop() { + let mut generator = OpenAiSseGenerator::new("gpt-4".to_string()); + let event = StreamEvent::MessageStop { + stop_reason: StopReason::EndTurn, + }; + + let sse = generator.generate(&event); + assert!(sse.is_some()); + let sse = sse.unwrap(); + assert!(sse.contains("\"finish_reason\":\"stop\"")); + assert!(sse.contains("[DONE]")); + } +} diff --git a/src-tauri/src/stream/mod.rs b/src-tauri/src/stream/mod.rs new file mode 100644 index 000000000..9c730f174 --- /dev/null +++ b/src-tauri/src/stream/mod.rs @@ -0,0 +1,36 @@ +//! 流式处理层 +//! +//! 提供统一的流式数据处理能力,包括: +//! - 事件类型定义 (events) +//! - 后端流格式解析 (parsers) +//! - 前端流格式生成 (generators) +//! +//! # 架构设计 +//! +//! ```text +//! 后端响应流 ──> [Parser] ──> StreamEvent ──> [Generator] ──> 前端 SSE +//! +//! 例如: +//! AWS Event Stream ──> [AwsEventStreamParser] ──> StreamEvent ──> [AnthropicSseGenerator] ──> Anthropic SSE +//! AWS Event Stream ──> [AwsEventStreamParser] ──> StreamEvent ──> [OpenAiSseGenerator] ──> OpenAI SSE +//! ``` +//! +//! # 模块结构 +//! +//! - `events`: 统一的流事件类型定义 (`StreamEvent`) +//! - `parsers`: 后端流格式解析器 +//! - `aws_event_stream`: AWS Event Stream 解析器 (Kiro/CodeWhisperer) +//! - `generators`: 前端流格式生成器 +//! - `openai_sse`: OpenAI SSE 格式生成器 +//! - `anthropic_sse`: Anthropic SSE 格式生成器 + +pub mod events; +pub mod generators; +pub mod parsers; +pub mod pipeline; + +// 重新导出核心类型 +pub use events::{ContentBlockType, StopReason, StreamContext, StreamEvent}; +pub use generators::{AnthropicSseGenerator, OpenAiSseGenerator}; +pub use parsers::{AwsEventStreamParser, ParserState}; +pub use pipeline::{create_sse_stream, BackendType, FrontendType, PipelineConfig, StreamPipeline}; diff --git a/src-tauri/src/stream/parsers/aws_event_stream.rs b/src-tauri/src/stream/parsers/aws_event_stream.rs new file mode 100644 index 000000000..01eb2266f --- /dev/null +++ b/src-tauri/src/stream/parsers/aws_event_stream.rs @@ -0,0 +1,500 @@ +//! AWS Event Stream 解析器 +//! +//! 解析 Kiro/CodeWhisperer 的 AWS Event Stream 二进制格式, +//! 输出统一的 `StreamEvent` 类型。 +//! +//! # 协议格式 +//! +//! CodeWhisperer 使用 AWS Event Stream 二进制格式,每个事件包含: +//! - `{"content": "文本内容"}` - 文本增量 +//! - `{"toolUseId": "id", "name": "tool_name"}` - 工具调用开始 +//! - `{"toolUseId": "id", "input": "部分JSON"}` - 工具参数增量 +//! - `{"toolUseId": "id", "stop": true}` - 工具调用结束 +//! - `{"stop": true}` - 流结束 +//! - `{"usage": 0.34}` - Credits 使用量 +//! - `{"contextUsagePercentage": 54.36}` - 上下文使用百分比 + +use crate::stream::events::{ContentBlockType, StopReason, StreamContext, StreamEvent}; +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 + } +} + +/// 工具调用累积器 +#[derive(Debug, Clone, Default)] +struct ToolAccumulator { + /// 工具名称 + name: String, + /// 累积的输入 + input: String, + /// 内容块索引 + block_index: u32, +} + +/// AWS Event Stream 解析器 +/// +/// 解析 CodeWhisperer 的 AWS Event Stream 格式,输出统一的 `StreamEvent`。 +#[derive(Debug)] +pub struct AwsEventStreamParser { + /// 缓冲区(用于处理部分 chunk) + buffer: Vec, + /// 当前状态 + state: ParserState, + /// 工具调用累积器 + tool_accumulators: HashMap, + /// 解析错误计数 + parse_error_count: u32, + /// 最大缓冲区大小(防止内存耗尽) + max_buffer_size: usize, + /// 流上下文 + context: StreamContext, + /// 是否已发送消息开始事件 + message_started: bool, + /// 是否在文本块中 + in_text_block: bool, + /// 当前文本块索引 + text_block_index: Option, +} + +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, + context: StreamContext::new(), + message_started: false, + in_text_block: false, + text_block_index: None, + } + } + + /// 创建带模型名称的解析器 + pub fn with_model(model: String) -> Self { + let mut parser = Self::new(); + parser.context.model = Some(model); + parser + } + + /// 获取当前状态 + 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; + self.context = StreamContext::new(); + self.message_started = false; + self.in_text_block = false; + self.text_block_index = None; + } + + /// 处理接收到的字节 + /// + /// # 返回 + /// + /// 解析出的 `StreamEvent` 列表 + 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![StreamEvent::Error { + error_type: "buffer_overflow".to_string(), + message: "缓冲区溢出".to_string(), + }]; + } + + // 将新数据添加到缓冲区 + 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(StreamEvent::ToolUseStop { id: id.clone() }); + events.push(StreamEvent::ContentBlockStop { + index: accumulator.block_index, + }); + } + } + + // 关闭文本块(如果有) + if let Some(index) = self.text_block_index.take() { + events.push(StreamEvent::ContentBlockStop { index }); + } + + // 更新状态 + 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) => { + self.parse_error_count += 1; + events.push(StreamEvent::Error { + error_type: "parse_error".to_string(), + message: e, + }); + } + } + pos = end_pos; + } + None => { + // JSON 对象不完整,等待更多数据 + break; + } + } + } + + // 移除已处理的数据 + if pos > 0 { + self.buffer.drain(..pos); + } + + events + } + + /// 查找 JSON 对象的开始位置 + fn find_json_start(&self, from: usize) -> Option { + self.buffer[from..] + .iter() + .position(|&b| b == b'{') + .map(|p| from + p) + } + + /// 从缓冲区中提取完整的 JSON 对象 + 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; + } + } + } + _ => {} + } + } + + None + } + + /// 解析 JSON 事件并生成 StreamEvent + 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(); + + // 如果还没发送消息开始事件,先发送 + if !self.message_started { + self.message_started = true; + let msg_id = format!("msg_{}", uuid::Uuid::new_v4().simple()); + self.context.message_id = Some(msg_id.clone()); + events.push(StreamEvent::MessageStart { + id: msg_id, + model: self + .context + .model + .clone() + .unwrap_or_else(|| "unknown".to_string()), + }); + } + + // 处理 content 事件 + if let Some(content) = value.get("content").and_then(|v| v.as_str()) { + // 跳过 followupPrompt + if value.get("followupPrompt").is_none() { + // 如果还没有文本块,创建一个 + if !self.in_text_block { + self.in_text_block = true; + let index = self.context.next_block_index(); + self.text_block_index = Some(index); + events.push(StreamEvent::ContentBlockStart { + index, + block_type: ContentBlockType::Text, + }); + } + + events.push(StreamEvent::TextDelta { + text: content.to_string(), + }); + } + } + // 处理 tool use 事件 (包含 toolUseId) + else if let Some(tool_use_id) = value.get("toolUseId").and_then(|v| v.as_str()) { + // 如果有文本块,先关闭它 + if let Some(index) = self.text_block_index.take() { + self.in_text_block = false; + events.push(StreamEvent::ContentBlockStop { index }); + } + + 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(); + accumulator.block_index = self.context.next_block_index(); + self.context.add_tool_call(tool_id.clone()); + + events.push(StreamEvent::ContentBlockStart { + index: accumulator.block_index, + block_type: ContentBlockType::ToolUse { + id: tool_id.clone(), + name: name.clone(), + }, + }); + + events.push(StreamEvent::ToolUseStart { + id: tool_id.clone(), + name, + }); + } + + // 如果有输入增量 + if !input_chunk.is_empty() { + accumulator.input.push_str(&input_chunk); + events.push(StreamEvent::ToolUseInputDelta { + id: tool_id.clone(), + partial_json: input_chunk, + }); + } + + // 如果是 stop 事件 + if is_stop { + if let Some(acc) = self.tool_accumulators.remove(&tool_id) { + self.context.remove_tool_call(&tool_id); + events.push(StreamEvent::ToolUseStop { id: tool_id }); + events.push(StreamEvent::ContentBlockStop { + index: acc.block_index, + }); + } + } + } + // 处理独立的 stop 事件 + else if value.get("stop").and_then(|v| v.as_bool()).unwrap_or(false) { + // 关闭文本块(如果有) + if let Some(index) = self.text_block_index.take() { + self.in_text_block = false; + events.push(StreamEvent::ContentBlockStop { index }); + } + + // 确定停止原因 + let stop_reason = if self.context.has_active_tool_calls() { + StopReason::ToolUse + } else { + StopReason::EndTurn + }; + + events.push(StreamEvent::MessageStop { stop_reason }); + } + // 处理 usage 事件 + else if let Some(usage) = value.get("usage").and_then(|v| v.as_f64()) { + events.push(StreamEvent::BackendUsage { + credits: usage, + context_percentage: 0.0, + }); + } + // 处理 contextUsagePercentage 事件 + else if let Some(ctx_usage) = value.get("contextUsagePercentage").and_then(|v| v.as_f64()) + { + events.push(StreamEvent::BackendUsage { + credits: 0.0, + context_percentage: ctx_usage, + }); + } + + Ok(events) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_parse_content_event() { + let mut parser = AwsEventStreamParser::with_model("test-model".to_string()); + let events = parser.process(br#"{"content":"Hello"}"#); + + assert!(events.len() >= 2); + assert!( + matches!(&events[0], StreamEvent::MessageStart { model, .. } if model == "test-model") + ); + assert!(matches!( + &events[1], + StreamEvent::ContentBlockStart { + block_type: ContentBlockType::Text, + .. + } + )); + assert!(matches!(&events[2], StreamEvent::TextDelta { text } if text == "Hello")); + } + + #[test] + fn test_parse_tool_use_event() { + let mut parser = AwsEventStreamParser::new(); + + // 工具调用开始 + let events = parser.process(br#"{"toolUseId":"tool_123","name":"read_file"}"#); + assert!(events.iter().any(|e| matches!(e, StreamEvent::ToolUseStart { id, name } if id == "tool_123" && name == "read_file"))); + + // 工具参数增量 + let events = parser.process(br#"{"toolUseId":"tool_123","input":"{\"path\":"}"#); + assert!(events.iter().any(|e| matches!(e, StreamEvent::ToolUseInputDelta { id, partial_json } if id == "tool_123" && partial_json == "{\"path\":"))); + + // 工具调用结束 + let events = parser.process(br#"{"toolUseId":"tool_123","stop":true}"#); + assert!(events + .iter() + .any(|e| matches!(e, StreamEvent::ToolUseStop { id } if id == "tool_123"))); + } + + #[test] + fn test_parse_stop_event() { + let mut parser = AwsEventStreamParser::new(); + + // 先发送一些内容 + let _ = parser.process(br#"{"content":"test"}"#); + + // 发送 stop 事件 + let events = parser.process(br#"{"stop":true}"#); + assert!(events.iter().any(|e| matches!( + e, + StreamEvent::MessageStop { + stop_reason: StopReason::EndTurn + } + ))); + } + + #[test] + fn test_incremental_parsing() { + let mut parser = AwsEventStreamParser::new(); + + // 发送部分数据 + let events1 = parser.process(br#"{"con"#); + assert!(events1.is_empty()); // 不完整,没有事件 + + // 发送剩余数据 + let events2 = parser.process(br#"tent":"Hello"}"#); + assert!(!events2.is_empty()); // 现在有事件了 + } +} diff --git a/src-tauri/src/stream/parsers/mod.rs b/src-tauri/src/stream/parsers/mod.rs new file mode 100644 index 000000000..c84f9a0f6 --- /dev/null +++ b/src-tauri/src/stream/parsers/mod.rs @@ -0,0 +1,13 @@ +//! 流式数据解析器 +//! +//! 解析不同后端的流式响应格式,输出统一的 `StreamEvent`。 +//! +//! # 支持的格式 +//! +//! - AWS Event Stream (Kiro/CodeWhisperer) +//! - OpenAI SSE (待实现) +//! - Anthropic SSE (待实现) + +pub mod aws_event_stream; + +pub use aws_event_stream::{AwsEventStreamParser, ParserState}; diff --git a/src-tauri/src/stream/pipeline.rs b/src-tauri/src/stream/pipeline.rs new file mode 100644 index 000000000..79b7b96a3 --- /dev/null +++ b/src-tauri/src/stream/pipeline.rs @@ -0,0 +1,327 @@ +//! 统一流处理管道 +//! +//! 封装完整的流式处理流程:后端字节流 → 解析 → 转换 → 前端 SSE +//! +//! # 使用示例 +//! +//! ```ignore +//! use proxycast::stream::pipeline::{StreamPipeline, PipelineConfig}; +//! +//! let config = PipelineConfig::kiro_to_anthropic("claude-sonnet-4-5".to_string()); +//! let pipeline = StreamPipeline::new(config); +//! +//! // 处理字节流 +//! let sse_stream = pipeline.process_stream(byte_stream); +//! ``` + +use crate::stream::events::StreamEvent; +use crate::stream::generators::{AnthropicSseGenerator, OpenAiSseGenerator}; +use crate::stream::parsers::AwsEventStreamParser; +use bytes::Bytes; +use futures::{Stream, StreamExt}; + +/// 后端类型 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum BackendType { + /// Kiro/CodeWhisperer (AWS Event Stream) + Kiro, + /// OpenAI (SSE) + OpenAi, + /// Anthropic (SSE) + Anthropic, +} + +/// 前端类型 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum FrontendType { + /// OpenAI SSE 格式 + OpenAi, + /// Anthropic SSE 格式 + Anthropic, +} + +/// 流处理管道配置 +#[derive(Debug, Clone)] +pub struct PipelineConfig { + /// 后端类型 + pub backend: BackendType, + /// 前端类型 + pub frontend: FrontendType, + /// 模型名称 + pub model: String, + /// 消息 ID(可选) + pub message_id: Option, +} + +impl PipelineConfig { + /// 创建 Kiro → Anthropic 配置 + pub fn kiro_to_anthropic(model: String) -> Self { + Self { + backend: BackendType::Kiro, + frontend: FrontendType::Anthropic, + model, + message_id: None, + } + } + + /// 创建 Kiro → OpenAI 配置 + pub fn kiro_to_openai(model: String) -> Self { + Self { + backend: BackendType::Kiro, + frontend: FrontendType::OpenAi, + model, + message_id: None, + } + } + + /// 设置消息 ID + pub fn with_message_id(mut self, id: String) -> Self { + self.message_id = Some(id); + self + } +} + +/// SSE 生成器封装 +enum SseGenerator { + Anthropic(AnthropicSseGenerator), + OpenAi(OpenAiSseGenerator), +} + +impl SseGenerator { + fn generate(&mut self, event: &StreamEvent) -> Vec { + match self { + SseGenerator::Anthropic(g) => g.generate(event), + SseGenerator::OpenAi(g) => g.generate(event).into_iter().collect(), + } + } +} + +/// 统一流处理管道 +/// +/// 将后端字节流转换为前端 SSE 字符串流 +pub struct StreamPipeline { + /// 配置 + config: PipelineConfig, + /// AWS Event Stream 解析器(用于 Kiro 后端) + aws_parser: Option, + /// SSE 生成器 + generator: SseGenerator, +} + +impl StreamPipeline { + /// 创建新的管道 + pub fn new(config: PipelineConfig) -> Self { + let aws_parser = match config.backend { + BackendType::Kiro => Some(AwsEventStreamParser::with_model(config.model.clone())), + _ => None, + }; + + let generator = match config.frontend { + FrontendType::Anthropic => { + if let Some(id) = &config.message_id { + SseGenerator::Anthropic(AnthropicSseGenerator::with_id( + id.clone(), + config.model.clone(), + )) + } else { + SseGenerator::Anthropic(AnthropicSseGenerator::new(config.model.clone())) + } + } + FrontendType::OpenAi => { + if let Some(id) = &config.message_id { + SseGenerator::OpenAi(OpenAiSseGenerator::with_id( + id.clone(), + config.model.clone(), + )) + } else { + SseGenerator::OpenAi(OpenAiSseGenerator::new(config.model.clone())) + } + } + }; + + Self { + config, + aws_parser, + generator, + } + } + + /// 处理单个字节块 + /// + /// # 返回 + /// + /// 生成的 SSE 字符串列表 + pub fn process_chunk(&mut self, bytes: &[u8]) -> Vec { + let events = self.parse_bytes(bytes); + self.generate_sse(&events) + } + + /// 完成处理 + /// + /// # 返回 + /// + /// 最终的 SSE 字符串列表 + pub fn finish(&mut self) -> Vec { + let events = self.finish_parsing(); + self.generate_sse(&events) + } + + /// 解析字节为 StreamEvent + fn parse_bytes(&mut self, bytes: &[u8]) -> Vec { + match &mut self.aws_parser { + Some(parser) => parser.process(bytes), + None => Vec::new(), // TODO: 支持其他后端格式的解析 + } + } + + /// 完成解析 + fn finish_parsing(&mut self) -> Vec { + match &mut self.aws_parser { + Some(parser) => parser.finish(), + None => Vec::new(), + } + } + + /// 将 StreamEvent 转换为 SSE 字符串 + fn generate_sse(&mut self, events: &[StreamEvent]) -> Vec { + let mut result = Vec::new(); + for event in events { + let sse_strings = self.generator.generate(event); + result.extend(sse_strings); + } + result + } + + /// 获取配置 + pub fn config(&self) -> &PipelineConfig { + &self.config + } + + /// 重置管道状态 + pub fn reset(&mut self) { + if let Some(ref mut parser) = self.aws_parser { + parser.reset(); + } + self.generator = match self.config.frontend { + FrontendType::Anthropic => { + SseGenerator::Anthropic(AnthropicSseGenerator::new(self.config.model.clone())) + } + FrontendType::OpenAi => { + SseGenerator::OpenAi(OpenAiSseGenerator::new(self.config.model.clone())) + } + }; + } +} + +/// 创建流式处理的异步流 +/// +/// 将字节流转换为 SSE 字符串流 +pub fn create_sse_stream( + byte_stream: S, + config: PipelineConfig, +) -> impl Stream> +where + S: Stream> + Send + 'static, + E: Send + 'static, +{ + async_stream::stream! { + let mut pipeline = StreamPipeline::new(config); + let mut byte_stream = std::pin::pin!(byte_stream); + + while let Some(result) = byte_stream.next().await { + match result { + Ok(bytes) => { + let sse_strings = pipeline.process_chunk(&bytes); + for sse in sse_strings { + yield Ok(sse); + } + } + Err(e) => { + yield Err(e); + return; + } + } + } + + // 完成处理 + let final_sse = pipeline.finish(); + for sse in final_sse { + yield Ok(sse); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_pipeline_config_kiro_to_anthropic() { + let config = PipelineConfig::kiro_to_anthropic("claude-sonnet-4-5".to_string()); + assert_eq!(config.backend, BackendType::Kiro); + assert_eq!(config.frontend, FrontendType::Anthropic); + assert_eq!(config.model, "claude-sonnet-4-5"); + } + + #[test] + fn test_pipeline_config_kiro_to_openai() { + let config = PipelineConfig::kiro_to_openai("gpt-4".to_string()); + assert_eq!(config.backend, BackendType::Kiro); + assert_eq!(config.frontend, FrontendType::OpenAi); + } + + #[test] + fn test_pipeline_process_content() { + let config = PipelineConfig::kiro_to_anthropic("claude-sonnet-4-5".to_string()); + let mut pipeline = StreamPipeline::new(config); + + // 模拟 Kiro 内容事件 + let bytes = br#"{"content":"Hello"}"#; + let sse = pipeline.process_chunk(bytes); + + // 应该生成 message_start, content_block_start, content_block_delta + assert!(!sse.is_empty()); + assert!(sse.iter().any(|s| s.contains("message_start"))); + assert!(sse.iter().any(|s| s.contains("content_block_start"))); + assert!(sse.iter().any(|s| s.contains("Hello"))); + } + + #[test] + fn test_pipeline_process_tool_use() { + let config = PipelineConfig::kiro_to_anthropic("claude-sonnet-4-5".to_string()); + let mut pipeline = StreamPipeline::new(config); + + // 工具调用开始 + let bytes = br#"{"toolUseId":"tool_123","name":"read_file"}"#; + let sse = pipeline.process_chunk(bytes); + + assert!(!sse.is_empty()); + assert!(sse.iter().any(|s| s.contains("tool_use"))); + assert!(sse.iter().any(|s| s.contains("read_file"))); + + // 工具参数 + let bytes = br#"{"toolUseId":"tool_123","input":"{\"path\":"}"#; + let sse = pipeline.process_chunk(bytes); + assert!(sse.iter().any(|s| s.contains("input_json_delta"))); + + // 工具结束 + let bytes = br#"{"toolUseId":"tool_123","stop":true}"#; + let sse = pipeline.process_chunk(bytes); + assert!(sse.iter().any(|s| s.contains("content_block_stop"))); + } + + #[test] + fn test_pipeline_openai_output() { + let config = PipelineConfig::kiro_to_openai("gpt-4".to_string()); + let mut pipeline = StreamPipeline::new(config); + + // 模拟 Kiro 内容事件 + let bytes = br#"{"content":"Hello"}"#; + let sse = pipeline.process_chunk(bytes); + + assert!(!sse.is_empty()); + // OpenAI 格式应该包含 data: 前缀和 choices + assert!(sse.iter().any(|s| s.starts_with("data: "))); + assert!(sse.iter().any(|s| s.contains("\"content\":\"Hello\""))); + } +} diff --git a/src-tauri/src/translator/kiro/anthropic/mod.rs b/src-tauri/src/translator/kiro/anthropic/mod.rs new file mode 100644 index 000000000..430de7866 --- /dev/null +++ b/src-tauri/src/translator/kiro/anthropic/mod.rs @@ -0,0 +1,13 @@ +//! Anthropic 协议 → Kiro 后端转换 +//! +//! 处理 Anthropic Messages API 格式与 CodeWhisperer 格式之间的转换。 +//! +//! # 重要说明 +//! +//! 这是 Claude Code 使用的协议,是最核心的转换路径。 + +pub mod request; +pub mod response; + +pub use request::AnthropicRequestTranslator; +pub use response::AnthropicResponseTranslator; diff --git a/src-tauri/src/translator/kiro/anthropic/request.rs b/src-tauri/src/translator/kiro/anthropic/request.rs new file mode 100644 index 000000000..58296ae44 --- /dev/null +++ b/src-tauri/src/translator/kiro/anthropic/request.rs @@ -0,0 +1,677 @@ +//! Anthropic 请求直接转换为 CodeWhisperer 请求 +//! +//! 直接将 Anthropic MessagesRequest 转换为 CodeWhisperer API 格式, +//! 无需经过 OpenAI 中间格式,减少转换开销。 + +use crate::models::anthropic::*; +use crate::models::codewhisperer::*; +use crate::translator::kiro::openai::request::{get_model_map, DEFAULT_MODEL}; +use crate::translator::traits::{RequestTranslator, TranslateError}; +use std::collections::HashSet; +use uuid::Uuid; + +/// Anthropic 到 Kiro 请求转换器 +#[derive(Debug, Clone)] +pub struct AnthropicRequestTranslator { + /// 可选的 Profile ARN (AWS CodeWhisperer) + pub profile_arn: Option, +} + +impl Default for AnthropicRequestTranslator { + fn default() -> Self { + Self::new() + } +} + +impl AnthropicRequestTranslator { + /// 创建新的转换器 + pub fn new() -> Self { + Self { profile_arn: None } + } + + /// 使用 Profile ARN 创建转换器 + pub fn with_profile_arn(profile_arn: String) -> Self { + Self { + profile_arn: Some(profile_arn), + } + } +} + +impl RequestTranslator for AnthropicRequestTranslator { + type Input = AnthropicMessagesRequest; + type Output = CodeWhispererRequest; + type Error = TranslateError; + + fn translate_request(&self, request: Self::Input) -> Result { + Ok(convert_anthropic_to_codewhisperer( + &request, + self.profile_arn.clone(), + )) + } +} + +// ============================================================================ +// 内部类型 +// ============================================================================ + +#[derive(Debug, Clone)] +struct ProcessedMessage { + role: String, + content: String, + tool_uses: Option>, + tool_results: Option>, +} + +// ============================================================================ +// 转换函数 +// ============================================================================ + +/// 将 Anthropic MessagesRequest 直接转换为 CodeWhisperer 请求 +pub fn convert_anthropic_to_codewhisperer( + request: &AnthropicMessagesRequest, + profile_arn: Option, +) -> CodeWhispererRequest { + let model_map = get_model_map(); + let cw_model = model_map + .get(request.model.as_str()) + .map(|s| s.to_string()) + .unwrap_or_else(|| DEFAULT_MODEL.to_string()); + + let conversation_id = Uuid::new_v4().to_string(); + + // 提取 system prompt + let mut system_prompt = extract_system_text(&request.system); + + // 处理 tool_choice: required - CodeWhisperer 不支持此参数,通过 prompt 注入强制 + if is_tool_choice_required(&request.tool_choice) && request.tools.is_some() { + let tool_instruction = "\n\n[CRITICAL INSTRUCTION] You MUST use one of the provided tools to respond. Do NOT respond with plain text. Call a tool function immediately."; + system_prompt.push_str(tool_instruction); + tracing::info!("[KIRO_TRANSLATE] tool_choice=required detected in Anthropic request, injected tool instruction"); + } + + // 预处理消息 + let messages = preprocess_anthropic_messages(&request.messages); + + // 构建历史记录 + let mut history: Vec = Vec::new(); + let mut start_idx = 0; + + // 处理 system prompt - 合并到第一条用户消息 + if !system_prompt.is_empty() && !messages.is_empty() && messages[0].role == "user" { + let first_content = &messages[0].content; + let combined = format!("{system_prompt}\n\n{first_content}"); + + let mut user_msg = UserInputMessage { + content: combined, + model_id: cw_model.clone(), + origin: "AI_EDITOR".to_string(), + images: None, + user_input_message_context: None, + }; + + if let Some(ref tool_results) = messages[0].tool_results { + user_msg.user_input_message_context = Some(UserInputMessageContext { + tools: None, + tool_results: Some(tool_results.clone()), + }); + } + + history.push(HistoryItem::User(UserHistoryItem { + user_input_message: user_msg, + })); + start_idx = 1; + } + + // 处理历史消息(除最后一条) + for msg in messages + .iter() + .take(messages.len().saturating_sub(1)) + .skip(start_idx) + { + match msg.role.as_str() { + "user" => { + let content = if msg.content.is_empty() { + if msg.tool_results.is_some() { + "Tool results provided.".to_string() + } else { + "Continue".to_string() + } + } else { + msg.content.clone() + }; + + let mut user_msg = UserInputMessage { + content, + model_id: cw_model.clone(), + origin: "AI_EDITOR".to_string(), + images: None, + user_input_message_context: None, + }; + + if let Some(ref tool_results) = msg.tool_results { + user_msg.user_input_message_context = Some(UserInputMessageContext { + tools: None, + tool_results: Some(tool_results.clone()), + }); + } + + history.push(HistoryItem::User(UserHistoryItem { + user_input_message: user_msg, + })); + } + "assistant" => { + let content = if msg.content.is_empty() { + "I understand.".to_string() + } else { + msg.content.clone() + }; + + history.push(HistoryItem::Assistant(AssistantHistoryItem { + assistant_response_message: AssistantResponseMessage { + content, + tool_uses: msg.tool_uses.clone(), + }, + })); + } + _ => {} + } + } + + // 修复历史记录交替顺序 + let history = fix_history_alternation(history, &cw_model); + + // 构建当前消息 + let (current_content, current_tool_results) = if let Some(last_msg) = messages.last() { + if last_msg.role == "assistant" { + ("Continue".to_string(), None) + } else { + let content = if last_msg.content.is_empty() { + if last_msg.tool_results.is_some() { + "Tool results provided.".to_string() + } else { + "Continue".to_string() + } + } else { + last_msg.content.clone() + }; + (content, last_msg.tool_results.clone()) + } + } else { + ("Continue".to_string(), None) + }; + + // 构建 tools + let tools = convert_anthropic_tools(&request.tools); + + let user_input_message_context = if tools.is_some() || current_tool_results.is_some() { + Some(UserInputMessageContext { + tools, + tool_results: current_tool_results, + }) + } else { + None + }; + + CodeWhispererRequest { + conversation_state: ConversationState { + chat_trigger_type: "MANUAL".to_string(), + conversation_id, + current_message: CurrentMessage { + user_input_message: UserInputMessage { + content: current_content, + model_id: cw_model, + origin: "AI_EDITOR".to_string(), + images: None, + user_input_message_context, + }, + }, + history: if history.is_empty() { + None + } else { + Some(history) + }, + }, + profile_arn, + } +} + +/// 提取 system prompt 文本 +fn extract_system_text(system: &Option) -> String { + match system { + Some(serde_json::Value::String(s)) => s.clone(), + Some(serde_json::Value::Array(arr)) => arr + .iter() + .filter_map(|item| { + if item.get("type") == Some(&serde_json::Value::String("text".to_string())) { + item.get("text") + .and_then(|t| t.as_str()) + .map(|s| s.to_string()) + } else { + None + } + }) + .collect::>() + .join("\n"), + _ => String::new(), + } +} + +/// 预处理 Anthropic 消息 +fn preprocess_anthropic_messages(messages: &[AnthropicMessage]) -> Vec { + let mut result: Vec = Vec::new(); + + for msg in messages { + let processed = convert_anthropic_message(msg); + result.extend(processed); + } + + // 合并连续的 user 消息中的 tool_results + let mut merged: Vec = Vec::new(); + let mut pending_tool_results: Vec = Vec::new(); + + for msg in result { + if msg.role == "user" { + if let Some(ref tr) = msg.tool_results { + pending_tool_results.extend(tr.clone()); + } + if !msg.content.is_empty() || msg.tool_results.is_none() { + // 去重 tool_results + let mut seen_ids = HashSet::new(); + pending_tool_results.retain(|tr| seen_ids.insert(tr.tool_use_id.clone())); + + merged.push(ProcessedMessage { + role: msg.role, + content: if msg.content.is_empty() && !pending_tool_results.is_empty() { + "Tool results provided.".to_string() + } else { + msg.content + }, + tool_uses: None, + tool_results: if pending_tool_results.is_empty() { + None + } else { + Some(pending_tool_results.clone()) + }, + }); + pending_tool_results.clear(); + } + } else { + // 如果有待处理的 tool_results,先创建 user 消息 + if !pending_tool_results.is_empty() { + let mut seen_ids = HashSet::new(); + pending_tool_results.retain(|tr| seen_ids.insert(tr.tool_use_id.clone())); + + merged.push(ProcessedMessage { + role: "user".to_string(), + content: "Tool results provided.".to_string(), + tool_uses: None, + tool_results: Some(pending_tool_results.clone()), + }); + pending_tool_results.clear(); + } + merged.push(msg); + } + } + + // 处理末尾的 tool_results + if !pending_tool_results.is_empty() { + let mut seen_ids = HashSet::new(); + pending_tool_results.retain(|tr| seen_ids.insert(tr.tool_use_id.clone())); + + merged.push(ProcessedMessage { + role: "user".to_string(), + content: "Tool results provided.".to_string(), + tool_uses: None, + tool_results: Some(pending_tool_results), + }); + } + + merged +} + +/// 转换单条 Anthropic 消息 +fn convert_anthropic_message(msg: &AnthropicMessage) -> Vec { + let mut result: Vec = Vec::new(); + + match &msg.content { + serde_json::Value::String(s) => { + result.push(ProcessedMessage { + role: msg.role.clone(), + content: s.clone(), + tool_uses: None, + tool_results: None, + }); + } + serde_json::Value::Array(parts) => { + let mut text_parts: Vec = Vec::new(); + let mut tool_uses: Vec = Vec::new(); + let mut tool_results: Vec = Vec::new(); + + for part in parts { + let part_type = part.get("type").and_then(|t| t.as_str()).unwrap_or(""); + + match part_type { + "text" => { + if let Some(text) = part.get("text").and_then(|t| t.as_str()) { + text_parts.push(text.to_string()); + } + } + "tool_use" => { + let default_id = format!("toolu_{}", &Uuid::new_v4().to_string()[..8]); + let id = part + .get("id") + .and_then(|i| i.as_str()) + .unwrap_or(&default_id); + let name = part.get("name").and_then(|n| n.as_str()).unwrap_or(""); + let input = part.get("input").cloned().unwrap_or(serde_json::json!({})); + + tool_uses.push(CWToolUse { + tool_use_id: id.to_string(), + name: name.to_string(), + input, + }); + } + "tool_result" => { + let tool_use_id = part + .get("tool_use_id") + .and_then(|i| i.as_str()) + .unwrap_or(""); + let content_text = extract_tool_result_content(part.get("content")); + let is_error = part + .get("is_error") + .and_then(|e| e.as_bool()) + .unwrap_or(false); + + tool_results.push(CWToolResult { + tool_use_id: tool_use_id.to_string(), + content: vec![CWTextContent { text: content_text }], + status: if is_error { + "error".to_string() + } else { + "success".to_string() + }, + }); + } + _ => {} + } + } + + // 处理 assistant 消息 + if msg.role == "assistant" { + result.push(ProcessedMessage { + role: "assistant".to_string(), + content: text_parts.join(""), + tool_uses: if tool_uses.is_empty() { + None + } else { + Some(tool_uses) + }, + tool_results: None, + }); + } + // 处理 user 消息 + else if msg.role == "user" { + // 先添加 tool results + if !tool_results.is_empty() { + result.push(ProcessedMessage { + role: "user".to_string(), + content: String::new(), + tool_uses: None, + tool_results: Some(tool_results), + }); + } + + // 添加文本内容 + if !text_parts.is_empty() { + result.push(ProcessedMessage { + role: "user".to_string(), + content: text_parts.join(""), + tool_uses: None, + tool_results: None, + }); + } + } + } + _ => {} + } + + result +} + +/// 提取 tool_result 内容 +fn extract_tool_result_content(content: Option<&serde_json::Value>) -> String { + match content { + Some(serde_json::Value::String(s)) => s.clone(), + Some(serde_json::Value::Array(arr)) => arr + .iter() + .filter_map(|item| { + if item.get("type") == Some(&serde_json::Value::String("text".to_string())) { + item.get("text") + .and_then(|t| t.as_str()) + .map(|s| s.to_string()) + } else { + None + } + }) + .collect::>() + .join("\n"), + _ => String::new(), + } +} + +/// 转换 Anthropic tools 为 CodeWhisperer tools +fn convert_anthropic_tools(tools: &Option>) -> Option> { + tools.as_ref().map(|tools| { + let mut cw_tools: Vec = Vec::new(); + let mut function_count = 0; + + for t in tools.iter() { + // 处理特殊工具类型 + if t.name == "web_search" || t.name == "web_search_20250305" { + cw_tools.push(CWToolItem::WebSearch(CWWebSearchTool { + tool_type: "web_search".to_string(), + })); + continue; + } + + // 限制最多 50 个函数工具 + if function_count >= 50 { + continue; + } + function_count += 1; + + let params = t + .input_schema + .clone() + .unwrap_or_else(|| serde_json::json!({"type": "object", "properties": {}})); + + let desc = t + .description + .clone() + .unwrap_or_else(|| format!("Tool: {}", t.name)); + + cw_tools.push(CWToolItem::Standard(CWTool { + tool_specification: ToolSpecification { + name: t.name.clone(), + description: if desc.len() > 500 { + let truncated: String = desc.chars().take(497).collect(); + format!("{}...", truncated) + } else { + desc + }, + input_schema: InputSchema { json: params }, + }, + })); + } + + cw_tools + }) +} + +/// 修复历史记录,确保 user/assistant 严格交替 +fn fix_history_alternation(history: Vec, model_id: &str) -> Vec { + if history.is_empty() { + return history; + } + + let mut fixed: Vec = Vec::new(); + + for item in history { + match &item { + HistoryItem::User(user_item) => { + if let Some(HistoryItem::User(last_user)) = fixed.last_mut() { + let has_tool_results = user_item + .user_input_message + .user_input_message_context + .as_ref() + .map(|ctx| ctx.tool_results.is_some()) + .unwrap_or(false); + + if has_tool_results { + let new_results = user_item + .user_input_message + .user_input_message_context + .as_ref() + .and_then(|ctx| ctx.tool_results.clone()) + .unwrap_or_default(); + + if let Some(ref mut ctx) = + last_user.user_input_message.user_input_message_context + { + if let Some(ref mut existing) = ctx.tool_results { + existing.extend(new_results); + } else { + ctx.tool_results = Some(new_results); + } + } else { + last_user.user_input_message.user_input_message_context = + Some(UserInputMessageContext { + tools: None, + tool_results: Some(new_results), + }); + } + continue; + } else { + fixed.push(HistoryItem::Assistant(AssistantHistoryItem { + assistant_response_message: AssistantResponseMessage { + content: "I understand.".to_string(), + tool_uses: None, + }, + })); + } + } + fixed.push(item); + } + HistoryItem::Assistant(_) => { + if let Some(HistoryItem::Assistant(_)) = fixed.last() { + fixed.push(HistoryItem::User(UserHistoryItem { + user_input_message: UserInputMessage { + content: "Continue".to_string(), + model_id: model_id.to_string(), + origin: "AI_EDITOR".to_string(), + images: None, + user_input_message_context: None, + }, + })); + } + if fixed.is_empty() { + fixed.push(HistoryItem::User(UserHistoryItem { + user_input_message: UserInputMessage { + content: "Continue".to_string(), + model_id: model_id.to_string(), + origin: "AI_EDITOR".to_string(), + images: None, + user_input_message_context: None, + }, + })); + } + fixed.push(item); + } + } + } + + // 确保以 assistant 结尾 + if let Some(HistoryItem::User(_)) = fixed.last() { + fixed.push(HistoryItem::Assistant(AssistantHistoryItem { + assistant_response_message: AssistantResponseMessage { + content: "I understand.".to_string(), + tool_uses: None, + }, + })); + } + + fixed +} + +/// 检查 tool_choice 是否为 required +/// +/// Anthropic tool_choice 可以是: +/// - {"type": "any"} - 必须调用工具 +/// - {"type": "tool", "name": "xxx"} - 必须调用指定工具 +fn is_tool_choice_required(tool_choice: &Option) -> bool { + match tool_choice { + Some(serde_json::Value::Object(obj)) => { + if let Some(serde_json::Value::String(t)) = obj.get("type") { + t == "any" || t == "tool" + } else { + false + } + } + // OpenAI 风格的 "required" 字符串 + Some(serde_json::Value::String(s)) => s == "required" || s == "any", + _ => false, + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_convert_simple_request() { + let request = AnthropicMessagesRequest { + model: "claude-sonnet-4-5".to_string(), + messages: vec![AnthropicMessage { + role: "user".to_string(), + content: serde_json::json!("Hello"), + }], + system: None, + max_tokens: Some(1024), + stream: true, + temperature: None, + tools: None, + tool_choice: None, + }; + + let translator = AnthropicRequestTranslator::new(); + let result = translator.translate_request(request); + assert!(result.is_ok()); + + let cw_request = result.unwrap(); + assert_eq!( + cw_request + .conversation_state + .current_message + .user_input_message + .model_id, + "CLAUDE_SONNET_4_5_20250929_V1_0" + ); + } + + #[test] + fn test_extract_system_text_string() { + let system = Some(serde_json::json!("You are a helpful assistant.")); + let text = extract_system_text(&system); + assert_eq!(text, "You are a helpful assistant."); + } + + #[test] + fn test_extract_system_text_array() { + let system = Some(serde_json::json!([ + {"type": "text", "text": "Line 1"}, + {"type": "text", "text": "Line 2"} + ])); + let text = extract_system_text(&system); + assert_eq!(text, "Line 1\nLine 2"); + } +} diff --git a/src-tauri/src/translator/kiro/anthropic/response.rs b/src-tauri/src/translator/kiro/anthropic/response.rs new file mode 100644 index 000000000..0b952b067 --- /dev/null +++ b/src-tauri/src/translator/kiro/anthropic/response.rs @@ -0,0 +1,189 @@ +//! Kiro 响应转换为 Anthropic SSE 格式 +//! +//! 将 `StreamEvent` 转换为 Anthropic Messages API 流式响应格式。 +//! 这是 Claude Code 使用的协议。 + +use crate::stream::{AnthropicSseGenerator, StreamEvent}; +use crate::translator::traits::{ResponseTranslator, SseResponseTranslator}; + +/// Anthropic 响应转换器 +/// +/// 将 `StreamEvent` 转换为 Anthropic SSE 格式 +#[derive(Debug)] +pub struct AnthropicResponseTranslator { + /// SSE 生成器 + generator: AnthropicSseGenerator, +} + +impl Default for AnthropicResponseTranslator { + fn default() -> Self { + Self::new("unknown".to_string()) + } +} + +impl AnthropicResponseTranslator { + /// 创建新的转换器 + pub fn new(model: String) -> Self { + Self { + generator: AnthropicSseGenerator::new(model), + } + } + + /// 使用指定的消息 ID 创建转换器 + pub fn with_id(id: String, model: String) -> Self { + Self { + generator: AnthropicSseGenerator::with_id(id, model), + } + } + + /// 获取消息 ID + pub fn message_id(&self) -> &str { + self.generator.message_id() + } + + /// 获取模型名称 + pub fn model(&self) -> &str { + self.generator.model() + } +} + +impl ResponseTranslator for AnthropicResponseTranslator { + type Output = Vec; + + fn translate_event(&mut self, event: &StreamEvent) -> Option { + let events = self.generator.generate(event); + if events.is_empty() { + None + } else { + Some(events) + } + } + + fn finalize(&mut self) -> Vec { + Vec::new() // Anthropic 生成器在 MessageStop 时已经发送了所有结束事件 + } + + fn reset(&mut self) { + self.generator = AnthropicSseGenerator::new("unknown".to_string()); + } +} + +impl SseResponseTranslator for AnthropicResponseTranslator { + fn translate_to_sse(&mut self, event: &StreamEvent) -> Vec { + self.generator.generate(event) + } + + fn finalize_sse(&mut self) -> Vec { + Vec::new() // Anthropic 生成器在 MessageStop 时已经发送了所有结束事件 + } + + fn reset(&mut self) { + self.generator = AnthropicSseGenerator::new("unknown".to_string()); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::stream::{ContentBlockType, StopReason}; + + #[test] + fn test_translate_message_start() { + let mut translator = AnthropicResponseTranslator::new("claude-3-sonnet".to_string()); + + let event = StreamEvent::MessageStart { + id: "msg_123".to_string(), + model: "claude-3-sonnet".to_string(), + }; + + let sse = translator.translate_event(&event); + assert!(sse.is_some()); + let sse = sse.unwrap(); + assert!(!sse.is_empty()); + assert!(sse[0].starts_with("event: message_start\ndata: ")); + } + + #[test] + fn test_translate_text_content() { + let mut translator = AnthropicResponseTranslator::new("claude-3-sonnet".to_string()); + + // 先发送 message_start + let _ = translator.translate_event(&StreamEvent::MessageStart { + id: "msg_123".to_string(), + model: "claude-3-sonnet".to_string(), + }); + + // 内容块开始 + let sse = translator.translate_event(&StreamEvent::ContentBlockStart { + index: 0, + block_type: ContentBlockType::Text, + }); + assert!(sse.is_some()); + let sse = sse.unwrap(); + assert!(sse[0].contains("content_block_start")); + + // 文本增量 + let sse = translator.translate_event(&StreamEvent::TextDelta { + text: "Hello".to_string(), + }); + assert!(sse.is_some()); + let sse = sse.unwrap(); + assert!(sse[0].contains("content_block_delta")); + assert!(sse[0].contains("Hello")); + } + + #[test] + fn test_translate_tool_use() { + let mut translator = AnthropicResponseTranslator::new("claude-3-sonnet".to_string()); + + // 发送 message_start + let _ = translator.translate_event(&StreamEvent::MessageStart { + id: "msg_123".to_string(), + model: "claude-3-sonnet".to_string(), + }); + + // 工具调用内容块开始 + let sse = translator.translate_event(&StreamEvent::ContentBlockStart { + index: 1, + block_type: ContentBlockType::ToolUse { + id: "tool_abc".to_string(), + name: "read_file".to_string(), + }, + }); + assert!(sse.is_some()); + let sse = sse.unwrap(); + assert!(sse[0].contains("content_block_start")); + assert!(sse[0].contains("tool_use")); + + // 工具参数增量 + let sse = translator.translate_event(&StreamEvent::ToolUseInputDelta { + id: "tool_abc".to_string(), + partial_json: "{\"path\":".to_string(), + }); + assert!(sse.is_some()); + let sse = sse.unwrap(); + assert!(sse[0].contains("input_json_delta")); + } + + #[test] + fn test_translate_message_stop() { + let mut translator = AnthropicResponseTranslator::new("claude-3-sonnet".to_string()); + + // 发送 message_start + let _ = translator.translate_event(&StreamEvent::MessageStart { + id: "msg_123".to_string(), + model: "claude-3-sonnet".to_string(), + }); + + // 消息结束 + let sse = translator.translate_event(&StreamEvent::MessageStop { + stop_reason: StopReason::EndTurn, + }); + assert!(sse.is_some()); + let sse = sse.unwrap(); + assert_eq!(sse.len(), 2); + assert!(sse[0].contains("message_delta")); + assert!(sse[0].contains("end_turn")); + assert!(sse[1].contains("message_stop")); + } +} diff --git a/src-tauri/src/translator/kiro/mod.rs b/src-tauri/src/translator/kiro/mod.rs new file mode 100644 index 000000000..590b4c6bf --- /dev/null +++ b/src-tauri/src/translator/kiro/mod.rs @@ -0,0 +1,43 @@ +//! Kiro/CodeWhisperer 后端协议转换 +//! +//! 处理与 AWS CodeWhisperer (Kiro) 后端的协议转换。 +//! +//! # 子模块 +//! +//! - `openai`: OpenAI 前端协议支持 +//! - `anthropic`: Anthropic 前端协议支持 +//! +//! # 调用链 +//! +//! ## OpenAI 协议 +//! ```text +//! OpenAI ChatCompletionRequest +//! → [openai/request.rs] translate_request +//! → CodeWhispererRequest +//! → [backends/kiro.rs] call_stream +//! → AWS Event Stream bytes +//! → [stream/parsers/aws_event_stream.rs] parse +//! → StreamEvent +//! → [openai/response.rs] translate_event +//! → OpenAI SSE +//! ``` +//! +//! ## Anthropic 协议 +//! ```text +//! Anthropic MessagesRequest +//! → [anthropic/request.rs] translate_request +//! → CodeWhispererRequest +//! → [backends/kiro.rs] call_stream +//! → AWS Event Stream bytes +//! → [stream/parsers/aws_event_stream.rs] parse +//! → StreamEvent +//! → [anthropic/response.rs] translate_event +//! → Anthropic SSE +//! ``` + +pub mod anthropic; +pub mod openai; + +// 重新导出常用类型 +pub use anthropic::{AnthropicRequestTranslator, AnthropicResponseTranslator}; +pub use openai::{OpenAiRequestTranslator, OpenAiResponseTranslator}; diff --git a/src-tauri/src/translator/kiro/openai/mod.rs b/src-tauri/src/translator/kiro/openai/mod.rs new file mode 100644 index 000000000..6bd8bc2e4 --- /dev/null +++ b/src-tauri/src/translator/kiro/openai/mod.rs @@ -0,0 +1,9 @@ +//! OpenAI 协议 → Kiro 后端转换 +//! +//! 处理 OpenAI ChatCompletion API 格式与 CodeWhisperer 格式之间的转换。 + +pub mod request; +pub mod response; + +pub use request::OpenAiRequestTranslator; +pub use response::OpenAiResponseTranslator; diff --git a/src-tauri/src/translator/kiro/openai/request.rs b/src-tauri/src/translator/kiro/openai/request.rs new file mode 100644 index 000000000..60aa5ebd8 --- /dev/null +++ b/src-tauri/src/translator/kiro/openai/request.rs @@ -0,0 +1,620 @@ +//! OpenAI 请求转换为 CodeWhisperer 请求 +//! +//! 将 OpenAI ChatCompletionRequest 转换为 CodeWhisperer API 格式。 +//! +//! # 模型映射 +//! +//! - claude-opus-4-5 → claude-opus-4.5 +//! - claude-sonnet-4-5 → CLAUDE_SONNET_4_5_20250929_V1_0 +//! - claude-sonnet-4-20250514 → CLAUDE_SONNET_4_20250514_V1_0 +//! - claude-haiku-4-5 → claude-haiku-4.5 + +use crate::models::codewhisperer::*; +use crate::models::openai::*; +use crate::translator::traits::{RequestTranslator, TranslateError}; +use std::collections::{HashMap, HashSet}; +use uuid::Uuid; + +/// OpenAI 到 Kiro 请求转换器 +#[derive(Debug, Clone)] +pub struct OpenAiRequestTranslator { + /// 可选的 Profile ARN (AWS CodeWhisperer) + pub profile_arn: Option, +} + +impl Default for OpenAiRequestTranslator { + fn default() -> Self { + Self::new() + } +} + +impl OpenAiRequestTranslator { + /// 创建新的转换器 + pub fn new() -> Self { + Self { profile_arn: None } + } + + /// 使用 Profile ARN 创建转换器 + pub fn with_profile_arn(profile_arn: String) -> Self { + Self { + profile_arn: Some(profile_arn), + } + } +} + +impl RequestTranslator for OpenAiRequestTranslator { + type Input = ChatCompletionRequest; + type Output = CodeWhispererRequest; + type Error = TranslateError; + + fn translate_request(&self, request: Self::Input) -> Result { + Ok(convert_openai_to_codewhisperer( + &request, + self.profile_arn.clone(), + )) + } +} + +// ============================================================================ +// 模型映射 +// ============================================================================ + +/// 模型映射表 +pub fn get_model_map() -> HashMap<&'static str, &'static str> { + let mut map = HashMap::new(); + // Opus 4.5 系列 + map.insert("claude-opus-4-5", "claude-opus-4.5"); + map.insert("claude-opus-4-5-20251101", "claude-opus-4.5"); + // Haiku 4.5 系列 + map.insert("claude-haiku-4-5", "claude-haiku-4.5"); + map.insert("claude-haiku-4-5-20251001", "claude-haiku-4.5"); + // Sonnet 4.5 系列 + map.insert("claude-sonnet-4-5", "CLAUDE_SONNET_4_5_20250929_V1_0"); + map.insert( + "claude-sonnet-4-5-20250929", + "CLAUDE_SONNET_4_5_20250929_V1_0", + ); + // Sonnet 4 系列 + map.insert("claude-sonnet-4-20250514", "CLAUDE_SONNET_4_20250514_V1_0"); + // Sonnet 3.7/3.5 系列(兼容旧版本) + map.insert( + "claude-3-7-sonnet-20250219", + "CLAUDE_3_7_SONNET_20250219_V1_0", + ); + map.insert( + "claude-3-5-sonnet-20241022", + "CLAUDE_3_7_SONNET_20250219_V1_0", + ); + map.insert( + "claude-3-5-sonnet-latest", + "CLAUDE_3_7_SONNET_20250219_V1_0", + ); + map +} + +/// 获取支持的模型列表 +pub fn get_supported_models() -> Vec<&'static str> { + vec![ + "claude-opus-4-5", + "claude-opus-4-5-20251101", + "claude-haiku-4-5", + "claude-haiku-4-5-20251001", + "claude-sonnet-4-5", + "claude-sonnet-4-5-20250929", + "claude-sonnet-4-20250514", + "claude-3-7-sonnet-20250219", + ] +} + +/// 默认模型 +pub const DEFAULT_MODEL: &str = "CLAUDE_SONNET_4_5_20250929_V1_0"; + +// ============================================================================ +// 内部类型 +// ============================================================================ + +#[derive(Debug, Clone)] +struct ProcessedMessage { + role: String, + content: String, + tool_calls: Option>, + tool_results: Option>, +} + +// ============================================================================ +// 转换函数 +// ============================================================================ + +/// 将 OpenAI ChatCompletionRequest 转换为 CodeWhisperer 请求 +pub fn convert_openai_to_codewhisperer( + request: &ChatCompletionRequest, + profile_arn: Option, +) -> CodeWhispererRequest { + let model_map = get_model_map(); + let cw_model = model_map + .get(request.model.as_str()) + .map(|s| s.to_string()) + .unwrap_or_else(|| DEFAULT_MODEL.to_string()); + + let conversation_id = Uuid::new_v4().to_string(); + + // 提取 system prompt 和消息 + let mut system_prompt = String::new(); + let mut raw_messages: Vec<&ChatMessage> = Vec::new(); + + for msg in &request.messages { + if msg.role == "system" { + system_prompt = msg.get_content_text(); + } else { + raw_messages.push(msg); + } + } + + // 调试日志:打印 tool_choice 和 tools 信息 + tracing::info!( + "[KIRO_TRANSLATE] 收到请求: tool_choice={:?}, has_tools={}, tools_count={}", + request.tool_choice, + request.tools.is_some(), + request.tools.as_ref().map(|t| t.len()).unwrap_or(0) + ); + + // 处理 tool_choice: required - CodeWhisperer 不支持此参数,通过 prompt 注入强制 + if is_tool_choice_required(&request.tool_choice) && request.tools.is_some() { + let tool_instruction = "\n\n[CRITICAL INSTRUCTION] You MUST use one of the provided tools to respond. Do NOT respond with plain text. Call a tool function immediately."; + system_prompt.push_str(tool_instruction); + tracing::info!("[KIRO_TRANSLATE] tool_choice=required detected, injected tool instruction"); + } + + // 预处理消息:合并 tool 消息 + let messages = preprocess_messages(&raw_messages); + + // 构建历史记录 + let mut history: Vec = Vec::new(); + let mut start_idx = 0; + + // 处理 system prompt - 合并到第一条用户消息 + if !system_prompt.is_empty() && !messages.is_empty() && messages[0].role == "user" { + let first_content = &messages[0].content; + let combined = format!("{system_prompt}\n\n{first_content}"); + + let mut user_msg = UserInputMessage { + content: combined, + model_id: cw_model.clone(), + origin: "AI_EDITOR".to_string(), + images: None, + user_input_message_context: None, + }; + + if let Some(ref tool_results) = messages[0].tool_results { + user_msg.user_input_message_context = Some(UserInputMessageContext { + tools: None, + tool_results: Some(tool_results.clone()), + }); + } + + history.push(HistoryItem::User(UserHistoryItem { + user_input_message: user_msg, + })); + start_idx = 1; + } + + // 处理历史消息(除最后一条) + for msg in messages + .iter() + .take(messages.len().saturating_sub(1)) + .skip(start_idx) + { + match msg.role.as_str() { + "user" => { + let content = if msg.content.is_empty() { + if msg.tool_results.is_some() { + "Tool results provided.".to_string() + } else { + "Continue".to_string() + } + } else { + msg.content.clone() + }; + + let mut user_msg = UserInputMessage { + content, + model_id: cw_model.clone(), + origin: "AI_EDITOR".to_string(), + images: None, + user_input_message_context: None, + }; + + if let Some(ref tool_results) = msg.tool_results { + user_msg.user_input_message_context = Some(UserInputMessageContext { + tools: None, + tool_results: Some(tool_results.clone()), + }); + } + + history.push(HistoryItem::User(UserHistoryItem { + user_input_message: user_msg, + })); + } + "assistant" => { + let content = if msg.content.is_empty() { + "I understand.".to_string() + } else { + msg.content.clone() + }; + + history.push(HistoryItem::Assistant(AssistantHistoryItem { + assistant_response_message: AssistantResponseMessage { + content, + tool_uses: msg.tool_calls.clone(), + }, + })); + } + _ => {} + } + } + + // 修复历史记录交替顺序 + let history = fix_history_alternation(history, &cw_model); + + // 构建当前消息 + let (current_content, current_tool_results) = if let Some(last_msg) = messages.last() { + if last_msg.role == "assistant" { + ("Continue".to_string(), None) + } else { + let content = if last_msg.content.is_empty() { + if last_msg.tool_results.is_some() { + "Tool results provided.".to_string() + } else { + "Continue".to_string() + } + } else { + last_msg.content.clone() + }; + (content, last_msg.tool_results.clone()) + } + } else { + ("Continue".to_string(), None) + }; + + // 构建 tools + let tools = convert_tools(&request.tools); + + let user_input_message_context = if tools.is_some() || current_tool_results.is_some() { + Some(UserInputMessageContext { + tools, + tool_results: current_tool_results, + }) + } else { + None + }; + + CodeWhispererRequest { + conversation_state: ConversationState { + chat_trigger_type: "MANUAL".to_string(), + conversation_id, + current_message: CurrentMessage { + user_input_message: UserInputMessage { + content: current_content, + model_id: cw_model, + origin: "AI_EDITOR".to_string(), + images: None, + user_input_message_context, + }, + }, + history: if history.is_empty() { + None + } else { + Some(history) + }, + }, + profile_arn, + } +} + +/// 预处理消息:合并连续的 tool 消息到前一个 assistant 消息后的 user 消息 +fn preprocess_messages(messages: &[&ChatMessage]) -> Vec { + let mut result: Vec = Vec::new(); + let mut pending_tool_results: Vec = Vec::new(); + + for msg in messages { + match msg.role.as_str() { + "tool" => { + let content = msg.get_content_text(); + let tool_id = msg.tool_call_id.clone().unwrap_or_default(); + pending_tool_results.push(CWToolResult { + content: vec![CWTextContent { text: content }], + status: "success".to_string(), + tool_use_id: tool_id, + }); + } + "user" => { + let content = msg.get_content_text(); + let mut tool_results = pending_tool_results.clone(); + pending_tool_results.clear(); + + // 去重 tool_results + let mut seen_ids = HashSet::new(); + tool_results.retain(|tr| seen_ids.insert(tr.tool_use_id.clone())); + + result.push(ProcessedMessage { + role: "user".to_string(), + content, + tool_calls: None, + tool_results: if tool_results.is_empty() { + None + } else { + Some(tool_results) + }, + }); + } + "assistant" => { + // 如果有待处理的 tool results,先创建一个 user 消息 + if !pending_tool_results.is_empty() { + let mut tool_results = pending_tool_results.clone(); + pending_tool_results.clear(); + + let mut seen_ids = HashSet::new(); + tool_results.retain(|tr| seen_ids.insert(tr.tool_use_id.clone())); + + result.push(ProcessedMessage { + role: "user".to_string(), + content: "Tool results provided.".to_string(), + tool_calls: None, + tool_results: Some(tool_results), + }); + } + + let content = msg.get_content_text(); + let tool_calls = msg.tool_calls.as_ref().map(|calls| { + calls + .iter() + .map(|tc| CWToolUse { + input: serde_json::from_str(&tc.function.arguments) + .unwrap_or(serde_json::json!({})), + name: tc.function.name.clone(), + tool_use_id: tc.id.clone(), + }) + .collect() + }); + + result.push(ProcessedMessage { + role: "assistant".to_string(), + content, + tool_calls, + tool_results: None, + }); + } + _ => {} + } + } + + // 处理末尾的 tool results + if !pending_tool_results.is_empty() { + let mut tool_results = pending_tool_results; + let mut seen_ids = HashSet::new(); + tool_results.retain(|tr| seen_ids.insert(tr.tool_use_id.clone())); + + result.push(ProcessedMessage { + role: "user".to_string(), + content: "Tool results provided.".to_string(), + tool_calls: None, + tool_results: Some(tool_results), + }); + } + + result +} + +/// 转换工具列表 +fn convert_tools(tools: &Option>) -> Option> { + tools.as_ref().map(|tools| { + let mut cw_tools: Vec = Vec::new(); + let mut function_count = 0; + + for t in tools.iter() { + match t { + Tool::Function { function } => { + // 限制最多 50 个函数工具 + if function_count >= 50 { + continue; + } + function_count += 1; + + let params = function + .parameters + .clone() + .unwrap_or_else(|| serde_json::json!({"type": "object", "properties": {}})); + + let desc = function + .description + .clone() + .unwrap_or_else(|| format!("Tool: {}", function.name)); + + cw_tools.push(CWToolItem::Standard(CWTool { + tool_specification: ToolSpecification { + name: function.name.clone(), + description: if desc.len() > 500 { + let truncated: String = desc.chars().take(497).collect(); + format!("{}...", truncated) + } else { + desc + }, + input_schema: InputSchema { json: params }, + }, + })); + } + Tool::WebSearch | Tool::WebSearch20250305 => { + cw_tools.push(CWToolItem::WebSearch(CWWebSearchTool { + tool_type: "web_search".to_string(), + })); + } + } + } + + cw_tools + }) +} + +/// 修复历史记录,确保 user/assistant 严格交替 +fn fix_history_alternation(history: Vec, model_id: &str) -> Vec { + if history.is_empty() { + return history; + } + + let mut fixed: Vec = Vec::new(); + + for item in history { + match &item { + HistoryItem::User(user_item) => { + if let Some(HistoryItem::User(last_user)) = fixed.last_mut() { + let has_tool_results = user_item + .user_input_message + .user_input_message_context + .as_ref() + .map(|ctx| ctx.tool_results.is_some()) + .unwrap_or(false); + + if has_tool_results { + let new_results = user_item + .user_input_message + .user_input_message_context + .as_ref() + .and_then(|ctx| ctx.tool_results.clone()) + .unwrap_or_default(); + + if let Some(ref mut ctx) = + last_user.user_input_message.user_input_message_context + { + if let Some(ref mut existing) = ctx.tool_results { + existing.extend(new_results); + } else { + ctx.tool_results = Some(new_results); + } + } else { + last_user.user_input_message.user_input_message_context = + Some(UserInputMessageContext { + tools: None, + tool_results: Some(new_results), + }); + } + continue; + } else { + fixed.push(HistoryItem::Assistant(AssistantHistoryItem { + assistant_response_message: AssistantResponseMessage { + content: "I understand.".to_string(), + tool_uses: None, + }, + })); + } + } + fixed.push(item); + } + HistoryItem::Assistant(_) => { + if let Some(HistoryItem::Assistant(_)) = fixed.last() { + fixed.push(HistoryItem::User(UserHistoryItem { + user_input_message: UserInputMessage { + content: "Continue".to_string(), + model_id: model_id.to_string(), + origin: "AI_EDITOR".to_string(), + images: None, + user_input_message_context: None, + }, + })); + } + if fixed.is_empty() { + fixed.push(HistoryItem::User(UserHistoryItem { + user_input_message: UserInputMessage { + content: "Continue".to_string(), + model_id: model_id.to_string(), + origin: "AI_EDITOR".to_string(), + images: None, + user_input_message_context: None, + }, + })); + } + fixed.push(item); + } + } + } + + // 确保以 assistant 结尾 + if let Some(HistoryItem::User(_)) = fixed.last() { + fixed.push(HistoryItem::Assistant(AssistantHistoryItem { + assistant_response_message: AssistantResponseMessage { + content: "I understand.".to_string(), + tool_uses: None, + }, + })); + } + + fixed +} + +/// 检查 tool_choice 是否为 required +/// +/// tool_choice 可以是: +/// - "required" 字符串 +/// - {"type": "any"} 或类似结构 +fn is_tool_choice_required(tool_choice: &Option) -> bool { + match tool_choice { + Some(serde_json::Value::String(s)) => s == "required" || s == "any", + Some(serde_json::Value::Object(obj)) => { + // 检查 {"type": "any"} 或 {"type": "tool", ...} + if let Some(serde_json::Value::String(t)) = obj.get("type") { + t == "any" || t == "tool" + } else { + false + } + } + _ => false, + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_model_mapping() { + let map = get_model_map(); + assert_eq!(map.get("claude-opus-4-5"), Some(&"claude-opus-4.5")); + assert_eq!( + map.get("claude-sonnet-4-5"), + Some(&"CLAUDE_SONNET_4_5_20250929_V1_0") + ); + } + + #[test] + fn test_convert_simple_request() { + let request = ChatCompletionRequest { + model: "claude-sonnet-4-5".to_string(), + messages: vec![ChatMessage { + role: "user".to_string(), + content: Some(MessageContent::Text("Hello".to_string())), + tool_calls: None, + tool_call_id: None, + }], + tools: None, + stream: false, + max_tokens: None, + temperature: None, + top_p: None, + tool_choice: None, + reasoning_effort: None, + }; + + let translator = OpenAiRequestTranslator::new(); + let result = translator.translate_request(request); + assert!(result.is_ok()); + + let cw_request = result.unwrap(); + assert_eq!( + cw_request + .conversation_state + .current_message + .user_input_message + .model_id, + "CLAUDE_SONNET_4_5_20250929_V1_0" + ); + } +} diff --git a/src-tauri/src/translator/kiro/openai/response.rs b/src-tauri/src/translator/kiro/openai/response.rs new file mode 100644 index 000000000..66539d6d9 --- /dev/null +++ b/src-tauri/src/translator/kiro/openai/response.rs @@ -0,0 +1,130 @@ +//! Kiro 响应转换为 OpenAI SSE 格式 +//! +//! 将 `StreamEvent` 转换为 OpenAI Chat Completions 流式响应格式。 + +use crate::stream::{OpenAiSseGenerator, StreamEvent}; +use crate::translator::traits::{ResponseTranslator, SseResponseTranslator}; + +/// OpenAI 响应转换器 +/// +/// 将 `StreamEvent` 转换为 OpenAI SSE 格式 +#[derive(Debug)] +pub struct OpenAiResponseTranslator { + /// SSE 生成器 + generator: OpenAiSseGenerator, +} + +impl Default for OpenAiResponseTranslator { + fn default() -> Self { + Self::new("unknown".to_string()) + } +} + +impl OpenAiResponseTranslator { + /// 创建新的转换器 + pub fn new(model: String) -> Self { + Self { + generator: OpenAiSseGenerator::new(model), + } + } + + /// 使用指定的响应 ID 创建转换器 + pub fn with_id(id: String, model: String) -> Self { + Self { + generator: OpenAiSseGenerator::with_id(id, model), + } + } + + /// 获取响应 ID + pub fn response_id(&self) -> &str { + self.generator.response_id() + } +} + +impl ResponseTranslator for OpenAiResponseTranslator { + type Output = String; + + fn translate_event(&mut self, event: &StreamEvent) -> Option { + self.generator.generate(event) + } + + fn finalize(&mut self) -> Vec { + vec![self.generator.generate_done()] + } + + fn reset(&mut self) { + self.generator = OpenAiSseGenerator::new("unknown".to_string()); + } +} + +impl SseResponseTranslator for OpenAiResponseTranslator { + fn translate_to_sse(&mut self, event: &StreamEvent) -> Vec { + self.generator.generate(event).into_iter().collect() + } + + fn finalize_sse(&mut self) -> Vec { + vec![self.generator.generate_done()] + } + + fn reset(&mut self) { + self.generator = OpenAiSseGenerator::new("unknown".to_string()); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::stream::{ContentBlockType, StopReason}; + + #[test] + fn test_translate_text_delta() { + let mut translator = OpenAiResponseTranslator::new("gpt-4".to_string()); + + let event = StreamEvent::TextDelta { + text: "Hello".to_string(), + }; + + let sse = translator.translate_event(&event); + assert!(sse.is_some()); + let sse = sse.unwrap(); + assert!(sse.starts_with("data: ")); + assert!(sse.contains("\"content\":\"Hello\"")); + } + + #[test] + fn test_translate_tool_use() { + let mut translator = OpenAiResponseTranslator::new("gpt-4".to_string()); + + // 工具调用开始 + let event = StreamEvent::ToolUseStart { + id: "call_123".to_string(), + name: "read_file".to_string(), + }; + let sse = translator.translate_event(&event); + assert!(sse.is_some()); + assert!(sse.unwrap().contains("\"tool_calls\"")); + + // 工具参数增量 + let event = StreamEvent::ToolUseInputDelta { + id: "call_123".to_string(), + partial_json: "{\"path\":".to_string(), + }; + let sse = translator.translate_event(&event); + assert!(sse.is_some()); + } + + #[test] + fn test_translate_message_stop() { + let mut translator = OpenAiResponseTranslator::new("gpt-4".to_string()); + + let event = StreamEvent::MessageStop { + stop_reason: StopReason::EndTurn, + }; + + let sse = translator.translate_event(&event); + assert!(sse.is_some()); + let sse = sse.unwrap(); + assert!(sse.contains("\"finish_reason\":\"stop\"")); + assert!(sse.contains("[DONE]")); + } +} diff --git a/src-tauri/src/translator/mod.rs b/src-tauri/src/translator/mod.rs new file mode 100644 index 000000000..e79d0789a --- /dev/null +++ b/src-tauri/src/translator/mod.rs @@ -0,0 +1,53 @@ +//! 协议转换层 +//! +//! 处理不同前端协议(OpenAI、Anthropic、Gemini CLI)与不同后端(Kiro、Codex、Claude) +//! 之间的请求和响应格式转换。 +//! +//! # 架构设计 +//! +//! ```text +//! translator/ +//! ├── traits.rs # 转换器 trait 定义 +//! └── kiro/ # Kiro/CodeWhisperer 后端 +//! ├── openai/ # OpenAI 前端协议 +//! │ ├── request.rs # OpenAI → Kiro 请求 +//! │ └── response.rs # StreamEvent → OpenAI SSE +//! └── anthropic/ # Anthropic 前端协议 +//! ├── request.rs # Anthropic → Kiro 请求 +//! └── response.rs # StreamEvent → Anthropic SSE +//! ``` +//! +//! # 使用示例 +//! +//! ```ignore +//! use proxycast::translator::kiro::{ +//! AnthropicRequestTranslator, AnthropicResponseTranslator, +//! }; +//! use proxycast::translator::traits::RequestTranslator; +//! +//! // 请求转换 +//! let translator = AnthropicRequestTranslator::new(); +//! let cw_request = translator.translate_request(anthropic_request)?; +//! +//! // 响应转换 +//! let mut response_translator = AnthropicResponseTranslator::new(model); +//! for event in stream_events { +//! let sse_events = response_translator.translate_to_sse(&event); +//! for sse in sse_events { +//! // 发送 SSE 到客户端 +//! } +//! } +//! ``` + +pub mod kiro; +pub mod traits; + +// 重新导出核心类型 +pub use kiro::{ + AnthropicRequestTranslator, AnthropicResponseTranslator, OpenAiRequestTranslator, + OpenAiResponseTranslator, +}; +pub use traits::{ + RequestTranslator, ResponseTranslator, SseResponseTranslator, TranslateError, + TranslateErrorKind, +}; diff --git a/src-tauri/src/translator/traits.rs b/src-tauri/src/translator/traits.rs new file mode 100644 index 000000000..9ff4fe5c5 --- /dev/null +++ b/src-tauri/src/translator/traits.rs @@ -0,0 +1,215 @@ +//! 协议转换器 Trait 定义 +//! +//! 定义请求和响应转换器的核心接口,用于在不同协议之间进行转换。 +//! +//! # 设计原则 +//! +//! - `RequestTranslator`: 将前端协议请求转换为后端协议请求 +//! - `ResponseTranslator`: 将后端响应/事件转换为前端格式 +//! - 使用 `StreamEvent` 作为流式响应的中间表示 + +use crate::stream::StreamEvent; + +/// 请求转换器 Trait +/// +/// 将前端协议的请求转换为后端协议的请求格式。 +/// +/// # 类型参数 +/// +/// - `Input`: 前端请求类型(如 OpenAI ChatCompletionRequest) +/// - `Output`: 后端请求类型(如 CodeWhispererRequest) +/// - `Error`: 转换错误类型 +pub trait RequestTranslator { + /// 前端请求类型 + type Input; + /// 后端请求类型 + type Output; + /// 转换错误类型 + type Error: std::error::Error + Send + Sync + 'static; + + /// 转换请求 + /// + /// # 参数 + /// + /// - `request`: 前端协议请求 + /// + /// # 返回 + /// + /// 转换后的后端协议请求 + fn translate_request(&self, request: Self::Input) -> Result; +} + +/// 响应转换器 Trait +/// +/// 将 `StreamEvent` 转换为目标前端协议的响应格式。 +/// +/// # 类型参数 +/// +/// - `Output`: 目标响应类型或 SSE 字符串 +pub trait ResponseTranslator { + /// 目标响应类型 + type Output; + + /// 转换单个流事件 + /// + /// # 参数 + /// + /// - `event`: 统一的流事件 + /// + /// # 返回 + /// + /// - `Some(output)`: 生成的响应数据 + /// - `None`: 该事件不需要生成输出 + fn translate_event(&mut self, event: &StreamEvent) -> Option; + + /// 完成转换 + /// + /// 用于生成流结束时的最终事件(如果需要) + fn finalize(&mut self) -> Vec { + Vec::new() + } + + /// 重置转换器状态 + fn reset(&mut self); +} + +/// SSE 响应转换器 Trait +/// +/// 专门用于将 `StreamEvent` 转换为 SSE 字符串格式。 +/// 这是 `ResponseTranslator` 的一个常见特化。 +pub trait SseResponseTranslator { + /// 将流事件转换为 SSE 字符串 + /// + /// # 返回 + /// + /// SSE 格式的字符串列表,每个字符串都是完整的 SSE 事件 + fn translate_to_sse(&mut self, event: &StreamEvent) -> Vec; + + /// 生成结束 SSE 事件 + fn finalize_sse(&mut self) -> Vec { + Vec::new() + } + + /// 重置状态 + fn reset(&mut self); +} + +/// 转换错误类型 +#[derive(Debug, Clone)] +pub struct TranslateError { + /// 错误类型 + pub kind: TranslateErrorKind, + /// 错误消息 + pub message: String, + /// 原始数据(用于调试) + pub source_data: Option, +} + +impl std::fmt::Display for TranslateError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "{}: {}", self.kind, self.message) + } +} + +impl std::error::Error for TranslateError {} + +/// 转换错误类型枚举 +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum TranslateErrorKind { + /// 无效的请求格式 + InvalidRequest, + /// 不支持的功能 + UnsupportedFeature, + /// 缺少必要字段 + MissingField, + /// 数据验证失败 + ValidationFailed, + /// 序列化/反序列化错误 + SerializationError, + /// 其他错误 + Other, +} + +impl std::fmt::Display for TranslateErrorKind { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::InvalidRequest => write!(f, "InvalidRequest"), + Self::UnsupportedFeature => write!(f, "UnsupportedFeature"), + Self::MissingField => write!(f, "MissingField"), + Self::ValidationFailed => write!(f, "ValidationFailed"), + Self::SerializationError => write!(f, "SerializationError"), + Self::Other => write!(f, "Other"), + } + } +} + +impl TranslateError { + /// 创建新的转换错误 + pub fn new(kind: TranslateErrorKind, message: impl Into) -> Self { + Self { + kind, + message: message.into(), + source_data: None, + } + } + + /// 带原始数据创建错误 + pub fn with_source( + kind: TranslateErrorKind, + message: impl Into, + source: impl Into, + ) -> Self { + Self { + kind, + message: message.into(), + source_data: Some(source.into()), + } + } + + /// 创建无效请求错误 + pub fn invalid_request(message: impl Into) -> Self { + Self::new(TranslateErrorKind::InvalidRequest, message) + } + + /// 创建不支持功能错误 + pub fn unsupported(message: impl Into) -> Self { + Self::new(TranslateErrorKind::UnsupportedFeature, message) + } + + /// 创建缺少字段错误 + pub fn missing_field(field: &str) -> Self { + Self::new( + TranslateErrorKind::MissingField, + format!("Missing required field: {}", field), + ) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_translate_error_display() { + let err = TranslateError::new(TranslateErrorKind::InvalidRequest, "test error"); + assert_eq!(format!("{}", err), "InvalidRequest: test error"); + } + + #[test] + fn test_translate_error_with_source() { + let err = TranslateError::with_source( + TranslateErrorKind::SerializationError, + "failed to parse", + "{invalid json}", + ); + assert!(err.source_data.is_some()); + assert_eq!(err.source_data.unwrap(), "{invalid json}"); + } + + #[test] + fn test_missing_field_error() { + let err = TranslateError::missing_field("model"); + assert_eq!(err.kind, TranslateErrorKind::MissingField); + assert!(err.message.contains("model")); + } +} diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index 2e081bfa7..be1163fa1 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.26.0", + "version": "0.27.0", "identifier": "com.proxycast.app", "build": { "beforeDevCommand": "npm run dev", diff --git a/src-tauri/tests/api_key_provider_tests.rs b/src-tauri/tests/api_key_provider_tests.rs new file mode 100644 index 000000000..fb8d17785 --- /dev/null +++ b/src-tauri/tests/api_key_provider_tests.rs @@ -0,0 +1,813 @@ +//! API Key Provider 属性测试 +//! +//! 使用 proptest 进行属性测试,验证 API Key Provider 服务的正确性。 +//! +//! **Feature: provider-ui-refactor** + +use proptest::prelude::*; +use std::collections::HashSet; +use std::sync::Arc; +use tempfile::TempDir; + +use proxycast_lib::database::dao::api_key_provider::{ + ApiKeyEntry, ApiKeyProvider, ApiKeyProviderDao, ApiProviderType, ProviderGroup, +}; +use proxycast_lib::database::{init_database, DbConnection}; +use proxycast_lib::services::api_key_provider_service::ApiKeyProviderService; +use rusqlite::Connection; + +/// 测试上下文 +struct TestContext { + pub temp_dir: TempDir, + pub db: DbConnection, + pub service: ApiKeyProviderService, +} + +impl TestContext { + /// 创建测试上下文 + pub fn new() -> Result> { + let temp_dir = TempDir::new()?; + let db_path = temp_dir.path().join("test.db"); + let conn = Connection::open(&db_path)?; + + // 创建表结构 + conn.execute( + "CREATE TABLE IF NOT EXISTS api_key_providers ( + id TEXT PRIMARY KEY, + name TEXT NOT NULL, + type TEXT NOT NULL, + api_host TEXT NOT NULL, + is_system INTEGER NOT NULL DEFAULT 0, + group_name TEXT NOT NULL, + enabled INTEGER NOT NULL DEFAULT 0, + sort_order INTEGER NOT NULL DEFAULT 0, + api_version TEXT, + project TEXT, + location TEXT, + region TEXT, + created_at TEXT NOT NULL, + updated_at TEXT NOT NULL + )", + [], + )?; + + conn.execute( + "CREATE TABLE IF NOT EXISTS api_keys ( + id TEXT PRIMARY KEY, + provider_id TEXT NOT NULL, + api_key_encrypted TEXT NOT NULL, + alias TEXT, + enabled INTEGER NOT NULL DEFAULT 1, + usage_count INTEGER NOT NULL DEFAULT 0, + error_count INTEGER NOT NULL DEFAULT 0, + last_used_at TEXT, + created_at TEXT NOT NULL, + FOREIGN KEY (provider_id) REFERENCES api_key_providers(id) ON DELETE CASCADE + )", + [], + )?; + + conn.execute( + "CREATE TABLE IF NOT EXISTS provider_ui_state ( + key TEXT PRIMARY KEY, + value TEXT NOT NULL + )", + [], + )?; + + let db = Arc::new(std::sync::Mutex::new(conn)); + let service = ApiKeyProviderService::new(); + + Ok(Self { + temp_dir, + db, + service, + }) + } + + /// 创建测试 Provider + pub fn create_test_provider(&self, id: &str) -> Result { + let now = chrono::Utc::now(); + let provider = ApiKeyProvider { + id: id.to_string(), + name: format!("Test Provider {}", id), + provider_type: ApiProviderType::Openai, + api_host: "https://api.test.com".to_string(), + is_system: false, + group: ProviderGroup::Custom, + enabled: true, + sort_order: 0, + api_version: None, + project: None, + location: None, + region: None, + created_at: now, + updated_at: now, + }; + + let conn = self.db.lock().map_err(|e| e.to_string())?; + ApiKeyProviderDao::insert_provider(&conn, &provider).map_err(|e| e.to_string())?; + + Ok(provider) + } + + /// 添加测试 API Key + pub fn add_test_api_key( + &self, + provider_id: &str, + api_key: &str, + ) -> Result { + self.service + .add_api_key(&self.db, provider_id, api_key, None) + } +} + +// ============================================================================ +// Property 12: 轮询负载均衡正确性 +// **Validates: Requirements 7.3** +// ============================================================================ + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// Property 12: 轮询负载均衡正确性 + /// + /// *对于任意* 拥有 N 个启用的 API Key 的 Provider,连续 N 次获取 API Key 应各返回不同的 Key + /// + /// **Feature: provider-ui-refactor, Property 12: 轮询负载均衡正确性** + /// **Validates: Requirements 7.3** + #[test] + fn test_round_robin_load_balancing(num_keys in 2usize..10) { + let ctx = TestContext::new().expect("Failed to create test context"); + + // 创建测试 Provider + let provider_id = format!("test-provider-{}", uuid::Uuid::new_v4()); + ctx.create_test_provider(&provider_id).expect("Failed to create provider"); + + // 添加 N 个 API Keys + let mut expected_keys = Vec::new(); + for i in 0..num_keys { + let api_key = format!("sk-test-key-{}-{}", provider_id, i); + ctx.add_test_api_key(&provider_id, &api_key).expect("Failed to add API key"); + expected_keys.push(api_key); + } + + // 连续获取 N 次 API Key + let mut retrieved_keys = Vec::new(); + for _ in 0..num_keys { + let key = ctx.service + .get_next_api_key(&ctx.db, &provider_id) + .expect("Failed to get next API key") + .expect("No API key returned"); + retrieved_keys.push(key); + } + + // 验证:连续 N 次获取应返回 N 个不同的 Key + let unique_keys: HashSet<_> = retrieved_keys.iter().collect(); + prop_assert_eq!( + unique_keys.len(), + num_keys, + "Expected {} unique keys, but got {}. Keys: {:?}", + num_keys, + unique_keys.len(), + retrieved_keys + ); + + // 验证:所有返回的 Key 都在预期列表中 + for key in &retrieved_keys { + prop_assert!( + expected_keys.contains(key), + "Unexpected key returned: {}", + key + ); + } + } + + /// Property 12 补充测试:轮询循环性 + /// + /// *对于任意* 拥有 N 个启用的 API Key 的 Provider,获取 2N 次应该循环使用所有 Key + /// + /// **Feature: provider-ui-refactor, Property 12: 轮询负载均衡正确性** + /// **Validates: Requirements 7.3** + #[test] + fn test_round_robin_cycling(num_keys in 2usize..8) { + let ctx = TestContext::new().expect("Failed to create test context"); + + // 创建测试 Provider + let provider_id = format!("test-provider-cycle-{}", uuid::Uuid::new_v4()); + ctx.create_test_provider(&provider_id).expect("Failed to create provider"); + + // 添加 N 个 API Keys + for i in 0..num_keys { + let api_key = format!("sk-cycle-key-{}-{}", provider_id, i); + ctx.add_test_api_key(&provider_id, &api_key).expect("Failed to add API key"); + } + + // 获取 2N 次 API Key + let mut first_cycle = Vec::new(); + let mut second_cycle = Vec::new(); + + for i in 0..(num_keys * 2) { + let key = ctx.service + .get_next_api_key(&ctx.db, &provider_id) + .expect("Failed to get next API key") + .expect("No API key returned"); + + if i < num_keys { + first_cycle.push(key); + } else { + second_cycle.push(key); + } + } + + // 验证:第一轮和第二轮应该返回相同的 Key 序列 + prop_assert_eq!( + first_cycle, + second_cycle, + "Round robin should cycle through keys in the same order" + ); + } +} + +// ============================================================================ +// Property 13: API Key 使用统计正确性 +// **Validates: Requirements 7.4** +// ============================================================================ + +proptest! { + #![proptest_config(ProptestConfig::with_cases(20))] + + /// Property 13: API Key 使用统计正确性 + /// + /// *对于任意* API Key 使用记录操作,使用次数应正确递增 + /// + /// **Feature: provider-ui-refactor, Property 13: API Key 使用统计正确性** + /// **Validates: Requirements 7.4** + #[test] + fn test_usage_count_increment(num_usages in 1usize..10) { + let ctx = TestContext::new().expect("Failed to create test context"); + + // 创建测试 Provider + let provider_id = format!("test-provider-usage-{}", uuid::Uuid::new_v4()); + ctx.create_test_provider(&provider_id).expect("Failed to create provider"); + + // 添加 API Key + let api_key = format!("sk-usage-test-{}", provider_id); + let entry = ctx.add_test_api_key(&provider_id, &api_key) + .expect("Failed to add API key"); + + // 初始使用次数应为 0 + prop_assert_eq!(entry.usage_count, 0, "Initial usage count should be 0"); + + // 记录 N 次使用 + for _ in 0..num_usages { + ctx.service.record_usage(&ctx.db, &entry.id) + .expect("Failed to record usage"); + } + + // 获取更新后的 API Key + let conn = ctx.db.lock().expect("Failed to lock db"); + let updated = ApiKeyProviderDao::get_api_key_by_id(&conn, &entry.id) + .expect("Failed to get API key") + .expect("API key not found"); + + // 验证:使用次数应等于记录次数 + prop_assert_eq!( + updated.usage_count as usize, + num_usages, + "Usage count should equal number of record_usage calls" + ); + + // 验证:最后使用时间应被更新 + prop_assert!( + updated.last_used_at.is_some(), + "last_used_at should be set after usage" + ); + } + + /// Property 13 补充测试:错误次数递增 + /// + /// *对于任意* API Key 错误记录操作,错误次数应正确递增 + /// + /// **Feature: provider-ui-refactor, Property 13: API Key 使用统计正确性** + /// **Validates: Requirements 7.4** + #[test] + fn test_error_count_increment(num_errors in 1usize..10) { + let ctx = TestContext::new().expect("Failed to create test context"); + + // 创建测试 Provider + let provider_id = format!("test-provider-error-{}", uuid::Uuid::new_v4()); + ctx.create_test_provider(&provider_id).expect("Failed to create provider"); + + // 添加 API Key + let api_key = format!("sk-error-test-{}", provider_id); + let entry = ctx.add_test_api_key(&provider_id, &api_key) + .expect("Failed to add API key"); + + // 初始错误次数应为 0 + prop_assert_eq!(entry.error_count, 0, "Initial error count should be 0"); + + // 记录 N 次错误 + for _ in 0..num_errors { + ctx.service.record_error(&ctx.db, &entry.id) + .expect("Failed to record error"); + } + + // 获取更新后的 API Key + let conn = ctx.db.lock().expect("Failed to lock db"); + let updated = ApiKeyProviderDao::get_api_key_by_id(&conn, &entry.id) + .expect("Failed to get API key") + .expect("API key not found"); + + // 验证:错误次数应等于记录次数 + prop_assert_eq!( + updated.error_count as usize, + num_errors, + "Error count should equal number of record_error calls" + ); + } + + /// Property 13 补充测试:使用和错误统计独立 + /// + /// *对于任意* API Key,使用次数和错误次数应独立递增 + /// + /// **Feature: provider-ui-refactor, Property 13: API Key 使用统计正确性** + /// **Validates: Requirements 7.4** + #[test] + fn test_usage_and_error_independent( + num_usages in 1usize..5, + num_errors in 1usize..5 + ) { + let ctx = TestContext::new().expect("Failed to create test context"); + + // 创建测试 Provider + let provider_id = format!("test-provider-mixed-{}", uuid::Uuid::new_v4()); + ctx.create_test_provider(&provider_id).expect("Failed to create provider"); + + // 添加 API Key + let api_key = format!("sk-mixed-test-{}", provider_id); + let entry = ctx.add_test_api_key(&provider_id, &api_key) + .expect("Failed to add API key"); + + // 交替记录使用和错误 + for i in 0..(num_usages + num_errors) { + if i < num_usages { + ctx.service.record_usage(&ctx.db, &entry.id) + .expect("Failed to record usage"); + } + if i < num_errors { + ctx.service.record_error(&ctx.db, &entry.id) + .expect("Failed to record error"); + } + } + + // 获取更新后的 API Key + let conn = ctx.db.lock().expect("Failed to lock db"); + let updated = ApiKeyProviderDao::get_api_key_by_id(&conn, &entry.id) + .expect("Failed to get API key") + .expect("API key not found"); + + // 验证:使用次数和错误次数应独立 + prop_assert_eq!( + updated.usage_count as usize, + num_usages, + "Usage count should equal number of record_usage calls" + ); + prop_assert_eq!( + updated.error_count as usize, + num_errors, + "Error count should equal number of record_error calls" + ); + } +} + +// ============================================================================ +// Property 16: 数据持久化 Round-Trip +// **Validates: Requirements 9.1** +// ============================================================================ + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// Property 16: 数据持久化 Round-Trip + /// + /// *对于任意* Provider 配置,保存后重新加载应得到等价的配置数据 + /// + /// **Feature: provider-ui-refactor, Property 16: 数据持久化 Round-Trip** + /// **Validates: Requirements 9.1** + #[test] + fn test_provider_persistence_round_trip( + name in "[a-zA-Z0-9 ]{3,30}", + api_host in "https://[a-z]{3,10}\\.[a-z]{2,5}/[a-z]{0,10}" + ) { + let ctx = TestContext::new().expect("Failed to create test context"); + + // 创建 Provider + let provider = ctx.service + .add_custom_provider( + &ctx.db, + name.clone(), + ApiProviderType::Openai, + api_host.clone(), + None, + None, + None, + None, + ) + .expect("Failed to create provider"); + + // 重新加载 Provider + let loaded = ctx.service + .get_provider(&ctx.db, &provider.id) + .expect("Failed to get provider") + .expect("Provider not found"); + + // 验证:加载的数据应与保存的数据等价 + prop_assert_eq!(&loaded.provider.id, &provider.id, "ID should match"); + prop_assert_eq!(&loaded.provider.name, &name, "Name should match"); + prop_assert_eq!(&loaded.provider.api_host, &api_host, "API host should match"); + prop_assert_eq!(loaded.provider.is_system, false, "Should not be system provider"); + prop_assert_eq!(loaded.provider.group, ProviderGroup::Custom, "Group should be Custom"); + } + + /// Property 16 补充测试:UI 状态持久化 Round-Trip + /// + /// *对于任意* UI 状态键值对,保存后重新加载应得到相同的值 + /// + /// **Feature: provider-ui-refactor, Property 16: 数据持久化 Round-Trip** + /// **Validates: Requirements 9.1, 8.4** + #[test] + fn test_ui_state_persistence_round_trip( + key in "[a-z_]{3,20}", + value in "[a-zA-Z0-9_,\\[\\]\"{}:]{1,100}" + ) { + let ctx = TestContext::new().expect("Failed to create test context"); + + // 保存 UI 状态 + ctx.service + .set_ui_state(&ctx.db, &key, &value) + .expect("Failed to set UI state"); + + // 重新加载 UI 状态 + let loaded = ctx.service + .get_ui_state(&ctx.db, &key) + .expect("Failed to get UI state") + .expect("UI state not found"); + + // 验证:加载的值应与保存的值相同 + prop_assert_eq!(&loaded, &value, "UI state value should match"); + } + + /// Property 16 补充测试:Provider 排序持久化 Round-Trip + /// + /// *对于任意* Provider 排序顺序,保存后重新加载应保持相同的顺序 + /// + /// **Feature: provider-ui-refactor, Property 16: 数据持久化 Round-Trip** + /// **Validates: Requirements 9.1, 8.4** + #[test] + fn test_provider_sort_order_persistence(num_providers in 2usize..6) { + let ctx = TestContext::new().expect("Failed to create test context"); + + // 创建多个 Provider + let mut provider_ids = Vec::new(); + for i in 0..num_providers { + let provider = ctx.service + .add_custom_provider( + &ctx.db, + format!("Provider {}", i), + ApiProviderType::Openai, + format!("https://api{}.test.com", i), + None, + None, + None, + None, + ) + .expect("Failed to create provider"); + provider_ids.push(provider.id); + } + + // 反转排序顺序 + let reversed_ids: Vec<_> = provider_ids.iter().rev().cloned().collect(); + let sort_orders: Vec<(String, i32)> = reversed_ids + .iter() + .enumerate() + .map(|(i, id)| (id.clone(), i as i32)) + .collect(); + + // 更新排序顺序 + ctx.service + .update_provider_sort_orders(&ctx.db, sort_orders) + .expect("Failed to update sort orders"); + + // 重新加载所有 Provider + let loaded = ctx.service + .get_all_providers(&ctx.db) + .expect("Failed to get providers"); + + // 过滤出我们创建的 Provider + let our_providers: Vec<_> = loaded + .iter() + .filter(|p| provider_ids.contains(&p.provider.id)) + .collect(); + + // 验证:排序顺序应与更新后的顺序一致 + for (i, expected_id) in reversed_ids.iter().enumerate() { + let provider = our_providers + .iter() + .find(|p| &p.provider.id == expected_id) + .expect("Provider not found"); + prop_assert_eq!( + provider.provider.sort_order, + i as i32, + "Sort order should match for provider {}", + expected_id + ); + } + } + + /// Property 16 补充测试:API Key 持久化 Round-Trip + /// + /// *对于任意* API Key,保存后重新加载应得到等价的数据(除了加密的 key) + /// + /// **Feature: provider-ui-refactor, Property 16: 数据持久化 Round-Trip** + /// **Validates: Requirements 9.1** + #[test] + fn test_api_key_persistence_round_trip( + api_key in "[a-zA-Z0-9_-]{20,50}", + alias in proptest::option::of("[a-zA-Z0-9 ]{3,20}") + ) { + let ctx = TestContext::new().expect("Failed to create test context"); + + // 创建 Provider + let provider_id = format!("test-provider-key-rt-{}", uuid::Uuid::new_v4()); + ctx.create_test_provider(&provider_id).expect("Failed to create provider"); + + // 添加 API Key + let entry = ctx.service + .add_api_key(&ctx.db, &provider_id, &api_key, alias.clone()) + .expect("Failed to add API key"); + + // 重新加载 Provider(包含 API Keys) + let loaded = ctx.service + .get_provider(&ctx.db, &provider_id) + .expect("Failed to get provider") + .expect("Provider not found"); + + // 找到我们添加的 API Key + let loaded_key = loaded.api_keys + .iter() + .find(|k| k.id == entry.id) + .expect("API Key not found"); + + // 验证:加载的数据应与保存的数据等价 + prop_assert_eq!(&loaded_key.id, &entry.id, "ID should match"); + prop_assert_eq!(&loaded_key.provider_id, &provider_id, "Provider ID should match"); + prop_assert_eq!(&loaded_key.alias, &alias, "Alias should match"); + prop_assert_eq!(loaded_key.enabled, true, "Should be enabled by default"); + prop_assert_eq!(loaded_key.usage_count, 0, "Usage count should be 0"); + prop_assert_eq!(loaded_key.error_count, 0, "Error count should be 0"); + + // 验证:解密后的 API Key 应与原始值相同 + let decrypted = ctx.service + .decrypt_api_key(&loaded_key.api_key_encrypted) + .expect("Failed to decrypt"); + prop_assert_eq!(&decrypted, &api_key, "Decrypted API key should match original"); + } +} + +// ============================================================================ +// Property 17: API Key 加密存储 +// **Validates: Requirements 9.2** +// ============================================================================ + +proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// Property 17: API Key 加密存储 + /// + /// *对于任意* 存储的 API Key,数据库中的值不应为明文 + /// + /// **Feature: provider-ui-refactor, Property 17: API Key 加密存储** + /// **Validates: Requirements 9.2** + #[test] + fn test_api_key_encryption(api_key in "[a-zA-Z0-9_-]{20,50}") { + let ctx = TestContext::new().expect("Failed to create test context"); + + // 创建测试 Provider + let provider_id = format!("test-provider-enc-{}", uuid::Uuid::new_v4()); + ctx.create_test_provider(&provider_id).expect("Failed to create provider"); + + // 添加 API Key + let entry = ctx.add_test_api_key(&provider_id, &api_key) + .expect("Failed to add API key"); + + // 验证:存储的值不是明文 + prop_assert_ne!( + &entry.api_key_encrypted, + &api_key, + "API Key should be encrypted, not stored as plaintext" + ); + + // 验证:加密后的值看起来像 Base64 + prop_assert!( + entry.api_key_encrypted.chars().all(|c| c.is_alphanumeric() || c == '+' || c == '/' || c == '='), + "Encrypted value should be Base64 encoded" + ); + + // 验证:可以正确解密 + let decrypted = ctx.service.decrypt_api_key(&entry.api_key_encrypted) + .expect("Failed to decrypt API key"); + prop_assert_eq!( + &decrypted, + &api_key, + "Decrypted key should match original" + ); + } + + /// Property 17 补充测试:加密 Round-Trip + /// + /// *对于任意* API Key,加密后解密应得到原始值 + /// + /// **Feature: provider-ui-refactor, Property 17: API Key 加密存储** + /// **Validates: Requirements 9.2** + #[test] + fn test_encryption_round_trip(api_key in "[a-zA-Z0-9_-]{10,100}") { + let service = ApiKeyProviderService::new(); + + // 加密 + let encrypted = service.encrypt_api_key(&api_key); + + // 验证:加密后不等于原文 + prop_assert_ne!( + &encrypted, + &api_key, + "Encrypted value should differ from original" + ); + + // 解密 + let decrypted = service.decrypt_api_key(&encrypted) + .expect("Failed to decrypt"); + + // 验证:解密后等于原文 + prop_assert_eq!( + &decrypted, + &api_key, + "Decrypted value should match original" + ); + } +} + +#[cfg(test)] +mod unit_tests { + use super::*; + + /// 单元测试:基本的 Provider CRUD 操作 + #[test] + fn test_provider_crud() { + let ctx = TestContext::new().expect("Failed to create test context"); + + // 创建 Provider + let provider = ctx + .service + .add_custom_provider( + &ctx.db, + "Test Provider".to_string(), + ApiProviderType::Openai, + "https://api.test.com".to_string(), + None, + None, + None, + None, + ) + .expect("Failed to create provider"); + + assert!(provider.id.starts_with("custom-")); + assert_eq!(provider.name, "Test Provider"); + assert!(!provider.is_system); + + // 获取 Provider + let retrieved = ctx + .service + .get_provider(&ctx.db, &provider.id) + .expect("Failed to get provider") + .expect("Provider not found"); + + assert_eq!(retrieved.provider.id, provider.id); + + // 更新 Provider + let updated = ctx + .service + .update_provider( + &ctx.db, + &provider.id, + Some("Updated Name".to_string()), + None, + Some(false), + None, + None, + None, + None, + None, + ) + .expect("Failed to update provider"); + + assert_eq!(updated.name, "Updated Name"); + assert!(!updated.enabled); + + // 删除 Provider + let deleted = ctx + .service + .delete_custom_provider(&ctx.db, &provider.id) + .expect("Failed to delete provider"); + + assert!(deleted); + + // 验证已删除 + let not_found = ctx + .service + .get_provider(&ctx.db, &provider.id) + .expect("Failed to get provider"); + + assert!(not_found.is_none()); + } + + /// 单元测试:API Key CRUD 操作 + #[test] + fn test_api_key_crud() { + let ctx = TestContext::new().expect("Failed to create test context"); + + // 创建 Provider + let provider_id = "test-provider-key-crud"; + ctx.create_test_provider(provider_id) + .expect("Failed to create provider"); + + // 添加 API Key + let key = ctx + .add_test_api_key(provider_id, "sk-test-key-123") + .expect("Failed to add API key"); + + assert!(!key.id.is_empty()); + assert_eq!(key.provider_id, provider_id); + assert!(key.enabled); + + // 切换启用状态 + let toggled = ctx + .service + .toggle_api_key(&ctx.db, &key.id, false) + .expect("Failed to toggle API key"); + + assert!(!toggled.enabled); + + // 更新别名 + let aliased = ctx + .service + .update_api_key_alias(&ctx.db, &key.id, Some("My Key".to_string())) + .expect("Failed to update alias"); + + assert_eq!(aliased.alias, Some("My Key".to_string())); + + // 删除 API Key + let deleted = ctx + .service + .delete_api_key(&ctx.db, &key.id) + .expect("Failed to delete API key"); + + assert!(deleted); + } + + /// 单元测试:系统 Provider 不能删除 + #[test] + fn test_system_provider_cannot_be_deleted() { + let ctx = TestContext::new().expect("Failed to create test context"); + + // 创建系统 Provider + let now = chrono::Utc::now(); + let provider = ApiKeyProvider { + id: "system-openai".to_string(), + name: "OpenAI".to_string(), + provider_type: ApiProviderType::Openai, + api_host: "https://api.openai.com".to_string(), + is_system: true, // 系统 Provider + group: ProviderGroup::Mainstream, + enabled: true, + sort_order: 1, + api_version: None, + project: None, + location: None, + region: None, + created_at: now, + updated_at: now, + }; + + { + let conn = ctx.db.lock().expect("Failed to lock db"); + ApiKeyProviderDao::insert_provider(&conn, &provider).expect("Failed to insert"); + } + + // 尝试删除系统 Provider + let result = ctx.service.delete_custom_provider(&ctx.db, "system-openai"); + + assert!(result.is_err()); + assert!(result.unwrap_err().contains("不允许删除系统 Provider")); + } +} diff --git a/src/components/agent/chat/components/MessageList.tsx b/src/components/agent/chat/components/MessageList.tsx index 21fd92a57..9828a0ec6 100644 --- a/src/components/agent/chat/components/MessageList.tsx +++ b/src/components/agent/chat/components/MessageList.tsx @@ -1,16 +1,6 @@ import React, { useState, useRef, useEffect } from "react"; -import { - User, - Bot, - Copy, - Edit2, - Trash2, - Lightbulb, - ChevronDown, - Check, -} from "lucide-react"; +import { User, Bot, Copy, Edit2, Trash2, Check } from "lucide-react"; import { Button } from "@/components/ui/button"; -import { cn } from "@/lib/utils"; import { toast } from "sonner"; import { MessageListContainer, @@ -23,9 +13,6 @@ import { TimeStamp, MessageBubble, MessageActions, - ThinkingBox, - ThinkingHeader, - ThinkingContent, } from "../styles"; import { MarkdownRenderer } from "./MarkdownRenderer"; import { StreamingRenderer } from "./StreamingRenderer"; @@ -44,9 +31,6 @@ export const MessageList: React.FC = ({ onEditMessage, }) => { const scrollRef = useRef(null); - const [expandedThinking, setExpandedThinking] = useState< - Record - >({}); const [copiedId, setCopiedId] = useState(null); const [editingId, setEditingId] = useState(null); const [editContent, setEditContent] = useState(""); @@ -57,10 +41,6 @@ export const MessageList: React.FC = ({ } }, [messages]); - const toggleThinking = (id: string) => { - setExpandedThinking((prev) => ({ ...prev, [id]: !prev[id] })); - }; - const formatTime = (date: Date) => { return date.toLocaleTimeString([], { hour: "2-digit", minute: "2-digit" }); }; @@ -128,25 +108,6 @@ export const MessageList: React.FC = ({ - {msg.isThinking && ( - - toggleThinking(msg.id)}> - - {msg.thinkingContent} - - - {expandedThinking[msg.id] && ( - 正在深度思考... - )} - - )} - {editingId === msg.id ? (