From a92b626ce19a8706abf0a9cdea9becdfc13c701d Mon Sep 17 00:00:00 2001 From: coso Date: Tue, 6 Jan 2026 10:10:05 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=B7=BB=E5=8A=A0=20GeminiFormStandalo?= =?UTF-8?q?ne=20=E7=BB=84=E4=BB=B6=E5=B9=B6=E6=9B=B4=E6=96=B0=20gemini-pro?= =?UTF-8?q?vider=20=E7=89=88=E6=9C=AC=E5=88=B0=20v0.4.0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增 GeminiFormStandalone 组件支持 OAuth 和 API Key 认证 - 在 plugin-components 中导出 GeminiFormStandalone - 添加 KeyRound 图标导出 - 更新 OAuthPluginTab 中 gemini-provider 版本配置 - 重构 ProviderModelList 将工具函数移到单独文件 --- src-tauri/src/commands/model_registry_cmd.rs | 4 +- .../src/converter/openai_to_antigravity.rs | 7 +- src-tauri/src/credential/registry.rs | 6 +- src-tauri/src/data/local_models.rs | 615 ++++++++++++------ src-tauri/src/models/model_registry.rs | 5 +- src-tauri/src/orchestrator/pool_builder.rs | 4 +- .../src/services/model_registry_service.rs | 34 +- .../agent/chat/components/ChatNavbar.tsx | 171 ++++- .../agent/chat/hooks/useAgentChat.ts | 7 +- src/components/api-server/ApiServerPage.tsx | 7 +- .../api-server/EnhancedModelsTab.tsx | 66 +- .../model-selector/EnhancedModelList.tsx | 37 +- .../model-selector/ProviderModelSelector.tsx | 413 ++++++++++++ src/components/model-selector/index.ts | 2 + .../provider-pool/ModelRegistryTab.tsx | 22 + .../provider-pool/OAuthPluginTab.tsx | 8 +- .../provider-pool/ProviderPoolPage.tsx | 19 +- .../api-key/AddCustomProviderModal.tsx | 275 +++++++- .../api-key/ProviderModelList.tsx | 209 ++++++ .../provider-pool/api-key/ProviderSetting.tsx | 9 + src/components/provider-pool/api-key/index.ts | 5 + .../api-key/providerTypeMapping.ts | 28 + .../credential-forms/GeminiFormStandalone.tsx | 474 ++++++++++++++ src/components/provider-pool/index.ts | 1 + src/hooks/useModelRegistry.ts | 12 +- src/lib/api/modelRegistry.ts | 6 +- src/lib/plugin-components/index.ts | 4 + src/lib/plugin-loader/PluginUIRenderer.tsx | 27 +- 28 files changed, 2158 insertions(+), 319 deletions(-) create mode 100644 src/components/model-selector/ProviderModelSelector.tsx create mode 100644 src/components/provider-pool/ModelRegistryTab.tsx create mode 100644 src/components/provider-pool/api-key/ProviderModelList.tsx create mode 100644 src/components/provider-pool/api-key/providerTypeMapping.ts create mode 100644 src/components/provider-pool/credential-forms/GeminiFormStandalone.tsx diff --git a/src-tauri/src/commands/model_registry_cmd.rs b/src-tauri/src/commands/model_registry_cmd.rs index 0324be79d..24e1e4b24 100644 --- a/src-tauri/src/commands/model_registry_cmd.rs +++ b/src-tauri/src/commands/model_registry_cmd.rs @@ -28,9 +28,7 @@ pub async fn get_model_registry( /// 刷新模型注册表(从 models.dev 获取最新数据) #[tauri::command] -pub async fn refresh_model_registry( - state: State<'_, ModelRegistryState>, -) -> Result<(), String> { +pub async fn refresh_model_registry(state: State<'_, ModelRegistryState>) -> Result<(), String> { let guard = state.read().await; let service = guard .as_ref() diff --git a/src-tauri/src/converter/openai_to_antigravity.rs b/src-tauri/src/converter/openai_to_antigravity.rs index 4a2539f14..c62988420 100644 --- a/src-tauri/src/converter/openai_to_antigravity.rs +++ b/src-tauri/src/converter/openai_to_antigravity.rs @@ -602,8 +602,7 @@ pub fn convert_openai_to_antigravity_with_context( .parameters .as_ref() .map(|p| { - let mut schema = - clean_parameters(Some(p.clone())).unwrap_or_default(); + let mut schema = clean_parameters(Some(p.clone())).unwrap_or_default(); // 确保有 type 和 properties if schema.get("type").is_none() { schema["type"] = serde_json::json!("object"); @@ -613,9 +612,7 @@ pub fn convert_openai_to_antigravity_with_context( } schema }) - .unwrap_or_else(|| { - serde_json::json!({"type": "object", "properties": {}}) - }); + .unwrap_or_else(|| serde_json::json!({"type": "object", "properties": {}})); if is_claude { // Claude 模型使用 parameters 字段(标准 Gemini 格式) diff --git a/src-tauri/src/credential/registry.rs b/src-tauri/src/credential/registry.rs index f902df922..26fd9a2bc 100644 --- a/src-tauri/src/credential/registry.rs +++ b/src-tauri/src/credential/registry.rs @@ -581,7 +581,11 @@ impl CredentialProviderRegistry { } } } else { - info!("No UI assets available for plugin: {} (HTTP {})", plugin_id, ui_response.status()); + info!( + "No UI assets available for plugin: {} (HTTP {})", + plugin_id, + ui_response.status() + ); } } else { info!("No UI assets available for plugin: {}", plugin_id); diff --git a/src-tauri/src/data/local_models.rs b/src-tauri/src/data/local_models.rs index 4d8928b03..0c6eb3b7c 100644 --- a/src-tauri/src/data/local_models.rs +++ b/src-tauri/src/data/local_models.rs @@ -34,24 +34,33 @@ fn get_dashscope_models() -> Vec { family: Some("qwen-coder".to_string()), tier: ModelTier::Pro, capabilities: ModelCapabilities { - vision: false, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: false, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(4.0), output_per_million: Some(16.0), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(4.0), + output_per_million: Some(16.0), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(131072), max_output_tokens: Some(8192), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(131072), + max_output_tokens: Some(8192), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2025-01-01".to_string()), is_latest: true, description: Some("阿里云通义千问代码模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, EnhancedModelMetadata { id: "qwen-max".to_string(), @@ -61,24 +70,33 @@ fn get_dashscope_models() -> Vec { family: Some("qwen".to_string()), tier: ModelTier::Max, capabilities: ModelCapabilities { - vision: true, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: true, + vision: true, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: true, }, pricing: Some(ModelPricing { - input_per_million: Some(20.0), output_per_million: Some(60.0), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(20.0), + output_per_million: Some(60.0), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(32768), max_output_tokens: Some(8192), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(32768), + max_output_tokens: Some(8192), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-12-01".to_string()), is_latest: true, description: Some("通义千问旗舰模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, EnhancedModelMetadata { id: "qwen-plus".to_string(), @@ -88,24 +106,33 @@ fn get_dashscope_models() -> Vec { family: Some("qwen".to_string()), tier: ModelTier::Pro, capabilities: ModelCapabilities { - vision: true, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: true, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(4.0), output_per_million: Some(12.0), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(4.0), + output_per_million: Some(12.0), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(131072), max_output_tokens: Some(8192), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(131072), + max_output_tokens: Some(8192), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-12-01".to_string()), is_latest: true, description: Some("通义千问增强版".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, EnhancedModelMetadata { id: "qwen-turbo".to_string(), @@ -115,24 +142,33 @@ fn get_dashscope_models() -> Vec { family: Some("qwen".to_string()), tier: ModelTier::Mini, capabilities: ModelCapabilities { - vision: false, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: false, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(0.3), output_per_million: Some(0.6), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(0.3), + output_per_million: Some(0.6), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(131072), max_output_tokens: Some(8192), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(131072), + max_output_tokens: Some(8192), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-12-01".to_string()), is_latest: true, description: Some("通义千问快速版".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, ] } @@ -149,24 +185,33 @@ fn get_zhipu_models() -> Vec { family: Some("glm-4".to_string()), tier: ModelTier::Max, capabilities: ModelCapabilities { - vision: true, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: true, + vision: true, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: true, }, pricing: Some(ModelPricing { - input_per_million: Some(50.0), output_per_million: Some(50.0), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(50.0), + output_per_million: Some(50.0), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(128000), max_output_tokens: Some(4096), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(128000), + max_output_tokens: Some(4096), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-08-01".to_string()), is_latest: true, description: Some("智谱旗舰模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, EnhancedModelMetadata { id: "glm-4-air".to_string(), @@ -176,24 +221,33 @@ fn get_zhipu_models() -> Vec { family: Some("glm-4".to_string()), tier: ModelTier::Pro, capabilities: ModelCapabilities { - vision: false, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: false, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(1.0), output_per_million: Some(1.0), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(1.0), + output_per_million: Some(1.0), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(128000), max_output_tokens: Some(4096), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(128000), + max_output_tokens: Some(4096), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-06-01".to_string()), is_latest: true, description: Some("智谱高性价比模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, EnhancedModelMetadata { id: "glm-4-flash".to_string(), @@ -203,24 +257,33 @@ fn get_zhipu_models() -> Vec { family: Some("glm-4".to_string()), tier: ModelTier::Mini, capabilities: ModelCapabilities { - vision: false, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: false, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(0.1), output_per_million: Some(0.1), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(0.1), + output_per_million: Some(0.1), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(128000), max_output_tokens: Some(4096), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(128000), + max_output_tokens: Some(4096), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-06-01".to_string()), is_latest: true, description: Some("智谱快速模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, ] } @@ -237,24 +300,33 @@ fn get_baichuan_models() -> Vec { family: Some("baichuan".to_string()), tier: ModelTier::Max, capabilities: ModelCapabilities { - vision: false, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: false, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(100.0), output_per_million: Some(100.0), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(100.0), + output_per_million: Some(100.0), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(32768), max_output_tokens: Some(4096), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(32768), + max_output_tokens: Some(4096), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-10-01".to_string()), is_latest: true, description: Some("百川旗舰模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, EnhancedModelMetadata { id: "Baichuan3-Turbo".to_string(), @@ -264,24 +336,33 @@ fn get_baichuan_models() -> Vec { family: Some("baichuan".to_string()), tier: ModelTier::Pro, capabilities: ModelCapabilities { - vision: false, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: false, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(12.0), output_per_million: Some(12.0), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(12.0), + output_per_million: Some(12.0), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(32768), max_output_tokens: Some(4096), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(32768), + max_output_tokens: Some(4096), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-06-01".to_string()), is_latest: true, description: Some("百川高性价比模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, ] } @@ -298,24 +379,33 @@ fn get_moonshot_models() -> Vec { family: Some("moonshot".to_string()), tier: ModelTier::Pro, capabilities: ModelCapabilities { - vision: false, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: false, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(60.0), output_per_million: Some(60.0), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(60.0), + output_per_million: Some(60.0), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(128000), max_output_tokens: Some(4096), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(128000), + max_output_tokens: Some(4096), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-03-01".to_string()), is_latest: true, description: Some("月之暗面长上下文模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, EnhancedModelMetadata { id: "moonshot-v1-32k".to_string(), @@ -325,24 +415,33 @@ fn get_moonshot_models() -> Vec { family: Some("moonshot".to_string()), tier: ModelTier::Pro, capabilities: ModelCapabilities { - vision: false, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: false, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(24.0), output_per_million: Some(24.0), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(24.0), + output_per_million: Some(24.0), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(32000), max_output_tokens: Some(4096), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(32000), + max_output_tokens: Some(4096), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-03-01".to_string()), is_latest: true, description: Some("月之暗面标准模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, EnhancedModelMetadata { id: "moonshot-v1-8k".to_string(), @@ -352,24 +451,33 @@ fn get_moonshot_models() -> Vec { family: Some("moonshot".to_string()), tier: ModelTier::Mini, capabilities: ModelCapabilities { - vision: false, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: false, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(12.0), output_per_million: Some(12.0), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(12.0), + output_per_million: Some(12.0), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(8000), max_output_tokens: Some(4096), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(8000), + max_output_tokens: Some(4096), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-03-01".to_string()), is_latest: true, description: Some("月之暗面快速模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, ] } @@ -386,24 +494,33 @@ fn get_deepseek_models() -> Vec { family: Some("deepseek".to_string()), tier: ModelTier::Pro, capabilities: ModelCapabilities { - vision: false, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: false, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(1.0), output_per_million: Some(2.0), - cache_read_per_million: Some(0.1), cache_write_per_million: Some(1.0), + input_per_million: Some(1.0), + output_per_million: Some(2.0), + cache_read_per_million: Some(0.1), + cache_write_per_million: Some(1.0), currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(64000), max_output_tokens: Some(8192), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(64000), + max_output_tokens: Some(8192), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-12-01".to_string()), is_latest: true, description: Some("DeepSeek V3 对话模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, EnhancedModelMetadata { id: "deepseek-reasoner".to_string(), @@ -413,24 +530,33 @@ fn get_deepseek_models() -> Vec { family: Some("deepseek".to_string()), tier: ModelTier::Max, capabilities: ModelCapabilities { - vision: false, tools: false, streaming: true, - json_mode: false, function_calling: false, reasoning: true, + vision: false, + tools: false, + streaming: true, + json_mode: false, + function_calling: false, + reasoning: true, }, pricing: Some(ModelPricing { - input_per_million: Some(4.0), output_per_million: Some(16.0), - cache_read_per_million: Some(0.4), cache_write_per_million: Some(4.0), + input_per_million: Some(4.0), + output_per_million: Some(16.0), + cache_read_per_million: Some(0.4), + cache_write_per_million: Some(4.0), currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(64000), max_output_tokens: Some(8192), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(64000), + max_output_tokens: Some(8192), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2025-01-01".to_string()), is_latest: true, description: Some("DeepSeek R1 推理模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, EnhancedModelMetadata { id: "deepseek-coder".to_string(), @@ -440,24 +566,33 @@ fn get_deepseek_models() -> Vec { family: Some("deepseek-coder".to_string()), tier: ModelTier::Pro, capabilities: ModelCapabilities { - vision: false, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: false, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(1.0), output_per_million: Some(2.0), - cache_read_per_million: Some(0.1), cache_write_per_million: Some(1.0), + input_per_million: Some(1.0), + output_per_million: Some(2.0), + cache_read_per_million: Some(0.1), + cache_write_per_million: Some(1.0), currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(64000), max_output_tokens: Some(8192), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(64000), + max_output_tokens: Some(8192), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-06-01".to_string()), is_latest: true, description: Some("DeepSeek 代码模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, ] } @@ -474,24 +609,33 @@ fn get_doubao_models() -> Vec { family: Some("doubao".to_string()), tier: ModelTier::Max, capabilities: ModelCapabilities { - vision: true, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: true, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(5.0), output_per_million: Some(9.0), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(5.0), + output_per_million: Some(9.0), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(256000), max_output_tokens: Some(4096), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(256000), + max_output_tokens: Some(4096), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-10-01".to_string()), is_latest: true, description: Some("豆包旗舰长上下文模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, EnhancedModelMetadata { id: "doubao-pro-32k".to_string(), @@ -501,24 +645,33 @@ fn get_doubao_models() -> Vec { family: Some("doubao".to_string()), tier: ModelTier::Pro, capabilities: ModelCapabilities { - vision: true, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: true, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(0.8), output_per_million: Some(2.0), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(0.8), + output_per_million: Some(2.0), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(32000), max_output_tokens: Some(4096), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(32000), + max_output_tokens: Some(4096), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-06-01".to_string()), is_latest: true, description: Some("豆包标准模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, EnhancedModelMetadata { id: "doubao-lite-32k".to_string(), @@ -528,24 +681,33 @@ fn get_doubao_models() -> Vec { family: Some("doubao".to_string()), tier: ModelTier::Mini, capabilities: ModelCapabilities { - vision: false, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: false, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(0.3), output_per_million: Some(0.6), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(0.3), + output_per_million: Some(0.6), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(32000), max_output_tokens: Some(4096), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(32000), + max_output_tokens: Some(4096), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-06-01".to_string()), is_latest: true, description: Some("豆包轻量模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, ] } @@ -553,35 +715,42 @@ fn get_doubao_models() -> Vec { /// MiniMax 系列模型 fn get_minimax_models() -> Vec { let now = chrono::Utc::now().timestamp(); - vec![ - EnhancedModelMetadata { - id: "abab6.5s-chat".to_string(), - display_name: "MiniMax abab6.5s".to_string(), - provider_id: "minimax".to_string(), - provider_name: "MiniMax".to_string(), - family: Some("abab".to_string()), - tier: ModelTier::Pro, - capabilities: ModelCapabilities { - vision: false, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, - }, - pricing: Some(ModelPricing { - input_per_million: Some(1.0), output_per_million: Some(1.0), - cache_read_per_million: None, cache_write_per_million: None, - currency: "CNY".to_string(), - }), - limits: ModelLimits { - context_length: Some(245760), max_output_tokens: Some(8192), - requests_per_minute: None, tokens_per_minute: None, - }, - status: ModelStatus::Active, - release_date: Some("2024-06-01".to_string()), - is_latest: true, - description: Some("MiniMax 长上下文模型".to_string()), - source: ModelSource::Local, - created_at: now, updated_at: now, + vec![EnhancedModelMetadata { + id: "abab6.5s-chat".to_string(), + display_name: "MiniMax abab6.5s".to_string(), + provider_id: "minimax".to_string(), + provider_name: "MiniMax".to_string(), + family: Some("abab".to_string()), + tier: ModelTier::Pro, + capabilities: ModelCapabilities { + vision: false, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, - ] + pricing: Some(ModelPricing { + input_per_million: Some(1.0), + output_per_million: Some(1.0), + cache_read_per_million: None, + cache_write_per_million: None, + currency: "CNY".to_string(), + }), + limits: ModelLimits { + context_length: Some(245760), + max_output_tokens: Some(8192), + requests_per_minute: None, + tokens_per_minute: None, + }, + status: ModelStatus::Active, + release_date: Some("2024-06-01".to_string()), + is_latest: true, + description: Some("MiniMax 长上下文模型".to_string()), + source: ModelSource::Local, + created_at: now, + updated_at: now, + }] } /// 零一万物 Yi 系列模型 @@ -596,24 +765,33 @@ fn get_yi_models() -> Vec { family: Some("yi".to_string()), tier: ModelTier::Max, capabilities: ModelCapabilities { - vision: false, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: false, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(20.0), output_per_million: Some(20.0), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(20.0), + output_per_million: Some(20.0), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(32768), max_output_tokens: Some(4096), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(32768), + max_output_tokens: Some(4096), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-05-01".to_string()), is_latest: true, description: Some("零一万物旗舰模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, EnhancedModelMetadata { id: "yi-medium".to_string(), @@ -623,24 +801,33 @@ fn get_yi_models() -> Vec { family: Some("yi".to_string()), tier: ModelTier::Pro, capabilities: ModelCapabilities { - vision: false, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: false, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(2.5), output_per_million: Some(2.5), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(2.5), + output_per_million: Some(2.5), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(16384), max_output_tokens: Some(4096), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(16384), + max_output_tokens: Some(4096), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-05-01".to_string()), is_latest: true, description: Some("零一万物标准模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, EnhancedModelMetadata { id: "yi-spark".to_string(), @@ -650,24 +837,33 @@ fn get_yi_models() -> Vec { family: Some("yi".to_string()), tier: ModelTier::Mini, capabilities: ModelCapabilities { - vision: false, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: false, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(1.0), output_per_million: Some(1.0), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(1.0), + output_per_million: Some(1.0), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(16384), max_output_tokens: Some(4096), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(16384), + max_output_tokens: Some(4096), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-05-01".to_string()), is_latest: true, description: Some("零一万物快速模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, ] } @@ -684,24 +880,33 @@ fn get_stepfun_models() -> Vec { family: Some("step".to_string()), tier: ModelTier::Max, capabilities: ModelCapabilities { - vision: true, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: true, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(38.0), output_per_million: Some(120.0), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(38.0), + output_per_million: Some(120.0), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(16384), max_output_tokens: Some(4096), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(16384), + max_output_tokens: Some(4096), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-09-01".to_string()), is_latest: true, description: Some("阶跃星辰旗舰模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, EnhancedModelMetadata { id: "step-1-128k".to_string(), @@ -711,24 +916,33 @@ fn get_stepfun_models() -> Vec { family: Some("step".to_string()), tier: ModelTier::Pro, capabilities: ModelCapabilities { - vision: false, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: false, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(40.0), output_per_million: Some(100.0), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(40.0), + output_per_million: Some(100.0), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(128000), max_output_tokens: Some(4096), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(128000), + max_output_tokens: Some(4096), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-06-01".to_string()), is_latest: true, description: Some("阶跃星辰长上下文模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, EnhancedModelMetadata { id: "step-1-flash".to_string(), @@ -738,24 +952,33 @@ fn get_stepfun_models() -> Vec { family: Some("step".to_string()), tier: ModelTier::Mini, capabilities: ModelCapabilities { - vision: false, tools: true, streaming: true, - json_mode: true, function_calling: true, reasoning: false, + vision: false, + tools: true, + streaming: true, + json_mode: true, + function_calling: true, + reasoning: false, }, pricing: Some(ModelPricing { - input_per_million: Some(1.0), output_per_million: Some(4.0), - cache_read_per_million: None, cache_write_per_million: None, + input_per_million: Some(1.0), + output_per_million: Some(4.0), + cache_read_per_million: None, + cache_write_per_million: None, currency: "CNY".to_string(), }), limits: ModelLimits { - context_length: Some(8192), max_output_tokens: Some(4096), - requests_per_minute: None, tokens_per_minute: None, + context_length: Some(8192), + max_output_tokens: Some(4096), + requests_per_minute: None, + tokens_per_minute: None, }, status: ModelStatus::Active, release_date: Some("2024-06-01".to_string()), is_latest: true, description: Some("阶跃星辰快速模型".to_string()), source: ModelSource::Local, - created_at: now, updated_at: now, + created_at: now, + updated_at: now, }, ] } diff --git a/src-tauri/src/models/model_registry.rs b/src-tauri/src/models/model_registry.rs index 38f0e9e64..be5505a9b 100644 --- a/src-tauri/src/models/model_registry.rs +++ b/src-tauri/src/models/model_registry.rs @@ -598,7 +598,10 @@ mod tests { #[test] fn test_model_status_parsing() { - assert_eq!("active".parse::().unwrap(), ModelStatus::Active); + assert_eq!( + "active".parse::().unwrap(), + ModelStatus::Active + ); assert_eq!( "deprecated".parse::().unwrap(), ModelStatus::Deprecated diff --git a/src-tauri/src/orchestrator/pool_builder.rs b/src-tauri/src/orchestrator/pool_builder.rs index d6ee2dd10..b4decdb84 100644 --- a/src-tauri/src/orchestrator/pool_builder.rs +++ b/src-tauri/src/orchestrator/pool_builder.rs @@ -109,7 +109,9 @@ impl ProviderDefinition { model_lower.starts_with(parts[0]) } else { // 复杂模式,回退到简单包含检查 - parts.iter().all(|p| p.is_empty() || model_lower.contains(p)) + parts + .iter() + .all(|p| p.is_empty() || model_lower.contains(p)) } } else { model_lower.contains(&pattern_lower) diff --git a/src-tauri/src/services/model_registry_service.rs b/src-tauri/src/services/model_registry_service.rs index 13bcc3aff..4cad6a25c 100644 --- a/src-tauri/src/services/model_registry_service.rs +++ b/src-tauri/src/services/model_registry_service.rs @@ -5,8 +5,8 @@ use crate::data::get_local_models; use crate::database::DbConnection; use crate::models::model_registry::{ - EnhancedModelMetadata, ModelSource, ModelStatus, - ModelSyncState, ModelTier, ModelsDevProvider, UserModelPreference, + EnhancedModelMetadata, ModelSource, ModelStatus, ModelSyncState, ModelTier, ModelsDevProvider, + UserModelPreference, }; use rusqlite::params; use std::collections::HashMap; @@ -135,10 +135,7 @@ impl ModelRegistryService { let local_models = get_local_models(); let merged = self.merge_models(models_dev_models, local_models); - tracing::info!( - "[ModelRegistry] 获取并合并了 {} 个模型", - merged.len() - ); + tracing::info!("[ModelRegistry] 获取并合并了 {} 个模型", merged.len()); // 更新缓存 { @@ -193,10 +190,7 @@ impl ModelRegistryService { .map_err(|e| format!("请求 models.dev 失败: {}", e))?; if !response.status().is_success() { - return Err(format!( - "models.dev 返回错误状态码: {}", - response.status() - )); + return Err(format!("models.dev 返回错误状态码: {}", response.status())); } let data: HashMap = response @@ -283,10 +277,8 @@ impl ModelRegistryService { provider_name: row.get(3)?, family: row.get(4)?, tier: tier_str.parse().unwrap_or(ModelTier::Pro), - capabilities: serde_json::from_str(&capabilities_json) - .unwrap_or_default(), - pricing: pricing_json - .and_then(|s| serde_json::from_str(&s).ok()), + capabilities: serde_json::from_str(&capabilities_json).unwrap_or_default(), + pricing: pricing_json.and_then(|s| serde_json::from_str(&s).ok()), limits: serde_json::from_str(&limits_json).unwrap_or_default(), status: status_str.parse().unwrap_or(ModelStatus::Active), release_date: row.get(10)?, @@ -361,8 +353,7 @@ impl ModelRegistryService { .map_err(|e| e.to_string())?; for model in models { - let capabilities_json = - serde_json::to_string(&model.capabilities).unwrap_or_default(); + let capabilities_json = serde_json::to_string(&model.capabilities).unwrap_or_default(); let pricing_json = model .pricing .as_ref() @@ -402,7 +393,11 @@ impl ModelRegistryService { 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()) + ( + state.last_sync_at, + state.model_count, + state.last_error.clone(), + ) }; let conn = self.db.lock().map_err(|e| e.to_string())?; @@ -485,10 +480,7 @@ impl ModelRegistryService { .collect(); // 按分数降序排序 - scored.sort_by(|a, b| { - b.0.partial_cmp(&a.0) - .unwrap_or(std::cmp::Ordering::Equal) - }); + scored.sort_by(|a, b| b.0.partial_cmp(&a.0).unwrap_or(std::cmp::Ordering::Equal)); scored .into_iter() diff --git a/src/components/agent/chat/components/ChatNavbar.tsx b/src/components/agent/chat/components/ChatNavbar.tsx index f54f2f90d..9320cd182 100644 --- a/src/components/agent/chat/components/ChatNavbar.tsx +++ b/src/components/agent/chat/components/ChatNavbar.tsx @@ -1,4 +1,4 @@ -import React, { useState } from "react"; +import React, { useState, useMemo, useEffect } from "react"; import { Bot, ChevronDown, Check, Box, Settings2 } from "lucide-react"; import { Button } from "@/components/ui/button"; import { @@ -8,8 +8,47 @@ import { } from "@/components/ui/popover"; import { ScrollArea } from "@/components/ui/scroll-area"; import { Navbar } from "../styles"; -import { PROVIDER_CONFIG } from "../types"; import { cn } from "@/lib/utils"; +import { useProviderPool } from "@/hooks/useProviderPool"; +import { useApiKeyProvider } from "@/hooks/useApiKeyProvider"; +import { useModelRegistry } from "@/hooks/useModelRegistry"; + +// OAuth 凭证类型到显示名称和 registry ID 的映射 +const CREDENTIAL_TYPE_CONFIG: Record< + string, + { label: string; registryId: string } +> = { + kiro: { label: "Kiro", registryId: "anthropic" }, + gemini: { label: "Gemini", registryId: "google" }, + qwen: { label: "通义千问", registryId: "alibaba" }, + antigravity: { label: "Antigravity", registryId: "google" }, + codex: { label: "Codex", registryId: "openai" }, + claude_oauth: { label: "Claude OAuth", registryId: "anthropic" }, + iflow: { label: "iFlow", registryId: "custom" }, + openai: { label: "OpenAI", registryId: "openai" }, + claude: { label: "Claude", registryId: "anthropic" }, + gemini_api_key: { label: "Gemini", registryId: "google" }, +}; + +// API Key Provider 类型到显示名称和 registry ID 的映射 +const API_KEY_PROVIDER_CONFIG: Record< + string, + { label: string; registryId: string } +> = { + anthropic: { label: "Anthropic", registryId: "anthropic" }, + openai: { label: "OpenAI", registryId: "openai" }, + gemini: { label: "Gemini", registryId: "google" }, + "azure-openai": { label: "Azure OpenAI", registryId: "openai" }, + vertexai: { label: "VertexAI", registryId: "google" }, + ollama: { label: "Ollama", registryId: "ollama" }, +}; + +/** 已配置的 Provider 信息 */ +interface ConfiguredProvider { + key: string; + label: string; + registryId: string; +} interface ChatNavbarProps { providerType: string; @@ -34,9 +73,79 @@ export const ChatNavbar: React.FC = ({ }) => { const [open, setOpen] = useState(false); - const selectedProviderLabel = - PROVIDER_CONFIG[providerType]?.label || providerType; - const currentModels = PROVIDER_CONFIG[providerType]?.models || []; + // 获取凭证池数据 + const { overview: oauthCredentials } = useProviderPool(); + const { providers: apiKeyProviders } = useApiKeyProvider(); + + // 获取模型注册表数据 + const { models: registryModels } = useModelRegistry({ autoLoad: true }); + + // 计算已配置的 Provider 列表 + const configuredProviders = useMemo(() => { + const providerMap = new Map(); + + // 从 OAuth 凭证提取 Provider + oauthCredentials.forEach((overview) => { + if (overview.credentials.length > 0) { + const config = CREDENTIAL_TYPE_CONFIG[overview.provider_type]; + if (config && !providerMap.has(overview.provider_type)) { + providerMap.set(overview.provider_type, { + key: overview.provider_type, + label: config.label, + registryId: config.registryId, + }); + } + } + }); + + // 从 API Key Provider 提取(只包含有 API Key 的) + apiKeyProviders + .filter((p) => p.api_key_count > 0 && p.enabled) + .forEach((provider) => { + const config = API_KEY_PROVIDER_CONFIG[provider.type]; + if (config && !providerMap.has(provider.type)) { + providerMap.set(provider.type, { + key: provider.type, + label: config.label, + registryId: config.registryId, + }); + } + }); + + return Array.from(providerMap.values()); + }, [oauthCredentials, apiKeyProviders]); + + // 获取当前选中 Provider 的配置 + const selectedProvider = useMemo(() => { + return configuredProviders.find((p) => p.key === providerType); + }, [configuredProviders, providerType]); + + // 获取当前 Provider 的模型列表(从 model_registry 获取) + const currentModels = useMemo(() => { + if (!selectedProvider) return []; + + // 从 model_registry 获取模型 + return registryModels + .filter((m) => m.provider_id === selectedProvider.registryId) + .map((m) => m.id); + }, [selectedProvider, registryModels]); + + // 如果当前选中的 Provider 不在已配置列表中,自动切换到第一个已配置的 + useEffect(() => { + if (configuredProviders.length > 0 && !selectedProvider) { + const firstProvider = configuredProviders[0]; + setProviderType(firstProvider.key); + } + }, [configuredProviders, selectedProvider, setProviderType]); + + // 当 Provider 切换时,自动选择第一个模型 + useEffect(() => { + if (currentModels.length > 0 && !currentModels.includes(model)) { + setModel(currentModels[0]); + } + }, [currentModels, model, setModel]); + + const selectedProviderLabel = selectedProvider?.label || providerType; return ( @@ -75,36 +184,36 @@ export const ChatNavbar: React.FC = ({ > {/* Provider/Model Selection */}
- {/* Left Column: Providers */} + {/* Left Column: Providers (只显示已配置的) */}
Providers
- {Object.entries(PROVIDER_CONFIG).map(([key, config]) => ( - - ))} + {configuredProviders.length === 0 ? ( +
+ 暂无已配置的 Provider +
+ ) : ( + configuredProviders.map((provider) => ( + + )) + )}
{/* Right Column: Models */} diff --git a/src/components/agent/chat/hooks/useAgentChat.ts b/src/components/agent/chat/hooks/useAgentChat.ts index 14ee571bd..9b48eba6c 100644 --- a/src/components/agent/chat/hooks/useAgentChat.ts +++ b/src/components/agent/chat/hooks/useAgentChat.ts @@ -141,9 +141,12 @@ export function useAgentChat() { // 如果不兼容,自动切换到新 provider 的第一个模型 useEffect(() => { const currentProviderModels = providerConfig[providerType]?.models || []; - if (currentProviderModels.length > 0 && !currentProviderModels.includes(model)) { + if ( + currentProviderModels.length > 0 && + !currentProviderModels.includes(model) + ) { console.log( - `[useAgentChat] 模型 ${model} 不在 ${providerType} 支持列表中,自动切换到 ${currentProviderModels[0]}` + `[useAgentChat] 模型 ${model} 不在 ${providerType} 支持列表中,自动切换到 ${currentProviderModels[0]}`, ); setModel(currentProviderModels[0]); } diff --git a/src/components/api-server/ApiServerPage.tsx b/src/components/api-server/ApiServerPage.tsx index 44ffc510b..08deecf35 100644 --- a/src/components/api-server/ApiServerPage.tsx +++ b/src/components/api-server/ApiServerPage.tsx @@ -9,7 +9,6 @@ import { } from "lucide-react"; import { LogsTab } from "./LogsTab"; import { RoutesTab } from "./RoutesTab"; -import { EnhancedModelsTab } from "./EnhancedModelsTab"; import { ProviderIcon } from "@/icons/providers"; import { startServer, @@ -41,7 +40,7 @@ interface TestState { httpStatus?: number; } -type TabId = "server" | "routes" | "models" | "logs"; +type TabId = "server" | "routes" | "logs"; // 可用的 Provider 信息(合并 OAuth 凭证池和 API Key Provider) interface AvailableProvider { @@ -671,7 +670,6 @@ export function ApiServerPage() { {[ { id: "server" as TabId, name: "服务器控制" }, { id: "routes" as TabId, name: "路由端点" }, - { id: "models" as TabId, name: "模型列表" }, { id: "logs" as TabId, name: "系统日志" }, ].map((tab) => (
diff --git a/src/components/api-server/EnhancedModelsTab.tsx b/src/components/api-server/EnhancedModelsTab.tsx index 195348564..60226e837 100644 --- a/src/components/api-server/EnhancedModelsTab.tsx +++ b/src/components/api-server/EnhancedModelsTab.tsx @@ -21,7 +21,10 @@ import { } from "lucide-react"; import { cn } from "@/lib/utils"; import { useModelRegistry } from "@/hooks/useModelRegistry"; -import type { EnhancedModelMetadata, ModelTier } from "@/lib/types/modelRegistry"; +import type { + EnhancedModelMetadata, + ModelTier, +} from "@/lib/types/modelRegistry"; export function EnhancedModelsTab() { const { @@ -113,10 +116,15 @@ export function EnhancedModelsTab() { "flex items-center gap-2 rounded-lg border px-4 py-2 text-sm font-medium transition-colors", showFavoritesOnly ? "bg-yellow-100 border-yellow-300 text-yellow-700 dark:bg-yellow-900 dark:border-yellow-700 dark:text-yellow-300" - : "hover:bg-muted" + : "hover:bg-muted", )} > - + 收藏 @@ -127,7 +135,9 @@ export function EnhancedModelsTab() { onClick={() => setSelectedProvider(null)} className={cn( "rounded-lg px-3 py-1.5 text-sm font-medium transition-colors", - !selectedProvider ? "bg-primary text-primary-foreground" : "bg-muted hover:bg-muted/80" + !selectedProvider + ? "bg-primary text-primary-foreground" + : "bg-muted hover:bg-muted/80", )} > 全部 ({models.length}) @@ -138,10 +148,16 @@ export function EnhancedModelsTab() { return ( @@ -336,10 +360,24 @@ function ModelRow({ /** 服务等级徽章 */ function TierBadge({ tier }: { tier: string }) { const config = { - mini: { label: "Mini", color: "bg-green-100 text-green-700 dark:bg-green-900 dark:text-green-300" }, - pro: { label: "Pro", color: "bg-blue-100 text-blue-700 dark:bg-blue-900 dark:text-blue-300" }, - max: { label: "Max", color: "bg-purple-100 text-purple-700 dark:bg-purple-900 dark:text-purple-300" }, - }[tier] || { label: tier, color: "bg-gray-100 text-gray-700 dark:bg-gray-800 dark:text-gray-300" }; + mini: { + label: "Mini", + color: + "bg-green-100 text-green-700 dark:bg-green-900 dark:text-green-300", + }, + pro: { + label: "Pro", + color: "bg-blue-100 text-blue-700 dark:bg-blue-900 dark:text-blue-300", + }, + max: { + label: "Max", + color: + "bg-purple-100 text-purple-700 dark:bg-purple-900 dark:text-purple-300", + }, + }[tier] || { + label: tier, + color: "bg-gray-100 text-gray-700 dark:bg-gray-800 dark:text-gray-300", + }; return ( diff --git a/src/components/model-selector/EnhancedModelList.tsx b/src/components/model-selector/EnhancedModelList.tsx index 0e7dcf80e..dc9a9c596 100644 --- a/src/components/model-selector/EnhancedModelList.tsx +++ b/src/components/model-selector/EnhancedModelList.tsx @@ -61,7 +61,7 @@ export function EnhancedModelList({ }: EnhancedModelListProps) { const [searchQuery, setSearchQuery] = useState(""); const [expandedGroups, setExpandedGroups] = useState>( - new Set(["favorites"]) + new Set(["favorites"]), ); // 过滤模型 @@ -74,7 +74,7 @@ export function EnhancedModelList({ m.id.toLowerCase().includes(query) || m.display_name.toLowerCase().includes(query) || m.provider_name.toLowerCase().includes(query) || - m.family?.toLowerCase().includes(query) + m.family?.toLowerCase().includes(query), ); }, [models, searchQuery]); @@ -130,7 +130,7 @@ export function EnhancedModelList({
@@ -248,7 +248,7 @@ function ModelItem({
@@ -257,7 +257,9 @@ function ModelItem({
{isSelected && } @@ -313,9 +315,7 @@ function ModelItem({ {showPricing && model.pricing && (
- - {model.pricing.input_per_million?.toFixed(2) || "?"} - + {model.pricing.input_per_million?.toFixed(2) || "?"}
)} @@ -334,7 +334,7 @@ function ModelItem({ "h-4 w-4", isFavorite ? "text-yellow-500 fill-yellow-500" - : "text-muted-foreground" + : "text-muted-foreground", )} /> @@ -346,9 +346,20 @@ function ModelItem({ /** 服务等级徽章 */ function TierBadge({ tier }: { tier: string }) { const config = { - mini: { label: "Mini", color: "bg-green-100 text-green-700 dark:bg-green-900 dark:text-green-300" }, - pro: { label: "Pro", color: "bg-blue-100 text-blue-700 dark:bg-blue-900 dark:text-blue-300" }, - max: { label: "Max", color: "bg-purple-100 text-purple-700 dark:bg-purple-900 dark:text-purple-300" }, + mini: { + label: "Mini", + color: + "bg-green-100 text-green-700 dark:bg-green-900 dark:text-green-300", + }, + pro: { + label: "Pro", + color: "bg-blue-100 text-blue-700 dark:bg-blue-900 dark:text-blue-300", + }, + max: { + label: "Max", + color: + "bg-purple-100 text-purple-700 dark:bg-purple-900 dark:text-purple-300", + }, }[tier] || { label: tier, color: "bg-gray-100 text-gray-700" }; return ( @@ -361,7 +372,7 @@ function TierBadge({ tier }: { tier: string }) { /** 获取分组名称 */ function getGroupName( groupId: string, - firstModel: EnhancedModelMetadata + firstModel: EnhancedModelMetadata, ): string { if (groupId === "favorites") return "收藏"; if (groupId === "all") return "全部模型"; diff --git a/src/components/model-selector/ProviderModelSelector.tsx b/src/components/model-selector/ProviderModelSelector.tsx new file mode 100644 index 000000000..1e0f62768 --- /dev/null +++ b/src/components/model-selector/ProviderModelSelector.tsx @@ -0,0 +1,413 @@ +/** + * @file ProviderModelSelector 组件 + * @description 双栏模型选择器:左侧 Provider 列表,右侧模型列表 + * @module components/model-selector/ProviderModelSelector + */ + +import React, { useState, useMemo, useCallback, useEffect } from "react"; +import { cn } from "@/lib/utils"; +import { useModelRegistry } from "@/hooks/useModelRegistry"; +import { useProviderPool } from "@/hooks/useProviderPool"; +import { useApiKeyProvider } from "@/hooks/useApiKeyProvider"; +import { + Check, + ChevronRight, + Eye, + Wrench, + Brain, + Loader2, + AlertCircle, +} from "lucide-react"; +import type { EnhancedModelMetadata } from "@/lib/types/modelRegistry"; + +// ============================================================================ +// 类型定义 +// ============================================================================ + +export interface ProviderModelSelectorProps { + /** 选择模型回调 */ + onSelect?: (model: EnhancedModelMetadata, providerId: string) => void; + /** 初始选中的 Provider */ + initialProviderId?: string; + /** 初始选中的模型 */ + initialModelId?: string; + /** 自定义类名 */ + className?: string; +} + +/** 已配置的 Provider 信息 */ +interface ConfiguredProvider { + id: string; + name: string; + registryId: string; + source: "oauth" | "apikey"; + credentialCount: number; +} + +// ============================================================================ +// 常量 +// ============================================================================ + +/** OAuth 凭证类型到 Provider ID 的映射 */ +const CREDENTIAL_TYPE_TO_PROVIDER_ID: Record = { + kiro: "anthropic", + gemini: "google", + qwen: "alibaba", + antigravity: "google", + codex: "openai", + claude_oauth: "anthropic", + iflow: "anthropic", + openai: "openai", + claude: "anthropic", + gemini_api_key: "google", +}; + +/** API Key Provider 类型到 Registry ID 的映射 */ +const PROVIDER_TYPE_TO_REGISTRY_ID: Record = { + anthropic: "anthropic", + openai: "openai", + "openai-response": "openai", + gemini: "google", + "azure-openai": "openai", + vertexai: "google", + "aws-bedrock": "anthropic", + ollama: "ollama", + "new-api": "custom", + gateway: "custom", +}; + +/** Provider 显示名称 */ +const PROVIDER_DISPLAY_NAMES: Record = { + anthropic: "Anthropic", + openai: "OpenAI", + google: "Google", + alibaba: "阿里云", + ollama: "Ollama", + custom: "自定义", +}; + +// ============================================================================ +// 子组件 +// ============================================================================ + +interface ProviderItemProps { + provider: ConfiguredProvider; + isSelected: boolean; + onClick: () => void; +} + +/** Provider 列表项 */ +const ProviderItem: React.FC = ({ + provider, + isSelected, + onClick, +}) => { + return ( + + ); +}; + +interface ModelItemProps { + model: EnhancedModelMetadata; + isSelected: boolean; + onClick: () => void; +} + +/** 模型列表项 */ +const ModelItem: React.FC = ({ + model, + isSelected, + onClick, +}) => { + return ( + + ); +}; + +// ============================================================================ +// 主组件 +// ============================================================================ + +/** + * 双栏模型选择器组件 + * + * 左侧显示已配置凭证的 Provider 列表(单选) + * 右侧显示选中 Provider 对应的模型列表(单选) + * + * @example + * ```tsx + * { + * console.log("选中模型:", model.display_name); + * }} + * /> + * ``` + */ +export const ProviderModelSelector: React.FC = ({ + onSelect, + initialProviderId, + initialModelId, + className, +}) => { + // 状态 + const [selectedProviderId, setSelectedProviderId] = useState( + initialProviderId || null, + ); + const [selectedModelId, setSelectedModelId] = useState( + initialModelId || null, + ); + + // 获取凭证池数据 + const { overview: oauthCredentials, loading: oauthLoading } = + useProviderPool(); + const { providers: apiKeyProviders, loading: apiKeyLoading } = + useApiKeyProvider(); + + // 获取模型注册表数据 + const { + models, + loading: modelsLoading, + error: modelsError, + } = useModelRegistry({ + autoLoad: true, + }); + + // 计算已配置的 Provider 列表 + const configuredProviders = useMemo(() => { + const providerMap = new Map(); + + // 从 OAuth 凭证提取 Provider + oauthCredentials.forEach((overview) => { + const registryId = CREDENTIAL_TYPE_TO_PROVIDER_ID[overview.provider_type]; + if (registryId && overview.credentials.length > 0) { + const existing = providerMap.get(registryId); + if (existing) { + existing.credentialCount += overview.credentials.length; + } else { + providerMap.set(registryId, { + id: registryId, + name: PROVIDER_DISPLAY_NAMES[registryId] || registryId, + registryId, + source: "oauth", + credentialCount: overview.credentials.length, + }); + } + } + }); + + // 从 API Key Provider 提取(只包含有 API Key 的) + apiKeyProviders + .filter((p) => p.api_key_count > 0 && p.enabled) + .forEach((provider) => { + const registryId = + PROVIDER_TYPE_TO_REGISTRY_ID[provider.type] || provider.type; + const existing = providerMap.get(registryId); + if (existing) { + existing.credentialCount += provider.api_key_count; + } else { + providerMap.set(registryId, { + id: registryId, + name: PROVIDER_DISPLAY_NAMES[registryId] || provider.name, + registryId, + source: "apikey", + credentialCount: provider.api_key_count, + }); + } + }); + + return Array.from(providerMap.values()).sort((a, b) => + a.name.localeCompare(b.name), + ); + }, [oauthCredentials, apiKeyProviders]); + + // 默认选中第一个 Provider + useEffect(() => { + if (!selectedProviderId && configuredProviders.length > 0) { + setSelectedProviderId(configuredProviders[0].registryId); + } + }, [selectedProviderId, configuredProviders]); + + // 过滤当前 Provider 的模型 + const filteredModels = useMemo(() => { + if (!selectedProviderId) return []; + return models.filter((m) => m.provider_id === selectedProviderId); + }, [models, selectedProviderId]); + + // 选择 Provider + const handleSelectProvider = useCallback((providerId: string) => { + setSelectedProviderId(providerId); + setSelectedModelId(null); // 切换 Provider 时清除模型选择 + }, []); + + // 选择模型 + const handleSelectModel = useCallback( + (model: EnhancedModelMetadata) => { + setSelectedModelId(model.id); + if (selectedProviderId) { + onSelect?.(model, selectedProviderId); + } + }, + [selectedProviderId, onSelect], + ); + + const isLoading = oauthLoading || apiKeyLoading || modelsLoading; + + // 空状态 + if (!isLoading && configuredProviders.length === 0) { + return ( +
+ +

暂无已配置的 Provider

+

请先在凭证池中添加凭证

+
+ ); + } + + return ( +
+ {/* 左侧:Provider 列表 */} +
+
+

Providers

+

已配置凭证的

+
+
+ {isLoading ? ( +
+ +
+ ) : ( + configuredProviders.map((provider) => ( + handleSelectProvider(provider.registryId)} + /> + )) + )} +
+
+ + {/* 右侧:模型列表 */} +
+
+

Models

+

+ {selectedProviderId + ? `${PROVIDER_DISPLAY_NAMES[selectedProviderId] || selectedProviderId} 的模型` + : "请选择 Provider"} +

+
+
+ {modelsLoading ? ( +
+ +
+ ) : modelsError ? ( +
+ +

{modelsError}

+
+ ) : filteredModels.length === 0 ? ( +
+

暂无模型数据

+
+ ) : ( + filteredModels.map((model) => ( + handleSelectModel(model)} + /> + )) + )} +
+
+
+ ); +}; + +export default ProviderModelSelector; diff --git a/src/components/model-selector/index.ts b/src/components/model-selector/index.ts index 7f5cac6fa..239c6dcd7 100644 --- a/src/components/model-selector/index.ts +++ b/src/components/model-selector/index.ts @@ -6,6 +6,8 @@ export { ModelSelector } from "./ModelSelector"; export { TierSelector, tierOptions } from "./TierSelector"; export { ModeToggle } from "./ModeToggle"; export { ModelList } from "./ModelList"; +export { ProviderModelSelector } from "./ProviderModelSelector"; export type { SelectionMode } from "./ModeToggle"; export type { TierOption } from "./TierSelector"; +export type { ProviderModelSelectorProps } from "./ProviderModelSelector"; diff --git a/src/components/provider-pool/ModelRegistryTab.tsx b/src/components/provider-pool/ModelRegistryTab.tsx new file mode 100644 index 000000000..97df8952b --- /dev/null +++ b/src/components/provider-pool/ModelRegistryTab.tsx @@ -0,0 +1,22 @@ +/** + * @file ModelRegistryTab 组件 + * @description 模型库 Tab,显示所有可用模型 + * @module components/provider-pool/ModelRegistryTab + */ + +import { EnhancedModelsTab } from "@/components/api-server/EnhancedModelsTab"; + +/** + * 模型库 Tab 组件 + * + * 复用 API Server 的 EnhancedModelsTab 组件 + */ +export function ModelRegistryTab() { + return ( +
+ +
+ ); +} + +export default ModelRegistryTab; diff --git a/src/components/provider-pool/OAuthPluginTab.tsx b/src/components/provider-pool/OAuthPluginTab.tsx index c7434b8f6..9e51b3639 100644 --- a/src/components/provider-pool/OAuthPluginTab.tsx +++ b/src/components/provider-pool/OAuthPluginTab.tsx @@ -147,10 +147,10 @@ const recommendedOAuthPlugins: RecommendedOAuthPlugin[] = [ type: "git_hub", owner: "aiclientproxy", repo: "droid-provider", - version: "v0.2.0", + version: "v0.3.0", }, downloadUrl: - "https://github.com/aiclientproxy/droid-provider/releases/download/v0.2.0/droid-provider-plugin.zip", + "https://github.com/aiclientproxy/droid-provider/releases/download/v0.3.0/droid-provider-plugin.zip", tags: ["anthropic", "openai"], recommended: false, available: true, @@ -165,10 +165,10 @@ const recommendedOAuthPlugins: RecommendedOAuthPlugin[] = [ type: "git_hub", owner: "aiclientproxy", repo: "gemini-provider", - version: "v0.3.0", + version: "v0.4.0", }, downloadUrl: - "https://github.com/aiclientproxy/gemini-provider/releases/download/v0.3.0/gemini-provider-plugin.zip", + "https://github.com/aiclientproxy/gemini-provider/releases/download/v0.4.0/gemini-provider-plugin.zip", tags: ["gemini", "API Key"], recommended: false, available: true, diff --git a/src/components/provider-pool/ProviderPoolPage.tsx b/src/components/provider-pool/ProviderPoolPage.tsx index a317e49dc..f494a0f5a 100644 --- a/src/components/provider-pool/ProviderPoolPage.tsx +++ b/src/components/provider-pool/ProviderPoolPage.tsx @@ -36,6 +36,7 @@ import { ProviderIcon } from "@/icons/providers"; import { ApiKeyProviderSection, AddCustomProviderModal } from "./api-key"; import { OAuthPluginTab } from "./OAuthPluginTab"; import { RelayProvidersSection } from "./RelayProvidersSection"; +import { ModelRegistryTab } from "./ModelRegistryTab"; import type { AddCustomProviderRequest } from "@/lib/api/apiKeyProvider"; import { getLocalKiroCredentialUuid, @@ -84,7 +85,7 @@ const isConfigTab = (tab: TabType): tab is ConfigTabType => { }; // 分类类型 -type CategoryType = "oauth" | "apikey" | "plugins" | "connect"; +type CategoryType = "oauth" | "apikey" | "plugins" | "connect" | "models"; export const ProviderPoolPage = forwardRef( (_props, ref) => { @@ -410,6 +411,19 @@ export const ProviderPoolPage = forwardRef( > Connect +
{/* OAuth 凭证分类 - Provider 选择图标网格 */} @@ -478,6 +492,9 @@ export const ProviderPoolPage = forwardRef(
)} + {/* 模型库分类 */} + {activeCategory === "models" && } + {/* OAuth 凭证内容 - 卡片布局 */} {activeCategory === "oauth" && !isConfigTab(activeTab) && diff --git a/src/components/provider-pool/api-key/AddCustomProviderModal.tsx b/src/components/provider-pool/api-key/AddCustomProviderModal.tsx index 0ef46488b..44e037a48 100644 --- a/src/components/provider-pool/api-key/AddCustomProviderModal.tsx +++ b/src/components/provider-pool/api-key/AddCustomProviderModal.tsx @@ -7,7 +7,13 @@ * **Validates: Requirements 6.1, 6.2** */ -import React, { useState, useCallback, useMemo } from "react"; +import React, { + useState, + useCallback, + useMemo, + useRef, + useEffect, +} from "react"; import { cn } from "@/lib/utils"; import { Modal, ModalHeader, ModalBody, ModalFooter } from "@/components/Modal"; import { Button } from "@/components/ui/button"; @@ -20,6 +26,8 @@ import { SelectTrigger, SelectValue, } from "@/components/ui/select"; +import { Search, X } from "lucide-react"; +import { useModelRegistry } from "@/hooks/useModelRegistry"; import type { ProviderType } from "@/lib/types/provider"; import type { AddCustomProviderRequest } from "@/lib/api/apiKeyProvider"; @@ -55,6 +63,114 @@ const PROVIDER_TYPE_EXTRA_FIELDS: Record = { gateway: [], }; +/** 已知厂商配置 */ +interface KnownProvider { + id: string; + name: string; + type: ProviderType; + apiHost?: string; +} + +/** 已知厂商列表(用于快速填充) */ +const KNOWN_PROVIDERS: KnownProvider[] = [ + { + id: "anthropic", + name: "Anthropic", + type: "anthropic", + apiHost: "https://api.anthropic.com", + }, + { + id: "openai", + name: "OpenAI", + type: "openai", + apiHost: "https://api.openai.com", + }, + { + id: "google", + name: "Google (Gemini)", + type: "gemini", + apiHost: "https://generativelanguage.googleapis.com", + }, + { + id: "alibaba", + name: "阿里云 (通义千问)", + type: "openai", + apiHost: "https://dashscope.aliyuncs.com/compatible-mode", + }, + { + id: "deepseek", + name: "DeepSeek", + type: "openai", + apiHost: "https://api.deepseek.com", + }, + { + id: "moonshot", + name: "Moonshot (月之暗面)", + type: "openai", + apiHost: "https://api.moonshot.cn", + }, + { + id: "zhipu", + name: "智谱 AI", + type: "openai", + apiHost: "https://open.bigmodel.cn/api/paas", + }, + { + id: "baichuan", + name: "百川智能", + type: "openai", + apiHost: "https://api.baichuan-ai.com", + }, + { + id: "minimax", + name: "MiniMax", + type: "openai", + apiHost: "https://api.minimax.chat", + }, + { + id: "groq", + name: "Groq", + type: "openai", + apiHost: "https://api.groq.com/openai", + }, + { + id: "together", + name: "Together AI", + type: "openai", + apiHost: "https://api.together.xyz", + }, + { + id: "fireworks", + name: "Fireworks AI", + type: "openai", + apiHost: "https://api.fireworks.ai/inference", + }, + { + id: "perplexity", + name: "Perplexity", + type: "openai", + apiHost: "https://api.perplexity.ai", + }, + { + id: "mistral", + name: "Mistral AI", + type: "openai", + apiHost: "https://api.mistral.ai", + }, + { + id: "cohere", + name: "Cohere", + type: "openai", + apiHost: "https://api.cohere.ai", + }, + { + id: "ollama", + name: "Ollama (本地)", + type: "ollama", + apiHost: "http://localhost:11434", + }, +]; + // ============================================================================ // 类型定义 // ============================================================================ @@ -204,6 +320,88 @@ export const AddCustomProviderModal: React.FC = ({ const [isSubmitting, setIsSubmitting] = useState(false); const [submitError, setSubmitError] = useState(null); + // 厂商搜索状态 + const [providerSearch, setProviderSearch] = useState(""); + const [showProviderDropdown, setShowProviderDropdown] = useState(false); + const [selectedKnownProvider, setSelectedKnownProvider] = + useState(null); + const searchInputRef = useRef(null); + const dropdownRef = useRef(null); + + // 从 model_registry 获取额外的 Provider 信息 + const { groupedByProvider } = useModelRegistry({ autoLoad: true }); + + // 合并已知厂商和 model_registry 中的厂商 + const allProviders = useMemo(() => { + const providers = [...KNOWN_PROVIDERS]; + const existingIds = new Set(providers.map((p) => p.id)); + + // 从 model_registry 添加额外的厂商 + groupedByProvider.forEach((models, providerId) => { + if (!existingIds.has(providerId) && models.length > 0) { + const firstModel = models[0]; + providers.push({ + id: providerId, + name: firstModel.provider_name, + type: "openai" as ProviderType, // 默认使用 OpenAI 兼容 + }); + } + }); + + return providers; + }, [groupedByProvider]); + + // 过滤厂商列表 + const filteredProviders = useMemo(() => { + if (!providerSearch.trim()) { + return allProviders; + } + const query = providerSearch.toLowerCase(); + return allProviders.filter( + (p) => + p.name.toLowerCase().includes(query) || + p.id.toLowerCase().includes(query), + ); + }, [allProviders, providerSearch]); + + // 点击外部关闭下拉框 + useEffect(() => { + const handleClickOutside = (event: MouseEvent) => { + if ( + dropdownRef.current && + !dropdownRef.current.contains(event.target as Node) && + searchInputRef.current && + !searchInputRef.current.contains(event.target as Node) + ) { + setShowProviderDropdown(false); + } + }; + + document.addEventListener("mousedown", handleClickOutside); + return () => document.removeEventListener("mousedown", handleClickOutside); + }, []); + + // 选择已知厂商 + const handleSelectKnownProvider = useCallback((provider: KnownProvider) => { + setSelectedKnownProvider(provider); + setProviderSearch(provider.name); + setShowProviderDropdown(false); + + // 自动填充表单 + setFormState((prev) => ({ + ...prev, + name: provider.name, + type: provider.type, + apiHost: provider.apiHost || "", + })); + }, []); + + // 清除选中的厂商 + const handleClearKnownProvider = useCallback(() => { + setSelectedKnownProvider(null); + setProviderSearch(""); + }, []); + // 获取当前类型需要的额外字段 const extraFields = useMemo( () => PROVIDER_TYPE_EXTRA_FIELDS[formState.type] || [], @@ -215,6 +413,9 @@ export const AddCustomProviderModal: React.FC = ({ setFormState(INITIAL_FORM_STATE); setErrors({}); setSubmitError(null); + setProviderSearch(""); + setSelectedKnownProvider(null); + setShowProviderDropdown(false); }, []); // 关闭模态框 @@ -291,6 +492,78 @@ export const AddCustomProviderModal: React.FC = ({ 添加自定义 Provider + {/* 搜索厂商(可选) */} +
+ +
+
+ + { + setProviderSearch(e.target.value); + setShowProviderDropdown(true); + if (selectedKnownProvider) { + setSelectedKnownProvider(null); + } + }} + onFocus={() => setShowProviderDropdown(true)} + placeholder="搜索厂商名称..." + disabled={isSubmitting} + className="pl-10 pr-8" + data-testid="provider-search-input" + /> + {(providerSearch || selectedKnownProvider) && ( + + )} +
+ + {/* 下拉列表 */} + {showProviderDropdown && filteredProviders.length > 0 && ( +
+ {filteredProviders.map((provider) => ( + + ))} +
+ )} +
+

+ 选择已知厂商可自动填充配置,或直接手动填写下方表单 +

+
+ +
+ {/* Provider 名称 */}
); diff --git a/src/components/provider-pool/api-key/index.ts b/src/components/provider-pool/api-key/index.ts index 9bf5e7f34..96542d5e5 100644 --- a/src/components/provider-pool/api-key/index.ts +++ b/src/components/provider-pool/api-key/index.ts @@ -45,3 +45,8 @@ export type { DeleteProviderDialogProps } from "./DeleteProviderDialog"; export { ImportExportDialog } from "./ImportExportDialog"; export type { ImportExportDialogProps } from "./ImportExportDialog"; + +export { ProviderModelList } from "./ProviderModelList"; +export type { ProviderModelListProps } from "./ProviderModelList"; + +export { mapProviderTypeToRegistryId } from "./providerTypeMapping"; diff --git a/src/components/provider-pool/api-key/providerTypeMapping.ts b/src/components/provider-pool/api-key/providerTypeMapping.ts new file mode 100644 index 000000000..7d5e0e4b0 --- /dev/null +++ b/src/components/provider-pool/api-key/providerTypeMapping.ts @@ -0,0 +1,28 @@ +/** + * @file Provider 类型映射工具 + * @description Provider 类型到 model_registry provider_id 的映射 + * @module components/provider-pool/api-key/providerTypeMapping + */ + +/** + * Provider 类型到 model_registry provider_id 的映射 + */ +const PROVIDER_TYPE_TO_REGISTRY_ID: Record = { + anthropic: "anthropic", + openai: "openai", + "openai-response": "openai", + gemini: "google", + "azure-openai": "openai", + vertexai: "google", + "aws-bedrock": "anthropic", + ollama: "ollama", + "new-api": "custom", + gateway: "custom", +}; + +/** + * 将 Provider 类型转换为 model_registry 的 provider_id + */ +export function mapProviderTypeToRegistryId(providerType: string): string { + return PROVIDER_TYPE_TO_REGISTRY_ID[providerType] || providerType; +} diff --git a/src/components/provider-pool/credential-forms/GeminiFormStandalone.tsx b/src/components/provider-pool/credential-forms/GeminiFormStandalone.tsx new file mode 100644 index 000000000..bd5eeb7f1 --- /dev/null +++ b/src/components/provider-pool/credential-forms/GeminiFormStandalone.tsx @@ -0,0 +1,474 @@ +/** + * Gemini 凭证添加表单(自包含版本) + * + * 支持两种认证方式: + * 1. Google OAuth - 使用 Google 账户授权 + * 2. API Key - 使用 Google AI Studio API Key + * + * @module components/provider-pool/credential-forms/GeminiFormStandalone + */ + +import { useState, useCallback, useEffect } from "react"; +import { open } from "@tauri-apps/plugin-dialog"; +import { listen } from "@tauri-apps/api/event"; +import { providerPoolApi } from "@/lib/api/providerPool"; +import { Button } from "@/components/ui/button"; +import { Input } from "@/components/ui/input"; +import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; +import { + Loader2, + Key, + KeyRound, + Copy, + Check, + ExternalLink, + Upload, +} from "lucide-react"; + +type AuthMethod = "oauth" | "api_key"; + +interface GeminiFormStandaloneProps { + /** 添加成功回调 */ + onSuccess: () => void; + /** 取消回调 */ + onCancel?: () => void; + /** 初始名称 */ + initialName?: string; + /** 初始认证方式 */ + initialAuthMethod?: AuthMethod; +} + +/** + * 自包含的 Gemini 凭证添加表单 + * + * 内部管理所有状态,只需要提供 onSuccess 和 onCancel 回调 + */ +export function GeminiFormStandalone({ + onSuccess, + onCancel, + initialName = "", + initialAuthMethod = "oauth", +}: GeminiFormStandaloneProps) { + const [authMethod, setAuthMethod] = useState(initialAuthMethod); + const [name, setName] = useState(initialName); + const [loading, setLoading] = useState(false); + const [error, setError] = useState(null); + + // OAuth 状态 + const [authUrl, setAuthUrl] = useState(null); + const [sessionId, setSessionId] = useState(null); + const [authCode, setAuthCode] = useState(""); + const [copied, setCopied] = useState(false); + const [exchanging, setExchanging] = useState(false); + + // 文件导入状态 + const [credsFilePath, setCredsFilePath] = useState(""); + const [projectId, setProjectId] = useState(""); + + // API Key 状态 + const [apiKey, setApiKey] = useState(""); + const [baseUrl, setBaseUrl] = useState(""); + + // 监听后端发送的授权 URL 事件 + useEffect(() => { + let unlisten: (() => void) | undefined; + + const setupListener = async () => { + unlisten = await listen<{ auth_url: string; session_id: string }>( + "gemini-auth-url", + (event) => { + console.log("[Gemini OAuth] 收到授权 URL 事件:", event.payload); + setAuthUrl(event.payload.auth_url); + setSessionId(event.payload.session_id); + }, + ); + }; + + setupListener(); + + return () => { + if (unlisten) unlisten(); + }; + }, []); + + // 获取授权 URL + const handleGetAuthUrl = useCallback(async () => { + setLoading(true); + setError(null); + setAuthUrl(null); + setSessionId(null); + setAuthCode(""); + + try { + await providerPoolApi.getGeminiAuthUrlAndWait(name.trim() || undefined); + } catch (e) { + const errorMsg = e instanceof Error ? e.message : String(e); + if (errorMsg.includes("AUTH_URL:")) { + const urlMatch = errorMsg.match(/AUTH_URL:(.+?)(?:\s|$)/); + if (urlMatch) { + setAuthUrl(urlMatch[1]); + } + } else { + setError(errorMsg); + } + } finally { + setLoading(false); + } + }, [name]); + + // 用 code 交换 token + const handleExchangeCode = useCallback(async () => { + if (!authCode.trim()) { + setError("请输入授权码"); + return; + } + + setExchanging(true); + setError(null); + + try { + await providerPoolApi.exchangeGeminiCode( + authCode.trim(), + sessionId || undefined, + name.trim() || undefined, + ); + onSuccess(); + } catch (e) { + setError(e instanceof Error ? e.message : String(e)); + } finally { + setExchanging(false); + } + }, [authCode, sessionId, name, onSuccess]); + + // 复制 URL + const handleCopyUrl = useCallback(async () => { + if (authUrl) { + await navigator.clipboard.writeText(authUrl); + setCopied(true); + setTimeout(() => setCopied(false), 2000); + } + }, [authUrl]); + + // 选择文件 + const handleSelectFile = useCallback(async () => { + try { + const selected = await open({ + multiple: false, + filters: [{ name: "JSON", extensions: ["json"] }], + }); + if (selected) { + setCredsFilePath(selected as string); + } + } catch (e) { + console.error("Failed to open file dialog:", e); + } + }, []); + + // 文件导入提交 + const handleFileSubmit = useCallback(async () => { + if (!credsFilePath) { + setError("请选择凭证文件"); + return; + } + + setLoading(true); + setError(null); + + try { + await providerPoolApi.addGeminiOAuth( + credsFilePath, + projectId.trim() || undefined, + name.trim() || undefined, + ); + onSuccess(); + } catch (e) { + setError(e instanceof Error ? e.message : String(e)); + } finally { + setLoading(false); + } + }, [credsFilePath, projectId, name, onSuccess]); + + // API Key 提交 + const handleApiKeySubmit = useCallback(async () => { + if (!apiKey.trim()) { + setError("请输入 API Key"); + return; + } + + setLoading(true); + setError(null); + + try { + await providerPoolApi.addGeminiApiKey( + apiKey.trim(), + baseUrl.trim() || undefined, + undefined, + name.trim() || undefined, + ); + onSuccess(); + } catch (e) { + setError(e instanceof Error ? e.message : String(e)); + } finally { + setLoading(false); + } + }, [apiKey, baseUrl, name, onSuccess]); + + return ( +
+ {/* 名称输入 */} +
+ + setName(e.target.value)} + placeholder="给这个凭证起个名字..." + disabled={loading || exchanging} + /> +
+ + {/* 认证方式选择 */} + setAuthMethod(v as AuthMethod)} + > + + + + Google OAuth + + + + API Key + + + + {/* OAuth 认证 */} + +
+

+ 点击下方按钮获取授权 URL,然后复制到浏览器完成 Google 登录。 +

+

+ 授权成功后,复制页面显示的授权码粘贴到下方输入框。 +

+
+ + {!authUrl ? ( + + ) : ( +
+
+ 授权 URL + +
+
+

+ {authUrl.length > 100 + ? `${authUrl.slice(0, 100)}...` + : authUrl} +

+
+ +
+ + setAuthCode(e.target.value)} + placeholder="粘贴浏览器页面显示的授权码..." + /> +

+ 在浏览器中完成授权后,复制页面显示的授权码 +

+
+ + +
+ )} + + {/* 文件导入选项 */} +
+

+ 或者导入已有的凭证文件: +

+
+
+ setCredsFilePath(e.target.value)} + placeholder="选择 oauth_creds.json..." + className="flex-1" + /> + +
+ setProjectId(e.target.value)} + placeholder="Project ID (可选)" + /> + +
+
+
+ + {/* API Key 认证 */} + +
+

+ 使用 Google AI Studio 的 API Key 进行认证。 +

+

+ 从{" "} + + Google AI Studio + {" "} + 获取 API Key。 +

+
+ +
+
+ + setApiKey(e.target.value)} + placeholder="AIzaSy..." + /> +
+ +
+ + setBaseUrl(e.target.value)} + placeholder="https://generativelanguage.googleapis.com" + /> +

+ 留空使用官方 API +

+
+
+
+
+ + {/* 错误提示 */} + {error && ( +
+ {error} +
+ )} + + {/* 按钮区域 */} +
+ {onCancel && ( + + )} + {authMethod === "api_key" && ( + + )} +
+
+ ); +} + +export default GeminiFormStandalone; diff --git a/src/components/provider-pool/index.ts b/src/components/provider-pool/index.ts index 96af6d860..183a3e000 100644 --- a/src/components/provider-pool/index.ts +++ b/src/components/provider-pool/index.ts @@ -9,3 +9,4 @@ export { IFlowSection } from "./IFlowSection"; export { AmpConfigSection } from "./AmpConfigSection"; export { UsageDisplay } from "./UsageDisplay"; export { OAuthPluginTab } from "./OAuthPluginTab"; +export { ModelRegistryTab } from "./ModelRegistryTab"; diff --git a/src/hooks/useModelRegistry.ts b/src/hooks/useModelRegistry.ts index e2737354b..308cafc1b 100644 --- a/src/hooks/useModelRegistry.ts +++ b/src/hooks/useModelRegistry.ts @@ -55,7 +55,7 @@ interface UseModelRegistryReturn { */ function sortModels( models: EnhancedModelMetadata[], - preferences: Map + preferences: Map, ): EnhancedModelMetadata[] { return [...models].sort((a, b) => { const prefA = preferences.get(a.id); @@ -88,7 +88,7 @@ function sortModels( */ function fuzzySearch( models: EnhancedModelMetadata[], - query: string + query: string, ): EnhancedModelMetadata[] { if (!query.trim()) { return models; @@ -140,7 +140,7 @@ function fuzzySearch( } export function useModelRegistry( - options: UseModelRegistryOptions = {} + options: UseModelRegistryOptions = {}, ): UseModelRegistryReturn { const { autoLoad = true, @@ -206,7 +206,7 @@ export function useModelRegistry( // 等级过滤 if (tierFilter && tierFilter.length > 0) { filtered = filtered.filter((m) => - tierFilter.includes(m.tier as ModelTier) + tierFilter.includes(m.tier as ModelTier), ); } @@ -227,7 +227,7 @@ export function useModelRegistry( (query: string): EnhancedModelMetadata[] => { return fuzzySearch(models, query); }, - [models] + [models], ); // 切换收藏 @@ -283,7 +283,7 @@ export function useModelRegistry( (modelId: string) => { return allModels.find((m) => m.id === modelId); }, - [allModels] + [allModels], ); // 按 Provider 分组 diff --git a/src/lib/api/modelRegistry.ts b/src/lib/api/modelRegistry.ts index f74c24b94..a4e3f2aab 100644 --- a/src/lib/api/modelRegistry.ts +++ b/src/lib/api/modelRegistry.ts @@ -33,7 +33,7 @@ export async function refreshModelRegistry(): Promise { */ export async function searchModels( query: string, - limit?: number + limit?: number, ): Promise { return invoke("search_models", { query, limit }); } @@ -82,7 +82,7 @@ export async function getModelSyncState(): Promise { * @param providerId Provider ID */ export async function getModelsForProvider( - providerId: string + providerId: string, ): Promise { return invoke("get_models_for_provider", { providerId }); } @@ -92,7 +92,7 @@ export async function getModelsForProvider( * @param tier 服务等级 */ export async function getModelsByTier( - tier: ModelTier + tier: ModelTier, ): Promise { return invoke("get_models_by_tier", { tier }); } diff --git a/src/lib/plugin-components/index.ts b/src/lib/plugin-components/index.ts index 06aeefa83..2b2f65b08 100644 --- a/src/lib/plugin-components/index.ts +++ b/src/lib/plugin-components/index.ts @@ -95,6 +95,9 @@ export { KiroForm } from "@/components/provider-pool/credential-forms/KiroForm"; // Antigravity 凭证表单(自包含版本,适合插件使用) export { AntigravityFormStandalone } from "@/components/provider-pool/credential-forms/AntigravityFormStandalone"; +// Gemini 凭证表单(自包含版本,适合插件使用) +export { GeminiFormStandalone } from "@/components/provider-pool/credential-forms/GeminiFormStandalone"; + // 浏览器模式选择器 export { BrowserModeSelector, @@ -162,6 +165,7 @@ export { ExternalLink, // 凭证相关 Key, + KeyRound, Lock, Unlock, Shield, diff --git a/src/lib/plugin-loader/PluginUIRenderer.tsx b/src/lib/plugin-loader/PluginUIRenderer.tsx index e7bd33356..e3c000101 100644 --- a/src/lib/plugin-loader/PluginUIRenderer.tsx +++ b/src/lib/plugin-loader/PluginUIRenderer.tsx @@ -118,23 +118,30 @@ export function PluginUIRenderer({ // 错误 if (error) { // 检查是否是文件不存在的错误 - const isFileNotFound = error.includes("读取插件 UI 文件失败") || - error.includes("No such file") || - error.includes("not found") || - error.includes("没有找到有效的组件导出") || - error.includes("插件加载失败"); - + const isFileNotFound = + error.includes("读取插件 UI 文件失败") || + error.includes("No such file") || + error.includes("not found") || + error.includes("没有找到有效的组件导出") || + error.includes("插件加载失败"); + if (isFileNotFound) { // UI 文件不存在时显示友好提示 - return fallback ? <>{fallback} : ( -
+ return fallback ? ( + <>{fallback} + ) : ( +

该插件暂无 UI 界面

-

请通过命令行或 API 使用此插件

+

+ 请通过命令行或 API 使用此插件 +

); } - + return (