diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index b2d5f4821..e32c9a61c 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -74,6 +74,31 @@ jobs: - name: Install frontend dependencies run: npm ci + - name: Download models data + run: | + mkdir -p src-tauri/resources/models/providers + mkdir -p src-tauri/resources/models/aliases + + # 下载 index.json + curl -sL https://raw.githubusercontent.com/aiclientproxy/models/main/index.json \ + -o src-tauri/resources/models/index.json + + # 解析 providers 列表并下载每个 provider 的数据 + for provider in $(cat src-tauri/resources/models/index.json | jq -r '.providers[]'); do + curl -sL "https://raw.githubusercontent.com/aiclientproxy/models/main/providers/${provider}.json" \ + -o "src-tauri/resources/models/providers/${provider}.json" + done + + # 下载别名配置 + for alias in kiro antigravity; do + curl -sL "https://raw.githubusercontent.com/aiclientproxy/models/main/aliases/${alias}.json" \ + -o "src-tauri/resources/models/aliases/${alias}.json" || true + done + + # 输出统计 + echo "Downloaded providers: $(ls -1 src-tauri/resources/models/providers | wc -l)" + echo "Downloaded aliases: $(ls -1 src-tauri/resources/models/aliases | wc -l)" + - name: Build Tauri app uses: tauri-apps/tauri-action@v0 env: diff --git a/.gitignore b/.gitignore index 2b01a18b4..9b340ecce 100644 --- a/.gitignore +++ b/.gitignore @@ -20,3 +20,6 @@ Thumbs.db # Kiro .kiro/ + +# Models 资源(构建时下载,不提交到 git) +src-tauri/resources/models/ diff --git a/package.json b/package.json index 760bc4161..ec01b500a 100644 --- a/package.json +++ b/package.json @@ -1,7 +1,7 @@ { "name": "proxycast", "private": true, - "version": "0.32.0", + "version": "0.33.0", "type": "module", "repository": { "type": "git", diff --git a/scripts/download-models.sh b/scripts/download-models.sh new file mode 100755 index 000000000..1c6b0f279 --- /dev/null +++ b/scripts/download-models.sh @@ -0,0 +1,44 @@ +#!/bin/bash +# 下载 models 仓库数据到 src-tauri/resources/models +# 用于本地开发环境 + +set -e + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +PROJECT_ROOT="$(dirname "$SCRIPT_DIR")" +MODELS_DIR="$PROJECT_ROOT/src-tauri/resources/models" + +BASE_URL="https://raw.githubusercontent.com/aiclientproxy/models/main" + +echo "📦 下载 models 数据..." +echo " 目标目录: $MODELS_DIR" + +# 创建目录结构 +mkdir -p "$MODELS_DIR/providers" "$MODELS_DIR/aliases" + +# 下载 index.json +echo " 下载 index.json..." +curl -sL "$BASE_URL/index.json" -o "$MODELS_DIR/index.json" + +# 解析 providers 列表并下载每个 provider 的数据 +echo " 下载 providers..." +for provider in $(cat "$MODELS_DIR/index.json" | jq -r '.providers[]'); do + echo " - $provider" + curl -sL "$BASE_URL/providers/${provider}.json" -o "$MODELS_DIR/providers/${provider}.json" +done + +# 下载别名配置 +echo " 下载 aliases..." +for alias in kiro antigravity; do + echo " - $alias" + curl -sL "$BASE_URL/aliases/${alias}.json" -o "$MODELS_DIR/aliases/${alias}.json" 2>/dev/null || echo " (跳过 $alias - 文件不存在)" +done + +# 统计 +provider_count=$(ls -1 "$MODELS_DIR/providers" 2>/dev/null | wc -l | tr -d ' ') +alias_count=$(ls -1 "$MODELS_DIR/aliases" 2>/dev/null | wc -l | tr -d ' ') + +echo "" +echo "✅ 下载完成!" +echo " Providers: $provider_count" +echo " Aliases: $alias_count" diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 96a174fa7..fba29324d 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -3720,7 +3720,7 @@ dependencies = [ [[package]] name = "proxycast" -version = "0.32.0" +version = "0.33.0" dependencies = [ "anyhow", "arboard", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index baa266f6b..f7552d693 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "proxycast" -version = "0.32.0" +version = "0.33.0" description = "AI API Proxy Desktop App" authors = ["you"] edition = "2021" diff --git a/src-tauri/build.rs b/src-tauri/build.rs index e878ef8d0..e9d7ed3be 100644 --- a/src-tauri/build.rs +++ b/src-tauri/build.rs @@ -3,8 +3,30 @@ fn main() { // 开发/CI 场景下可能只跑 `cargo check/test` 而未先构建前端,从而导致宏 panic。 // 这里提前创建配置中的 `../dist` 目录,避免无关的编译阻塞。 if let Ok(manifest_dir) = std::env::var("CARGO_MANIFEST_DIR") { - let dist_dir = std::path::PathBuf::from(manifest_dir).join("../dist"); + let manifest_path = std::path::PathBuf::from(&manifest_dir); + let dist_dir = manifest_path.join("../dist"); let _ = std::fs::create_dir_all(dist_dir); + + // 检查 models 资源是否存在 + check_models_resources(&manifest_path); } tauri_build::build() } + +/// 检查 models 资源目录是否存在 +/// 如果不存在,输出警告提示用户运行下载脚本 +fn check_models_resources(manifest_dir: &std::path::Path) { + let models_dir = manifest_dir.join("resources/models"); + let index_file = models_dir.join("index.json"); + + if !index_file.exists() { + println!("cargo:warning======================================================="); + println!("cargo:warning=Models 资源不存在!请运行以下命令下载:"); + println!("cargo:warning= ./scripts/download-models.sh"); + println!("cargo:warning======================================================="); + + // 创建空目录结构,避免 Tauri 构建失败 + let _ = std::fs::create_dir_all(models_dir.join("providers")); + let _ = std::fs::create_dir_all(models_dir.join("aliases")); + } +} diff --git a/src-tauri/src/app/bootstrap.rs b/src-tauri/src/app/bootstrap.rs index dfee72e8c..40a1e0e1b 100644 --- a/src-tauri/src/app/bootstrap.rs +++ b/src-tauri/src/app/bootstrap.rs @@ -117,6 +117,7 @@ pub struct AppStates { pub resilience_config: ResilienceConfigState, pub plugin_manager: PluginManagerState, pub plugin_installer: PluginInstallerState, + pub plugin_rpc_manager: crate::commands::plugin_rpc_cmd::PluginRpcManagerState, pub telemetry: crate::commands::telemetry_cmd::TelemetryState, pub flow_monitor: FlowMonitorState, pub flow_query_service: FlowQueryServiceState, @@ -180,6 +181,9 @@ pub fn init_states(config: &Config) -> Result { // 插件安装器 let plugin_installer_state = init_plugin_installer()?; + // 插件 RPC 管理器 + let plugin_rpc_manager_state = crate::commands::plugin_rpc_cmd::PluginRpcManagerState::new(); + // 遥测系统 let (telemetry_state, shared_stats, shared_tokens, shared_logger) = init_telemetry(config)?; @@ -235,6 +239,7 @@ pub fn init_states(config: &Config) -> Result { resilience_config: resilience_config_state, plugin_manager: plugin_manager_state, plugin_installer: plugin_installer_state, + plugin_rpc_manager: plugin_rpc_manager_state, telemetry: telemetry_state, flow_monitor: flow_monitor_state, flow_query_service: flow_query_service_state, diff --git a/src-tauri/src/app/runner.rs b/src-tauri/src/app/runner.rs index fd72883ec..b1a05bcf2 100644 --- a/src-tauri/src/app/runner.rs +++ b/src-tauri/src/app/runner.rs @@ -56,6 +56,7 @@ pub fn run() { resilience_config: resilience_config_state, plugin_manager: plugin_manager_state, plugin_installer: plugin_installer_state, + plugin_rpc_manager: plugin_rpc_manager_state, telemetry: telemetry_state, flow_monitor: flow_monitor_state, flow_query_service: flow_query_service_state, @@ -132,6 +133,7 @@ pub fn run() { .manage(telemetry_state) .manage(plugin_manager_state) .manage(plugin_installer_state) + .manage(plugin_rpc_manager_state) .manage(flow_monitor_state) .manage(flow_query_service_state) .manage(flow_interceptor_state) @@ -233,9 +235,13 @@ pub fn run() { { let app_handle = app.handle().clone(); let db_clone = db_clone.clone(); + // 获取资源目录路径 + let resource_dir = app.path().resource_dir().unwrap_or_else(|_| std::path::PathBuf::from(".")); tauri::async_runtime::spawn(async move { // 创建 ModelRegistryService - let service = crate::services::model_registry_service::ModelRegistryService::new(db_clone); + let mut service = crate::services::model_registry_service::ModelRegistryService::new(db_clone); + // 设置资源目录路径 + service.set_resource_dir(resource_dir); // 初始化服务 match service.initialize().await { @@ -748,6 +754,10 @@ pub fn run() { commands::plugin_install_cmd::is_plugin_installed, // Plugin UI commands commands::plugin_cmd::get_plugins_with_ui, + // Plugin RPC commands + commands::plugin_rpc_cmd::plugin_rpc_connect, + commands::plugin_rpc_cmd::plugin_rpc_disconnect, + commands::plugin_rpc_cmd::plugin_rpc_call, // Flow Monitor commands commands::flow_monitor_cmd::query_flows, commands::flow_monitor_cmd::get_flow_detail, @@ -1008,7 +1018,7 @@ pub fn run() { commands::connect_cmd::send_connect_callback, // Model Registry commands commands::model_registry_cmd::get_model_registry, - commands::model_registry_cmd::refresh_model_registry, + // commands::model_registry_cmd::refresh_model_registry, // TODO: 暂时禁用 commands::model_registry_cmd::search_models, commands::model_registry_cmd::get_model_preferences, commands::model_registry_cmd::toggle_model_favorite, diff --git a/src-tauri/src/commands/mod.rs b/src-tauri/src/commands/mod.rs index 0f1634236..c79946bb2 100644 --- a/src-tauri/src/commands/mod.rs +++ b/src-tauri/src/commands/mod.rs @@ -18,6 +18,7 @@ pub mod oauth_plugin_cmd; pub mod orchestrator_cmd; pub mod plugin_cmd; pub mod plugin_install_cmd; +pub mod plugin_rpc_cmd; pub mod prompt_cmd; pub mod provider_pool_cmd; pub mod resilience_cmd; diff --git a/src-tauri/src/commands/model_registry_cmd.rs b/src-tauri/src/commands/model_registry_cmd.rs index 703a71614..62322a5ba 100644 --- a/src-tauri/src/commands/model_registry_cmd.rs +++ b/src-tauri/src/commands/model_registry_cmd.rs @@ -26,17 +26,6 @@ pub async fn get_model_registry( Ok(service.get_all_models().await) } -/// 刷新模型注册表(从 models.dev 获取最新数据) -#[tauri::command] -pub async fn refresh_model_registry(state: State<'_, ModelRegistryState>) -> Result<(), String> { - let guard = state.read().await; - let service = guard - .as_ref() - .ok_or_else(|| "模型注册服务未初始化".to_string())?; - - service.refresh_from_repo().await -} - /// 搜索模型 #[tauri::command] pub async fn search_models( diff --git a/src-tauri/src/commands/plugin_rpc_cmd.rs b/src-tauri/src/commands/plugin_rpc_cmd.rs new file mode 100644 index 000000000..dc417d5cf --- /dev/null +++ b/src-tauri/src/commands/plugin_rpc_cmd.rs @@ -0,0 +1,227 @@ +//! 插件 RPC 通信命令 +//! +//! 提供插件与其 Binary 后端进程的 JSON-RPC 通信功能: +//! - plugin_rpc_connect: 启动插件进程并建立连接 +//! - plugin_rpc_disconnect: 关闭插件进程 +//! - plugin_rpc_call: 发送 RPC 请求并等待响应 +//! +//! _需求: 插件 RPC 通信_ + +use crate::commands::plugin_install_cmd::PluginInstallerState; +use serde::{Deserialize, Serialize}; +use serde_json::Value; +use std::collections::HashMap; +use std::io::{BufRead, BufReader, Write}; +use std::process::{Child, Command, Stdio}; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::sync::Arc; +use tokio::sync::{Mutex, RwLock}; + +/// RPC 请求 ID 生成器 +static REQUEST_ID: AtomicU64 = AtomicU64::new(1); + +/// JSON-RPC 请求 +#[derive(Debug, Serialize)] +struct JsonRpcRequest { + jsonrpc: &'static str, + method: String, + params: Option, + id: u64, +} + +/// JSON-RPC 响应 +#[derive(Debug, Deserialize)] +struct JsonRpcResponse { + #[allow(dead_code)] + jsonrpc: String, + result: Option, + error: Option, + #[allow(dead_code)] + id: Option, +} + +/// JSON-RPC 错误 +#[derive(Debug, Deserialize)] +struct JsonRpcError { + code: i32, + message: String, + #[allow(dead_code)] + data: Option, +} + +/// 插件进程信息 +struct PluginProcess { + child: Child, + #[allow(dead_code)] + plugin_id: String, +} + +/// 插件 RPC 管理器状态 +pub struct PluginRpcManagerState { + /// 运行中的插件进程 + processes: RwLock>>>, +} + +impl PluginRpcManagerState { + pub fn new() -> Self { + Self { + processes: RwLock::new(HashMap::new()), + } + } +} + +impl Default for PluginRpcManagerState { + fn default() -> Self { + Self::new() + } +} + +/// 启动插件进程并建立 RPC 连接 +#[tauri::command] +pub async fn plugin_rpc_connect( + plugin_id: String, + installer_state: tauri::State<'_, PluginInstallerState>, + rpc_state: tauri::State<'_, PluginRpcManagerState>, +) -> Result<(), String> { + // 检查是否已连接 + { + let processes = rpc_state.processes.read().await; + if processes.contains_key(&plugin_id) { + return Ok(()); // 已连接 + } + } + + // 获取插件信息 + let installer = installer_state.0.read().await; + let plugins = installer.list_installed().map_err(|e| e.to_string())?; + let plugin = plugins + .iter() + .find(|p| p.id == plugin_id) + .ok_or_else(|| format!("插件 {} 未安装", plugin_id))?; + + // 读取插件 manifest + let manifest_path = plugin.install_path.join("plugin.json"); + let manifest_content = std::fs::read_to_string(&manifest_path) + .map_err(|e| format!("读取 manifest 失败: {}", e))?; + let manifest: Value = serde_json::from_str(&manifest_content) + .map_err(|e| format!("解析 manifest 失败: {}", e))?; + + // 获取二进制文件路径 + let binary_name = manifest["binary"]["binary_name"] + .as_str() + .ok_or("manifest 中缺少 binary.binary_name")?; + + // 根据平台选择二进制文件 + let platform_key = match (std::env::consts::ARCH, std::env::consts::OS) { + ("aarch64", "macos") => "macos-arm64", + ("x86_64", "macos") => "macos-x64", + ("x86_64", "linux") => "linux-x64", + ("aarch64", "linux") => "linux-arm64", + ("x86_64", "windows") => "windows-x64", + _ => return Err("不支持的平台".to_string()), + }; + + let binary_filename = manifest["binary"]["platform_binaries"][platform_key] + .as_str() + .ok_or_else(|| format!("manifest 中缺少 {} 平台的二进制文件", platform_key))?; + + let binary_path = plugin.install_path.join(binary_filename); + if !binary_path.exists() { + return Err(format!("二进制文件不存在: {:?}", binary_path)); + } + + // 启动进程 + let child = Command::new(&binary_path) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn() + .map_err(|e| format!("启动插件进程失败: {}", e))?; + + tracing::info!("插件 {} 进程已启动, PID: {:?}", plugin_id, child.id()); + + let process = PluginProcess { + child, + plugin_id: plugin_id.clone(), + }; + + // 保存进程 + let mut processes = rpc_state.processes.write().await; + processes.insert(plugin_id, Arc::new(Mutex::new(process))); + + Ok(()) +} + +/// 关闭插件 RPC 连接 +#[tauri::command] +pub async fn plugin_rpc_disconnect( + plugin_id: String, + rpc_state: tauri::State<'_, PluginRpcManagerState>, +) -> Result<(), String> { + let mut processes = rpc_state.processes.write().await; + + if let Some(process_arc) = processes.remove(&plugin_id) { + let mut process = process_arc.lock().await; + if let Err(e) = process.child.kill() { + tracing::warn!("关闭插件 {} 进程失败: {}", plugin_id, e); + } + tracing::info!("插件 {} 进程已关闭", plugin_id); + } + + Ok(()) +} + +/// 发送 RPC 请求 +#[tauri::command] +pub async fn plugin_rpc_call( + plugin_id: String, + method: String, + params: Option, + rpc_state: tauri::State<'_, PluginRpcManagerState>, +) -> Result { + let processes = rpc_state.processes.read().await; + let process_arc = processes + .get(&plugin_id) + .ok_or_else(|| format!("插件 {} 未连接", plugin_id))? + .clone(); + drop(processes); + + let mut process = process_arc.lock().await; + + // 构建请求 + let request_id = REQUEST_ID.fetch_add(1, Ordering::SeqCst); + let request = JsonRpcRequest { + jsonrpc: "2.0", + method, + params, + id: request_id, + }; + + let request_json = + serde_json::to_string(&request).map_err(|e| format!("序列化请求失败: {}", e))?; + + // 发送请求 + let stdin = process.child.stdin.as_mut().ok_or("无法获取进程 stdin")?; + writeln!(stdin, "{}", request_json).map_err(|e| format!("发送请求失败: {}", e))?; + stdin + .flush() + .map_err(|e| format!("刷新 stdin 失败: {}", e))?; + + // 读取响应 + let stdout = process.child.stdout.as_mut().ok_or("无法获取进程 stdout")?; + let mut reader = BufReader::new(stdout); + let mut response_line = String::new(); + reader + .read_line(&mut response_line) + .map_err(|e| format!("读取响应失败: {}", e))?; + + // 解析响应 + let response: JsonRpcResponse = + serde_json::from_str(&response_line).map_err(|e| format!("解析响应失败: {}", e))?; + + if let Some(error) = response.error { + return Err(format!("RPC 错误 [{}]: {}", error.code, error.message)); + } + + Ok(response.result.unwrap_or(Value::Null)) +} diff --git a/src-tauri/src/models/model_registry.rs b/src-tauri/src/models/model_registry.rs index 5f1ca9e95..770658ea2 100644 --- a/src-tauri/src/models/model_registry.rs +++ b/src-tauri/src/models/model_registry.rs @@ -159,7 +159,9 @@ impl std::str::FromStr for ModelTier { #[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] #[serde(rename_all = "lowercase")] pub enum ModelSource { - /// 从 models.dev API 获取 + /// 从内嵌资源加载(构建时打包) + Embedded, + /// 从 models.dev API 获取(已弃用) ModelsDev, /// 本地硬编码(国内模型等) Local, @@ -176,6 +178,7 @@ impl Default for ModelSource { impl std::fmt::Display for ModelSource { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { + Self::Embedded => write!(f, "embedded"), Self::ModelsDev => write!(f, "models.dev"), Self::Local => write!(f, "local"), Self::Custom => write!(f, "custom"), @@ -188,6 +191,7 @@ impl std::str::FromStr for ModelSource { fn from_str(s: &str) -> Result { match s.to_lowercase().as_str() { + "embedded" => Ok(Self::Embedded), "models.dev" | "modelsdev" => Ok(Self::ModelsDev), "local" => Ok(Self::Local), "custom" => Ok(Self::Custom), diff --git a/src-tauri/src/processor/mod.rs b/src-tauri/src/processor/mod.rs index 8df726cb4..aa27520b4 100644 --- a/src-tauri/src/processor/mod.rs +++ b/src-tauri/src/processor/mod.rs @@ -190,8 +190,8 @@ impl RequestProcessor { /// * `model` - 模型名称(应该是解析后的实际模型名) /// /// # Returns - /// 选择的 Provider 类型和是否使用默认 Provider - pub async fn route_model(&self, model: &str) -> (crate::ProviderType, bool) { + /// 选择的 Provider 类型(如果设置了)和是否使用默认 Provider + pub async fn route_model(&self, model: &str) -> (Option, bool) { let router = self.router.read().await; let result = router.route(model); (result.provider, result.is_default) @@ -203,18 +203,26 @@ impl RequestProcessor { /// * `ctx` - 请求上下文 /// /// # Returns - /// 选择的 Provider 类型 - pub async fn route_for_context(&self, ctx: &mut RequestContext) -> crate::ProviderType { + /// 选择的 Provider 类型,如果未设置默认 Provider 则返回 None + pub async fn route_for_context(&self, ctx: &mut RequestContext) -> Option { let (provider, is_default) = self.route_model(&ctx.resolved_model).await; - ctx.set_provider(provider); - tracing::info!( - "[ROUTE] request_id={} model={} provider={} is_default={}", - ctx.request_id, - ctx.resolved_model, - provider, - is_default - ); + if let Some(p) = provider { + ctx.set_provider(p); + tracing::info!( + "[ROUTE] request_id={} model={} provider={} is_default={}", + ctx.request_id, + ctx.resolved_model, + p, + is_default + ); + } else { + tracing::warn!( + "[ROUTE] request_id={} model={} 未设置默认 Provider", + ctx.request_id, + ctx.resolved_model + ); + } provider } @@ -227,8 +235,8 @@ impl RequestProcessor { /// * `ctx` - 请求上下文 /// /// # Returns - /// 选择的 Provider 类型 - pub async fn resolve_and_route(&self, ctx: &mut RequestContext) -> crate::ProviderType { + /// 选择的 Provider 类型,如果未设置默认 Provider 则返回 None + pub async fn resolve_and_route(&self, ctx: &mut RequestContext) -> Option { // 1. 解析模型别名 self.resolve_model_for_context(ctx).await; diff --git a/src-tauri/src/processor/steps/routing.rs b/src-tauri/src/processor/steps/routing.rs index bb0f00e38..d290207ec 100644 --- a/src-tauri/src/processor/steps/routing.rs +++ b/src-tauri/src/processor/steps/routing.rs @@ -50,7 +50,11 @@ impl RoutingStep { // 使用路由规则(如果没有匹配的规则,会返回默认 Provider) let result = router.route(model); - Ok(result.provider) + + // 如果没有设置默认 Provider,返回错误 + result.provider.ok_or_else(|| { + StepError::Routing("未设置默认 Provider,请先在设置中选择一个默认 Provider".to_string()) + }) } } diff --git a/src-tauri/src/processor/tests.rs b/src-tauri/src/processor/tests.rs index 9256479c0..fdfb65ed5 100644 --- a/src-tauri/src/processor/tests.rs +++ b/src-tauri/src/processor/tests.rs @@ -30,7 +30,8 @@ async fn test_request_processor_components() { // 验证路由器可以正常使用 { let router = processor.router.read().await; - assert_eq!(router.default_provider(), ProviderType::Kiro); + // 默认使用 Kiro + assert_eq!(router.default_provider(), Some(ProviderType::Kiro)); } // 验证映射器可以正常使用 @@ -118,11 +119,11 @@ async fn test_route_model_returns_default() { // 所有模型都应返回默认 Provider let (provider, is_default) = processor.route_model("gemini-2.5-flash").await; - assert_eq!(provider, ProviderType::Kiro); + assert_eq!(provider, Some(ProviderType::Kiro)); assert!(is_default); let (provider, is_default) = processor.route_model("claude-sonnet-4-5").await; - assert_eq!(provider, ProviderType::Kiro); + assert_eq!(provider, Some(ProviderType::Kiro)); assert!(is_default); } @@ -138,7 +139,7 @@ async fn test_route_for_context() { // 路由并更新上下文 let provider = processor.route_for_context(&mut ctx).await; - assert_eq!(provider, ProviderType::Kiro); + assert_eq!(provider, Some(ProviderType::Kiro)); assert_eq!(ctx.provider, Some(ProviderType::Kiro)); } @@ -160,7 +161,7 @@ async fn test_resolve_and_route() { // gpt-4 -> claude-sonnet-4-5 -> Kiro (默认) assert_eq!(ctx.original_model, "gpt-4"); assert_eq!(ctx.resolved_model, "claude-sonnet-4-5"); - assert_eq!(provider, ProviderType::Kiro); + assert_eq!(provider, Some(ProviderType::Kiro)); assert_eq!(ctx.provider, Some(ProviderType::Kiro)); } diff --git a/src-tauri/src/router/rules.rs b/src-tauri/src/router/rules.rs index f58be531c..2a57cabf9 100644 --- a/src-tauri/src/router/rules.rs +++ b/src-tauri/src/router/rules.rs @@ -7,8 +7,8 @@ use crate::ProviderType; /// 路由结果 #[derive(Debug, Clone)] pub struct RouteResult { - /// 目标 Provider - pub provider: ProviderType, + /// 目标 Provider(如果未设置默认 Provider 则为 None) + pub provider: Option, /// 是否使用默认 Provider pub is_default: bool, } @@ -16,29 +16,43 @@ pub struct RouteResult { /// 路由器 - 根据默认 Provider 路由请求 #[derive(Debug, Clone)] pub struct Router { - /// 默认 Provider - default_provider: ProviderType, + /// 默认 Provider(可选,未设置时为 None) + default_provider: Option, } impl Router { /// 创建新的路由器 pub fn new(default_provider: ProviderType) -> Self { - Self { default_provider } + Self { + default_provider: Some(default_provider), + } + } + + /// 创建没有默认 Provider 的路由器 + pub fn new_empty() -> Self { + Self { + default_provider: None, + } } /// 设置默认 Provider pub fn set_default_provider(&mut self, provider: ProviderType) { - self.default_provider = provider; + self.default_provider = Some(provider); } /// 获取默认 Provider - pub fn default_provider(&self) -> ProviderType { + pub fn default_provider(&self) -> Option { self.default_provider } + /// 检查是否设置了默认 Provider + pub fn has_default_provider(&self) -> bool { + self.default_provider.is_some() + } + /// 路由请求到 Provider /// - /// 直接返回默认 Provider + /// 返回默认 Provider,如果未设置则返回 None pub fn route(&self, _model: &str) -> RouteResult { RouteResult { provider: self.default_provider, @@ -49,7 +63,7 @@ impl Router { impl Default for Router { fn default() -> Self { - Self::new(ProviderType::Kiro) + Self::new_empty() } } @@ -60,21 +74,44 @@ mod tests { #[test] fn test_new_router() { let router = Router::new(ProviderType::Kiro); - assert_eq!(router.default_provider(), ProviderType::Kiro); + assert_eq!(router.default_provider(), Some(ProviderType::Kiro)); + } + + #[test] + fn test_new_empty_router() { + let router = Router::new_empty(); + assert_eq!(router.default_provider(), None); + assert!(!router.has_default_provider()); + } + + #[test] + fn test_default_router_is_empty() { + let router = Router::default(); + assert_eq!(router.default_provider(), None); } #[test] fn test_route_returns_default() { let router = Router::new(ProviderType::Antigravity); let result = router.route("any-model"); - assert_eq!(result.provider, ProviderType::Antigravity); + assert_eq!(result.provider, Some(ProviderType::Antigravity)); + assert!(result.is_default); + } + + #[test] + fn test_route_returns_none_when_no_default() { + let router = Router::new_empty(); + let result = router.route("any-model"); + assert_eq!(result.provider, None); assert!(result.is_default); } #[test] fn test_set_default_provider() { - let mut router = Router::new(ProviderType::Kiro); + let mut router = Router::new_empty(); + assert!(!router.has_default_provider()); router.set_default_provider(ProviderType::Gemini); - assert_eq!(router.default_provider(), ProviderType::Gemini); + assert_eq!(router.default_provider(), Some(ProviderType::Gemini)); + assert!(router.has_default_provider()); } } diff --git a/src-tauri/src/server/handlers/api.rs b/src-tauri/src/server/handlers/api.rs index 5e1aef323..682729543 100644 --- a/src-tauri/src/server/handlers/api.rs +++ b/src-tauri/src/server/handlers/api.rs @@ -23,6 +23,7 @@ use axum::{ Json, }; use chrono::Utc; +use serde_json::json; use std::collections::HashMap; use crate::converter::anthropic_to_openai::convert_anthropic_to_openai; @@ -660,6 +661,29 @@ pub async fn chat_completions( provider, ctx.resolved_model ); + // 如果没有设置默认 Provider,返回错误 + let provider = match provider { + Some(p) => p, + None => { + eprintln!("[CHAT_COMPLETIONS] 未设置默认 Provider,返回错误"); + state.logs.write().await.add( + "error", + &format!("[ROUTE] request_id={} 未设置默认 Provider", ctx.request_id), + ); + return ( + StatusCode::SERVICE_UNAVAILABLE, + Json(json!({ + "error": { + "message": "未设置默认 Provider,请先在设置中选择一个默认 Provider", + "type": "configuration_error", + "code": "no_default_provider" + } + })), + ) + .into_response(); + } + }; + // 更新请求中的模型名为解析后的模型 if ctx.resolved_model != ctx.original_model { request.model = ctx.resolved_model.clone(); @@ -723,58 +747,105 @@ pub async fn chat_completions( ), ); + // 从请求头提取 X-Provider-Id(用于精确路由) + let provider_id_header = headers + .get("x-provider-id") + .and_then(|v| v.to_str().ok()) + .map(|s| s.to_lowercase()); + // 尝试从凭证池中选择凭证 - // 优先使用路由规则选择的 provider,如果找不到再回退到 selected_provider + // 如果指定了 X-Provider-Id,优先使用它(不降级) + // 否则使用路由规则选择的 provider,如果找不到再回退到 selected_provider eprintln!("[CHAT_COMPLETIONS] 开始选择凭证..."); let credential = match &state.db { Some(db) => { - // 首先尝试使用路由规则选择的 provider - let provider_str = provider.to_string(); - eprintln!( - "[CHAT_COMPLETIONS] 尝试从凭证池选择: provider={}, model={}", - provider_str, request.model - ); - let cred = state - .pool_service - .select_credential(db, &provider_str, Some(&request.model)) - .ok() - .flatten(); - - if cred.is_some() { - eprintln!("[CHAT_COMPLETIONS] 找到凭证: provider={}", provider_str); - } else { - eprintln!("[CHAT_COMPLETIONS] 未找到凭证: provider={}", provider_str); - } - - // 如果路由规则的 provider 没有找到凭证,回退到 selected_provider - if cred.is_none() && provider_str != selected_provider { + // 如果指定了 X-Provider-Id,优先使用它(不降级) + if let Some(ref explicit_provider_id) = provider_id_header { eprintln!( - "[CHAT_COMPLETIONS] 回退到 selected_provider: {}", - selected_provider + "[CHAT_COMPLETIONS] 使用 X-Provider-Id 指定的 provider: {}", + explicit_provider_id ); - state.logs.write().await.add( - "debug", - &format!( - "[ROUTE] No credential found for routed provider '{}', trying selected_provider '{}'", - provider_str, selected_provider - ), - ); - let fallback_cred = state + let cred = state .pool_service - .select_credential(db, &selected_provider, Some(&request.model)) + .select_credential(db, explicit_provider_id, Some(&request.model)) .ok() .flatten(); - if fallback_cred.is_some() { + + if cred.is_none() { eprintln!( - "[CHAT_COMPLETIONS] 回退凭证找到: provider={}", + "[CHAT_COMPLETIONS] X-Provider-Id '{}' 没有可用凭证,不进行降级", + explicit_provider_id + ); + state.logs.write().await.add( + "error", + &format!( + "[ROUTE] No available credentials for explicitly specified provider '{}', refusing to fallback", + explicit_provider_id + ), + ); + // 返回错误,不降级 + return ( + StatusCode::SERVICE_UNAVAILABLE, + Json(json!({ + "error": { + "message": format!("No available credentials for provider '{}'", explicit_provider_id), + "type": "provider_unavailable", + "code": "no_credentials" + } + })), + ) + .into_response(); + } + cred + } else { + // 原有逻辑:使用路由规则选择的 provider + let provider_str = provider.to_string(); + eprintln!( + "[CHAT_COMPLETIONS] 尝试从凭证池选择: provider={}, model={}", + provider_str, request.model + ); + let cred = state + .pool_service + .select_credential(db, &provider_str, Some(&request.model)) + .ok() + .flatten(); + + if cred.is_some() { + eprintln!("[CHAT_COMPLETIONS] 找到凭证: provider={}", provider_str); + } else { + eprintln!("[CHAT_COMPLETIONS] 未找到凭证: provider={}", provider_str); + } + + // 如果路由规则的 provider 没有找到凭证,回退到 selected_provider + if cred.is_none() && provider_str != selected_provider { + eprintln!( + "[CHAT_COMPLETIONS] 回退到 selected_provider: {}", selected_provider ); + state.logs.write().await.add( + "debug", + &format!( + "[ROUTE] No credential found for routed provider '{}', trying selected_provider '{}'", + provider_str, selected_provider + ), + ); + let fallback_cred = state + .pool_service + .select_credential(db, &selected_provider, Some(&request.model)) + .ok() + .flatten(); + if fallback_cred.is_some() { + eprintln!( + "[CHAT_COMPLETIONS] 回退凭证找到: provider={}", + selected_provider + ); + } else { + eprintln!("[CHAT_COMPLETIONS] 回退凭证也未找到!"); + } + fallback_cred } else { - eprintln!("[CHAT_COMPLETIONS] 回退凭证也未找到!"); + cred } - fallback_cred - } else { - cred } } None => { @@ -1676,6 +1747,27 @@ pub async fn anthropic_messages( // 使用 RequestProcessor 解析模型别名和路由 let provider = state.processor.resolve_and_route(&mut ctx).await; + // 如果没有设置默认 Provider,返回错误 + let provider = match provider { + Some(p) => p, + None => { + state.logs.write().await.add( + "error", + &format!("[ROUTE] request_id={} 未设置默认 Provider", ctx.request_id), + ); + return ( + StatusCode::SERVICE_UNAVAILABLE, + Json(json!({ + "error": { + "type": "configuration_error", + "message": "未设置默认 Provider,请先在设置中选择一个默认 Provider" + } + })), + ) + .into_response(); + } + }; + // 更新请求中的模型名为解析后的模型 if ctx.resolved_model != ctx.original_model { request.model = ctx.resolved_model.clone(); @@ -1757,24 +1849,69 @@ pub async fn anthropic_messages( ), ); - // 尝试从凭证池中选择凭证(带智能降级) - // 优先使用路由结果 provider,而不是 selected_provider - // 这确保了路由规则(如 claude-* → Kiro)能够正确生效 + // 从请求头提取 X-Provider-Id(用于精确路由) + let provider_id_header = headers + .get("x-provider-id") + .and_then(|v| v.to_str().ok()) + .map(|s| s.to_lowercase()); + + // 尝试从凭证池中选择凭证 + // 如果指定了 X-Provider-Id,优先使用它(不降级) + // 否则使用路由结果 provider(带智能降级) let credential_provider = provider.to_string().to_lowercase(); let credential = match &state.db { Some(db) => { - // 根据路由结果选择凭证 - state - .pool_service - .select_credential_with_fallback( - db, - &state.api_key_service, - &credential_provider, - Some(&request.model), - None, // provider_id_hint 可从路由或请求头提取 - ) - .ok() - .flatten() + // 如果指定了 X-Provider-Id,优先使用它(不降级) + if let Some(ref explicit_provider_id) = provider_id_header { + eprintln!( + "[AMP] 使用 X-Provider-Id 指定的 provider: {}", + explicit_provider_id + ); + let cred = state + .pool_service + .select_credential(db, explicit_provider_id, Some(&request.model)) + .ok() + .flatten(); + + if cred.is_none() { + eprintln!( + "[AMP] X-Provider-Id '{}' 没有可用凭证,不进行降级", + explicit_provider_id + ); + state.logs.write().await.add( + "error", + &format!( + "[ROUTE] No available credentials for explicitly specified provider '{}', refusing to fallback", + explicit_provider_id + ), + ); + // 返回错误,不降级 + return ( + StatusCode::SERVICE_UNAVAILABLE, + Json(json!({ + "error": { + "type": "provider_unavailable", + "message": format!("No available credentials for provider '{}'", explicit_provider_id) + } + })), + ) + .into_response(); + } + cred + } else { + // 原有逻辑:根据路由结果选择凭证(带智能降级) + state + .pool_service + .select_credential_with_fallback( + db, + &state.api_key_service, + &credential_provider, + Some(&request.model), + None, // provider_id_hint 可从路由或请求头提取 + ) + .ok() + .flatten() + } } None => None, }; diff --git a/src-tauri/src/server/handlers/websocket.rs b/src-tauri/src/server/handlers/websocket.rs index b588ccbe4..50cd21250 100644 --- a/src-tauri/src/server/handlers/websocket.rs +++ b/src-tauri/src/server/handlers/websocket.rs @@ -459,17 +459,11 @@ async fn handle_ws_chat_completions( // 获取默认 provider let default_provider = state.default_provider.read().await.clone(); - // 尝试从凭证池中选择凭证(带智能降级) + // 尝试从凭证池中选择凭证(不降级,指定什么就用什么) let credential = match &state.db { Some(db) => state .pool_service - .select_credential_with_fallback( - db, - &state.api_key_service, - &default_provider, - Some(&request.model), - None, // provider_id_hint - ) + .select_credential(db, &default_provider, Some(&request.model)) .ok() .flatten(), None => None, @@ -487,81 +481,14 @@ async fn handle_ws_chat_completions( Err(e) => WsProtoMessage::Error(WsError::upstream(Some(request_id.to_string()), e)), } } else { - // 回退到 Kiro provider - let kiro = state.kiro.read().await; - match kiro.call_api(&request).await { - Ok(resp) => { - if resp.status().is_success() { - match resp.text().await { - Ok(body) => { - let parsed = parse_cw_response(&body); - let has_tool_calls = !parsed.tool_calls.is_empty(); - - let message = if has_tool_calls { - serde_json::json!({ - "role": "assistant", - "content": if parsed.content.is_empty() { serde_json::Value::Null } else { serde_json::json!(parsed.content) }, - "tool_calls": parsed.tool_calls.iter().map(|tc| { - serde_json::json!({ - "id": tc.id, - "type": "function", - "function": { - "name": tc.function.name, - "arguments": tc.function.arguments - } - }) - }).collect::>() - }) - } else { - serde_json::json!({ - "role": "assistant", - "content": parsed.content - }) - }; - - let response = serde_json::json!({ - "id": format!("chatcmpl-{}", uuid::Uuid::new_v4()), - "object": "chat.completion", - "created": std::time::SystemTime::now() - .duration_since(std::time::UNIX_EPOCH) - .unwrap_or_default() - .as_secs(), - "model": request.model, - "choices": [{ - "index": 0, - "message": message, - "finish_reason": if has_tool_calls { "tool_calls" } else { "stop" } - }], - "usage": { - "prompt_tokens": 0, - "completion_tokens": 0, - "total_tokens": 0 - } - }); - - WsProtoMessage::Response(WsApiResponse { - request_id: request_id.to_string(), - payload: response, - }) - } - Err(e) => WsProtoMessage::Error(WsError::internal( - Some(request_id.to_string()), - e.to_string(), - )), - } - } else { - let body = resp.text().await.unwrap_or_default(); - WsProtoMessage::Error(WsError::upstream( - Some(request_id.to_string()), - format!("Upstream error: {}", body), - )) - } - } - Err(e) => WsProtoMessage::Error(WsError::internal( - Some(request_id.to_string()), - e.to_string(), - )), - } + // 不再回退到 Kiro provider,直接返回错误 + WsProtoMessage::Error(WsError::internal( + Some(request_id.to_string()), + format!( + "No available credentials for provider '{}'. Please add credentials in the Provider Pool.", + default_provider + ), + )) } } @@ -602,13 +529,7 @@ async fn handle_ws_anthropic_messages( let credential = match &state.db { Some(db) => state .pool_service - .select_credential_with_fallback( - db, - &state.api_key_service, - &default_provider, - Some(&request.model), - None, // provider_id_hint - ) + .select_credential(db, &default_provider, Some(&request.model)) .ok() .flatten(), None => None, @@ -624,59 +545,14 @@ async fn handle_ws_anthropic_messages( Err(e) => WsProtoMessage::Error(WsError::upstream(Some(request_id.to_string()), e)), } } else { - // 回退到 Kiro provider - let kiro = state.kiro.read().await; - - // 转换为 OpenAI 格式 - let openai_request = convert_anthropic_to_openai(&request); - - match kiro.call_api(&openai_request).await { - Ok(resp) => { - if resp.status().is_success() { - match resp.text().await { - Ok(body) => { - let parsed = parse_cw_response(&body); - - // 转换为 Anthropic 格式响应 - let response = serde_json::json!({ - "id": format!("msg_{}", uuid::Uuid::new_v4()), - "type": "message", - "role": "assistant", - "content": [{ - "type": "text", - "text": parsed.content - }], - "model": request.model, - "stop_reason": "end_turn", - "usage": { - "input_tokens": 0, - "output_tokens": 0 - } - }); - - WsProtoMessage::Response(WsApiResponse { - request_id: request_id.to_string(), - payload: response, - }) - } - Err(e) => WsProtoMessage::Error(WsError::internal( - Some(request_id.to_string()), - e.to_string(), - )), - } - } else { - let body = resp.text().await.unwrap_or_default(); - WsProtoMessage::Error(WsError::upstream( - Some(request_id.to_string()), - format!("Upstream error: {}", body), - )) - } - } - Err(e) => WsProtoMessage::Error(WsError::internal( - Some(request_id.to_string()), - e.to_string(), - )), - } + // 不再回退到 Kiro provider,直接返回错误 + WsProtoMessage::Error(WsError::internal( + Some(request_id.to_string()), + format!( + "No available credentials for provider '{}'. Please add credentials in the Provider Pool.", + default_provider + ), + )) } } diff --git a/src-tauri/src/server/mod.rs b/src-tauri/src/server/mod.rs index 5488d12f9..a786a0153 100644 --- a/src-tauri/src/server/mod.rs +++ b/src-tauri/src/server/mod.rs @@ -1029,17 +1029,11 @@ async fn gemini_generate_content( // 获取默认 provider let default_provider = state.default_provider.read().await.clone(); - // 尝试从凭证池中选择凭证(带智能降级) + // 尝试从凭证池中选择凭证(不降级,指定什么就用什么) let credential = match &state.db { Some(db) => state .pool_service - .select_credential_with_fallback( - db, - &state.api_key_service, - &default_provider, - Some(model), - None, // provider_id_hint - ) + .select_credential(db, &default_provider, Some(model)) .ok() .flatten(), None => None, @@ -1049,10 +1043,10 @@ async fn gemini_generate_content( Some(c) => c, None => { return ( - StatusCode::NOT_FOUND, + StatusCode::SERVICE_UNAVAILABLE, Json(serde_json::json!({ "error": { - "message": "没有可用的凭证。您可以在 API Key Provider 中配置 API Key 作为降级选项。" + "message": format!("No available credentials for provider '{}'. Please add credentials in the Provider Pool.", default_provider) } })), ) @@ -1311,7 +1305,7 @@ async fn anthropic_messages_with_selector( ), ); - // 尝试解析凭证(带智能降级) + // 尝试解析凭证(不降级,指定什么就用什么) let credential = match &state.db { Some(db) => { // 首先尝试按名称查找 @@ -1322,14 +1316,12 @@ async fn anthropic_messages_with_selector( else if let Ok(Some(cred)) = state.pool_service.get_by_uuid(db, &selector) { Some(cred) } - // 最后尝试按 provider 类型轮询(带智能降级) - else if let Ok(Some(cred)) = state.pool_service.select_credential_with_fallback( - db, - &state.api_key_service, - &selector, - Some(&request.model), - None, // provider_id_hint - ) { + // 最后尝试按 provider 类型选择(不降级) + else if let Ok(Some(cred)) = + state + .pool_service + .select_credential(db, &selector, Some(&request.model)) + { Some(cred) } else { None @@ -1355,16 +1347,24 @@ async fn anthropic_messages_with_selector( handlers::call_provider_anthropic(&state, &cred, &request, None).await } None => { - // 回退到默认 Kiro provider + // 不再回退到默认 provider,直接返回错误 state.logs.write().await.add( - "warn", + "error", &format!( - "[ROUTE] Credential not found for selector '{}', falling back to default", + "[ROUTE] No available credentials for selector '{}', refusing to fallback", selector ), ); - // 调用原有的 Kiro 处理逻辑 - anthropic_messages_internal(&state, &request).await + ( + StatusCode::SERVICE_UNAVAILABLE, + Json(serde_json::json!({ + "error": { + "type": "provider_unavailable", + "message": format!("No available credentials for selector '{}'", selector) + } + })), + ) + .into_response() } } } @@ -1392,20 +1392,18 @@ async fn chat_completions_with_selector( ), ); - // 尝试解析凭证(带智能降级) + // 尝试解析凭证(不降级,指定什么就用什么) let credential = match &state.db { Some(db) => { if let Ok(Some(cred)) = state.pool_service.get_by_name(db, &selector) { Some(cred) } else if let Ok(Some(cred)) = state.pool_service.get_by_uuid(db, &selector) { Some(cred) - } else if let Ok(Some(cred)) = state.pool_service.select_credential_with_fallback( - db, - &state.api_key_service, - &selector, - Some(&request.model), - None, // provider_id_hint - ) { + } else if let Ok(Some(cred)) = + state + .pool_service + .select_credential(db, &selector, Some(&request.model)) + { Some(cred) } else { None @@ -1430,14 +1428,25 @@ async fn chat_completions_with_selector( handlers::call_provider_openai(&state, &cred, &request, None).await } None => { + // 不再回退到默认 provider,直接返回错误 state.logs.write().await.add( - "warn", + "error", &format!( - "[ROUTE] Credential not found for selector '{}', falling back to default", + "[ROUTE] No available credentials for selector '{}', refusing to fallback", selector ), ); - chat_completions_internal(&state, &request).await + ( + StatusCode::SERVICE_UNAVAILABLE, + Json(serde_json::json!({ + "error": { + "message": format!("No available credentials for selector '{}'", selector), + "type": "provider_unavailable", + "code": "no_credentials" + } + })), + ) + .into_response() } } } @@ -1487,7 +1496,7 @@ async fn amp_chat_completions( ), ); - // 尝试根据 provider 名称选择凭证(带智能降级) + // 尝试根据 provider 名称选择凭证(不降级,指定什么就用什么) eprintln!( "[AMP] 开始查找凭证: provider={}, model={}, db={}", provider, @@ -1497,21 +1506,16 @@ async fn amp_chat_completions( let credential = match &state.db { Some(db) => { eprintln!( - "[AMP] 调用 select_credential_with_fallback, provider_id_hint={}", + "[AMP] 使用 select_credential 查找凭证(不降级): provider={}", provider ); - // 首先尝试按 provider 类型选择(带智能降级) - if let Ok(Some(cred)) = state.pool_service.select_credential_with_fallback( - db, - &state.api_key_service, - &provider, - Some(&request.model), - Some(&provider), // provider_id_hint 使用路由中的 provider 名称 - ) { - eprintln!( - "[AMP] select_credential_with_fallback 找到凭证: {:?}", - cred.name - ); + // 使用 select_credential 而不是 select_credential_with_fallback,禁止降级 + if let Ok(Some(cred)) = + state + .pool_service + .select_credential(db, &provider, Some(&request.model)) + { + eprintln!("[AMP] select_credential 找到凭证: {:?}", cred.name); Some(cred) } // 然后尝试按名称查找 @@ -1524,7 +1528,10 @@ async fn amp_chat_completions( eprintln!("[AMP] get_by_uuid 找到凭证: {:?}", cred.name); Some(cred) } else { - eprintln!("[AMP] 未找到任何凭证 for provider '{}'", provider); + eprintln!( + "[AMP] 未找到任何凭证 for provider '{}',不进行降级", + provider + ); None } } @@ -1549,14 +1556,25 @@ async fn amp_chat_completions( handlers::call_provider_openai(&state, &cred, &request, None).await } None => { + // 不再回退到默认 provider,直接返回错误 state.logs.write().await.add( - "warn", + "error", &format!( - "[AMP] Credential not found for provider '{}', falling back to default", + "[AMP] No available credentials for provider '{}', refusing to fallback", provider ), ); - chat_completions_internal(&state, &request).await + ( + StatusCode::SERVICE_UNAVAILABLE, + Json(serde_json::json!({ + "error": { + "message": format!("No available credentials for provider '{}'", provider), + "type": "provider_unavailable", + "code": "no_credentials" + } + })), + ) + .into_response() } } } @@ -1605,17 +1623,15 @@ async fn amp_messages( ), ); - // 尝试根据 provider 名称选择凭证(带智能降级) + // 尝试根据 provider 名称选择凭证(不降级,指定什么就用什么) let credential = match &state.db { Some(db) => { - // 首先尝试按 provider 类型选择(带智能降级) - if let Ok(Some(cred)) = state.pool_service.select_credential_with_fallback( - db, - &state.api_key_service, - &provider, - Some(&request.model), - Some(&provider), // provider_id_hint 使用路由中的 provider 名称 - ) { + // 使用 select_credential 而不是 select_credential_with_fallback,禁止降级 + if let Ok(Some(cred)) = + state + .pool_service + .select_credential(db, &provider, Some(&request.model)) + { Some(cred) } // 然后尝试按名称查找 @@ -1647,14 +1663,24 @@ async fn amp_messages( handlers::call_provider_anthropic(&state, &cred, &request, None).await } None => { + // 不再回退到默认 provider,直接返回错误 state.logs.write().await.add( - "warn", + "error", &format!( - "[AMP] Credential not found for provider '{}', falling back to default", + "[AMP] No available credentials for provider '{}', refusing to fallback", provider ), ); - anthropic_messages_internal(&state, &request).await + ( + StatusCode::SERVICE_UNAVAILABLE, + Json(serde_json::json!({ + "error": { + "type": "provider_unavailable", + "message": format!("No available credentials for provider '{}'", provider) + } + })), + ) + .into_response() } } } diff --git a/src-tauri/src/services/model_registry_service.rs b/src-tauri/src/services/model_registry_service.rs index 6b25443de..34dc8e0fa 100644 --- a/src-tauri/src/services/model_registry_service.rs +++ b/src-tauri/src/services/model_registry_service.rs @@ -1,6 +1,7 @@ //! 模型注册服务 //! -//! 负责从 aiclientproxy/models 仓库获取模型数据、管理本地缓存、提供模型搜索等功能 +//! 从内嵌资源加载模型数据,管理本地缓存,提供模型搜索等功能 +//! 模型数据在构建时从 aiclientproxy/models 仓库打包进应用 use crate::database::DbConnection; use crate::models::model_registry::{ @@ -13,9 +14,8 @@ use std::collections::HashMap; use std::sync::Arc; use tokio::sync::RwLock; -/// GitHub 仓库 raw 文件基础 URL -const MODELS_REPO_BASE_URL: &str = "https://raw.githubusercontent.com/aiclientproxy/models/main"; -const CACHE_DURATION_SECS: i64 = 3600; // 1 小时 +/// 内嵌的模型资源目录名 +const MODELS_RESOURCE_DIR: &str = "models"; /// 仓库索引文件结构 #[derive(Debug, Deserialize)] @@ -96,6 +96,8 @@ pub struct ModelRegistryService { aliases_cache: Arc>>, /// 同步状态 sync_state: Arc>, + /// 资源目录路径 + resource_dir: Option, } impl ModelRegistryService { @@ -106,254 +108,141 @@ impl ModelRegistryService { models_cache: Arc::new(RwLock::new(Vec::new())), aliases_cache: Arc::new(RwLock::new(HashMap::new())), sync_state: Arc::new(RwLock::new(ModelSyncState::default())), + resource_dir: None, } } - /// 初始化服务 + /// 设置资源目录路径 + pub fn set_resource_dir(&mut self, path: std::path::PathBuf) { + self.resource_dir = Some(path); + } + + /// 初始化服务 - 从内嵌资源加载模型数据 pub async fn initialize(&self) -> Result<(), String> { tracing::info!("[ModelRegistry] 初始化模型注册服务"); - // 1. 尝试从数据库加载缓存 + // 1. 首先尝试从内嵌资源加载 + match self.load_from_embedded_resources().await { + Ok((models, aliases)) => { + tracing::info!( + "[ModelRegistry] 从内嵌资源加载了 {} 个模型, {} 个别名配置", + models.len(), + aliases.len() + ); + + // 更新缓存 + { + let mut cache = self.models_cache.write().await; + *cache = models.clone(); + } + { + let mut cache = self.aliases_cache.write().await; + *cache = aliases; + } + + // 更新同步状态 + { + let mut state = self.sync_state.write().await; + state.model_count = models.len() as u32; + state.last_sync_at = Some(chrono::Utc::now().timestamp()); + state.is_syncing = false; + state.last_error = None; + } + + // 保存到数据库 + if let Err(e) = self.save_models_to_db(&models).await { + tracing::warn!("[ModelRegistry] 保存模型到数据库失败: {}", e); + } + + return Ok(()); + } + Err(e) => { + tracing::warn!("[ModelRegistry] 从内嵌资源加载失败: {}", e); + } + } + + // 2. 回退到从数据库加载 match self.load_from_db().await { Ok(models) if !models.is_empty() => { tracing::info!("[ModelRegistry] 从数据库加载了 {} 个模型", models.len()); let mut cache = self.models_cache.write().await; *cache = models; - - // 检查是否需要后台刷新 - if self.should_refresh().await { - tracing::info!("[ModelRegistry] 缓存已过期,启动后台刷新"); - self.spawn_background_refresh(); - } - return Ok(()); } Ok(_) => { - tracing::info!("[ModelRegistry] 数据库中没有缓存数据"); + tracing::warn!("[ModelRegistry] 数据库中没有模型数据"); } Err(e) => { - tracing::warn!("[ModelRegistry] 从数据库加载失败: {}", e); + tracing::error!("[ModelRegistry] 从数据库加载失败: {}", e); } } - // 2. 后台获取 models 仓库数据 - self.spawn_background_refresh(); - Ok(()) } - /// 检查是否需要刷新 - async fn should_refresh(&self) -> bool { - let state = self.sync_state.read().await; - match state.last_sync_at { - Some(last_sync) => { - let now = chrono::Utc::now().timestamp(); - now - last_sync > CACHE_DURATION_SECS - } - None => true, - } - } - - /// 启动后台刷新任务 - fn spawn_background_refresh(&self) { - let db = self.db.clone(); - let models_cache = self.models_cache.clone(); - let aliases_cache = self.aliases_cache.clone(); - let sync_state = self.sync_state.clone(); - - tokio::spawn(async move { - let service = ModelRegistryService { - db, - models_cache, - aliases_cache, - sync_state, - }; - if let Err(e) = service.refresh_from_repo().await { - tracing::error!("[ModelRegistry] 后台刷新失败: {}", e); - } - }); - } - - /// 从 aiclientproxy/models 仓库刷新数据 - pub async fn refresh_from_repo(&self) -> Result<(), String> { - tracing::info!("[ModelRegistry] 开始从 models 仓库获取数据"); - - // 设置同步状态 - { - let mut state = self.sync_state.write().await; - state.is_syncing = true; - state.last_error = None; - } - - // 获取模型数据 - let models_result = self.fetch_models_from_repo().await; - - // 获取别名数据 - let aliases_result = self.fetch_aliases_from_repo().await; - - match models_result { - Ok(models) => { - tracing::info!("[ModelRegistry] 获取了 {} 个模型", models.len()); - - // 更新模型缓存 - { - let mut cache = self.models_cache.write().await; - *cache = models.clone(); - } - - // 更新别名缓存 - if let Ok(aliases) = aliases_result { - tracing::info!( - "[ModelRegistry] 获取了 {} 个 Provider 别名配置", - aliases.len() - ); - let mut cache = self.aliases_cache.write().await; - *cache = aliases; - } - - // 保存到数据库 - self.save_models_to_db(&models).await?; - - // 更新同步状态 - { - let mut state = self.sync_state.write().await; - state.is_syncing = false; - state.last_sync_at = Some(chrono::Utc::now().timestamp()); - state.model_count = models.len() as u32; - state.last_error = None; - } - - // 保存同步状态到数据库 - self.save_sync_state().await?; - - Ok(()) - } - Err(e) => { - tracing::error!("[ModelRegistry] 从 models 仓库获取数据失败: {}", e); - - // 更新同步状态 - { - let mut state = self.sync_state.write().await; - state.is_syncing = false; - state.last_error = Some(e.clone()); - } - - Err(e) - } - } - } - - /// 从 models 仓库获取别名配置 - async fn fetch_aliases_from_repo( + /// 从内嵌资源加载模型数据 + async fn load_from_embedded_resources( &self, - ) -> Result, String> { - let client = reqwest::Client::builder() - .timeout(std::time::Duration::from_secs(30)) - .build() - .map_err(|e| format!("创建 HTTP 客户端失败: {}", e))?; + ) -> Result< + ( + Vec, + HashMap, + ), + String, + > { + let resource_dir = self + .resource_dir + .as_ref() + .ok_or_else(|| "资源目录未设置".to_string())?; - let mut aliases = HashMap::new(); + let models_dir = resource_dir.join(MODELS_RESOURCE_DIR); + let index_file = models_dir.join("index.json"); - // 已知的别名文件列表 - let alias_files = ["kiro", "antigravity"]; - - for alias_name in alias_files { - let alias_url = format!("{}/aliases/{}.json", MODELS_REPO_BASE_URL, alias_name); - - match client - .get(&alias_url) - .header("User-Agent", "ProxyCast/1.0") - .send() - .await - { - Ok(response) => { - if response.status().is_success() { - match response.json::().await { - Ok(config) => { - tracing::info!( - "[ModelRegistry] 加载别名配置: {} ({} 个模型)", - config.provider, - config.models.len() - ); - aliases.insert(config.provider.clone(), config); - } - Err(e) => { - tracing::warn!( - "[ModelRegistry] 解析别名配置 {} 失败: {}", - alias_name, - e - ); - } - } - } - } - Err(e) => { - tracing::warn!("[ModelRegistry] 获取别名配置 {} 失败: {}", alias_name, e); - } - } + if !index_file.exists() { + return Err(format!("索引文件不存在: {:?}", index_file)); } - Ok(aliases) - } - - /// 从 models 仓库获取数据 - async fn fetch_models_from_repo(&self) -> Result, String> { - let client = reqwest::Client::builder() - .timeout(std::time::Duration::from_secs(30)) - .build() - .map_err(|e| format!("创建 HTTP 客户端失败: {}", e))?; - - // 1. 获取索引文件 - let index_url = format!("{}/index.json", MODELS_REPO_BASE_URL); - let index: RepoIndex = client - .get(&index_url) - .header("User-Agent", "ProxyCast/1.0") - .send() - .await - .map_err(|e| format!("请求 index.json 失败: {}", e))? - .json() - .await - .map_err(|e| format!("解析 index.json 失败: {}", e))?; + // 1. 读取索引文件 + let index_content = + std::fs::read_to_string(&index_file).map_err(|e| format!("读取索引文件失败: {}", e))?; + let index: RepoIndex = + serde_json::from_str(&index_content).map_err(|e| format!("解析索引文件失败: {}", e))?; tracing::info!( "[ModelRegistry] 索引包含 {} 个 providers", index.providers.len() ); - // 2. 并发获取所有 provider 数据 + // 2. 加载所有 provider 数据 let mut models = Vec::new(); let now = chrono::Utc::now().timestamp(); + let providers_dir = models_dir.join("providers"); for provider_id in &index.providers { - let provider_url = format!("{}/providers/{}.json", MODELS_REPO_BASE_URL, provider_id); + let provider_file = providers_dir.join(format!("{}.json", provider_id)); + if !provider_file.exists() { + tracing::warn!("[ModelRegistry] Provider 文件不存在: {:?}", provider_file); + continue; + } - match client - .get(&provider_url) - .header("User-Agent", "ProxyCast/1.0") - .send() - .await - { - Ok(response) => { - if response.status().is_success() { - match response.json::().await { - Ok(provider_data) => { - for model in provider_data.models { - let enhanced = self.convert_repo_model( - model, - &provider_data.provider.id, - &provider_data.provider.name, - now, - ); - models.push(enhanced); - } - } - Err(e) => { - tracing::warn!("[ModelRegistry] 解析 {} 失败: {}", provider_id, e); - } + match std::fs::read_to_string(&provider_file) { + Ok(content) => match serde_json::from_str::(&content) { + Ok(provider_data) => { + for model in provider_data.models { + let enhanced = self.convert_repo_model( + model, + &provider_data.provider.id, + &provider_data.provider.name, + now, + ); + models.push(enhanced); } } - } + Err(e) => { + tracing::warn!("[ModelRegistry] 解析 {} 失败: {}", provider_id, e); + } + }, Err(e) => { - tracing::warn!("[ModelRegistry] 获取 {} 失败: {}", provider_id, e); + tracing::warn!("[ModelRegistry] 读取 {} 失败: {}", provider_id, e); } } } @@ -365,12 +254,40 @@ impl ModelRegistryService { .then(a.display_name.cmp(&b.display_name)) }); - tracing::info!( - "[ModelRegistry] 从 models 仓库获取了 {} 个模型", - models.len() - ); + // 3. 加载别名配置 + let mut aliases = HashMap::new(); + let aliases_dir = models_dir.join("aliases"); + let alias_files = ["kiro", "antigravity"]; - Ok(models) + for alias_name in alias_files { + let alias_file = aliases_dir.join(format!("{}.json", alias_name)); + if !alias_file.exists() { + continue; + } + + match std::fs::read_to_string(&alias_file) { + Ok(content) => match serde_json::from_str::(&content) { + Ok(config) => { + tracing::info!( + "[ModelRegistry] 加载别名配置: {} ({} 个模型)", + config.provider, + config.models.len() + ); + aliases.insert(config.provider.clone(), config); + } + Err(e) => { + tracing::warn!("[ModelRegistry] 解析别名配置 {} 失败: {}", alias_name, e); + } + }, + Err(e) => { + tracing::warn!("[ModelRegistry] 读取别名配置 {} 失败: {}", alias_name, e); + } + } + } + + tracing::info!("[ModelRegistry] 从内嵌资源加载了 {} 个模型", models.len()); + + Ok((models, aliases)) } /// 转换仓库模型格式为内部格式 @@ -421,7 +338,7 @@ impl ModelRegistryService { release_date: model.release_date, is_latest: model.is_latest.unwrap_or(false), description: model.description_zh.or(model.description), - source: ModelSource::ModelsDev, + source: ModelSource::Embedded, created_at: now, updated_at: now, } @@ -569,43 +486,6 @@ impl ModelRegistryService { Ok(()) } - /// 保存同步状态 - async fn save_sync_state(&self) -> Result<(), String> { - let (last_sync_at, model_count, last_error) = { - let state = self.sync_state.read().await; - ( - state.last_sync_at, - state.model_count, - state.last_error.clone(), - ) - }; - - let conn = self.db.lock().map_err(|e| e.to_string())?; - let now = chrono::Utc::now().timestamp(); - - let mut stmt = conn - .prepare( - "INSERT OR REPLACE INTO model_sync_state (key, value, updated_at) - VALUES (?, ?, ?)", - ) - .map_err(|e| e.to_string())?; - - if let Some(last_sync) = last_sync_at { - stmt.execute(params!["last_sync_at", last_sync.to_string(), now]) - .map_err(|e| e.to_string())?; - } - - stmt.execute(params!["model_count", model_count.to_string(), now]) - .map_err(|e| e.to_string())?; - - if let Some(ref error) = last_error { - stmt.execute(params!["last_error", error, now]) - .map_err(|e| e.to_string())?; - } - - Ok(()) - } - /// 获取所有模型 pub async fn get_all_models(&self) -> Vec { self.models_cache.read().await.clone() diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index efd998500..8c380a6aa 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.32.0", + "version": "0.33.0", "identifier": "com.proxycast.app", "build": { "beforeDevCommand": "npm run dev", @@ -39,7 +39,8 @@ ], "resources": [ "icons/tray/*", - "../scripts/playwright-login/**/*" + "../scripts/playwright-login/**/*", + "resources/models/**/*" ], "macOS": { "entitlements": null, diff --git a/src/App.tsx b/src/App.tsx index 56092c0d2..7e5202642 100644 --- a/src/App.tsx +++ b/src/App.tsx @@ -65,6 +65,17 @@ const PageWrapper = styled.div` overflow: auto; `; +/** + * 全屏页面容器(无 padding) + * 用于终端等需要全屏显示的插件 + */ +const FullscreenWrapper = styled.div` + flex: 1; + overflow: hidden; + display: flex; + flex-direction: column; +`; + function App() { const [showSplash, setShowSplash] = useState(true); const [currentPage, setCurrentPage] = useState("agent"); @@ -130,6 +141,19 @@ function App() { // 检查是否为动态插件页面 (plugin:xxx 格式) if (currentPage.startsWith("plugin:")) { const pluginId = currentPage.slice(7); // 移除 "plugin:" 前缀 + + // 需要全屏显示的插件列表 + const fullscreenPlugins = ["terminal-plugin"]; + const isFullscreen = fullscreenPlugins.includes(pluginId); + + if (isFullscreen) { + return ( + + + + ); + } + return ( diff --git a/src/components/agent/chat/components/ChatNavbar.tsx b/src/components/agent/chat/components/ChatNavbar.tsx index 7116b0e26..58bf77ad7 100644 --- a/src/components/agent/chat/components/ChatNavbar.tsx +++ b/src/components/agent/chat/components/ChatNavbar.tsx @@ -13,8 +13,11 @@ import { useProviderPool } from "@/hooks/useProviderPool"; import { useApiKeyProvider } from "@/hooks/useApiKeyProvider"; import { useModelRegistry } from "@/hooks/useModelRegistry"; import { getDefaultProvider } from "@/hooks/useTauri"; +import { getProviderAliasConfig } from "@/lib/api/modelRegistry"; +import type { ProviderAliasConfig } from "@/lib/types/modelRegistry"; // Provider type 到 registry ID 的映射(用于获取模型列表) +// 注意:antigravity 和 kiro 使用别名配置,需要单独处理 const getRegistryIdFromType = (providerType: string): string => { const typeMap: Record = { openai: "openai", @@ -23,23 +26,26 @@ const getRegistryIdFromType = (providerType: string): string => { "azure-openai": "openai", vertexai: "google", ollama: "ollama", - kiro: "anthropic", + kiro: "kiro", // 使用别名配置 claude: "anthropic", claude_oauth: "anthropic", qwen: "alibaba", codex: "openai", - antigravity: "google", + antigravity: "antigravity", // 使用别名配置 iflow: "openai", gemini_api_key: "google", }; return typeMap[providerType.toLowerCase()] || providerType.toLowerCase(); }; +// 需要使用别名配置的 Provider 列表 +const ALIAS_PROVIDERS = ["antigravity", "kiro"]; + // 生成 Provider 的显示标签 const getProviderLabel = (providerType: string): string => { const labelMap: Record = { kiro: "Kiro", - gemini: "Gemini", + gemini: "Gemini OAuth", qwen: "通义千问", antigravity: "Antigravity", codex: "Codex", @@ -50,7 +56,7 @@ const getProviderLabel = (providerType: string): string => { "azure-openai": "Azure OpenAI", vertexai: "VertexAI", ollama: "Ollama", - gemini_api_key: "Gemini", + gemini_api_key: "Gemini API Key", iflow: "iFlow", }; // 如果在映射表中,使用映射;否则首字母大写 @@ -98,6 +104,11 @@ export const ChatNavbar: React.FC = ({ // 用于防止无限循环 const hasInitialized = useRef(false); + // 别名配置缓存(用于 Antigravity/Kiro 等中转服务) + const [aliasConfig, setAliasConfig] = useState( + null, + ); + // 获取凭证池数据 const { overview: oauthCredentials } = useProviderPool(); const { providers: apiKeyProviders } = useApiKeyProvider(); @@ -139,16 +150,26 @@ export const ChatNavbar: React.FC = ({ // 从 API Key Provider 提取(动态,支持所有自定义 Provider) // 使用 provider.id 作为 key,确保每个 Provider 单独显示 + // 特殊处理:如果与 OAuth 凭证冲突,使用带后缀的 key apiKeyProviders .filter((p) => p.api_key_count > 0 && p.enabled) .forEach((provider) => { - const key = provider.id; // 使用 provider.id 而不是 type 映射 + let key = provider.id; + let label = provider.name; + + // 如果 key 与 OAuth 凭证冲突,添加 "_api_key" 后缀 + // 例如:Gemini OAuth 的 key 是 "gemini",Gemini API Key 的 key 变成 "gemini_api_key" + if (providerMap.has(key)) { + key = `${provider.id}_api_key`; + label = `${provider.name} API Key`; + } + if (!providerMap.has(key)) { // 优先使用 provider.id 作为 registryId(适用于系统预设的 Provider,如 deepseek, moonshot) // 如果模型注册表中没有该 id 的模型,则回退到使用 type 映射(适用于自定义 Provider) providerMap.set(key, { key, - label: provider.name, // 使用 Provider 的 name 作为显示名称 + label, registryId: provider.id, // 先尝试用 id fallbackRegistryId: getRegistryIdFromType(provider.type), // 回退用 type type: provider.type, @@ -164,11 +185,32 @@ export const ChatNavbar: React.FC = ({ return configuredProviders.find((p) => p.key === providerType); }, [configuredProviders, providerType]); + // 当选中别名 Provider 时,加载别名配置 + useEffect(() => { + if (selectedProvider && ALIAS_PROVIDERS.includes(selectedProvider.key)) { + getProviderAliasConfig(selectedProvider.key) + .then((config) => { + setAliasConfig(config); + }) + .catch((error) => { + console.error("加载别名配置失败:", error); + setAliasConfig(null); + }); + } else { + setAliasConfig(null); + } + }, [selectedProvider]); + // 获取当前 Provider 的模型列表(从 model_registry 获取) // 按照模型版本排序,最新的在前面 const currentModels = useMemo(() => { if (!selectedProvider) return []; + // 对于别名 Provider(Antigravity、Kiro),使用别名配置中的模型列表 + if (ALIAS_PROVIDERS.includes(selectedProvider.key) && aliasConfig) { + return aliasConfig.models; + } + // 从 model_registry 获取模型 // 优先使用 registryId,如果没有模型则回退到 fallbackRegistryId let models = registryModels @@ -211,7 +253,7 @@ export const ChatNavbar: React.FC = ({ // 其他情况按字母降序(通常版本号大的在前) return b.localeCompare(a); }); - }, [selectedProvider, registryModels]); + }, [selectedProvider, registryModels, aliasConfig]); // 初始化:优先选择服务器默认 Provider,否则选择第一个已配置的 useEffect(() => { @@ -246,15 +288,29 @@ export const ChatNavbar: React.FC = ({ ]); // 当 Provider 切换或模型列表变化时,自动选择第一个模型 + // 注意:使用 ref 跟踪 model 避免将其放入依赖中导致无限循环 + const modelRef = useRef(model); + modelRef.current = model; + useEffect(() => { + // 对于别名 Provider,等待别名配置加载完成 + if ( + selectedProvider && + ALIAS_PROVIDERS.includes(selectedProvider.key) && + !aliasConfig + ) { + return; + } + // 如果模型列表不为空,且当前模型为空或不在列表中,选择第一个模型 + const currentModel = modelRef.current; if ( currentModels.length > 0 && - (!model || !currentModels.includes(model)) + (!currentModel || !currentModels.includes(currentModel)) ) { setModel(currentModels[0]); } - }, [currentModels, model, setModel]); + }, [currentModels, setModel, selectedProvider, aliasConfig]); const selectedProviderLabel = selectedProvider?.label || providerType; diff --git a/src/components/agent/chat/hooks/useAgentChat.ts b/src/components/agent/chat/hooks/useAgentChat.ts index a35e2ff4e..d8aa79708 100644 --- a/src/components/agent/chat/hooks/useAgentChat.ts +++ b/src/components/agent/chat/hooks/useAgentChat.ts @@ -139,18 +139,23 @@ export function useAgentChat() { // 当 provider 改变时,检查当前模型是否兼容 // 如果不兼容,自动切换到新 provider 的第一个模型 + // 注意:model 不能放在依赖中,否则会导致无限循环 useEffect(() => { const currentProviderModels = providerConfig[providerType]?.models || []; - if ( - currentProviderModels.length > 0 && - !currentProviderModels.includes(model) - ) { - console.log( - `[useAgentChat] 模型 ${model} 不在 ${providerType} 支持列表中,自动切换到 ${currentProviderModels[0]}`, - ); - setModel(currentProviderModels[0]); + // 只有当模型列表非空时才检查兼容性 + if (currentProviderModels.length > 0) { + // 使用 setModel 的函数形式来访问当前 model 值,避免将 model 放入依赖 + setModel((currentModel) => { + if (!currentProviderModels.includes(currentModel)) { + console.log( + `[useAgentChat] 模型 ${currentModel} 不在 ${providerType} 支持列表中,自动切换到 ${currentProviderModels[0]}`, + ); + return currentProviderModels[0]; + } + return currentModel; + }); } - }, [providerType, providerConfig, model]); + }, [providerType, providerConfig]); useEffect(() => { saveTransient("agent_curr_sessionId", sessionId); diff --git a/src/components/model-selector/ProviderModelSelector.tsx b/src/components/model-selector/ProviderModelSelector.tsx index 42346a899..394fc78cb 100644 --- a/src/components/model-selector/ProviderModelSelector.tsx +++ b/src/components/model-selector/ProviderModelSelector.tsx @@ -55,7 +55,7 @@ interface ConfiguredProvider { /** OAuth 凭证类型到 Provider ID 的映射 */ const CREDENTIAL_TYPE_TO_PROVIDER_ID: Record = { kiro: "kiro", - gemini: "google", + gemini: "gemini_oauth", qwen: "alibaba", antigravity: "antigravity", codex: "openai", @@ -63,7 +63,7 @@ const CREDENTIAL_TYPE_TO_PROVIDER_ID: Record = { iflow: "iflow", openai: "openai", claude: "anthropic", - gemini_api_key: "google", + gemini_api_key: "gemini_api_key", }; /** API Key Provider 类型到 Registry ID 的映射 */ @@ -85,6 +85,8 @@ const PROVIDER_DISPLAY_NAMES: Record = { anthropic: "Anthropic", openai: "OpenAI", google: "Google", + gemini_oauth: "Gemini OAuth", + gemini_api_key: "Gemini API Key", alibaba: "阿里云", ollama: "Ollama", custom: "自定义", @@ -96,6 +98,12 @@ const PROVIDER_DISPLAY_NAMES: Record = { /** 别名 Provider 列表(使用别名配置而非标准模型注册表) */ const ALIAS_PROVIDERS = ["antigravity", "kiro"]; +/** Provider ID 到模型注册表 Provider ID 的映射(用于过滤模型) */ +const PROVIDER_TO_REGISTRY_MAPPING: Record = { + gemini_oauth: "google", + gemini_api_key: "google", +}; + // ============================================================================ // 子组件 // ============================================================================ @@ -378,7 +386,10 @@ export const ProviderModelSelector: React.FC = ({ } // 对于标准 Provider,从模型注册表过滤 - return models.filter((m) => m.provider_id === selectedProviderId); + // 使用映射表将 UI Provider ID 转换为模型注册表 Provider ID + const registryProviderId = + PROVIDER_TO_REGISTRY_MAPPING[selectedProviderId] || selectedProviderId; + return models.filter((m) => m.provider_id === registryProviderId); }, [models, selectedProviderId, aliasConfig]); // 选择 Provider diff --git a/src/components/plugins/PluginUIRenderer.tsx b/src/components/plugins/PluginUIRenderer.tsx index fdfb123f6..30379132b 100644 --- a/src/components/plugins/PluginUIRenderer.tsx +++ b/src/components/plugins/PluginUIRenderer.tsx @@ -2,17 +2,20 @@ * 插件 UI 渲染器组件 * * 根据 pluginId 渲染对应的插件 UI 组件 - * 支持内置插件组件映射和错误处理 + * 支持内置插件组件映射、动态加载外部插件和错误处理 * * _需求: 3.2_ */ -import React from "react"; -import { AlertCircle, Package } from "lucide-react"; +import React, { useState, useEffect } from "react"; +import { AlertCircle, Package, Loader2 } from "lucide-react"; +import { invoke } from "@tauri-apps/api/core"; import { MachineIdTool } from "@/components/tools/machine-id/MachineIdTool"; import { BrowserInterceptorTool } from "@/components/tools/browser-interceptor/BrowserInterceptorTool"; import { FlowMonitorPage } from "@/pages"; import { ConfigManagementPage } from "@/components/config/ConfigManagementPage"; +import { PluginUIRenderer as DynamicPluginRenderer } from "@/lib/plugin-loader/PluginUIRenderer"; +import { usePluginSDK } from "@/lib/plugin-sdk"; /** * 页面类型定义 @@ -39,8 +42,6 @@ interface PluginUIRendererProps { /** * 插件 UI 加载错误组件 - * - * 当插件 UI 组件加载失败时显示友好的错误提示 */ function PluginUIError({ pluginId, @@ -69,8 +70,6 @@ function PluginUIError({ /** * 插件未找到组件 - * - * 当请求的插件不存在时显示提示 */ function PluginNotFound({ pluginId }: { pluginId: string }) { return ( @@ -93,11 +92,20 @@ function PluginNotFound({ pluginId }: { pluginId: string }) { ); } +/** + * 加载中组件 + */ +function PluginLoading() { + return ( +
+ +

加载插件中...

+
+ ); +} + /** * 内置插件组件映射 - * - * 将插件 ID 映射到对应的 React 组件 - * 支持 machine-id-tool 和 browser-interception 插件 */ const builtinPluginComponents: Record< string, @@ -109,34 +117,146 @@ const builtinPluginComponents: Record< "config-switch": ConfigManagementPage, }; +/** + * 已安装插件信息 + */ +interface InstalledPlugin { + id: string; + name: string; + install_path: string; + has_ui: boolean; + ui_entry?: string; +} + +/** + * 动态插件渲染器 + * 用于加载外部安装的插件 UI + */ +function DynamicPluginUIRenderer({ + pluginId, + pluginsDir, + uiEntry, +}: { + pluginId: string; + pluginsDir: string; + uiEntry?: string; +}) { + const { sdk } = usePluginSDK(pluginId); + + return ( + } + /> + ); +} + /** * 插件 UI 渲染器 * * 根据 pluginId 渲染对应的插件 UI 组件 * - 对于内置插件,直接渲染对应的 React 组件 + * - 对于外部安装的插件,动态加载其 UI * - 对于未知插件,显示错误提示 - * - * @param pluginId - 插件 ID - * @param onNavigate - 页面导航回调 */ export function PluginUIRenderer({ pluginId, onNavigate, }: PluginUIRendererProps) { - // 查找内置插件组件 - const Component = builtinPluginComponents[pluginId]; + const [loading, setLoading] = useState(true); + const [pluginInfo, setPluginInfo] = useState(null); + const [pluginsDir, setPluginsDir] = useState(""); + const [error, setError] = useState(null); - if (Component) { + // 查找内置插件组件 + const BuiltinComponent = builtinPluginComponents[pluginId]; + + // 对于非内置插件,检查是否已安装并有 UI + useEffect(() => { + // 如果是内置插件,跳过检查 + if (BuiltinComponent) { + setLoading(false); + return; + } + + async function checkPlugin() { + setLoading(true); + setError(null); + + try { + // 获取插件目录 + const dir = await invoke("get_plugins_dir"); + setPluginsDir(dir); + + // 检查插件是否已安装 + const installed = await invoke("is_plugin_installed", { + pluginId, + }); + + if (!installed) { + setPluginInfo(null); + setLoading(false); + return; + } + + // 获取插件信息 + const plugins = await invoke( + "list_installed_plugins", + ); + const plugin = plugins.find((p) => p.id === pluginId); + + if (plugin) { + setPluginInfo(plugin); + } else { + setPluginInfo(null); + } + } catch (err) { + console.error("检查插件失败:", err); + setError(err instanceof Error ? err.message : String(err)); + } finally { + setLoading(false); + } + } + + checkPlugin(); + }, [pluginId, BuiltinComponent]); + + // 加载中 + if (loading) { + return ; + } + + // 如果是内置插件,直接渲染 + if (BuiltinComponent) { try { - return ; - } catch (error) { - const errorMessage = error instanceof Error ? error.message : "未知错误"; + return ; + } catch (err) { + const errorMessage = err instanceof Error ? err.message : "未知错误"; return ; } } - // 插件未找到 - return ; + // 错误 + if (error) { + return ; + } + + // 插件未安装 + if (!pluginInfo) { + return ; + } + + // 动态加载插件 UI + return ( + + ); } export default PluginUIRenderer; diff --git a/src/lib/plugin-loader/index.ts b/src/lib/plugin-loader/index.ts index 8df6a7cff..e72ad2be6 100644 --- a/src/lib/plugin-loader/index.ts +++ b/src/lib/plugin-loader/index.ts @@ -42,6 +42,7 @@ const PLUGIN_GLOBAL_NAMES: Record = { "gemini-provider": "GeminiProviderPlugin", "antigravity-provider": "AntigravityProviderPlugin", "codex-provider": "CodexProviderPlugin", + "terminal-plugin": "TerminalPlugin", }; /** @@ -51,25 +52,37 @@ const PLUGIN_GLOBAL_NAMES: Record = { function getPluginGlobalName(pluginPath: string): string { // 从路径中提取插件 ID const parts = pluginPath.split("/"); - const pluginIdIndex = parts.findIndex((p) => p.endsWith("-provider")); - const pluginId = pluginIdIndex >= 0 ? parts[pluginIdIndex] : null; + + // 查找插件 ID(在 plugins 目录后的那个目录名) + const pluginsIndex = parts.findIndex((p) => p === "plugins"); + const pluginId = + pluginsIndex >= 0 && pluginsIndex + 1 < parts.length + ? parts[pluginsIndex + 1] + : null; + + console.log(`[PluginLoader] 从路径提取插件 ID: ${pluginId}`); // 查找预定义的全局变量名 if (pluginId && PLUGIN_GLOBAL_NAMES[pluginId]) { + console.log( + `[PluginLoader] 使用预定义全局变量名: ${PLUGIN_GLOBAL_NAMES[pluginId]}`, + ); return PLUGIN_GLOBAL_NAMES[pluginId]; } // 尝试从插件 ID 推断全局变量名 - // 例如: my-plugin -> MyPluginPlugin + // 例如: my-plugin -> MyPlugin if (pluginId) { const camelCase = pluginId .split("-") .map((part) => part.charAt(0).toUpperCase() + part.slice(1)) .join(""); + console.log(`[PluginLoader] 推断全局变量名: ${camelCase}`); return camelCase; } // 默认回退 + console.log(`[PluginLoader] 使用默认全局变量名: KiroProviderPlugin`); return "KiroProviderPlugin"; } diff --git a/src/lib/plugin-sdk/index.ts b/src/lib/plugin-sdk/index.ts index 23226ed23..cfcb00302 100644 --- a/src/lib/plugin-sdk/index.ts +++ b/src/lib/plugin-sdk/index.ts @@ -35,6 +35,7 @@ export type { StorageApi, CredentialApi, PluginConfigApi, + RpcApi, // 数据类型 QueryResult, @@ -44,6 +45,7 @@ export type { EventCallback, Unsubscribe, CredentialInfo, + RpcNotificationCallback, // 主 SDK 类型 ProxyCastPluginSDK, @@ -61,6 +63,7 @@ export { clearSDKCache, subscribeNotifications, getGlobalEventBus, + handleRpcNotification, } from "./sdk"; // Hook 导出 diff --git a/src/lib/plugin-sdk/sdk.ts b/src/lib/plugin-sdk/sdk.ts index 3e12ee59c..2150b6407 100644 --- a/src/lib/plugin-sdk/sdk.ts +++ b/src/lib/plugin-sdk/sdk.ts @@ -16,6 +16,7 @@ import type { StorageApi, CredentialApi, PluginConfigApi, + RpcApi, QueryResult, ExecuteResult, HttpRequestOptions, @@ -24,6 +25,7 @@ import type { Unsubscribe, CredentialInfo, CredentialId, + RpcNotificationCallback, } from "./types"; // ============================================================================ @@ -462,6 +464,164 @@ function createPluginConfigApi(pluginId: PluginId): PluginConfigApi { }; } +// ============================================================================ +// RPC API 实现(用于 Binary 插件通信) +// ============================================================================ + +/** + * RPC 通知处理器管理 + */ +class RpcNotificationManager { + private handlers = new Map>(); + + on( + event: string, + callback: RpcNotificationCallback, + ): Unsubscribe { + if (!this.handlers.has(event)) { + this.handlers.set(event, new Set()); + } + const eventHandlers = this.handlers.get(event)!; + eventHandlers.add(callback as RpcNotificationCallback); + + return () => { + eventHandlers.delete(callback as RpcNotificationCallback); + if (eventHandlers.size === 0) { + this.handlers.delete(event); + } + }; + } + + off(event: string, callback: RpcNotificationCallback): void { + const eventHandlers = this.handlers.get(event); + if (eventHandlers) { + eventHandlers.delete(callback as RpcNotificationCallback); + if (eventHandlers.size === 0) { + this.handlers.delete(event); + } + } + } + + emit(event: string, params: unknown): void { + const eventHandlers = this.handlers.get(event); + if (eventHandlers) { + eventHandlers.forEach((handler) => { + try { + handler(params); + } catch (error) { + console.error( + `[RPC] Error in notification handler for '${event}':`, + error, + ); + } + }); + } + } + + clear(): void { + this.handlers.clear(); + } +} + +// 每个插件的 RPC 通知管理器 +const rpcNotificationManagers = new Map(); + +function getRpcNotificationManager(pluginId: PluginId): RpcNotificationManager { + let manager = rpcNotificationManagers.get(pluginId); + if (!manager) { + manager = new RpcNotificationManager(); + rpcNotificationManagers.set(pluginId, manager); + } + return manager; +} + +// 连接状态跟踪 +const rpcConnectionStatus = new Map(); + +/** + * 创建 RPC API + * + * 用于与 Binary 类型插件进行 JSON-RPC 通信 + */ +function createRpcApi(pluginId: PluginId): RpcApi { + const notificationManager = getRpcNotificationManager(pluginId); + + return { + async call(method: string, params?: unknown): Promise { + try { + const result = await invoke("plugin_rpc_call", { + pluginId, + method, + params: params ?? null, + }); + return result; + } catch (error) { + console.error( + `[Plugin ${pluginId}] RPC call error (${method}):`, + error, + ); + throw error; + } + }, + + on( + event: string, + callback: RpcNotificationCallback, + ): Unsubscribe { + return notificationManager.on(event, callback); + }, + + off( + event: string, + callback: RpcNotificationCallback, + ): void { + notificationManager.off(event, callback); + }, + + isConnected(): boolean { + return rpcConnectionStatus.get(pluginId) ?? false; + }, + + async connect(): Promise { + try { + await invoke("plugin_rpc_connect", { pluginId }); + rpcConnectionStatus.set(pluginId, true); + console.log(`[Plugin ${pluginId}] RPC connected`); + } catch (error) { + console.error(`[Plugin ${pluginId}] RPC connect error:`, error); + throw error; + } + }, + + async disconnect(): Promise { + try { + await invoke("plugin_rpc_disconnect", { pluginId }); + rpcConnectionStatus.set(pluginId, false); + notificationManager.clear(); + console.log(`[Plugin ${pluginId}] RPC disconnected`); + } catch (error) { + console.error(`[Plugin ${pluginId}] RPC disconnect error:`, error); + throw error; + } + }, + }; +} + +/** + * 处理来自后端的 RPC 通知 + * 由 Tauri 事件系统调用 + */ +export function handleRpcNotification( + pluginId: PluginId, + event: string, + params: unknown, +): void { + const manager = rpcNotificationManagers.get(pluginId); + if (manager) { + manager.emit(event, params); + } +} + // ============================================================================ // SDK 工厂函数 // ============================================================================ @@ -498,6 +658,7 @@ export function createPluginSDK(pluginId: PluginId): ProxyCastPluginSDK { storage: createStorageApi(pluginId), credential: createCredentialApi(pluginId), config: createPluginConfigApi(pluginId), + rpc: createRpcApi(pluginId), }; } diff --git a/src/lib/plugin-sdk/types.ts b/src/lib/plugin-sdk/types.ts index a3e458e60..6d490be07 100644 --- a/src/lib/plugin-sdk/types.ts +++ b/src/lib/plugin-sdk/types.ts @@ -177,6 +177,58 @@ export interface EventsApi { once(event: string, callback: EventCallback): void; } +// ============================================================================ +// RPC 操作(用于 Binary 插件通信) +// ============================================================================ + +/** RPC 通知回调 */ +export type RpcNotificationCallback = (params: T) => void; + +/** RPC 操作接口 */ +export interface RpcApi { + /** + * 发送 RPC 请求并等待响应 + * @param method RPC 方法名 + * @param params 请求参数 + * @returns 响应结果 + */ + call(method: string, params?: unknown): Promise; + + /** + * 订阅 RPC 通知 + * @param event 通知事件名 + * @param callback 回调函数 + * @returns 取消订阅函数 + */ + on( + event: string, + callback: RpcNotificationCallback, + ): Unsubscribe; + + /** + * 取消订阅 RPC 通知 + * @param event 通知事件名 + * @param callback 回调函数 + */ + off(event: string, callback: RpcNotificationCallback): void; + + /** + * 检查 RPC 连接状态 + * @returns 是否已连接 + */ + isConnected(): boolean; + + /** + * 初始化 RPC 连接(启动插件进程) + */ + connect(): Promise; + + /** + * 关闭 RPC 连接(停止插件进程) + */ + disconnect(): Promise; +} + // ============================================================================ // 存储操作 // ============================================================================ @@ -350,6 +402,9 @@ export interface ProxyCastPluginSDK { /** 插件配置操作 */ readonly config: PluginConfigApi; + + /** RPC 操作(用于 Binary 插件通信) */ + readonly rpc: RpcApi; } // ============================================================================