chore: bump version to v0.32.0

🤖 Generated with [Claude Code](https://claude.com/claude-code)

Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com>
This commit is contained in:
coso
2026-01-06 22:25:46 +08:00
co-authored by Claude Opus 4.5
parent 717abfd732
commit d29efe737c
36 changed files with 1589 additions and 2518 deletions
+1 -1
View File
@@ -1,7 +1,7 @@
{
"name": "proxycast",
"private": true,
"version": "0.31.0",
"version": "0.32.0",
"type": "module",
"repository": {
"type": "git",
+324 -197
View File
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -1,6 +1,6 @@
[package]
name = "proxycast"
version = "0.31.0"
version = "0.32.0"
description = "AI API Proxy Desktop App"
authors = ["you"]
edition = "2021"
+121 -6
View File
@@ -44,6 +44,8 @@ pub struct NativeAgent {
provider_type: ProviderType,
/// 协议处理器
protocol: Box<dyn Protocol>,
/// Provider ID,用于自定义 Provider 路由(如 "moonshot")
provider_id: Option<String>,
}
impl NativeAgent {
@@ -51,6 +53,7 @@ impl NativeAgent {
base_url: String,
api_key: String,
provider_type: ProviderType,
provider_id: Option<String>,
) -> Result<Self, String> {
let client = Client::builder()
.timeout(Duration::from_secs(300))
@@ -61,21 +64,27 @@ impl NativeAgent {
let protocol = create_protocol(provider_type);
// 保存原始 base_url,provider_id 将在构建请求 URL 时使用
let effective_base_url = base_url.clone();
info!(
"[NativeAgent] 创建 Agent: base_url={}, provider={:?}, protocol_endpoint={}",
"[NativeAgent] 创建 Agent: base_url={}, effective_base_url={}, provider={:?}, provider_id={:?}, protocol_endpoint={}",
base_url,
effective_base_url,
provider_type,
provider_id,
protocol.endpoint()
);
Ok(Self {
client,
base_url,
base_url: effective_base_url,
api_key,
sessions: Arc::new(RwLock::new(HashMap::new())),
config: AgentConfig::default(),
provider_type,
protocol,
provider_id,
})
}
@@ -89,6 +98,57 @@ impl NativeAgent {
self
}
/// 获取 API 请求的有效 base_url
///
/// 对于自定义 Provider(如 moonshot),返回 `{base_url}/api/provider/{provider_id}`
/// 对于内置 Provider,返回原始 base_url
fn get_effective_base_url(&self) -> String {
if let Some(ref pid) = self.provider_id {
// 检查 provider_id 是否是已知的内置类型
let is_builtin = matches!(
pid.to_lowercase().as_str(),
"openai"
| "claude"
| "anthropic"
| "gemini"
| "kiro"
| "qwen"
| "codex"
| "antigravity"
| "iflow"
);
if is_builtin {
self.base_url.clone()
} else {
// 自定义 Provider,使用 provider 特定路由
// 例如:http://127.0.0.1:8999/api/provider/moonshot
format!("{}/api/provider/{}", self.base_url, pid)
}
} else {
self.base_url.clone()
}
}
/// 检查是否是自定义 Provider
fn is_custom_provider(&self) -> bool {
if let Some(ref pid) = self.provider_id {
!matches!(
pid.to_lowercase().as_str(),
"openai"
| "claude"
| "anthropic"
| "gemini"
| "kiro"
| "qwen"
| "codex"
| "antigravity"
| "iflow"
)
} else {
false
}
}
/// 发送聊天请求(非流式,用于简单场景)
pub async fn chat(&self, request: NativeChatRequest) -> Result<NativeChatResponse, String> {
let model = request.model.unwrap_or_else(|| self.config.model.clone());
@@ -126,7 +186,12 @@ impl NativeAgent {
reasoning_effort: None,
};
let url = format!("{}/v1/chat/completions", self.base_url);
// 对于自定义 Provider,使用 provider 特定路由
let url = if self.is_custom_provider() {
format!("{}/chat/completions", self.get_effective_base_url())
} else {
format!("{}/v1/chat/completions", self.base_url)
};
let response = self
.client
@@ -241,11 +306,13 @@ impl NativeAgent {
};
// 使用协议策略发送请求
// 对于自定义 Provider,使用 provider 特定路由
let effective_base_url = self.get_effective_base_url();
let result = self
.protocol
.chat_stream(
&self.client,
&self.base_url,
&effective_base_url,
&self.api_key,
&history,
&request.message,
@@ -408,11 +475,13 @@ impl NativeAgent {
};
// 使用协议策略继续对话
// 对于自定义 Provider,使用 provider 特定路由
let effective_base_url = self.get_effective_base_url();
let result = self
.protocol
.chat_stream_continue(
&self.client,
&self.base_url,
&effective_base_url,
&self.api_key,
&session.messages,
&model,
@@ -676,8 +745,45 @@ impl NativeAgentState {
base_url: String,
api_key: String,
provider_type: ProviderType,
provider_id: Option<String>,
) -> Result<(), String> {
let agent = NativeAgent::new(base_url, api_key, provider_type)?;
let agent = NativeAgent::new(base_url, api_key, provider_type, provider_id)?;
*self.agent.write() = Some(agent);
Ok(())
}
/// 使用配置初始化 Agent
///
/// 从 NativeAgentConfig 加载系统提示词等配置
pub fn init_with_config(
&self,
base_url: String,
api_key: String,
provider_type: ProviderType,
provider_id: Option<String>,
agent_config: &crate::config::NativeAgentConfig,
) -> Result<(), String> {
let mut agent = NativeAgent::new(base_url, api_key, provider_type, provider_id)?;
// 从配置加载系统提示词
let system_prompt = agent_config.get_effective_system_prompt().or_else(|| {
// 如果配置启用了默认提示词,使用内置默认值
if agent_config.use_default_system_prompt {
Some(super::types::DEFAULT_SYSTEM_PROMPT.to_string())
} else {
None
}
});
if let Some(prompt) = system_prompt {
agent.config.system_prompt = Some(prompt);
}
// 从配置加载其他参数
agent.config.model = agent_config.default_model.clone();
agent.config.temperature = Some(agent_config.temperature);
agent.config.max_tokens = Some(agent_config.max_tokens);
*self.agent.write() = Some(agent);
Ok(())
}
@@ -691,6 +797,14 @@ impl NativeAgentState {
self.agent.read().as_ref().map(|a| a.provider_type)
}
/// 获取当前 Agent 的 provider ID
pub fn get_provider_id(&self) -> Option<String> {
self.agent
.read()
.as_ref()
.and_then(|a| a.provider_id.clone())
}
pub fn reset(&self) {
*self.agent.write() = None;
}
@@ -724,6 +838,7 @@ impl NativeAgentState {
config: agent.config.clone(),
provider_type: agent.provider_type,
protocol,
provider_id: agent.provider_id.clone(),
})
}
+88
View File
@@ -176,6 +176,52 @@ impl OpenAIProtocol {
);
buffer.push_str(&text);
// 检查是否是非流式响应(直接返回完整 JSON)
// 非流式响应以 { 开头,不是 SSE 格式
if buffer.trim().starts_with('{') && !buffer.contains("data: ") {
// 尝试解析为完整的 ChatCompletionResponse
if let Ok(response) = serde_json::from_str::<
crate::models::openai::ChatCompletionResponse,
>(&buffer)
{
eprintln!("[OpenAIProtocol] 检测到非流式响应,直接解析");
let content = response
.choices
.first()
.and_then(|c| c.message.content.clone())
.unwrap_or_default();
// 发送完整内容作为 TextDelta
if !content.is_empty() {
let _ = tx
.send(StreamEvent::TextDelta {
text: content.clone(),
})
.await;
}
let usage = Some(crate::agent::types::TokenUsage {
input_tokens: response.usage.prompt_tokens,
output_tokens: response.usage.completion_tokens,
});
if send_done {
let _ = tx
.send(StreamEvent::Done {
usage: usage.clone(),
})
.await;
}
return Ok(StreamResult {
content,
tool_calls: None,
usage,
});
}
}
// 处理完整的 SSE 事件(以 \n\n 分隔)
while let Some(pos) = buffer.find("\n\n") {
let event = buffer[..pos].to_string();
@@ -233,6 +279,48 @@ impl OpenAIProtocol {
}
// 流正常结束但没有收到 [DONE]
// 检查 buffer 中是否还有未处理的非流式响应
if !buffer.trim().is_empty() && buffer.trim().starts_with('{') {
if let Ok(response) =
serde_json::from_str::<crate::models::openai::ChatCompletionResponse>(&buffer)
{
eprintln!("[OpenAIProtocol] 流结束时检测到非流式响应");
let content = response
.choices
.first()
.and_then(|c| c.message.content.clone())
.unwrap_or_default();
if !content.is_empty() {
let _ = tx
.send(StreamEvent::TextDelta {
text: content.clone(),
})
.await;
}
let usage = Some(crate::agent::types::TokenUsage {
input_tokens: response.usage.prompt_tokens,
output_tokens: response.usage.completion_tokens,
});
if send_done {
let _ = tx
.send(StreamEvent::Done {
usage: usage.clone(),
})
.await;
}
return Ok(StreamResult {
content,
tool_calls: None,
usage,
});
}
}
let full_content = parser.get_full_content();
let tool_calls = if parser.has_tool_calls() {
Some(parser.finalize_tool_calls())
+62 -1
View File
@@ -213,7 +213,7 @@ impl Default for AgentConfig {
fn default() -> Self {
Self {
model: "claude-sonnet-4-20250514".to_string(),
system_prompt: None,
system_prompt: Some(DEFAULT_SYSTEM_PROMPT.to_string()),
temperature: Some(0.7),
max_tokens: Some(4096),
tools: Vec::new(),
@@ -221,6 +221,67 @@ impl Default for AgentConfig {
}
}
/// 默认系统提示词
///
/// 参考 Manus Agent 的模块化设计,使用结构化的提示词组织
/// 支持通过配置文件覆盖
pub const DEFAULT_SYSTEM_PROMPT: &str = r#"你是 ProxyCast 内置的 AI 助手。
<identity>
- 你是一个友好、专业的 AI 助手
- 擅长编程、文件操作和系统任务
- 使用中文与用户交流
</identity>
<core_principles>
1. **自然交流优先**:问候、闲聊、问答类对话,直接用文字回复
2. **显式授权操作**:只有当用户明确提供路径或命令时,才能执行工具
3. **不主动探索**:不要未经请求就读取文件或执行命令
</core_principles>
<tool_use_rules>
## 何时使用工具
✅ **使用工具的情况**:
- 用户明确提供了文件路径(如 "读取 /path/to/file")
- 用户明确要求执行命令(如 "运行 npm install")
- 用户要求创建或修改文件
❌ **禁止使用工具的情况**:
- 用户说 "你好"、"嗨"、"hello" 等问候语
- 用户进行闲聊或一般性提问
- 用户没有提供具体路径时猜测路径
- 为了 "了解环境" 或 "打招呼" 而读取文件
## 可用工具
- **read_file**:读取用户指定的文件或目录内容
- **write_file**:创建或覆盖用户指定的文件
- **edit_file**:修改用户指定文件的特定内容
- **bash**:执行用户要求的 shell 命令
</tool_use_rules>
<response_examples>
## 正确示例
用户: "你好"
助手: "你好!有什么我可以帮助你的吗?"
(直接文字回复,不调用任何工具)
用户: "看看 /tmp/test.txt"
助手: 调用 read_file 工具读取 /tmp/test.txt
用户: "帮我列出当前目录"
助手: "请告诉我你想查看哪个目录?"
(询问具体路径,不要猜测)
</response_examples>
<output_format>
- 使用 Markdown 格式
- 回复简洁明了
- 使用中文
</output_format>"#;
/// 聊天请求
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct NativeChatRequest {
+18 -75
View File
@@ -3,6 +3,7 @@
//! 包含 API 测试、模型列表和兼容性检查命令。
use crate::app::types::{AppState, LogState, ProviderType};
use crate::commands::model_registry_cmd::ModelRegistryState;
/// 测试结果
#[derive(serde::Serialize)]
@@ -279,82 +280,24 @@ pub async fn check_api_compatibility(
/// 获取可用模型列表
#[tauri::command]
pub async fn get_available_models() -> Result<Vec<ModelInfo>, String> {
Ok(vec![
// Kiro/Claude models
ModelInfo {
id: "claude-sonnet-4-5".to_string(),
pub async fn get_available_models(
state: tauri::State<'_, ModelRegistryState>,
) -> Result<Vec<ModelInfo>, String> {
let guard = state.read().await;
let service = guard
.as_ref()
.ok_or_else(|| "模型注册服务未初始化".to_string())?;
let models = service.get_all_models().await;
Ok(models
.into_iter()
.map(|m| ModelInfo {
id: m.id,
object: "model".to_string(),
owned_by: "anthropic".to_string(),
},
ModelInfo {
id: "claude-sonnet-4-5-20250514".to_string(),
object: "model".to_string(),
owned_by: "anthropic".to_string(),
},
ModelInfo {
id: "claude-sonnet-4-5-20250929".to_string(),
object: "model".to_string(),
owned_by: "anthropic".to_string(),
},
ModelInfo {
id: "claude-3-7-sonnet-20250219".to_string(),
object: "model".to_string(),
owned_by: "anthropic".to_string(),
},
ModelInfo {
id: "claude-3-5-sonnet-latest".to_string(),
object: "model".to_string(),
owned_by: "anthropic".to_string(),
},
ModelInfo {
id: "claude-opus-4-5-20250514".to_string(),
object: "model".to_string(),
owned_by: "anthropic".to_string(),
},
ModelInfo {
id: "claude-haiku-4-5-20250514".to_string(),
object: "model".to_string(),
owned_by: "anthropic".to_string(),
},
// Gemini models
ModelInfo {
id: "gemini-2.5-flash".to_string(),
object: "model".to_string(),
owned_by: "google".to_string(),
},
ModelInfo {
id: "gemini-2.5-flash-lite".to_string(),
object: "model".to_string(),
owned_by: "google".to_string(),
},
ModelInfo {
id: "gemini-2.5-pro".to_string(),
object: "model".to_string(),
owned_by: "google".to_string(),
},
ModelInfo {
id: "gemini-2.5-pro-preview-06-05".to_string(),
object: "model".to_string(),
owned_by: "google".to_string(),
},
ModelInfo {
id: "gemini-3-pro-preview".to_string(),
object: "model".to_string(),
owned_by: "google".to_string(),
},
// Qwen models
ModelInfo {
id: "qwen3-coder-plus".to_string(),
object: "model".to_string(),
owned_by: "alibaba".to_string(),
},
ModelInfo {
id: "qwen3-coder-flash".to_string(),
object: "model".to_string(),
owned_by: "alibaba".to_string(),
},
])
owned_by: m.provider_id,
})
.collect())
}
/// 测试 API
+6 -10
View File
@@ -2,7 +2,7 @@
//!
//! 包含配置读取、保存、Provider 设置等命令。
use crate::app::types::{AppState, LogState, ProviderType};
use crate::app::types::{AppState, LogState};
use crate::config;
/// 获取配置
@@ -51,8 +51,8 @@ pub async fn set_default_provider(
logs: tauri::State<'_, LogState>,
provider: String,
) -> Result<String, String> {
// 使用枚举验证 provider
let provider_type: ProviderType = provider.parse().map_err(|e: String| e)?;
// 允许任意 Provider ID(包括自定义 Provider 的 UUID)
// 不再强制验证为已知的 ProviderType
let mut s = state.write().await;
s.config.default_provider = provider.clone();
@@ -66,7 +66,7 @@ pub async fn set_default_provider(
config::save_config(&s.config).map_err(|e| e.to_string())?;
logs.write()
.await
.add("info", &format!("默认 Provider 已切换为: {provider_type}"));
.add("info", &format!("默认 Provider 已切换为: {provider}"));
Ok(provider)
}
@@ -95,12 +95,8 @@ pub async fn set_endpoint_provider(
endpoint: String,
provider: Option<String>,
) -> Result<String, String> {
// 验证 provider(如果提供)
if let Some(ref p) = provider {
if !p.is_empty() {
let _: ProviderType = p.parse().map_err(|e: String| e)?;
}
}
// 允许任意 Provider ID(包括自定义 Provider 的 UUID)
// 不再强制验证为已知的 ProviderType
let mut s = state.write().await;
-5
View File
@@ -687,9 +687,6 @@ pub fn run() {
commands::route_cmd::get_available_routes,
commands::route_cmd::get_route_curl_examples,
// Router config commands
commands::router_cmd::get_model_aliases,
commands::router_cmd::add_model_alias,
commands::router_cmd::remove_model_alias,
commands::router_cmd::get_routing_rules,
commands::router_cmd::add_routing_rule,
commands::router_cmd::remove_routing_rule,
@@ -698,8 +695,6 @@ pub fn run() {
commands::router_cmd::add_exclusion,
commands::router_cmd::remove_exclusion,
commands::router_cmd::set_router_default_provider,
commands::router_cmd::get_recommended_presets,
commands::router_cmd::apply_recommended_preset,
commands::router_cmd::clear_all_routing_config,
// Resilience config commands
commands::resilience_cmd::get_retry_config,
+8 -3
View File
@@ -52,7 +52,12 @@ pub async fn agent_start_process(
let base_url = format!("http://127.0.0.1:{}", port);
let provider_type = ProviderType::from_str(&default_provider);
agent_state.init(base_url.clone(), api_key, provider_type)?;
agent_state.init(
base_url.clone(),
api_key,
provider_type,
Some(default_provider),
)?;
Ok(AgentProcessStatus {
running: true,
@@ -137,7 +142,7 @@ pub async fn agent_create_session(
let api_key = api_key.ok_or_else(|| "未配置 API Key".to_string())?;
let base_url = format!("http://127.0.0.1:{}", port);
let provider_type = ProviderType::from_str(&default_provider);
agent_state.init(base_url, api_key, provider_type)?;
agent_state.init(base_url, api_key, provider_type, Some(default_provider))?;
}
// 构建包含 Skills 的 System Prompt
@@ -243,7 +248,7 @@ pub async fn agent_send_message(
let api_key = api_key.ok_or_else(|| "未配置 API Key".to_string())?;
let base_url = format!("http://127.0.0.1:{}", port);
let provider_type = ProviderType::from_str(&default_provider);
agent_state.init(base_url, api_key, provider_type)?;
agent_state.init(base_url, api_key, provider_type, Some(default_provider))?;
}
// 根据启用的模式构建最终消息
+1 -1
View File
@@ -34,7 +34,7 @@ pub async fn refresh_model_registry(state: State<'_, ModelRegistryState>) -> Res
.as_ref()
.ok_or_else(|| "模型注册服务未初始化".to_string())?;
service.refresh_from_models_dev().await
service.refresh_from_repo().await
}
/// 搜索模型
+43 -12
View File
@@ -6,6 +6,8 @@ use crate::agent::{
AgentSession, ImageData, NativeAgentState, NativeChatRequest, NativeChatResponse, ProviderType,
StreamEvent, ToolLoopEngine,
};
use crate::database::dao::api_key_provider::ApiKeyProviderDao;
use crate::database::DbConnection;
use crate::AppState;
use serde::{Deserialize, Serialize};
use tauri::{Emitter, State};
@@ -24,13 +26,14 @@ pub async fn native_agent_init(
) -> Result<NativeAgentStatus, String> {
tracing::info!("[NativeAgent] 初始化 Agent");
let (port, api_key, running, default_provider) = {
let (port, api_key, running, default_provider, agent_config) = {
let state = app_state.read().await;
(
state.config.server.port,
state.running_api_key.clone(),
state.running,
state.config.routing.default_provider.clone(),
state.config.agent.clone(),
)
};
@@ -44,12 +47,20 @@ pub async fn native_agent_init(
let provider_type = ProviderType::from_str(&default_provider);
tracing::info!(
"[NativeAgent] 初始化 Agent: base_url={}, provider={:?}",
"[NativeAgent] 初始化 Agent: base_url={}, provider={:?}, use_default_prompt={}",
base_url,
provider_type
provider_type,
agent_config.use_default_system_prompt
);
agent_state.init(base_url.clone(), api_key, provider_type)?;
// 使用带配置的初始化方法
agent_state.init_with_config(
base_url.clone(),
api_key,
provider_type,
Some(default_provider),
&agent_config,
)?;
tracing::info!("[NativeAgent] Agent 初始化成功: {}", base_url);
@@ -115,7 +126,7 @@ pub async fn native_agent_chat(
let api_key = api_key.ok_or_else(|| "未配置 API Key".to_string())?;
let base_url = format!("http://127.0.0.1:{}", port);
let provider_type = ProviderType::from_str(&default_provider);
agent_state.init(base_url, api_key, provider_type)?;
agent_state.init(base_url, api_key, provider_type, Some(default_provider))?;
}
let request = NativeChatRequest {
@@ -142,6 +153,7 @@ pub async fn native_agent_chat_stream(
app_handle: tauri::AppHandle,
agent_state: State<'_, NativeAgentState>,
app_state: State<'_, AppState>,
db: State<'_, DbConnection>,
message: String,
event_name: String,
session_id: Option<String>,
@@ -177,7 +189,25 @@ pub async fn native_agent_chat_stream(
// 使用前端传递的 provider,如果没有则使用默认值
let provider_str = provider.unwrap_or(default_provider);
let provider_type = ProviderType::from_str(&provider_str);
// 尝试从数据库查询 Provider 的类型(用于确定协议)
let provider_type = {
let conn = db.lock().map_err(|e| format!("数据库锁定失败: {}", e))?;
if let Ok(Some(api_provider)) = ApiKeyProviderDao::get_provider_by_id(&conn, &provider_str)
{
// 根据 API Key Provider 的 type 确定协议
let api_type = api_provider.provider_type.to_string();
tracing::info!(
"[NativeAgent] 从数据库获取 Provider 类型: {} -> {}",
provider_str,
api_type
);
ProviderType::from_str(&api_type)
} else {
// 数据库中没有找到,使用默认解析
ProviderType::from_str(&provider_str)
}
};
tracing::info!(
"[NativeAgent] 使用 provider: {:?} (原始值: {})",
@@ -186,15 +216,16 @@ pub async fn native_agent_chat_stream(
);
// 如果 Agent 未初始化,或者 provider 发生变化,重新初始化
// 使用 provider_str 而不是 provider_type 来判断,因为自定义 Provider 的 type 都是 OpenAI
let need_reinit = if !agent_state.is_initialized() {
tracing::info!("[NativeAgent] Agent 未初始化,需要初始化");
true
} else if let Some(current_provider) = agent_state.get_provider_type() {
if current_provider != provider_type {
} else if let Some(current_provider_id) = agent_state.get_provider_id() {
if current_provider_id != provider_str {
tracing::info!(
"[NativeAgent] Provider 发生变化: {:?} -> {:?},需要重新初始化",
current_provider,
provider_type
"[NativeAgent] Provider 发生变化: {} -> {},需要重新初始化",
current_provider_id,
provider_str
);
true
} else {
@@ -206,7 +237,7 @@ pub async fn native_agent_chat_stream(
if need_reinit {
let base_url = format!("http://127.0.0.1:{}", port);
agent_state.init(base_url, api_key, provider_type)?;
agent_state.init(base_url, api_key, provider_type, Some(provider_str.clone()))?;
}
// 获取工具注册表(用于创建 ToolLoopEngine)
+12 -3
View File
@@ -925,9 +925,18 @@ pub async fn plugin_config_set(
pub async fn read_plugin_ui_file(path: String) -> Result<String, String> {
use std::fs;
// 安全检查:确保路径在插件目录内
let path = std::path::PathBuf::from(&path);
// 展开 ~ 为用户主目录
let expanded_path = if path.starts_with("~/") {
if let Some(home) = dirs::home_dir() {
home.join(&path[2..])
} else {
std::path::PathBuf::from(&path)
}
} else {
std::path::PathBuf::from(&path)
};
// 读取文件内容
fs::read_to_string(&path).map_err(|e| format!("读取插件 UI 文件失败: {}", e))
fs::read_to_string(&expanded_path)
.map_err(|e| format!("读取插件 UI 文件失败: {} (路径: {:?})", e, expanded_path))
}
-440
View File
@@ -6,13 +6,6 @@ use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
/// 模型别名
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelAlias {
pub alias: String,
pub actual: String,
}
/// 路由规则
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RoutingRuleDto {
@@ -24,7 +17,6 @@ pub struct RoutingRuleDto {
/// 路由配置状态
pub struct RouterConfigState {
pub aliases: Arc<RwLock<HashMap<String, String>>>,
pub rules: Arc<RwLock<Vec<RoutingRuleDto>>>,
pub exclusions: Arc<RwLock<HashMap<ProviderType, Vec<String>>>>,
}
@@ -32,51 +24,12 @@ pub struct RouterConfigState {
impl Default for RouterConfigState {
fn default() -> Self {
Self {
aliases: Arc::new(RwLock::new(HashMap::new())),
rules: Arc::new(RwLock::new(Vec::new())),
exclusions: Arc::new(RwLock::new(HashMap::new())),
}
}
}
/// 获取所有模型别名
#[tauri::command]
pub async fn get_model_aliases(
state: tauri::State<'_, RouterConfigState>,
) -> Result<Vec<ModelAlias>, String> {
let aliases = state.aliases.read().await;
Ok(aliases
.iter()
.map(|(alias, actual)| ModelAlias {
alias: alias.clone(),
actual: actual.clone(),
})
.collect())
}
/// 添加模型别名
#[tauri::command]
pub async fn add_model_alias(
state: tauri::State<'_, RouterConfigState>,
alias: String,
actual: String,
) -> Result<(), String> {
let mut aliases = state.aliases.write().await;
aliases.insert(alias, actual);
Ok(())
}
/// 移除模型别名
#[tauri::command]
pub async fn remove_model_alias(
state: tauri::State<'_, RouterConfigState>,
alias: String,
) -> Result<(), String> {
let mut aliases = state.aliases.write().await;
aliases.remove(&alias);
Ok(())
}
/// 获取所有路由规则
#[tauri::command]
pub async fn get_routing_rules(
@@ -178,404 +131,11 @@ pub async fn set_router_default_provider(_provider: ProviderType) -> Result<(),
Ok(())
}
/// 推荐配置预设
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RecommendedPreset {
pub id: String,
pub name: String,
pub description: String,
pub aliases: Vec<ModelAlias>,
pub rules: Vec<RoutingRuleDto>,
/// 客户端路由配置
#[serde(default, skip_serializing_if = "Option::is_none")]
pub endpoint_providers: Option<EndpointProvidersConfigDto>,
}
/// 端点 Provider 配置 DTO
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct EndpointProvidersConfigDto {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cursor: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub claude_code: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub codex: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub windsurf: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub kiro: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub other: Option<String>,
}
/// 获取推荐配置列表
#[tauri::command]
pub async fn get_recommended_presets() -> Result<Vec<RecommendedPreset>, String> {
Ok(vec![
RecommendedPreset {
id: "claude-optimized".to_string(),
name: "Claude 优化配置".to_string(),
description: "将所有 Claude 模型请求路由到 Kiro,适合主要使用 Claude 的用户"
.to_string(),
aliases: vec![
// Claude 4.5 系列 (最新)
ModelAlias {
alias: "claude".to_string(),
actual: "claude-opus-4-5".to_string(),
},
ModelAlias {
alias: "opus".to_string(),
actual: "claude-opus-4-5".to_string(),
},
ModelAlias {
alias: "sonnet".to_string(),
actual: "claude-sonnet-4-5".to_string(),
},
ModelAlias {
alias: "haiku".to_string(),
actual: "claude-haiku-4-5".to_string(),
},
// Claude 4 系列
ModelAlias {
alias: "opus-4".to_string(),
actual: "claude-opus-4".to_string(),
},
ModelAlias {
alias: "sonnet-4".to_string(),
actual: "claude-sonnet-4".to_string(),
},
// Claude 3.7/3.5 系列 (旧版)
ModelAlias {
alias: "sonnet-3.7".to_string(),
actual: "claude-3-7-sonnet-latest".to_string(),
},
ModelAlias {
alias: "sonnet-3.5".to_string(),
actual: "claude-3-5-sonnet-latest".to_string(),
},
],
rules: vec![
RoutingRuleDto {
pattern: "claude-*".to_string(),
target_provider: ProviderType::Kiro,
priority: 1,
enabled: true,
},
RoutingRuleDto {
pattern: "*sonnet*".to_string(),
target_provider: ProviderType::Kiro,
priority: 2,
enabled: true,
},
RoutingRuleDto {
pattern: "*opus*".to_string(),
target_provider: ProviderType::Kiro,
priority: 2,
enabled: true,
},
RoutingRuleDto {
pattern: "*haiku*".to_string(),
target_provider: ProviderType::Kiro,
priority: 2,
enabled: true,
},
],
endpoint_providers: None,
},
RecommendedPreset {
id: "gemini-optimized".to_string(),
name: "Gemini 优化配置".to_string(),
description: "将 Gemini 模型请求路由到 Gemini Provider,适合主要使用 Google AI 的用户"
.to_string(),
aliases: vec![
// Gemini 3 系列 (最新)
ModelAlias {
alias: "gemini".to_string(),
actual: "gemini-3-pro".to_string(),
},
ModelAlias {
alias: "gemini-pro".to_string(),
actual: "gemini-3-pro".to_string(),
},
ModelAlias {
alias: "gemini-3".to_string(),
actual: "gemini-3-pro".to_string(),
},
// Gemini 2.5 系列
ModelAlias {
alias: "flash".to_string(),
actual: "gemini-2.5-flash".to_string(),
},
ModelAlias {
alias: "flash-lite".to_string(),
actual: "gemini-2.5-flash-lite".to_string(),
},
ModelAlias {
alias: "gemini-2.5".to_string(),
actual: "gemini-2.5-pro".to_string(),
},
],
rules: vec![
RoutingRuleDto {
pattern: "gemini-*".to_string(),
target_provider: ProviderType::Gemini,
priority: 1,
enabled: true,
},
RoutingRuleDto {
pattern: "*flash*".to_string(),
target_provider: ProviderType::Gemini,
priority: 2,
enabled: true,
},
],
endpoint_providers: None,
},
RecommendedPreset {
id: "multi-provider".to_string(),
name: "多 Provider 均衡配置".to_string(),
description: "根据模型名称自动路由到对应的 Provider,适合同时使用多个 AI 服务的用户"
.to_string(),
aliases: vec![
// Claude (最新)
ModelAlias {
alias: "claude".to_string(),
actual: "claude-opus-4-5".to_string(),
},
ModelAlias {
alias: "sonnet".to_string(),
actual: "claude-sonnet-4-5".to_string(),
},
// Gemini (最新)
ModelAlias {
alias: "gemini".to_string(),
actual: "gemini-3-pro".to_string(),
},
ModelAlias {
alias: "flash".to_string(),
actual: "gemini-2.5-flash".to_string(),
},
// Qwen
ModelAlias {
alias: "qwen".to_string(),
actual: "qwen3-coder-plus".to_string(),
},
// OpenAI (最新)
ModelAlias {
alias: "gpt".to_string(),
actual: "gpt-5.2".to_string(),
},
ModelAlias {
alias: "gpt-5".to_string(),
actual: "gpt-5.2".to_string(),
},
ModelAlias {
alias: "gpt-4".to_string(),
actual: "gpt-4o".to_string(),
},
ModelAlias {
alias: "o1".to_string(),
actual: "o1".to_string(),
},
ModelAlias {
alias: "o3".to_string(),
actual: "o3".to_string(),
},
],
rules: vec![
RoutingRuleDto {
pattern: "claude-*".to_string(),
target_provider: ProviderType::Kiro,
priority: 1,
enabled: true,
},
RoutingRuleDto {
pattern: "gemini-*".to_string(),
target_provider: ProviderType::Gemini,
priority: 1,
enabled: true,
},
RoutingRuleDto {
pattern: "qwen*".to_string(),
target_provider: ProviderType::Qwen,
priority: 1,
enabled: true,
},
RoutingRuleDto {
pattern: "gpt-*".to_string(),
target_provider: ProviderType::OpenAI,
priority: 1,
enabled: true,
},
RoutingRuleDto {
pattern: "o1*".to_string(),
target_provider: ProviderType::OpenAI,
priority: 1,
enabled: true,
},
RoutingRuleDto {
pattern: "o3*".to_string(),
target_provider: ProviderType::OpenAI,
priority: 1,
enabled: true,
},
],
endpoint_providers: None,
},
RecommendedPreset {
id: "coding-assistant".to_string(),
name: "编程助手配置".to_string(),
description:
"针对编程场景优化,Claude Opus 4.5 用于复杂代码,Gemini Flash 用于快速响应"
.to_string(),
aliases: vec![
ModelAlias {
alias: "code".to_string(),
actual: "claude-opus-4-5".to_string(),
},
ModelAlias {
alias: "coder".to_string(),
actual: "qwen3-coder-plus".to_string(),
},
ModelAlias {
alias: "fast".to_string(),
actual: "gemini-2.5-flash".to_string(),
},
ModelAlias {
alias: "think".to_string(),
actual: "claude-sonnet-4-5".to_string(),
},
],
rules: vec![
RoutingRuleDto {
pattern: "*coder*".to_string(),
target_provider: ProviderType::Qwen,
priority: 1,
enabled: true,
},
RoutingRuleDto {
pattern: "claude-*".to_string(),
target_provider: ProviderType::Kiro,
priority: 2,
enabled: true,
},
RoutingRuleDto {
pattern: "gemini-*".to_string(),
target_provider: ProviderType::Gemini,
priority: 2,
enabled: true,
},
],
endpoint_providers: None,
},
RecommendedPreset {
id: "cost-effective".to_string(),
name: "性价比优先配置".to_string(),
description: "优先使用免费或低成本的模型,适合预算有限的用户".to_string(),
aliases: vec![
ModelAlias {
alias: "default".to_string(),
actual: "gemini-2.5-flash".to_string(),
},
ModelAlias {
alias: "cheap".to_string(),
actual: "gemini-2.5-flash-lite".to_string(),
},
ModelAlias {
alias: "free".to_string(),
actual: "gemini-2.5-flash".to_string(),
},
],
rules: vec![
// 默认路由到 Gemini(免费额度高)
RoutingRuleDto {
pattern: "*".to_string(),
target_provider: ProviderType::Gemini,
priority: 100,
enabled: true,
},
// Claude 请求仍然路由到 Kiro
RoutingRuleDto {
pattern: "claude-*".to_string(),
target_provider: ProviderType::Kiro,
priority: 1,
enabled: true,
},
],
endpoint_providers: None,
},
// 客户端路由预设
RecommendedPreset {
id: "client-routing".to_string(),
name: "客户端路由配置".to_string(),
description: "为不同的 IDE 客户端配置不同的 Provider,Cursor/Windsurf 使用 Kiro,Claude Code 使用 Kiro,Codex 使用 OpenAI"
.to_string(),
aliases: vec![],
rules: vec![],
endpoint_providers: Some(EndpointProvidersConfigDto {
cursor: Some("kiro".to_string()),
claude_code: Some("kiro".to_string()),
codex: Some("openai".to_string()),
windsurf: Some("kiro".to_string()),
kiro: Some("kiro".to_string()),
other: None,
}),
},
])
}
/// 应用推荐配置
#[tauri::command]
pub async fn apply_recommended_preset(
state: tauri::State<'_, RouterConfigState>,
preset_id: String,
merge: bool,
) -> Result<(), String> {
let presets = get_recommended_presets().await?;
let preset = presets
.into_iter()
.find(|p| p.id == preset_id)
.ok_or_else(|| format!("未找到预设配置: {}", preset_id))?;
// 应用别名
{
let mut aliases = state.aliases.write().await;
if !merge {
aliases.clear();
}
for alias in preset.aliases {
aliases.insert(alias.alias, alias.actual);
}
}
// 应用规则
{
let mut rules = state.rules.write().await;
if !merge {
rules.clear();
}
for rule in preset.rules {
// 避免重复
if !rules.iter().any(|r| r.pattern == rule.pattern) {
rules.push(rule);
}
}
// 按优先级排序
rules.sort_by(|a, b| a.priority.cmp(&b.priority));
}
Ok(())
}
/// 清空所有路由配置
#[tauri::command]
pub async fn clear_all_routing_config(
state: tauri::State<'_, RouterConfigState>,
) -> Result<(), String> {
{
let mut aliases = state.aliases.write().await;
aliases.clear();
}
{
let mut rules = state.rules.write().await;
rules.clear();
+3 -3
View File
@@ -22,9 +22,9 @@ pub use types::{
generate_secure_api_key, AmpConfig, AmpModelMapping, ApiKeyEntry, Config, CredentialEntry,
CredentialPoolConfig, CustomProviderConfig, EndpointProvidersConfig, GeminiApiKeyEntry,
IFlowCredentialEntry, InjectionRuleConfig, InjectionSettings, LoggingConfig, ModelInfo,
ModelsConfig, ProviderConfig, ProviderModelsConfig, ProvidersConfig, QuotaExceededConfig,
RemoteManagementConfig, RetrySettings, RoutingConfig, RoutingRuleConfig, ServerConfig,
TlsConfig, VertexApiKeyEntry, VertexModelAlias, DEFAULT_API_KEY,
ModelsConfig, NativeAgentConfig, ProviderConfig, ProviderModelsConfig, ProvidersConfig,
QuotaExceededConfig, RemoteManagementConfig, RetrySettings, RoutingConfig, RoutingRuleConfig,
ServerConfig, TlsConfig, VertexApiKeyEntry, VertexModelAlias, DEFAULT_API_KEY,
};
pub use yaml::{load_config, save_config, ConfigError, ConfigManager, YamlService};
+3
View File
@@ -220,6 +220,7 @@ fn arb_config() -> impl Strategy<Value = Config> {
endpoint_providers: crate::config::EndpointProvidersConfig::default(),
minimize_to_tray: true,
models: crate::config::ModelsConfig::default(),
agent: crate::config::NativeAgentConfig::default(),
})
}
@@ -494,6 +495,7 @@ fn arb_valid_config() -> impl Strategy<Value = Config> {
endpoint_providers: crate::config::EndpointProvidersConfig::default(),
minimize_to_tray: true,
models: crate::config::ModelsConfig::default(),
agent: crate::config::NativeAgentConfig::default(),
})
}
@@ -540,6 +542,7 @@ fn arb_invalid_config() -> impl Strategy<Value = Config> {
endpoint_providers: crate::config::EndpointProvidersConfig::default(),
minimize_to_tray: true,
models: crate::config::ModelsConfig::default(),
agent: crate::config::NativeAgentConfig::default(),
};
// 根据类型使配置无效
match invalid_type {
+97
View File
@@ -314,6 +314,102 @@ pub struct Config {
/// 模型配置(动态加载 Provider 和模型列表)
#[serde(default)]
pub models: ModelsConfig,
/// Native Agent 配置
#[serde(default)]
pub agent: NativeAgentConfig,
}
// ============ Native Agent 配置类型 ============
/// Native Agent 配置
///
/// 配置内置 Agent 的行为,包括系统提示词、工具使用规则等
/// 参考 Manus Agent 的模块化设计,支持灵活配置
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct NativeAgentConfig {
/// 是否使用默认系统提示词
/// 当 custom_system_prompt 为空时,如果此项为 true 则使用内置默认提示词
#[serde(default = "default_use_default_prompt")]
pub use_default_system_prompt: bool,
/// 自定义系统提示词
/// 如果设置了此项,将覆盖默认系统提示词
#[serde(default, skip_serializing_if = "Option::is_none")]
pub custom_system_prompt: Option<String>,
/// 系统提示词模板文件路径(支持 ~ 展开)
/// 可以将系统提示词存储在外部文件中,便于管理和版本控制
#[serde(default, skip_serializing_if = "Option::is_none")]
pub system_prompt_file: Option<String>,
/// 默认模型
#[serde(default = "default_agent_model")]
pub default_model: String,
/// 默认温度参数
#[serde(default = "default_temperature")]
pub temperature: f32,
/// 默认最大 token 数
#[serde(default = "default_max_tokens")]
pub max_tokens: u32,
}
fn default_use_default_prompt() -> bool {
true
}
fn default_agent_model() -> String {
"claude-sonnet-4-20250514".to_string()
}
fn default_temperature() -> f32 {
0.7
}
fn default_max_tokens() -> u32 {
4096
}
impl Default for NativeAgentConfig {
fn default() -> Self {
Self {
use_default_system_prompt: default_use_default_prompt(),
custom_system_prompt: None,
system_prompt_file: None,
default_model: default_agent_model(),
temperature: default_temperature(),
max_tokens: default_max_tokens(),
}
}
}
impl NativeAgentConfig {
/// 获取有效的系统提示词
///
/// 优先级:
/// 1. system_prompt_file(外部文件)
/// 2. custom_system_prompt(配置中的自定义提示词)
/// 3. 如果 use_default_system_prompt 为 true,返回 None 让调用方使用默认提示词
/// 4. 否则返回 None(不使用任何系统提示词)
pub fn get_effective_system_prompt(&self) -> Option<String> {
// 优先从文件加载
if let Some(file_path) = &self.system_prompt_file {
let expanded_path = crate::config::expand_tilde(file_path);
if let Ok(content) = std::fs::read_to_string(&expanded_path) {
let trimmed = content.trim();
if !trimmed.is_empty() {
return Some(trimmed.to_string());
}
}
}
// 其次使用配置中的自定义提示词
if let Some(prompt) = &self.custom_system_prompt {
let trimmed = prompt.trim();
if !trimmed.is_empty() {
return Some(trimmed.to_string());
}
}
// 返回 None,让调用方根据 use_default_system_prompt 决定是否使用默认提示词
None
}
}
fn default_minimize_to_tray() -> bool {
@@ -1116,6 +1212,7 @@ impl Default for Config {
endpoint_providers: EndpointProvidersConfig::default(),
minimize_to_tray: default_minimize_to_tray(),
models: ModelsConfig::default(),
agent: NativeAgentConfig::default(),
}
}
}
+72 -6
View File
@@ -754,13 +754,56 @@ impl CredentialProviderRegistry {
/// 检查插件更新
pub async fn check_updates(&self) -> OAuthPluginResult<Vec<PluginUpdate>> {
// TODO: 实现更新检查逻辑
// 1. 遍历所有插件
// 2. 检查 GitHub Release 或其他来源
// 3. 比较版本号
// 4. 返回有更新的插件列表
// 已知插件的最新版本(与前端 OAuthPluginTab.tsx 保持同步)
let latest_versions: std::collections::HashMap<&str, &str> = [
("kiro-provider", "0.3.0"),
("antigravity-provider", "0.4.0"),
("claude-provider", "0.3.0"),
("droid-provider", "0.3.0"),
("gemini-provider", "0.4.0"),
("codex-provider", "0.1.0"),
]
.into_iter()
.collect();
Ok(vec![])
let mut updates = Vec::new();
// 扫描已安装的插件
if let Ok(entries) = std::fs::read_dir(&self.plugins_dir) {
for entry in entries.flatten() {
let path = entry.path();
if path.is_dir() {
let plugin_json = path.join("plugin.json");
if plugin_json.exists() {
if let Ok(content) = std::fs::read_to_string(&plugin_json) {
if let Ok(manifest) =
serde_json::from_str::<serde_json::Value>(&content)
{
let plugin_id = manifest["name"].as_str().unwrap_or_default();
let current_version =
manifest["version"].as_str().unwrap_or("0.0.0");
// 检查是否有更新
if let Some(&latest) = latest_versions.get(plugin_id) {
if version_compare(current_version, latest)
== std::cmp::Ordering::Less
{
updates.push(PluginUpdate {
plugin_id: plugin_id.to_string(),
current_version: current_version.to_string(),
latest_version: latest.to_string(),
changelog: None,
});
}
}
}
}
}
}
}
}
Ok(updates)
}
// ========================================================================
@@ -809,6 +852,29 @@ pub fn get_global_registry() -> Option<Arc<CredentialProviderRegistry>> {
GLOBAL_REGISTRY.get().cloned()
}
/// 比较语义化版本号
fn version_compare(v1: &str, v2: &str) -> std::cmp::Ordering {
let parse = |v: &str| -> Vec<u32> {
v.trim_start_matches('v')
.split('.')
.filter_map(|s| s.parse().ok())
.collect()
};
let v1_parts = parse(v1);
let v2_parts = parse(v2);
for i in 0..std::cmp::max(v1_parts.len(), v2_parts.len()) {
let p1 = v1_parts.get(i).copied().unwrap_or(0);
let p2 = v2_parts.get(i).copied().unwrap_or(0);
match p1.cmp(&p2) {
std::cmp::Ordering::Equal => continue,
other => return other,
}
}
std::cmp::Ordering::Equal
}
/// 递归复制目录
fn copy_dir_all(src: &Path, dst: &Path) -> std::io::Result<()> {
std::fs::create_dir_all(dst)?;
-984
View File
@@ -1,984 +0,0 @@
//! 本地硬编码的国内模型数据
//!
//! 这些模型数据用于补充 models.dev API 未覆盖的国内模型
use crate::models::model_registry::{
EnhancedModelMetadata, ModelCapabilities, ModelLimits, ModelPricing, ModelSource, ModelStatus,
ModelTier,
};
/// 获取所有本地硬编码的国内模型
pub fn get_local_models() -> Vec<EnhancedModelMetadata> {
let mut models = Vec::new();
models.extend(get_dashscope_models());
models.extend(get_zhipu_models());
models.extend(get_baichuan_models());
models.extend(get_moonshot_models());
models.extend(get_deepseek_models());
models.extend(get_doubao_models());
models.extend(get_minimax_models());
models.extend(get_yi_models());
models.extend(get_stepfun_models());
models
}
/// 通义千问系列模型 (阿里云百炼)
fn get_dashscope_models() -> Vec<EnhancedModelMetadata> {
let now = chrono::Utc::now().timestamp();
vec![
EnhancedModelMetadata {
id: "qwen3-coder-plus".to_string(),
display_name: "通义千问 Coder Plus".to_string(),
provider_id: "dashscope".to_string(),
provider_name: "阿里云百炼".to_string(),
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,
},
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,
currency: "CNY".to_string(),
}),
limits: ModelLimits {
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,
},
EnhancedModelMetadata {
id: "qwen-max".to_string(),
display_name: "通义千问 Max".to_string(),
provider_id: "dashscope".to_string(),
provider_name: "阿里云百炼".to_string(),
family: Some("qwen".to_string()),
tier: ModelTier::Max,
capabilities: ModelCapabilities {
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,
currency: "CNY".to_string(),
}),
limits: ModelLimits {
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,
},
EnhancedModelMetadata {
id: "qwen-plus".to_string(),
display_name: "通义千问 Plus".to_string(),
provider_id: "dashscope".to_string(),
provider_name: "阿里云百炼".to_string(),
family: Some("qwen".to_string()),
tier: ModelTier::Pro,
capabilities: ModelCapabilities {
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,
currency: "CNY".to_string(),
}),
limits: ModelLimits {
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,
},
EnhancedModelMetadata {
id: "qwen-turbo".to_string(),
display_name: "通义千问 Turbo".to_string(),
provider_id: "dashscope".to_string(),
provider_name: "阿里云百炼".to_string(),
family: Some("qwen".to_string()),
tier: ModelTier::Mini,
capabilities: ModelCapabilities {
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,
currency: "CNY".to_string(),
}),
limits: ModelLimits {
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,
},
]
}
/// 智谱 GLM 系列模型
fn get_zhipu_models() -> Vec<EnhancedModelMetadata> {
let now = chrono::Utc::now().timestamp();
vec![
EnhancedModelMetadata {
id: "glm-4-plus".to_string(),
display_name: "GLM-4 Plus".to_string(),
provider_id: "zhipu".to_string(),
provider_name: "智谱 AI".to_string(),
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,
},
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,
currency: "CNY".to_string(),
}),
limits: ModelLimits {
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,
},
EnhancedModelMetadata {
id: "glm-4-air".to_string(),
display_name: "GLM-4 Air".to_string(),
provider_id: "zhipu".to_string(),
provider_name: "智谱 AI".to_string(),
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,
},
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(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,
},
EnhancedModelMetadata {
id: "glm-4-flash".to_string(),
display_name: "GLM-4 Flash".to_string(),
provider_id: "zhipu".to_string(),
provider_name: "智谱 AI".to_string(),
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,
},
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,
currency: "CNY".to_string(),
}),
limits: ModelLimits {
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,
},
]
}
/// 百川系列模型
fn get_baichuan_models() -> Vec<EnhancedModelMetadata> {
let now = chrono::Utc::now().timestamp();
vec![
EnhancedModelMetadata {
id: "Baichuan4".to_string(),
display_name: "百川 4".to_string(),
provider_id: "baichuan".to_string(),
provider_name: "百川智能".to_string(),
family: Some("baichuan".to_string()),
tier: ModelTier::Max,
capabilities: ModelCapabilities {
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,
currency: "CNY".to_string(),
}),
limits: ModelLimits {
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,
},
EnhancedModelMetadata {
id: "Baichuan3-Turbo".to_string(),
display_name: "百川 3 Turbo".to_string(),
provider_id: "baichuan".to_string(),
provider_name: "百川智能".to_string(),
family: Some("baichuan".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(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,
},
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,
},
]
}
/// 月之暗面 Moonshot 系列模型
fn get_moonshot_models() -> Vec<EnhancedModelMetadata> {
let now = chrono::Utc::now().timestamp();
vec![
EnhancedModelMetadata {
id: "moonshot-v1-128k".to_string(),
display_name: "Moonshot V1 128K".to_string(),
provider_id: "moonshot".to_string(),
provider_name: "月之暗面".to_string(),
family: Some("moonshot".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(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,
},
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,
},
EnhancedModelMetadata {
id: "moonshot-v1-32k".to_string(),
display_name: "Moonshot V1 32K".to_string(),
provider_id: "moonshot".to_string(),
provider_name: "月之暗面".to_string(),
family: Some("moonshot".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(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,
},
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,
},
EnhancedModelMetadata {
id: "moonshot-v1-8k".to_string(),
display_name: "Moonshot V1 8K".to_string(),
provider_id: "moonshot".to_string(),
provider_name: "月之暗面".to_string(),
family: Some("moonshot".to_string()),
tier: ModelTier::Mini,
capabilities: ModelCapabilities {
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,
currency: "CNY".to_string(),
}),
limits: ModelLimits {
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,
},
]
}
/// DeepSeek 系列模型
fn get_deepseek_models() -> Vec<EnhancedModelMetadata> {
let now = chrono::Utc::now().timestamp();
vec![
EnhancedModelMetadata {
id: "deepseek-chat".to_string(),
display_name: "DeepSeek Chat".to_string(),
provider_id: "deepseek".to_string(),
provider_name: "DeepSeek".to_string(),
family: Some("deepseek".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(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,
},
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,
},
EnhancedModelMetadata {
id: "deepseek-reasoner".to_string(),
display_name: "DeepSeek Reasoner".to_string(),
provider_id: "deepseek".to_string(),
provider_name: "DeepSeek".to_string(),
family: Some("deepseek".to_string()),
tier: ModelTier::Max,
capabilities: ModelCapabilities {
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),
currency: "CNY".to_string(),
}),
limits: ModelLimits {
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,
},
EnhancedModelMetadata {
id: "deepseek-coder".to_string(),
display_name: "DeepSeek Coder".to_string(),
provider_id: "deepseek".to_string(),
provider_name: "DeepSeek".to_string(),
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,
},
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),
currency: "CNY".to_string(),
}),
limits: ModelLimits {
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,
},
]
}
/// 字节豆包系列模型
fn get_doubao_models() -> Vec<EnhancedModelMetadata> {
let now = chrono::Utc::now().timestamp();
vec![
EnhancedModelMetadata {
id: "doubao-pro-256k".to_string(),
display_name: "豆包 Pro 256K".to_string(),
provider_id: "doubao".to_string(),
provider_name: "字节跳动".to_string(),
family: Some("doubao".to_string()),
tier: ModelTier::Max,
capabilities: ModelCapabilities {
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,
currency: "CNY".to_string(),
}),
limits: ModelLimits {
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,
},
EnhancedModelMetadata {
id: "doubao-pro-32k".to_string(),
display_name: "豆包 Pro 32K".to_string(),
provider_id: "doubao".to_string(),
provider_name: "字节跳动".to_string(),
family: Some("doubao".to_string()),
tier: ModelTier::Pro,
capabilities: ModelCapabilities {
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,
currency: "CNY".to_string(),
}),
limits: ModelLimits {
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,
},
EnhancedModelMetadata {
id: "doubao-lite-32k".to_string(),
display_name: "豆包 Lite 32K".to_string(),
provider_id: "doubao".to_string(),
provider_name: "字节跳动".to_string(),
family: Some("doubao".to_string()),
tier: ModelTier::Mini,
capabilities: ModelCapabilities {
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,
currency: "CNY".to_string(),
}),
limits: ModelLimits {
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,
},
]
}
/// MiniMax 系列模型
fn get_minimax_models() -> Vec<EnhancedModelMetadata> {
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,
}]
}
/// 零一万物 Yi 系列模型
fn get_yi_models() -> Vec<EnhancedModelMetadata> {
let now = chrono::Utc::now().timestamp();
vec![
EnhancedModelMetadata {
id: "yi-large".to_string(),
display_name: "Yi Large".to_string(),
provider_id: "yi".to_string(),
provider_name: "零一万物".to_string(),
family: Some("yi".to_string()),
tier: ModelTier::Max,
capabilities: ModelCapabilities {
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,
currency: "CNY".to_string(),
}),
limits: ModelLimits {
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,
},
EnhancedModelMetadata {
id: "yi-medium".to_string(),
display_name: "Yi Medium".to_string(),
provider_id: "yi".to_string(),
provider_name: "零一万物".to_string(),
family: Some("yi".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(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,
},
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,
},
EnhancedModelMetadata {
id: "yi-spark".to_string(),
display_name: "Yi Spark".to_string(),
provider_id: "yi".to_string(),
provider_name: "零一万物".to_string(),
family: Some("yi".to_string()),
tier: ModelTier::Mini,
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(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,
},
]
}
/// 阶跃星辰 Step 系列模型
fn get_stepfun_models() -> Vec<EnhancedModelMetadata> {
let now = chrono::Utc::now().timestamp();
vec![
EnhancedModelMetadata {
id: "step-2-16k".to_string(),
display_name: "Step 2 16K".to_string(),
provider_id: "stepfun".to_string(),
provider_name: "阶跃星辰".to_string(),
family: Some("step".to_string()),
tier: ModelTier::Max,
capabilities: ModelCapabilities {
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,
currency: "CNY".to_string(),
}),
limits: ModelLimits {
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,
},
EnhancedModelMetadata {
id: "step-1-128k".to_string(),
display_name: "Step 1 128K".to_string(),
provider_id: "stepfun".to_string(),
provider_name: "阶跃星辰".to_string(),
family: Some("step".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(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,
},
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,
},
EnhancedModelMetadata {
id: "step-1-flash".to_string(),
display_name: "Step 1 Flash".to_string(),
provider_id: "stepfun".to_string(),
provider_name: "阶跃星辰".to_string(),
family: Some("step".to_string()),
tier: ModelTier::Mini,
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(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,
},
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,
},
]
}
+2 -5
View File
@@ -1,7 +1,4 @@
//! 静态数据模块
//!
//! 包含本地硬编码的模型数据等
pub mod local_models;
pub use local_models::get_local_models;
//! 模型数据现在从 aiclientproxy/models 仓库获取
//! 本地硬编码数据已迁移到独立仓库: https://github.com/aiclientproxy/models
+117 -74
View File
@@ -787,101 +787,144 @@ pub async fn chat_completions(
let credential = if credential.is_none() {
eprintln!("[CHAT_COMPLETIONS] Provider Pool 中未找到凭证,尝试 API Key Provider...");
// 根据 selected_provider 映射到 ApiProviderType
use crate::database::dao::api_key_provider::ApiProviderType;
let api_provider_type = match selected_provider.to_lowercase().as_str() {
"anthropic" | "claude" => Some(ApiProviderType::Anthropic),
"openai" => Some(ApiProviderType::Openai),
"gemini" => Some(ApiProviderType::Gemini),
// 以下都是 OpenAI 兼容的 Provider
"deepseek" | "moonshot" | "groq" | "grok" | "mistral" | "perplexity" | "cohere"
| "openrouter" | "silicon" => Some(ApiProviderType::Openai),
_ => None,
};
let provider_id_lower = selected_provider.to_lowercase();
if let (Some(db), Some(api_type)) = (&state.db, api_provider_type) {
// 策略 1: 优先按 provider_id 直接查找(支持 deepseek, moonshot 等 60+ Provider)
// 这些 Provider 在 API Key Provider 中有独立配置
let mut found_credential: Option<crate::models::provider_pool_model::ProviderCredential> =
None;
if let Some(db) = &state.db {
// 先尝试按 provider_id 直接查找
eprintln!(
"[CHAT_COMPLETIONS] 尝试从 API Key Provider 类型 '{:?}' 获取凭证",
api_type
"[CHAT_COMPLETIONS] 尝试按 provider_id '{}' 直接查找凭证",
provider_id_lower
);
// 使用按类型获取的方法(包括自定义 Provider)
match state.api_key_service.get_next_api_key_by_type(db, api_type) {
Ok(Some((_key_id, api_key, provider_info))) => {
match state.api_key_service.get_fallback_credential(
db,
&crate::models::provider_pool_model::PoolProviderType::OpenAI,
Some(&provider_id_lower),
) {
Ok(Some(cred)) => {
eprintln!(
"[CHAT_COMPLETIONS] 从 API Key Provider 获取到凭证: provider={}, api_host={}",
provider_info.name,
provider_info.api_host
"[CHAT_COMPLETIONS] 通过 provider_id '{}' 找到凭证: name={:?}",
provider_id_lower, cred.name
);
let base_url = if provider_info.api_host.is_empty() {
None
} else {
Some(provider_info.api_host.clone())
};
let provider_type = match provider_info.provider_type {
ApiProviderType::Anthropic => crate::ProviderType::Anthropic,
ApiProviderType::Openai | ApiProviderType::OpenaiResponse => {
crate::ProviderType::OpenAI
}
ApiProviderType::Gemini => crate::ProviderType::GeminiApiKey,
_ => crate::ProviderType::OpenAI,
};
// 根据 provider_type 创建对应的 CredentialData
let credential_data = match provider_type {
crate::ProviderType::Anthropic => {
crate::models::provider_pool_model::CredentialData::AnthropicKey {
api_key: api_key.clone(),
base_url,
}
}
crate::ProviderType::GeminiApiKey => {
crate::models::provider_pool_model::CredentialData::GeminiApiKey {
api_key: api_key.clone(),
base_url,
excluded_models: vec![],
}
}
_ => crate::models::provider_pool_model::CredentialData::OpenAIKey {
api_key: api_key.clone(),
base_url,
},
};
// 构建 ProviderCredential
let mut cred = crate::models::provider_pool_model::ProviderCredential::new(
provider_type,
credential_data,
);
cred.name = Some(provider_info.name.clone());
state.logs.write().await.add(
"info",
&format!(
"[ROUTE] Using API Key Provider credential: provider={}, type={:?}",
provider_info.name, provider_info.provider_type
"[ROUTE] Using API Key Provider credential by provider_id: {}",
provider_id_lower
),
);
Some(cred)
found_credential = Some(cred);
}
Ok(None) => {
eprintln!(
"[CHAT_COMPLETIONS] API Key Provider 类型 '{:?}' 没有可用的 API Key",
api_type
"[CHAT_COMPLETIONS] provider_id '{}' 未找到凭证,尝试按类型查找",
provider_id_lower
);
None
}
Err(e) => {
eprintln!("[CHAT_COMPLETIONS] 从 API Key Provider 获取凭证失败: {}", e);
None
eprintln!("[CHAT_COMPLETIONS] 按 provider_id 查找凭证失败: {}", e);
}
}
// 策略 2: 如果按 provider_id 未找到,按类型查找
if found_credential.is_none() {
let api_provider_type = match provider_id_lower.as_str() {
"anthropic" | "claude" => Some(ApiProviderType::Anthropic),
"openai" => Some(ApiProviderType::Openai),
"gemini" => Some(ApiProviderType::Gemini),
// 以下都是 OpenAI 兼容的 Provider,但优先按 provider_id 查找已在上面处理
"deepseek" | "moonshot" | "groq" | "grok" | "mistral" | "perplexity"
| "cohere" | "openrouter" | "silicon" => Some(ApiProviderType::Openai),
_ => None,
};
if let Some(api_type) = api_provider_type {
eprintln!(
"[CHAT_COMPLETIONS] 尝试从 API Key Provider 类型 '{:?}' 获取凭证",
api_type
);
match state.api_key_service.get_next_api_key_by_type(db, api_type) {
Ok(Some((_key_id, api_key, provider_info))) => {
eprintln!(
"[CHAT_COMPLETIONS] 从 API Key Provider 获取到凭证: provider={}, api_host={}",
provider_info.name,
provider_info.api_host
);
let base_url = if provider_info.api_host.is_empty() {
None
} else {
Some(provider_info.api_host.clone())
};
let provider_type = match provider_info.provider_type {
ApiProviderType::Anthropic => crate::ProviderType::Anthropic,
ApiProviderType::Openai | ApiProviderType::OpenaiResponse => {
crate::ProviderType::OpenAI
}
ApiProviderType::Gemini => crate::ProviderType::GeminiApiKey,
_ => crate::ProviderType::OpenAI,
};
let credential_data = match provider_type {
crate::ProviderType::Anthropic => {
crate::models::provider_pool_model::CredentialData::AnthropicKey {
api_key: api_key.clone(),
base_url,
}
}
crate::ProviderType::GeminiApiKey => {
crate::models::provider_pool_model::CredentialData::GeminiApiKey {
api_key: api_key.clone(),
base_url,
excluded_models: vec![],
}
}
_ => crate::models::provider_pool_model::CredentialData::OpenAIKey {
api_key: api_key.clone(),
base_url,
},
};
let mut cred =
crate::models::provider_pool_model::ProviderCredential::new(
provider_type,
credential_data,
);
cred.name = Some(provider_info.name.clone());
state.logs.write().await.add(
"info",
&format!(
"[ROUTE] Using API Key Provider credential: provider={}, type={:?}",
provider_info.name, provider_info.provider_type
),
);
found_credential = Some(cred);
}
Ok(None) => {
eprintln!(
"[CHAT_COMPLETIONS] API Key Provider 类型 '{:?}' 没有可用的 API Key",
api_type
);
}
Err(e) => {
eprintln!("[CHAT_COMPLETIONS] 从 API Key Provider 获取凭证失败: {}", e);
}
}
}
}
} else {
None
}
found_credential
} else {
credential
};
+21 -1
View File
@@ -1467,8 +1467,18 @@ async fn amp_chat_completions(
);
// 尝试根据 provider 名称选择凭证(带智能降级)
eprintln!(
"[AMP] 开始查找凭证: provider={}, model={}, db={}",
provider,
request.model,
state.db.is_some()
);
let credential = match &state.db {
Some(db) => {
eprintln!(
"[AMP] 调用 select_credential_with_fallback, provider_id_hint={}",
provider
);
// 首先尝试按 provider 类型选择(带智能降级)
if let Ok(Some(cred)) = state.pool_service.select_credential_with_fallback(
db,
@@ -1477,20 +1487,30 @@ async fn amp_chat_completions(
Some(&request.model),
Some(&provider), // provider_id_hint 使用路由中的 provider 名称
) {
eprintln!(
"[AMP] select_credential_with_fallback 找到凭证: {:?}",
cred.name
);
Some(cred)
}
// 然后尝试按名称查找
else if let Ok(Some(cred)) = state.pool_service.get_by_name(db, &provider) {
eprintln!("[AMP] get_by_name 找到凭证: {:?}", cred.name);
Some(cred)
}
// 最后尝试按 UUID 查找
else if let Ok(Some(cred)) = state.pool_service.get_by_uuid(db, &provider) {
eprintln!("[AMP] get_by_uuid 找到凭证: {:?}", cred.name);
Some(cred)
} else {
eprintln!("[AMP] 未找到任何凭证 for provider '{}'", provider);
None
}
}
None => None,
None => {
eprintln!("[AMP] 数据库未初始化");
None
}
};
match credential {
@@ -316,6 +316,8 @@ impl ApiKeyProviderService {
// ==================== API Key 操作 ====================
/// 添加 API Key
///
/// 当添加第一个 API Key 时,会自动启用 Provider
pub fn add_api_key(
&self,
db: &DbConnection,
@@ -325,10 +327,15 @@ impl ApiKeyProviderService {
) -> Result<ApiKeyEntry, String> {
// 验证 Provider 存在
let conn = db.lock().map_err(|e| e.to_string())?;
let _ = ApiKeyProviderDao::get_provider_by_id(&conn, provider_id)
let provider = ApiKeyProviderDao::get_provider_by_id(&conn, provider_id)
.map_err(|e| e.to_string())?
.ok_or_else(|| format!("Provider not found: {}", provider_id))?;
// 检查是否是第一个 API Key,如果是则自动启用 Provider
let existing_keys = ApiKeyProviderDao::get_api_keys_by_provider(&conn, provider_id)
.map_err(|e| e.to_string())?;
let should_enable_provider = existing_keys.is_empty() && !provider.enabled;
// 加密 API Key
let encrypted_key = self.encryption.encrypt(api_key);
@@ -347,6 +354,19 @@ impl ApiKeyProviderService {
ApiKeyProviderDao::insert_api_key(&conn, &key).map_err(|e| e.to_string())?;
// 如果是第一个 API Key,自动启用 Provider
if should_enable_provider {
let mut updated_provider = provider;
updated_provider.enabled = true;
updated_provider.updated_at = now;
ApiKeyProviderDao::update_provider(&conn, &updated_provider)
.map_err(|e| e.to_string())?;
tracing::info!(
"[ApiKeyProviderService] 自动启用 Provider: {} (添加了第一个 API Key)",
provider_id
);
}
Ok(key)
}
@@ -696,26 +716,49 @@ impl ApiKeyProviderService {
pool_type: &PoolProviderType,
provider_id_hint: Option<&str>,
) -> Result<Option<ProviderCredential>, String> {
// 策略 1: 通过类型映射查找
if let Some(api_type) = self.map_pool_type_to_api_type(pool_type) {
tracing::debug!("[智能降级] 尝试类型映射: {:?} -> {:?}", pool_type, api_type);
if let Some(cred) = self.find_by_api_type(db, pool_type, &api_type)? {
return Ok(Some(cred));
}
}
eprintln!(
"[get_fallback_credential] 开始查找: pool_type={:?}, provider_id_hint={:?}",
pool_type, provider_id_hint
);
// 策略 2: 通过 provider_id 直接查找 (支持 60+ Provider)
// 策略 1: 优先通过 provider_id 直接查找 (支持 deepseek, moonshot 等 60+ Provider)
// 这些 Provider 在 API Key Provider 中有独立配置,应该优先使用
if let Some(provider_id) = provider_id_hint {
tracing::debug!("[智能降级] 尝试 provider_id 查找: {}", provider_id);
eprintln!(
"[get_fallback_credential] 尝试按 provider_id '{}' 查找",
provider_id
);
if let Some(cred) = self.find_by_provider_id(db, provider_id)? {
eprintln!(
"[get_fallback_credential] 通过 provider_id '{}' 找到凭证: {:?}",
provider_id, cred.name
);
return Ok(Some(cred));
}
eprintln!(
"[get_fallback_credential] provider_id '{}' 未找到凭证",
provider_id
);
}
// 策略 2: 通过类型映射查找(降级方案)
if let Some(api_type) = self.map_pool_type_to_api_type(pool_type) {
eprintln!(
"[get_fallback_credential] 尝试类型映射: {:?} -> {:?}",
pool_type, api_type
);
if let Some(cred) = self.find_by_api_type(db, pool_type, &api_type)? {
eprintln!(
"[get_fallback_credential] 通过类型映射找到凭证: {:?}",
cred.name
);
return Ok(Some(cred));
}
}
tracing::debug!(
"[智能降级] 未找到 {:?} 的降级凭证 (provider_id_hint: {:?})",
pool_type,
provider_id_hint
eprintln!(
"[get_fallback_credential] 未找到 {:?} 的降级凭证 (provider_id_hint: {:?})",
pool_type, provider_id_hint
);
Ok(None)
}
@@ -831,8 +874,24 @@ impl ApiKeyProviderService {
ApiKeyProviderDao::get_provider_by_id(&conn, provider_id).map_err(|e| e.to_string())?;
let provider = match provider {
Some(p) if p.enabled => p,
_ => return Ok(None),
Some(p) if p.enabled => {
eprintln!(
"[find_by_provider_id] 找到已启用的 provider: id={}, name={}, api_host={}",
p.id, p.name, p.api_host
);
p
}
Some(_p) => {
eprintln!(
"[find_by_provider_id] provider '{}' 存在但未启用",
provider_id
);
return Ok(None);
}
None => {
eprintln!("[find_by_provider_id] provider '{}' 不存在", provider_id);
return Ok(None);
}
};
// 获取启用的 API Key
@@ -840,9 +899,19 @@ impl ApiKeyProviderService {
.map_err(|e| e.to_string())?;
if keys.is_empty() {
eprintln!(
"[find_by_provider_id] provider '{}' 没有启用的 API Key",
provider_id
);
return Ok(None);
}
eprintln!(
"[find_by_provider_id] provider '{}' 有 {} 个启用的 API Key",
provider_id,
keys.len()
);
// 轮询选择 API Key
let index = {
let mut indices = self.round_robin_index.write().map_err(|e| e.to_string())?;
+192 -84
View File
@@ -1,21 +1,90 @@
//! 模型注册服务
//!
//! 负责从 models.dev API 获取模型数据、管理本地缓存、提供模型搜索等功能
//! 负责从 aiclientproxy/models 仓库获取模型数据、管理本地缓存、提供模型搜索等功能
use crate::data::get_local_models;
use crate::database::DbConnection;
use crate::models::model_registry::{
EnhancedModelMetadata, ModelSource, ModelStatus, ModelSyncState, ModelTier, ModelsDevProvider,
UserModelPreference,
EnhancedModelMetadata, ModelCapabilities, ModelLimits, ModelPricing, ModelSource, ModelStatus,
ModelSyncState, ModelTier, UserModelPreference,
};
use rusqlite::params;
use std::collections::HashMap;
use serde::Deserialize;
use std::sync::Arc;
use tokio::sync::RwLock;
const MODELS_DEV_API_URL: &str = "https://models.dev/api.json";
/// GitHub 仓库 raw 文件基础 URL
const MODELS_REPO_BASE_URL: &str = "https://raw.githubusercontent.com/aiclientproxy/models/main";
const CACHE_DURATION_SECS: i64 = 3600; // 1 小时
/// 仓库索引文件结构
#[derive(Debug, Deserialize)]
struct RepoIndex {
providers: Vec<String>,
#[allow(dead_code)]
total_models: u32,
}
/// 仓库中的 Provider 数据结构
#[derive(Debug, Deserialize)]
struct RepoProviderData {
provider: RepoProvider,
models: Vec<RepoModel>,
}
#[derive(Debug, Deserialize)]
struct RepoProvider {
id: String,
name: String,
}
#[derive(Debug, Deserialize)]
struct RepoModel {
id: String,
name: String,
family: Option<String>,
tier: Option<String>,
capabilities: Option<RepoCapabilities>,
pricing: Option<RepoPricing>,
limits: Option<RepoLimits>,
status: Option<String>,
release_date: Option<String>,
is_latest: Option<bool>,
description: Option<String>,
#[serde(default)]
description_zh: Option<String>,
}
#[derive(Debug, Deserialize, Default)]
struct RepoCapabilities {
#[serde(default)]
vision: bool,
#[serde(default)]
tools: bool,
#[serde(default)]
streaming: bool,
#[serde(default)]
json_mode: bool,
#[serde(default)]
function_calling: bool,
#[serde(default)]
reasoning: bool,
}
#[derive(Debug, Deserialize)]
struct RepoPricing {
input: Option<f64>,
output: Option<f64>,
cache_read: Option<f64>,
cache_write: Option<f64>,
currency: Option<String>,
}
#[derive(Debug, Deserialize)]
struct RepoLimits {
context: Option<u32>,
max_output: Option<u32>,
}
/// 模型注册服务
pub struct ModelRegistryService {
/// 数据库连接
@@ -62,24 +131,7 @@ impl ModelRegistryService {
}
}
// 2. 使用本地硬编码数据作为初始数据
let local_models = get_local_models();
tracing::info!(
"[ModelRegistry] 使用 {} 个本地硬编码模型作为初始数据",
local_models.len()
);
{
let mut cache = self.models_cache.write().await;
*cache = local_models.clone();
}
// 保存到数据库
if let Err(e) = self.save_models_to_db(&local_models).await {
tracing::warn!("[ModelRegistry] 保存本地模型到数据库失败: {}", e);
}
// 3. 后台获取 models.dev 数据
// 2. 后台获取 models 仓库数据
self.spawn_background_refresh();
Ok(())
@@ -109,15 +161,15 @@ impl ModelRegistryService {
models_cache,
sync_state,
};
if let Err(e) = service.refresh_from_models_dev().await {
if let Err(e) = service.refresh_from_repo().await {
tracing::error!("[ModelRegistry] 后台刷新失败: {}", e);
}
});
}
/// 从 models.dev API 刷新数据
pub async fn refresh_from_models_dev(&self) -> Result<(), String> {
tracing::info!("[ModelRegistry] 开始从 models.dev 获取数据");
/// 从 aiclientproxy/models 仓库刷新数据
pub async fn refresh_from_repo(&self) -> Result<(), String> {
tracing::info!("[ModelRegistry] 开始从 models 仓库获取数据");
// 设置同步状态
{
@@ -127,31 +179,27 @@ impl ModelRegistryService {
}
// 获取数据
let result = self.fetch_models_dev_data().await;
let result = self.fetch_models_from_repo().await;
match result {
Ok(models_dev_models) => {
// 合并本地模型
let local_models = get_local_models();
let merged = self.merge_models(models_dev_models, local_models);
tracing::info!("[ModelRegistry] 获取并合并了 {} 个模型", merged.len());
Ok(models) => {
tracing::info!("[ModelRegistry] 获取了 {} 个模型", models.len());
// 更新缓存
{
let mut cache = self.models_cache.write().await;
*cache = merged.clone();
*cache = models.clone();
}
// 保存到数据库
self.save_models_to_db(&merged).await?;
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 = merged.len() as u32;
state.model_count = models.len() as u32;
state.last_error = None;
}
@@ -161,7 +209,7 @@ impl ModelRegistryService {
Ok(())
}
Err(e) => {
tracing::error!("[ModelRegistry] 从 models.dev 获取数据失败: {}", e);
tracing::error!("[ModelRegistry] 从 models 仓库获取数据失败: {}", e);
// 更新同步状态
{
@@ -175,76 +223,136 @@ impl ModelRegistryService {
}
}
/// 从 models.dev API 获取数据
async fn fetch_models_dev_data(&self) -> Result<Vec<EnhancedModelMetadata>, String> {
/// 从 models 仓库获取数据
async fn fetch_models_from_repo(&self) -> Result<Vec<EnhancedModelMetadata>, String> {
let client = reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(30))
.build()
.map_err(|e| format!("创建 HTTP 客户端失败: {}", e))?;
let response = client
.get(MODELS_DEV_API_URL)
// 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!("请求 models.dev 失败: {}", e))?;
if !response.status().is_success() {
return Err(format!("models.dev 返回错误状态码: {}", response.status()));
}
let data: HashMap<String, ModelsDevProvider> = response
.map_err(|e| format!("请求 index.json 失败: {}", e))?
.json()
.await
.map_err(|e| format!("解析 models.dev 响应失败: {}", e))?;
.map_err(|e| format!("解析 index.json 失败: {}", e))?;
// 转换为内部格式
tracing::info!(
"[ModelRegistry] 索引包含 {} 个 providers",
index.providers.len()
);
// 2. 并发获取所有 provider 数据
let mut models = Vec::new();
for (provider_id, provider) in data {
for (_, model) in provider.models {
let enhanced = model.to_enhanced_metadata(&provider_id, &provider.name);
models.push(enhanced);
let now = chrono::Utc::now().timestamp();
for provider_id in &index.providers {
let provider_url = format!("{}/providers/{}.json", MODELS_REPO_BASE_URL, provider_id);
match client
.get(&provider_url)
.header("User-Agent", "ProxyCast/1.0")
.send()
.await
{
Ok(response) => {
if response.status().is_success() {
match response.json::<RepoProviderData>().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);
}
}
}
}
Err(e) => {
tracing::warn!("[ModelRegistry] 获取 {} 失败: {}", provider_id, e);
}
}
}
// 按 provider_id 和 display_name 排序
models.sort_by(|a, b| {
a.provider_id
.cmp(&b.provider_id)
.then(a.display_name.cmp(&b.display_name))
});
tracing::info!(
"[ModelRegistry] 从 models.dev 获取了 {} 个模型",
"[ModelRegistry] 从 models 仓库获取了 {} 个模型",
models.len()
);
Ok(models)
}
/// 合并 models.dev 数据和本地数据
fn merge_models(
/// 转换仓库模型格式为内部格式
fn convert_repo_model(
&self,
models_dev: Vec<EnhancedModelMetadata>,
local: Vec<EnhancedModelMetadata>,
) -> Vec<EnhancedModelMetadata> {
let mut merged: HashMap<String, EnhancedModelMetadata> = HashMap::new();
model: RepoModel,
provider_id: &str,
provider_name: &str,
now: i64,
) -> EnhancedModelMetadata {
let caps = model.capabilities.unwrap_or_default();
// 先添加 models.dev 数据
for model in models_dev {
merged.insert(model.id.clone(), model);
EnhancedModelMetadata {
id: model.id,
display_name: model.name,
provider_id: provider_id.to_string(),
provider_name: provider_name.to_string(),
family: model.family,
tier: model
.tier
.and_then(|t| t.parse().ok())
.unwrap_or(ModelTier::Pro),
capabilities: ModelCapabilities {
vision: caps.vision,
tools: caps.tools,
streaming: caps.streaming,
json_mode: caps.json_mode,
function_calling: caps.function_calling,
reasoning: caps.reasoning,
},
pricing: model.pricing.map(|p| ModelPricing {
input_per_million: p.input,
output_per_million: p.output,
cache_read_per_million: p.cache_read,
cache_write_per_million: p.cache_write,
currency: p.currency.unwrap_or_else(|| "USD".to_string()),
}),
limits: ModelLimits {
context_length: model.limits.as_ref().and_then(|l| l.context),
max_output_tokens: model.limits.as_ref().and_then(|l| l.max_output),
requests_per_minute: None,
tokens_per_minute: None,
},
status: model
.status
.and_then(|s| s.parse().ok())
.unwrap_or(ModelStatus::Active),
release_date: model.release_date,
is_latest: model.is_latest.unwrap_or(false),
description: model.description_zh.or(model.description),
source: ModelSource::ModelsDev,
created_at: now,
updated_at: now,
}
// 本地数据覆盖或补充
for model in local {
// 如果 models.dev 没有这个模型,或者本地数据更新,则使用本地数据
if !merged.contains_key(&model.id) {
merged.insert(model.id.clone(), model);
}
}
let mut result: Vec<_> = merged.into_values().collect();
// 按 provider_id 和 display_name 排序
result.sort_by(|a, b| {
a.provider_id
.cmp(&b.provider_id)
.then(a.display_name.cmp(&b.display_name))
});
result
}
/// 从数据库加载模型
+40 -21
View File
@@ -223,7 +223,18 @@ impl ProviderPoolService {
provider_type: &str,
model: Option<&str>,
) -> Result<Option<ProviderCredential>, String> {
let pt: PoolProviderType = provider_type.parse().map_err(|e: String| e)?;
// 对于未知的 provider_type,直接返回 None(不是错误)
// 这样可以让 select_credential_with_fallback 继续尝试智能降级
let pt: PoolProviderType = match provider_type.parse() {
Ok(pt) => pt,
Err(_) => {
eprintln!(
"[SELECT_CREDENTIAL] 未知的 provider_type '{}', 返回 None 以便智能降级",
provider_type
);
return Ok(None);
}
};
let conn = db.lock().map_err(|e| e.to_string())?;
// 获取凭证,对于 Anthropic 类型,也查找 Claude 类型的凭证
@@ -340,34 +351,42 @@ impl ProviderPoolService {
model: Option<&str>,
provider_id_hint: Option<&str>,
) -> Result<Option<ProviderCredential>, String> {
eprintln!(
"[select_credential_with_fallback] 开始: provider_type={}, model={:?}, provider_id_hint={:?}",
provider_type, model, provider_id_hint
);
// Step 1: 尝试从 Provider Pool 选择 (OAuth + API Key)
if let Some(cred) = self.select_credential(db, provider_type, model)? {
tracing::debug!(
"[凭证选择] 从 Provider Pool 找到 '{}' 凭证: {:?}",
provider_type,
eprintln!(
"[select_credential_with_fallback] 从 Provider Pool 找到凭证: {:?}",
cred.name
);
return Ok(Some(cred));
}
eprintln!("[select_credential_with_fallback] Provider Pool 未找到凭证,尝试智能降级");
// Step 2: 智能降级到 API Key Provider
let pt: PoolProviderType = provider_type.parse().unwrap_or(PoolProviderType::OpenAI);
eprintln!(
"[select_credential_with_fallback] 解析 provider_type '{}' -> {:?}",
provider_type, pt
);
// 传入 provider_id_hint 支持 60+ Provider
eprintln!("[select_credential_with_fallback] 调用 get_fallback_credential");
if let Some(cred) = api_key_service.get_fallback_credential(db, &pt, provider_id_hint)? {
eprintln!(
"[select_credential_with_fallback] 智能降级成功: {:?}",
cred.name
);
return Ok(Some(cred));
}
// Step 2: 智能降级到 API Key Provider
let pt: PoolProviderType = provider_type.parse().unwrap_or(PoolProviderType::OpenAI);
// 传入 provider_id_hint 支持 60+ Provider
if let Some(cred) = api_key_service.get_fallback_credential(db, &pt, provider_id_hint)? {
tracing::info!(
"[智能降级] Provider Pool 无 '{}' 凭证,使用 API Key Provider 降级 (provider_id: {:?})",
provider_type,
provider_id_hint
);
return Ok(Some(cred));
}
// Step 3: 都没有找到
tracing::warn!(
"[凭证选择] 未找到 '{}' 的任何可用凭证 (provider_id_hint: {:?})",
provider_type,
provider_id_hint
eprintln!(
"[select_credential_with_fallback] 未找到任何凭证 for provider_type='{}'",
provider_type
);
Ok(None)
}
+1 -1
View File
@@ -1,7 +1,7 @@
{
"$schema": "https://schema.tauri.app/config/2",
"productName": "ProxyCast",
"version": "0.31.0",
"version": "0.32.0",
"identifier": "com.proxycast.app",
"build": {
"beforeDevCommand": "npm run dev",
@@ -14,34 +14,50 @@ import { useApiKeyProvider } from "@/hooks/useApiKeyProvider";
import { useModelRegistry } from "@/hooks/useModelRegistry";
import { getDefaultProvider } from "@/hooks/useTauri";
// 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" },
// Provider type 到 registry ID 的映射(用于获取模型列表)
const getRegistryIdFromType = (providerType: string): string => {
const typeMap: Record<string, string> = {
openai: "openai",
anthropic: "anthropic",
gemini: "google",
"azure-openai": "openai",
vertexai: "google",
ollama: "ollama",
kiro: "anthropic",
claude: "anthropic",
claude_oauth: "anthropic",
qwen: "alibaba",
codex: "openai",
antigravity: "google",
iflow: "openai",
gemini_api_key: "google",
};
return typeMap[providerType.toLowerCase()] || providerType.toLowerCase();
};
// 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 的显示标签
const getProviderLabel = (providerType: string): string => {
const labelMap: Record<string, string> = {
kiro: "Kiro",
gemini: "Gemini",
qwen: "通义千问",
antigravity: "Antigravity",
codex: "Codex",
claude_oauth: "Claude OAuth",
claude: "Claude",
openai: "OpenAI",
anthropic: "Anthropic",
"azure-openai": "Azure OpenAI",
vertexai: "VertexAI",
ollama: "Ollama",
gemini_api_key: "Gemini",
iflow: "iFlow",
};
// 如果在映射表中,使用映射;否则首字母大写
return (
labelMap[providerType.toLowerCase()] ||
providerType.charAt(0).toUpperCase() + providerType.slice(1)
);
};
/** 已配置的 Provider 信息 */
@@ -49,6 +65,8 @@ interface ConfiguredProvider {
key: string;
label: string;
registryId: string;
fallbackRegistryId?: string; // 当 registryId 没有模型时的回退
type: string; // 原始 provider type,用于确定 API 协议
}
interface ChatNavbarProps {
@@ -79,7 +97,6 @@ export const ChatNavbar: React.FC<ChatNavbarProps> = ({
// 用于防止无限循环
const hasInitialized = useRef(false);
const prevProviderType = useRef(providerType);
// 获取凭证池数据
const { overview: oauthCredentials } = useProviderPool();
@@ -101,34 +118,40 @@ export const ChatNavbar: React.FC<ChatNavbarProps> = ({
// 获取模型注册表数据
const { models: registryModels } = useModelRegistry({ autoLoad: true });
// 计算已配置的 Provider 列表
// 计算已配置的 Provider 列表(完全动态,无白名单限制)
const configuredProviders = useMemo(() => {
const providerMap = new Map<string, ConfiguredProvider>();
// 从 OAuth 凭证提取 Provider
// 从 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,
const key = overview.provider_type;
if (!providerMap.has(key)) {
providerMap.set(key, {
key,
label: getProviderLabel(key),
registryId: getRegistryIdFromType(key),
type: key,
});
}
}
});
// 从 API Key Provider 提取(只包含有 API Key 的)
// 从 API Key Provider 提取(动态,支持所有自定义 Provider)
// 使用 provider.id 作为 key,确保每个 Provider 单独显示
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,
const key = provider.id; // 使用 provider.id 而不是 type 映射
if (!providerMap.has(key)) {
// 优先使用 provider.id 作为 registryId(适用于系统预设的 Provider,如 deepseek, moonshot)
// 如果模型注册表中没有该 id 的模型,则回退到使用 type 映射(适用于自定义 Provider)
providerMap.set(key, {
key,
label: provider.name, // 使用 Provider 的 name 作为显示名称
registryId: provider.id, // 先尝试用 id
fallbackRegistryId: getRegistryIdFromType(provider.type), // 回退用 type
type: provider.type,
});
}
});
@@ -142,13 +165,52 @@ export const ChatNavbar: React.FC<ChatNavbarProps> = ({
}, [configuredProviders, providerType]);
// 获取当前 Provider 的模型列表(从 model_registry 获取)
// 按照模型版本排序,最新的在前面
const currentModels = useMemo(() => {
if (!selectedProvider) return [];
// 从 model_registry 获取模型
return registryModels
// 优先使用 registryId,如果没有模型则回退到 fallbackRegistryId
let models = registryModels
.filter((m) => m.provider_id === selectedProvider.registryId)
.map((m) => m.id);
// 如果没有找到模型,尝试使用 fallbackRegistryId
if (models.length === 0 && selectedProvider.fallbackRegistryId) {
models = registryModels
.filter((m) => m.provider_id === selectedProvider.fallbackRegistryId)
.map((m) => m.id);
}
// 按照模型名称排序,优先显示最新版本
// 排序规则:
// 1. 带日期后缀的模型(如 claude-opus-4-5-20251101)按日期降序
// 2. 带 "latest" 后缀的模型排在最前面
// 3. 其他模型按字母顺序
return models.sort((a, b) => {
const aIsLatest = a.includes("-latest");
const bIsLatest = b.includes("-latest");
// latest 版本排在最前面
if (aIsLatest && !bIsLatest) return -1;
if (!aIsLatest && bIsLatest) return 1;
// 提取日期后缀(如 20251101)
const dateRegex = /-(\d{8})$/;
const aMatch = a.match(dateRegex);
const bMatch = b.match(dateRegex);
if (aMatch && bMatch) {
// 两个都有日期,按日期降序(最新的在前)
return bMatch[1].localeCompare(aMatch[1]);
}
if (aMatch && !bMatch) return -1; // 有日期的排在前面
if (!aMatch && bMatch) return 1;
// 其他情况按字母降序(通常版本号大的在前)
return b.localeCompare(a);
});
}, [selectedProvider, registryModels]);
// 初始化:优先选择服务器默认 Provider,否则选择第一个已配置的
@@ -183,16 +245,16 @@ export const ChatNavbar: React.FC<ChatNavbarProps> = ({
providerType,
]);
// 当 Provider 切换时,自动选择第一个模型
// 当 Provider 切换或模型列表变化时,自动选择第一个模型
useEffect(() => {
// 只在 Provider 真正变化时触发
if (providerType === prevProviderType.current) return;
prevProviderType.current = providerType;
if (currentModels.length > 0 && !currentModels.includes(model)) {
// 如果模型列表不为空,且当前模型为空或不在列表中,选择第一个模型
if (
currentModels.length > 0 &&
(!model || !currentModels.includes(model))
) {
setModel(currentModels[0]);
}
}, [providerType, currentModels, model, setModel]);
}, [currentModels, model, setModel]);
const selectedProviderLabel = selectedProvider?.label || providerType;
+30 -3
View File
@@ -227,6 +227,19 @@ export function ApiServerPage() {
ollama: "ollama",
};
// 根据 Provider type 获取图标类型(用于自定义 Provider)
const getIconTypeFromProviderType = (providerType: string): string => {
const typeIconMap: Record<string, string> = {
openai: "openai",
anthropic: "claude",
gemini: "gemini",
"azure-openai": "openai",
vertexai: "gemini",
ollama: "ollama",
};
return typeIconMap[providerType.toLowerCase()] || "openai";
};
const [poolOverview, setPoolOverview] = useState<ProviderPoolOverview[]>([]);
const [apiKeyProviders, setApiKeyProviders] = useState<
ProviderWithKeysDisplay[]
@@ -292,11 +305,12 @@ export function ApiServerPage() {
});
// 添加 API Key Provider 中有 API Key 的 Provider
// 使用 provider.id 作为 key,确保每个 Provider 单独显示
apiKeyProviders.forEach((provider) => {
const enabledKeys = provider.api_keys.filter((k) => k.enabled);
if (enabledKeys.length > 0 && provider.enabled) {
// 将 API Key Provider 类型映射到统一的 ID
const id = mapApiKeyProviderToId(provider.type);
// 使用 provider.id 而不是 type 映射,确保自定义 Provider 单独显示
const id = provider.id;
const existing = providerMap.get(id);
if (existing) {
existing.apiKeyCount = enabledKeys.length;
@@ -306,10 +320,15 @@ export function ApiServerPage() {
? "both"
: "api_key";
} else {
// 根据 provider.type 确定图标类型(优先使用 id 映射,否则使用 type 映射)
const iconType =
providerIconMap[id] ||
providerIconMap[provider.type] ||
getIconTypeFromProviderType(provider.type);
providerMap.set(id, {
id,
label: providerLabels[id] || provider.name,
iconType: providerIconMap[id] || "openai",
iconType,
source: "api_key",
oauthCount: 0,
apiKeyCount: enabledKeys.length,
@@ -803,7 +822,15 @@ export function ApiServerPage() {
).filter((cred) => !cred.is_disabled);
// 获取 API Key 凭证 - 查找所有映射到当前 defaultProvider 的 API Key Provider
// 支持两种匹配方式:
// 1. 通过 provider.id 直接匹配(用于自定义 Provider)
// 2. 通过 type 映射匹配(用于内置 Provider)
const matchingApiKeyProviders = apiKeyProviders.filter((p) => {
// 首先尝试直接通过 id 匹配
if (p.id === defaultProvider && p.enabled) {
return true;
}
// 然后尝试通过 type 映射匹配
const mappedId = mapApiKeyProviderToId(p.type);
return mappedId === defaultProvider && p.enabled;
});
+92 -103
View File
@@ -1,60 +1,49 @@
import { useState, useEffect } from "react";
import { useState, useEffect, useMemo } from "react";
import { Cpu, RefreshCw, Copy, Check, Search } from "lucide-react";
import { getAvailableModels, ModelInfo } from "@/hooks/useTauri";
// 模型分组配置
const MODEL_GROUPS: Record<
string,
{ name: string; color: string; models: string[] }
> = {
kiro: {
name: "Kiro Claude",
color: "bg-purple-100 text-purple-700",
models: [
"claude-sonnet-4-5",
"claude-sonnet-4-5-20250514",
"claude-sonnet-4-5-20250929",
"claude-3-7-sonnet-20250219",
"claude-3-5-sonnet-latest",
"claude-3-5-sonnet-20241022",
"claude-opus-4-5-20250514",
"claude-haiku-4-5-20250514",
],
// 根据 provider_id 获取分组配置
const PROVIDER_GROUPS: Record<string, { name: string; color: string }> = {
anthropic: {
name: "Anthropic",
color:
"bg-purple-100 text-purple-700 dark:bg-purple-900/30 dark:text-purple-300",
},
gemini: {
name: "Gemini CLI",
color: "bg-blue-100 text-blue-700",
models: [
"gemini-2.5-flash",
"gemini-2.5-flash-lite",
"gemini-2.5-pro",
"gemini-2.5-pro-preview-06-05",
"gemini-3-pro-preview",
"gemini-2.0-flash-exp",
],
},
qwen: {
name: "通义千问",
color: "bg-orange-100 text-orange-700",
models: [
"qwen3-coder-plus",
"qwen3-coder-flash",
"qwen-coder-plus",
"qwen-coder-turbo",
],
google: {
name: "Google",
color: "bg-blue-100 text-blue-700 dark:bg-blue-900/30 dark:text-blue-300",
},
openai: {
name: "OpenAI",
color: "bg-green-100 text-green-700",
models: [
"gpt-4o",
"gpt-4o-mini",
"gpt-4-turbo",
"gpt-4",
"gpt-3.5-turbo",
"o1-preview",
"o1-mini",
],
color:
"bg-green-100 text-green-700 dark:bg-green-900/30 dark:text-green-300",
},
dashscope: {
name: "阿里云",
color:
"bg-orange-100 text-orange-700 dark:bg-orange-900/30 dark:text-orange-300",
},
deepseek: {
name: "DeepSeek",
color: "bg-cyan-100 text-cyan-700 dark:bg-cyan-900/30 dark:text-cyan-300",
},
zhipu: {
name: "智谱",
color:
"bg-indigo-100 text-indigo-700 dark:bg-indigo-900/30 dark:text-indigo-300",
},
moonshot: {
name: "月之暗面",
color:
"bg-yellow-100 text-yellow-700 dark:bg-yellow-900/30 dark:text-yellow-300",
},
mistral: {
name: "Mistral",
color: "bg-red-100 text-red-700 dark:bg-red-900/30 dark:text-red-300",
},
cohere: {
name: "Cohere",
color: "bg-pink-100 text-pink-700 dark:bg-pink-900/30 dark:text-pink-300",
},
};
@@ -64,7 +53,7 @@ export function ModelsTab() {
const [error, setError] = useState<string | null>(null);
const [copied, setCopied] = useState<string | null>(null);
const [search, setSearch] = useState("");
const [selectedGroup, setSelectedGroup] = useState<string | null>(null);
const [selectedProvider, setSelectedProvider] = useState<string | null>(null);
useEffect(() => {
fetchModels();
@@ -91,61 +80,59 @@ export function ModelsTab() {
setTimeout(() => setCopied(null), 2000);
};
const getModelGroup = (modelId: string): string | null => {
for (const [groupId, group] of Object.entries(MODEL_GROUPS)) {
if (
group.models.some((m) =>
modelId.toLowerCase().includes(m.toLowerCase().split("-")[0]),
)
) {
return groupId;
}
const getProviderBadge = (providerId: string) => {
const config = PROVIDER_GROUPS[providerId];
if (!config) {
return (
<span className="rounded px-2 py-0.5 text-xs font-medium bg-gray-100 text-gray-700 dark:bg-gray-800 dark:text-gray-300">
{providerId}
</span>
);
}
// 根据 owned_by 判断
const model = models.find((m) => m.id === modelId);
if (model?.owned_by === "anthropic") return "kiro";
if (model?.owned_by === "google") return "gemini";
if (model?.owned_by === "alibaba") return "qwen";
if (model?.owned_by === "openai") return "openai";
return null;
};
const getGroupBadge = (groupId: string | null) => {
if (!groupId || !MODEL_GROUPS[groupId]) return null;
const group = MODEL_GROUPS[groupId];
return (
<span
className={`rounded px-2 py-0.5 text-xs font-medium ${group.color}`}
className={`rounded px-2 py-0.5 text-xs font-medium ${config.color}`}
>
{group.name}
{config.name}
</span>
);
};
// 过滤模型
const filteredModels = models.filter((model) => {
const matchesSearch = model.id.toLowerCase().includes(search.toLowerCase());
const matchesGroup =
!selectedGroup || getModelGroup(model.id) === selectedGroup;
return matchesSearch && matchesGroup;
});
// 按 provider 分组统计
const groupCounts = models.reduce(
(acc, model) => {
const group = getModelGroup(model.id);
if (group) {
acc[group] = (acc[group] || 0) + 1;
}
return acc;
},
{} as Record<string, number>,
);
const providerCounts = useMemo(() => {
return models.reduce(
(acc, model) => {
const provider = model.owned_by;
acc[provider] = (acc[provider] || 0) + 1;
return acc;
},
{} as Record<string, number>,
);
}, [models]);
// 获取所有 provider 列表(按数量排序)
const providers = useMemo(() => {
return Object.entries(providerCounts)
.sort((a, b) => b[1] - a[1])
.map(([id]) => id);
}, [providerCounts]);
// 过滤模型
const filteredModels = useMemo(() => {
return models.filter((model) => {
const matchesSearch = model.id
.toLowerCase()
.includes(search.toLowerCase());
const matchesProvider =
!selectedProvider || model.owned_by === selectedProvider;
return matchesSearch && matchesProvider;
});
}, [models, search, selectedProvider]);
return (
<div className="space-y-6">
{error && (
<div className="rounded-lg border border-red-500 bg-red-50 p-4 text-red-700">
<div className="rounded-lg border border-red-500 bg-red-50 p-4 text-red-700 dark:bg-red-950/30">
{error}
</div>
)}
@@ -175,31 +162,33 @@ export function ModelsTab() {
{/* Provider 过滤标签 */}
<div className="flex flex-wrap gap-2">
<button
onClick={() => setSelectedGroup(null)}
onClick={() => setSelectedProvider(null)}
className={`rounded-lg px-3 py-1.5 text-sm font-medium transition-colors ${
!selectedGroup
!selectedProvider
? "bg-primary text-primary-foreground"
: "bg-muted hover:bg-muted/80"
}`}
>
全部 ({models.length})
</button>
{Object.entries(MODEL_GROUPS).map(([groupId, group]) => {
const count = groupCounts[groupId] || 0;
if (count === 0) return null;
{providers.map((providerId) => {
const count = providerCounts[providerId] || 0;
const config = PROVIDER_GROUPS[providerId];
return (
<button
key={groupId}
key={providerId}
onClick={() =>
setSelectedGroup(selectedGroup === groupId ? null : groupId)
setSelectedProvider(
selectedProvider === providerId ? null : providerId,
)
}
className={`rounded-lg px-3 py-1.5 text-sm font-medium transition-colors ${
selectedGroup === groupId
selectedProvider === providerId
? "bg-primary text-primary-foreground"
: "bg-muted hover:bg-muted/80"
}`}
>
{group.name} ({count})
{config?.name || providerId} ({count})
</button>
);
})}
@@ -237,7 +226,7 @@ export function ModelsTab() {
<div>
<div className="flex items-center gap-2">
<code className="font-medium">{model.id}</code>
{getGroupBadge(getModelGroup(model.id))}
{getProviderBadge(model.owned_by)}
</div>
<p className="text-xs text-muted-foreground">
{model.owned_by}
-171
View File
@@ -1,171 +0,0 @@
import { useState } from "react";
import { Plus, Trash2, ArrowRight, Check, X } from "lucide-react";
import type { ModelAlias } from "@/lib/api/router";
interface ModelMappingProps {
aliases: ModelAlias[];
onAdd: (alias: string, actual: string) => Promise<void>;
onRemove: (alias: string) => Promise<void>;
loading?: boolean;
}
export function ModelMapping({
aliases,
onAdd,
onRemove,
loading,
}: ModelMappingProps) {
const [isAdding, setIsAdding] = useState(false);
const [newAlias, setNewAlias] = useState("");
const [newActual, setNewActual] = useState("");
const [addError, setAddError] = useState<string | null>(null);
const [deletingAlias, setDeletingAlias] = useState<string | null>(null);
const handleAdd = async () => {
if (!newAlias.trim() || !newActual.trim()) {
setAddError("别名和实际模型名都不能为空");
return;
}
// Check for duplicate alias
if (aliases.some((a) => a.alias === newAlias.trim())) {
setAddError("该别名已存在");
return;
}
try {
await onAdd(newAlias.trim(), newActual.trim());
setNewAlias("");
setNewActual("");
setIsAdding(false);
setAddError(null);
} catch (e) {
setAddError(e instanceof Error ? e.message : String(e));
}
};
const handleRemove = async (alias: string) => {
setDeletingAlias(alias);
try {
await onRemove(alias);
} finally {
setDeletingAlias(null);
}
};
const handleCancel = () => {
setIsAdding(false);
setNewAlias("");
setNewActual("");
setAddError(null);
};
return (
<div className="space-y-4">
<div className="flex items-center justify-between">
<div>
<h3 className="text-lg font-semibold">模型别名映射</h3>
<p className="text-sm text-muted-foreground">
定义模型别名,使用熟悉的名称映射到实际模型
</p>
</div>
{!isAdding && (
<button
onClick={() => setIsAdding(true)}
disabled={loading}
className="flex items-center gap-1 rounded-lg bg-primary px-3 py-1.5 text-sm text-primary-foreground hover:bg-primary/90 disabled:opacity-50"
>
<Plus className="h-4 w-4" />
添加别名
</button>
)}
</div>
{/* Add new alias form */}
{isAdding && (
<div className="rounded-lg border border-primary/50 bg-primary/5 p-4 space-y-3">
<div className="flex items-center gap-3">
<div className="flex-1">
<label className="text-xs text-muted-foreground mb-1 block">
别名
</label>
<input
type="text"
value={newAlias}
onChange={(e) => setNewAlias(e.target.value)}
placeholder="例如: gpt-4"
className="w-full rounded-md border bg-background px-3 py-2 text-sm focus:border-primary focus:outline-none"
autoFocus
/>
</div>
<ArrowRight className="h-5 w-5 text-muted-foreground mt-5" />
<div className="flex-1">
<label className="text-xs text-muted-foreground mb-1 block">
实际模型
</label>
<input
type="text"
value={newActual}
onChange={(e) => setNewActual(e.target.value)}
placeholder="例如: claude-sonnet-4-5-20250514"
className="w-full rounded-md border bg-background px-3 py-2 text-sm focus:border-primary focus:outline-none"
/>
</div>
</div>
{addError && <p className="text-sm text-red-500">{addError}</p>}
<div className="flex justify-end gap-2">
<button
onClick={handleCancel}
className="flex items-center gap-1 rounded-lg border px-3 py-1.5 text-sm hover:bg-muted"
>
<X className="h-4 w-4" />
取消
</button>
<button
onClick={handleAdd}
className="flex items-center gap-1 rounded-lg bg-primary px-3 py-1.5 text-sm text-primary-foreground hover:bg-primary/90"
>
<Check className="h-4 w-4" />
确认添加
</button>
</div>
</div>
)}
{/* Aliases list */}
{aliases.length === 0 && !isAdding ? (
<div className="flex flex-col items-center justify-center rounded-lg border border-dashed py-8 text-muted-foreground">
<p className="text-sm">暂无模型别名</p>
<p className="text-xs mt-1">点击"添加别名"创建第一个映射</p>
</div>
) : (
<div className="space-y-2">
{aliases.map((alias) => (
<div
key={alias.alias}
className="flex items-center justify-between rounded-lg border p-3 hover:bg-muted/50 transition-colors"
>
<div className="flex items-center gap-3 flex-1 min-w-0">
<span className="font-mono text-sm font-medium truncate bg-muted px-2 py-1 rounded">
{alias.alias}
</span>
<ArrowRight className="h-4 w-4 text-muted-foreground shrink-0" />
<span className="font-mono text-sm text-muted-foreground truncate">
{alias.actual}
</span>
</div>
<button
onClick={() => handleRemove(alias.alias)}
disabled={deletingAlias === alias.alias}
className="rounded-lg p-2 text-red-500 hover:bg-red-100 dark:hover:bg-red-900/30 disabled:opacity-50 transition-colors shrink-0"
title="删除"
>
<Trash2 className="h-4 w-4" />
</button>
</div>
))}
</div>
)}
</div>
);
}
+7 -191
View File
@@ -1,7 +1,5 @@
import { useState, useEffect, forwardRef, useImperativeHandle } from "react";
import { RefreshCw, Route, Sparkles, Check, Trash2 } from "lucide-react";
import { Modal } from "@/components/Modal";
import { ModelMapping } from "./ModelMapping";
import { RefreshCw, Route, Trash2 } from "lucide-react";
import { RoutingRules } from "./RoutingRules";
import { ExclusionList } from "./ExclusionList";
import { InjectionRules } from "./InjectionRules";
@@ -9,20 +7,14 @@ import { ClientRouting } from "./ClientRouting";
import { HelpTip } from "@/components/HelpTip";
import { routerApi } from "@/lib/api/router";
import { injectionApi } from "@/lib/api/injection";
import { setEndpointProvider } from "@/hooks/useTauri";
import type {
ModelAlias,
RoutingRule,
ProviderType,
RecommendedPreset,
} from "@/lib/api/router";
import type { RoutingRule, ProviderType } from "@/lib/api/router";
import type { InjectionRule } from "@/lib/api/injection";
export interface RoutingPageRef {
refresh: () => void;
}
type TabType = "aliases" | "rules" | "exclusions" | "injection" | "clients";
type TabType = "rules" | "exclusions" | "injection" | "clients";
interface RoutingPageProps {
hideHeader?: boolean;
@@ -35,7 +27,6 @@ export const RoutingPage = forwardRef<RoutingPageRef, RoutingPageProps>(
const [error, setError] = useState<string | null>(null);
// Data state
const [aliases, setAliases] = useState<ModelAlias[]>([]);
const [rules, setRules] = useState<RoutingRule[]>([]);
const [exclusions, setExclusions] = useState<
Record<ProviderType, string[]>
@@ -43,34 +34,19 @@ export const RoutingPage = forwardRef<RoutingPageRef, RoutingPageProps>(
const [injectionRules, setInjectionRules] = useState<InjectionRule[]>([]);
const [injectionEnabled, setInjectionEnabled] = useState(false);
// Presets state
const [presets, setPresets] = useState<RecommendedPreset[]>([]);
const [showPresets, setShowPresets] = useState(false);
const [applyingPreset, setApplyingPreset] = useState<string | null>(null);
const refresh = async () => {
setLoading(true);
setError(null);
try {
const [
aliasesData,
rulesData,
exclusionsData,
injectionConfig,
presetsData,
] = await Promise.all([
routerApi.getModelAliases(),
const [rulesData, exclusionsData, injectionConfig] = await Promise.all([
routerApi.getRoutingRules(),
routerApi.getExclusions(),
injectionApi.getInjectionConfig(),
routerApi.getRecommendedPresets(),
]);
setAliases(aliasesData);
setRules(rulesData);
setExclusions(exclusionsData);
setInjectionRules(injectionConfig.rules);
setInjectionEnabled(injectionConfig.enabled);
setPresets(presetsData);
} catch (e) {
setError(e instanceof Error ? e.message : String(e));
} finally {
@@ -78,47 +54,6 @@ export const RoutingPage = forwardRef<RoutingPageRef, RoutingPageProps>(
}
};
const handleApplyPreset = async (
presetId: string,
merge: boolean = false,
) => {
setApplyingPreset(presetId);
try {
// 先找到预设配置
const preset = presets.find((p) => p.id === presetId);
// 应用别名和规则
await routerApi.applyRecommendedPreset(presetId, merge);
// 如果预设包含客户端路由配置,也应用它
if (preset?.endpoint_providers) {
const ep = preset.endpoint_providers;
const clientTypes = [
"cursor",
"claude_code",
"codex",
"windsurf",
"kiro",
"other",
] as const;
for (const clientType of clientTypes) {
const provider = ep[clientType];
// 只有在非合并模式或有值时才设置
if (!merge || provider) {
await setEndpointProvider(clientType, provider || null);
}
}
}
await refresh();
setShowPresets(false);
} catch (e) {
setError(e instanceof Error ? e.message : String(e));
} finally {
setApplyingPreset(null);
}
};
const handleClearAll = async () => {
if (!confirm("确定要清空所有路由配置吗?此操作不可撤销。")) return;
try {
@@ -137,17 +72,6 @@ export const RoutingPage = forwardRef<RoutingPageRef, RoutingPageProps>(
refresh();
}, []);
// Alias handlers
const handleAddAlias = async (alias: string, actual: string) => {
await routerApi.addModelAlias(alias, actual);
await refresh();
};
const handleRemoveAlias = async (alias: string) => {
await routerApi.removeModelAlias(alias);
await refresh();
};
// Rule handlers
const handleAddRule = async (rule: RoutingRule) => {
await routerApi.addRoutingRule(rule);
@@ -207,7 +131,6 @@ export const RoutingPage = forwardRef<RoutingPageRef, RoutingPageProps>(
const tabs: { id: TabType; label: string; count: number }[] = [
{ id: "clients", label: "客户端路由", count: 0 },
{ id: "aliases", label: "模型别名", count: aliases.length },
{ id: "rules", label: "路由规则", count: rules.length },
{
id: "exclusions",
@@ -226,23 +149,12 @@ export const RoutingPage = forwardRef<RoutingPageRef, RoutingPageProps>(
<Route className="h-6 w-6" />
智能路由
</h2>
<p className="text-muted-foreground">
配置模型映射、路由规则和排除列表
</p>
<p className="text-muted-foreground">配置路由规则和排除列表</p>
</div>
<div className="flex items-center gap-2">
<button
onClick={() => setShowPresets(true)}
className="flex items-center gap-2 rounded-lg bg-primary px-3 py-2 text-sm text-primary-foreground hover:bg-primary/90"
>
<Sparkles className="h-4 w-4" />
推荐配置
</button>
<button
onClick={handleClearAll}
disabled={
loading || (aliases.length === 0 && rules.length === 0)
}
disabled={loading || rules.length === 0}
className="flex items-center gap-2 rounded-lg border border-red-300 px-3 py-2 text-sm text-red-600 hover:bg-red-50 disabled:opacity-50 dark:border-red-800 dark:text-red-400 dark:hover:bg-red-950/30"
>
<Trash2 className="h-4 w-4" />
@@ -265,16 +177,9 @@ export const RoutingPage = forwardRef<RoutingPageRef, RoutingPageProps>(
{/* 当隐藏标题时,显示操作按钮 */}
{hideHeader && (
<div className="flex items-center justify-end gap-2">
<button
onClick={() => setShowPresets(true)}
className="flex items-center gap-2 rounded-lg bg-primary px-3 py-2 text-sm text-primary-foreground hover:bg-primary/90"
>
<Sparkles className="h-4 w-4" />
推荐配置
</button>
<button
onClick={handleClearAll}
disabled={loading || (aliases.length === 0 && rules.length === 0)}
disabled={loading || rules.length === 0}
className="flex items-center gap-2 rounded-lg border border-red-300 px-3 py-2 text-sm text-red-600 hover:bg-red-50 disabled:opacity-50 dark:border-red-800 dark:text-red-400 dark:hover:bg-red-950/30"
>
<Trash2 className="h-4 w-4" />
@@ -293,89 +198,8 @@ export const RoutingPage = forwardRef<RoutingPageRef, RoutingPageProps>(
</div>
)}
{/* Presets Modal */}
<Modal
isOpen={showPresets}
onClose={() => setShowPresets(false)}
maxWidth="max-w-2xl"
className="max-h-[80vh] overflow-y-auto"
>
<div className="p-6">
<div className="flex items-center justify-between mb-4">
<h3 className="text-lg font-semibold flex items-center gap-2">
<Sparkles className="h-5 w-5 text-primary" />
推荐配置
</h3>
</div>
<p className="text-sm text-muted-foreground mb-4">
选择一个预设配置快速设置路由规则和模型别名
</p>
<div className="space-y-3">
{presets.map((preset) => (
<div
key={preset.id}
className="rounded-lg border p-4 hover:border-primary/50 transition-colors"
>
<div className="flex items-start justify-between">
<div className="flex-1">
<h4 className="font-medium">{preset.name}</h4>
<p className="text-sm text-muted-foreground mt-1">
{preset.description}
</p>
<div className="flex gap-4 mt-2 text-xs text-muted-foreground">
<span>{preset.aliases.length} 个别名</span>
<span>{preset.rules.length} 条规则</span>
{preset.endpoint_providers && (
<span>
{
Object.values(preset.endpoint_providers).filter(
(v) => v,
).length
}{" "}
个客户端路由
</span>
)}
</div>
</div>
<div className="flex gap-2 ml-4">
<button
onClick={() => handleApplyPreset(preset.id, true)}
disabled={applyingPreset !== null}
className="flex items-center gap-1 rounded px-3 py-1.5 text-sm border hover:bg-muted disabled:opacity-50"
>
{applyingPreset === preset.id ? (
<RefreshCw className="h-3 w-3 animate-spin" />
) : (
<Check className="h-3 w-3" />
)}
合并
</button>
<button
onClick={() => handleApplyPreset(preset.id, false)}
disabled={applyingPreset !== null}
className="flex items-center gap-1 rounded bg-primary px-3 py-1.5 text-sm text-primary-foreground hover:bg-primary/90 disabled:opacity-50"
>
{applyingPreset === preset.id ? (
<RefreshCw className="h-3 w-3 animate-spin" />
) : (
<Check className="h-3 w-3" />
)}
应用
</button>
</div>
</div>
</div>
))}
</div>
</div>
</Modal>
<HelpTip title="智能路由说明" variant="blue">
<ul className="list-disc list-inside space-y-1 text-sm text-blue-700 dark:text-blue-400">
<li>
<span className="font-medium">模型别名</span>
:使用熟悉的模型名(如 gpt-4)映射到实际模型
</li>
<li>
<span className="font-medium">路由规则</span>
:将特定模型路由到指定 Provider,支持通配符匹配
@@ -423,14 +247,6 @@ export const RoutingPage = forwardRef<RoutingPageRef, RoutingPageProps>(
</div>
) : (
<div className="py-4">
{activeTab === "aliases" && (
<ModelMapping
aliases={aliases}
onAdd={handleAddAlias}
onRemove={handleRemoveAlias}
loading={loading}
/>
)}
{activeTab === "rules" && (
<RoutingRules
rules={rules}
-1
View File
@@ -1,4 +1,3 @@
export { ModelMapping } from "./ModelMapping";
export { RoutingRules } from "./RoutingRules";
export { ExclusionList } from "./ExclusionList";
export { InjectionRules } from "./InjectionRules";
-49
View File
@@ -9,12 +9,6 @@ export type ProviderType =
| "openai"
| "claude";
// Model alias mapping
export interface ModelAlias {
alias: string;
actual: string;
}
// Routing rule
export interface RoutingRule {
pattern: string;
@@ -29,27 +23,9 @@ export interface ExclusionPattern {
pattern: string;
}
// Recommended preset
export interface RecommendedPreset {
id: string;
name: string;
description: string;
aliases: ModelAlias[];
rules: RoutingRule[];
endpoint_providers?: {
cursor?: string;
claude_code?: string;
codex?: string;
windsurf?: string;
kiro?: string;
other?: string;
};
}
// Router configuration
export interface RouterConfig {
default_provider: ProviderType;
aliases: ModelAlias[];
rules: RoutingRule[];
exclusions: Record<ProviderType, string[]>;
}
@@ -60,19 +36,6 @@ export const routerApi = {
return invoke("get_router_config");
},
// Model aliases
async addModelAlias(alias: string, actual: string): Promise<void> {
return invoke("add_model_alias", { alias, actual });
},
async removeModelAlias(alias: string): Promise<void> {
return invoke("remove_model_alias", { alias });
},
async getModelAliases(): Promise<ModelAlias[]> {
return invoke("get_model_aliases");
},
// Routing rules
async addRoutingRule(rule: RoutingRule): Promise<void> {
return invoke("add_routing_rule", { rule });
@@ -111,18 +74,6 @@ export const routerApi = {
return invoke("set_router_default_provider", { provider });
},
// Recommended presets
async getRecommendedPresets(): Promise<RecommendedPreset[]> {
return invoke("get_recommended_presets");
},
async applyRecommendedPreset(
presetId: string,
merge: boolean = false,
): Promise<void> {
return invoke("apply_recommended_preset", { presetId, merge });
},
async clearAllRoutingConfig(): Promise<void> {
return invoke("clear_all_routing_config");
},
+13
View File
@@ -11,6 +11,19 @@ if (typeof window !== "undefined") {
(window as unknown as Record<string, unknown>).React = React;
(window as unknown as Record<string, unknown>).ProxyCastPluginComponents =
PluginComponents;
// 调试:检查所有导出
console.log("[PluginComponents] 已暴露到全局变量");
console.log("[PluginComponents] 导出的键:", Object.keys(PluginComponents));
// 检查是否有 undefined 的导出
const undefinedExports = Object.entries(PluginComponents)
.filter(([, value]) => value === undefined)
.map(([key]) => key);
if (undefinedExports.length > 0) {
console.error("[PluginComponents] 以下导出是 undefined:", undefinedExports);
}
}
export {};
+2
View File
@@ -204,6 +204,8 @@ export {
Sparkles,
Cookie,
FileJson,
Code,
Bot,
} from "lucide-react";
// ============================================================================
+15
View File
@@ -143,6 +143,21 @@ export async function loadPluginUI(
`[PluginLoader] 全局变量检查: React=${typeof (window as unknown as Record<string, unknown>).React}, ProxyCastPluginComponents=${typeof (window as unknown as Record<string, unknown>).ProxyCastPluginComponents}`,
);
// 检查 ProxyCastPluginComponents 中的所有导出
const components = (window as unknown as Record<string, unknown>)
.ProxyCastPluginComponents as Record<string, unknown> | undefined;
if (components) {
const undefinedKeys = Object.keys(components).filter(
(key) => components[key] === undefined,
);
if (undefinedKeys.length > 0) {
console.error(
`[PluginLoader] ProxyCastPluginComponents 中有 undefined 的导出:`,
undefinedKeys,
);
}
}
// 执行插件代码
await executeScript(content);