feat: 从 Provider API 获取模型列表 & 修复 UTF-8 切片 panic

主要更新:
- 新增从 Provider /v1/models API 获取模型列表功能
- 当本地模型注册表为空时自动从 API 获取
- 修复 OpenAI 协议中 UTF-8 字符串切片导致的 panic
- 统一 Provider ID 与 JSON 文件名一致
- 添加数据库迁移逻辑处理旧版 Provider ID
- 清理未使用的模块 (injection, proxy, resilience, telemetry)
- 重构 workspace crates 结构

版本: 0.47.4
This commit is contained in:
coso
2026-01-16 12:46:47 +08:00
parent 99453292a9
commit 6dd714f110
69 changed files with 5001 additions and 467 deletions
+1 -1
View File
@@ -1,7 +1,7 @@
{
"name": "proxycast",
"private": true,
"version": "0.47.3",
"version": "0.47.4",
"type": "module",
"repository": {
"type": "git",
+39 -77
View File
@@ -19,21 +19,6 @@ dependencies = [
"cpufeatures",
]
[[package]]
name = "agent"
version = "0.1.0"
dependencies = [
"anyhow",
"aster",
"core",
"futures",
"serde",
"serde_json",
"thiserror 1.0.69",
"tokio",
"tracing",
]
[[package]]
name = "ahash"
version = "0.8.12"
@@ -152,22 +137,6 @@ version = "1.0.100"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a23eb6b1614318a8071c9b2521f36b424b2c83db5eb3a0fead4a6c0809af6e61"
[[package]]
name = "app"
version = "0.1.0"
dependencies = [
"agent",
"anyhow",
"core",
"providers",
"serde",
"serde_json",
"server",
"thiserror 1.0.69",
"tokio",
"tracing",
]
[[package]]
name = "arboard"
version = "3.6.1"
@@ -1812,19 +1781,6 @@ dependencies = [
"url",
]
[[package]]
name = "core"
version = "0.1.0"
dependencies = [
"anyhow",
"chrono",
"serde",
"serde_json",
"thiserror 1.0.69",
"tracing",
"uuid",
]
[[package]]
name = "core-foundation"
version = "0.9.4"
@@ -6101,24 +6057,9 @@ dependencies = [
"syn 2.0.114",
]
[[package]]
name = "providers"
version = "0.1.0"
dependencies = [
"anyhow",
"async-trait",
"core",
"reqwest 0.12.28",
"serde",
"serde_json",
"thiserror 1.0.69",
"tokio",
"tracing",
]
[[package]]
name = "proxycast"
version = "0.47.3"
version = "0.47.4"
dependencies = [
"anyhow",
"arboard",
@@ -6150,6 +6091,8 @@ dependencies = [
"parking_lot",
"portable-pty",
"proptest",
"proxycast-core",
"proxycast-infra",
"rand 0.8.5",
"regex",
"reqwest 0.12.28",
@@ -6193,6 +6136,42 @@ dependencies = [
"zip",
]
[[package]]
name = "proxycast-core"
version = "0.47.4"
dependencies = [
"chrono",
"dirs 5.0.1",
"indexmap 2.13.0",
"parking_lot",
"proptest",
"serde",
"serde_json",
"sha2",
"tracing",
"uuid",
]
[[package]]
name = "proxycast-infra"
version = "0.47.4"
dependencies = [
"chrono",
"dashmap 5.5.3",
"dirs 5.0.1",
"parking_lot",
"proptest",
"proxycast-core",
"reqwest 0.12.28",
"serde",
"serde_json",
"thiserror 1.0.69",
"tiktoken-rs",
"tokio",
"tracing",
"uuid",
]
[[package]]
name = "psl-types"
version = "2.0.11"
@@ -7361,23 +7340,6 @@ dependencies = [
"syn 2.0.114",
]
[[package]]
name = "server"
version = "0.1.0"
dependencies = [
"anyhow",
"axum 0.7.9",
"core",
"providers",
"serde",
"serde_json",
"thiserror 1.0.69",
"tokio",
"tower 0.4.13",
"tower-http 0.5.2",
"tracing",
]
[[package]]
name = "servo_arc"
version = "0.2.0"
+244 -93
View File
@@ -3,12 +3,155 @@ members = ["crates/*"]
resolver = "2"
[workspace.package]
version = "0.47.3"
version = "0.47.4"
edition = "2021"
authors = ["you"]
repository = "https://github.com/aiclientproxy/proxycast"
homepage = "https://github.com/aiclientproxy/proxycast"
[workspace.dependencies]
# 项目内 crate 依赖
proxycast-core = { path = "crates/core" }
proxycast-infra = { path = "crates/infra" }
# 序列化
serde = { version = "1", features = ["derive"] }
serde_json = "1"
serde_yaml = "0.9"
serde_urlencoded = "0.7"
# 异步运行时
tokio = { version = "1", features = ["full"] }
tokio-util = "0.7"
futures = "0.3"
async-stream = "0.3"
async-trait = "0.1"
# 错误处理
anyhow = "1"
thiserror = "1"
# 日志
tracing = "0.1"
tracing-subscriber = "0.3"
# HTTP 服务器
axum = { version = "0.7", features = ["ws"] }
axum-server = { version = "0.7", features = ["tls-rustls"] }
tower = "0.4"
tower-http = { version = "0.5", features = ["limit", "cors"] }
# HTTP 客户端
reqwest = { version = "0.12", features = ["json", "stream", "gzip", "brotli", "deflate"] }
# 数据库
rusqlite = { version = "0.31", features = ["bundled", "backup"] }
# 时间和 UUID
chrono = { version = "0.4", features = ["serde"] }
uuid = { version = "1", features = ["v4"] }
# 工具库
dirs = "5"
regex = "1"
md5 = "0.7"
urlencoding = "2"
subtle = "2.5"
flate2 = "1"
tar = "0.4"
fs2 = "0.4"
indexmap = { version = "2", features = ["serde"] }
zip = "0.6"
dashmap = "5"
notify = { version = "6", default-features = false, features = ["macos_fsevent"] }
parking_lot = "0.12"
tiktoken-rs = "0.6"
base64 = "0.22"
bytes = "1"
rand = "0.8"
sha2 = "0.10"
open = "5"
url = "2"
once_cell = "1"
arboard = "3"
glob = "0.3.3"
hex = "0.4.3"
scopeguard = "1"
sysinfo = "0.32"
whoami = "1"
# TLS
rustls-pemfile = "2"
# 终端
portable-pty = "0.8"
# SSH
ssh2 = "0.9"
openssl = { version = "0.10", features = ["vendored"] }
# 系统交互
mouse_position = "0.1.4"
window-vibrancy = "0.7.1"
if-addrs = "0.13"
# Aster Agent Framework
aster = { git = "https://github.com/astercloud/aster-rust", tag = "v0.3.0" }
# Tauri
tauri = { version = "2.5", features = ["tray-icon", "image-png", "unstable", "macos-private-api"] }
tauri-build = { version = "2", features = [] }
tauri-plugin-shell = "2.3"
tauri-plugin-autostart = "2.3"
tauri-plugin-dialog = "2.5"
tauri-plugin-single-instance = "2.3"
tauri-plugin-global-shortcut = "2.3"
# 测试
proptest = "1"
tempfile = "3"
# Windows 平台依赖
[workspace.dependencies.windows]
version = "0.56"
features = [
"Win32_Foundation",
"Win32_System_Registry",
"Win32_System_Threading",
"Win32_System_ProcessStatus",
"Win32_UI_Shell",
"Win32_UI_WindowsAndMessaging",
"Win32_System_LibraryLoader",
"Win32_System_Memory",
"Win32_System_Diagnostics_ToolHelp",
"Win32_Security",
]
[workspace.dependencies.winapi]
version = "0.3"
features = [
"winuser",
"winreg",
"processthreadsapi",
"handleapi",
"shellapi",
"psapi",
"tlhelp32",
]
[workspace.dependencies.winreg]
version = "0.52"
# macOS 平台依赖
[workspace.dependencies.cocoa]
version = "0.26"
[workspace.dependencies.objc]
version = "0.2"
[workspace.dependencies.tauri-plugin-deep-link]
version = "2.4"
[package]
name = "proxycast"
version.workspace = true
@@ -23,110 +166,118 @@ name = "proxycast_lib"
crate-type = ["lib", "cdylib", "staticlib"]
[build-dependencies]
tauri-build = { version = "2", features = [] }
tauri-build.workspace = true
[dependencies]
tauri = { version = "2.5", features = ["tray-icon", "image-png", "unstable", "macos-private-api"] }
tauri-plugin-shell = "2.3"
tauri-plugin-autostart = "2.3"
tauri-plugin-dialog = "2.5"
tauri-plugin-single-instance = "2.3"
tauri-plugin-global-shortcut = "2.3"
serde = { version = "1", features = ["derive"] }
serde_json = "1"
tokio = { version = "1", features = ["full"] }
axum = { version = "0.7", features = ["ws"] }
axum-server = { version = "0.7", features = ["tls-rustls"] }
rustls-pemfile = "2"
tower = "0.4"
tower-http = { version = "0.5", features = ["limit", "cors"] }
reqwest = { version = "0.12", features = ["json", "stream", "gzip", "brotli", "deflate"] }
uuid = { version = "1", features = ["v4"] }
chrono = { version = "0.4", features = ["serde"] }
dirs = "5"
tracing = "0.1"
tracing-subscriber = "0.3"
futures = "0.3"
async-stream = "0.3"
regex = "1"
md5 = "0.7"
urlencoding = "2"
subtle = "2.5"
flate2 = "1"
tar = "0.4"
fs2 = "0.4"
rusqlite = { version = "0.31", features = ["bundled", "backup"] }
serde_yaml = "0.9"
indexmap = { version = "2", features = ["serde"] }
zip = "0.6"
anyhow = "1"
dashmap = "5"
notify = { version = "6", default-features = false, features = ["macos_fsevent"] }
parking_lot = "0.12"
tiktoken-rs = "0.6"
async-trait = "0.1"
thiserror = "1"
base64 = "0.22"
bytes = "1"
rand = "0.8"
sha2 = "0.10"
serde_urlencoded = "0.7"
open = "5"
url = "2"
once_cell = "1"
tokio-util = "0.7"
arboard = "3"
glob = "0.3.3"
hex = "0.4.3"
portable-pty = "0.8"
scopeguard = "1"
ssh2 = "0.9"
openssl = { version = "0.10", features = ["vendored"] }
sysinfo = "0.32"
whoami = "1"
mouse_position = "0.1.4"
window-vibrancy = "0.7.1"
if-addrs = "0.13"
# 项目内 crate
proxycast-core.workspace = true
proxycast-infra.workspace = true
# Aster Agent Framework - 使用固定 tag 避免每次重新拉取
aster = { git = "https://github.com/astercloud/aster-rust", tag = "v0.3.0" }
# Tauri
tauri.workspace = true
tauri-plugin-shell.workspace = true
tauri-plugin-autostart.workspace = true
tauri-plugin-dialog.workspace = true
tauri-plugin-single-instance.workspace = true
tauri-plugin-global-shortcut.workspace = true
# Platform specific dependencies for browser interceptor
# 序列化
serde.workspace = true
serde_json.workspace = true
serde_yaml.workspace = true
serde_urlencoded.workspace = true
# 异步运行时
tokio.workspace = true
tokio-util.workspace = true
futures.workspace = true
async-stream.workspace = true
async-trait.workspace = true
# 错误处理
anyhow.workspace = true
thiserror.workspace = true
# 日志
tracing.workspace = true
tracing-subscriber.workspace = true
# HTTP 服务器
axum.workspace = true
axum-server.workspace = true
tower.workspace = true
tower-http.workspace = true
rustls-pemfile.workspace = true
# HTTP 客户端
reqwest.workspace = true
# 数据库
rusqlite.workspace = true
# 时间和 UUID
chrono.workspace = true
uuid.workspace = true
# 工具库
dirs.workspace = true
regex.workspace = true
md5.workspace = true
urlencoding.workspace = true
subtle.workspace = true
flate2.workspace = true
tar.workspace = true
fs2.workspace = true
indexmap.workspace = true
zip.workspace = true
dashmap.workspace = true
notify.workspace = true
parking_lot.workspace = true
tiktoken-rs.workspace = true
base64.workspace = true
bytes.workspace = true
rand.workspace = true
sha2.workspace = true
open.workspace = true
url.workspace = true
once_cell.workspace = true
arboard.workspace = true
glob.workspace = true
hex.workspace = true
scopeguard.workspace = true
sysinfo.workspace = true
whoami.workspace = true
# 终端
portable-pty.workspace = true
# SSH
ssh2.workspace = true
openssl.workspace = true
# 系统交互
mouse_position.workspace = true
window-vibrancy.workspace = true
if-addrs.workspace = true
# Aster Agent Framework
aster.workspace = true
# Windows specific dependencies for browser interceptor and machine ID management
[target.'cfg(windows)'.dependencies]
windows = { version = "0.56", features = [
"Win32_Foundation",
"Win32_System_Registry",
"Win32_System_Threading",
"Win32_System_ProcessStatus",
"Win32_UI_Shell",
"Win32_UI_WindowsAndMessaging",
"Win32_System_LibraryLoader",
"Win32_System_Memory",
"Win32_System_Diagnostics_ToolHelp",
"Win32_Security",
] }
winapi = { version = "0.3", features = [
"winuser",
"winreg",
"processthreadsapi",
"handleapi",
"shellapi",
"psapi",
"tlhelp32",
] }
winreg = "0.52"
windows.workspace = true
winapi.workspace = true
winreg.workspace = true
# macOS specific dependencies for browser interceptor
[target.'cfg(target_os = "macos")'.dependencies]
cocoa = "0.26"
objc = "0.2"
tauri-plugin-deep-link = "2.4"
cocoa.workspace = true
objc.workspace = true
tauri-plugin-deep-link.workspace = true
[dev-dependencies]
proptest = "1"
tempfile = "3"
proptest.workspace = true
tempfile.workspace = true
[features]
default = ["custom-protocol"]
-17
View File
@@ -1,17 +0,0 @@
[package]
name = "agent"
version = "0.1.0"
edition = "2021"
[dependencies]
core = { path = "../core" }
serde = { version = "1", features = ["derive"] }
serde_json = "1"
anyhow = "1"
thiserror = "1"
tracing = "0.1"
tokio = { version = "1", features = ["full"] }
futures = "0.3"
# Aster Agent Framework
aster = { git = "https://github.com/astercloud/aster-rust", tag = "v0.3.0" }
-7
View File
@@ -1,7 +0,0 @@
//! Aster Agent 集成模块
//!
//! 包含 agent 相关功能,依赖 aster 框架
pub fn version() -> &'static str {
env!("CARGO_PKG_VERSION")
}
-16
View File
@@ -1,16 +0,0 @@
[package]
name = "app"
version = "0.1.0"
edition = "2021"
[dependencies]
core = { path = "../core" }
providers = { path = "../providers" }
server = { path = "../server" }
agent = { path = "../agent" }
serde = { version = "1", features = ["derive"] }
serde_json = "1"
anyhow = "1"
thiserror = "1"
tracing = "0.1"
tokio = { version = "1", features = ["full"] }
-7
View File
@@ -1,7 +0,0 @@
//! Tauri 应用入口模块
//!
//! 包含 app, commands, tray, services 等功能
pub fn version() -> &'static str {
env!("CARGO_PKG_VERSION")
}
+24 -10
View File
@@ -1,13 +1,27 @@
[package]
name = "core"
version = "0.1.0"
edition = "2021"
name = "proxycast-core"
version.workspace = true
edition.workspace = true
authors.workspace = true
repository.workspace = true
[dependencies]
serde = { version = "1", features = ["derive"] }
serde_json = "1"
anyhow = "1"
thiserror = "1"
tracing = "0.1"
chrono = { version = "0.4", features = ["serde"] }
uuid = { version = "1", features = ["v4"] }
# 序列化
serde.workspace = true
serde_json.workspace = true
# 日志
tracing.workspace = true
# 时间和 UUID
chrono.workspace = true
uuid.workspace = true
# 工具库
indexmap.workspace = true
parking_lot.workspace = true
dirs.workspace = true
sha2.workspace = true
[dev-dependencies]
proptest.workspace = true
+4
View File
@@ -0,0 +1,4 @@
//! 静态数据模块
//!
//! 模型数据现在从 aiclientproxy/models 仓库获取
//! 本地硬编码数据已迁移到独立仓库: https://github.com/aiclientproxy/models
+13 -2
View File
@@ -1,6 +1,17 @@
//! 核心类型和工具模块
//! 核心类型模块
//!
//! 包含 models, config, database, logger 等基础功能
//! 包含纯数据类型(models)、静态数据(data)、日志配置(logger)
//!
//! 本 crate 不包含任何业务逻辑,只提供基础类型定义。
pub mod data;
pub mod logger;
pub mod models;
// 重新导出常用类型
pub use logger::{LogEntry, LogStore, LogStoreConfig, SharedLogStore};
pub use models::provider_type::ProviderType;
pub use models::*;
pub fn version() -> &'static str {
env!("CARGO_PKG_VERSION")
+235
View File
@@ -0,0 +1,235 @@
//! 日志管理模块
use chrono::{Duration, Local, Utc};
use serde::{Deserialize, Serialize};
use std::collections::VecDeque;
use std::fs::{self, OpenOptions};
use std::io::Write;
use std::path::PathBuf;
use std::sync::Arc;
#[derive(Debug, Clone)]
pub struct LogStoreConfig {
pub max_logs: usize,
pub retention_days: u32,
pub max_file_size: u64,
pub enable_file_logging: bool,
}
impl Default for LogStoreConfig {
fn default() -> Self {
Self {
max_logs: 1000,
retention_days: 7,
max_file_size: 10 * 1024 * 1024,
enable_file_logging: true,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct LogEntry {
pub timestamp: String,
pub level: String,
pub message: String,
}
pub struct LogStore {
logs: VecDeque<LogEntry>,
max_logs: usize,
config: LogStoreConfig,
log_file_path: Option<PathBuf>,
}
impl Default for LogStore {
fn default() -> Self {
// 默认日志文件路径: ~/.proxycast/logs/proxycast.log
let log_dir = dirs::home_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join(".proxycast")
.join("logs");
// 创建日志目录
let _ = fs::create_dir_all(&log_dir);
let log_file = log_dir.join("proxycast.log");
let config = LogStoreConfig::default();
Self {
logs: VecDeque::new(),
max_logs: config.max_logs,
config,
log_file_path: Some(log_file),
}
}
}
impl LogStore {
pub fn new() -> Self {
Self::default()
}
/// 使用自定义配置创建 LogStore
pub fn with_custom_config(retention_days: u32, enabled: bool) -> Self {
let mut store = Self::default();
store.config.retention_days = retention_days;
store.config.enable_file_logging = enabled;
store.max_logs = store.config.max_logs;
store
}
pub fn add(&mut self, level: &str, message: &str) {
let sanitized = sanitize_log_message(message);
let now = Utc::now();
let entry = LogEntry {
timestamp: now.to_rfc3339(),
level: level.to_string(),
message: sanitized.clone(),
};
self.logs.push_back(entry.clone());
// 写入日志文件
if self.config.enable_file_logging {
if let Some(ref path) = self.log_file_path {
self.rotate_log_file_if_needed(path);
let local_time = Local::now().format("%Y-%m-%d %H:%M:%S%.3f");
let log_line = format!("{} [{}] {}\n", local_time, level.to_uppercase(), sanitized);
if let Ok(mut file) = OpenOptions::new().create(true).append(true).open(path) {
let _ = file.write_all(log_line.as_bytes());
}
self.prune_old_logs(path);
}
}
// 保持日志数量在限制内
if self.logs.len() > self.max_logs {
self.logs.pop_front();
}
}
/// 记录原始响应到单独的文件(用于调试)
pub fn log_raw_response(&self, request_id: &str, body: &str) {
if let Some(ref log_path) = self.log_file_path {
let log_dir = log_path.parent().unwrap_or(std::path::Path::new("."));
let raw_file = log_dir.join(format!("raw_response_{request_id}.txt"));
let sanitized = sanitize_log_message(body);
if let Ok(mut file) = OpenOptions::new()
.create(true)
.truncate(true)
.write(true)
.open(&raw_file)
{
let _ = file.write_all(sanitized.as_bytes());
}
}
}
pub fn get_logs(&self) -> Vec<LogEntry> {
self.logs.iter().cloned().collect()
}
pub fn clear(&mut self) {
self.logs.clear();
}
pub fn get_log_file_path(&self) -> Option<String> {
self.log_file_path
.as_ref()
.map(|p| p.to_string_lossy().to_string())
}
fn rotate_log_file_if_needed(&self, path: &PathBuf) {
let Ok(metadata) = fs::metadata(path) else {
return;
};
if metadata.len() <= self.config.max_file_size {
return;
}
let suffix = Local::now().format("%Y%m%d-%H%M%S");
let rotated = path.with_file_name(format!(
"{}.{}",
path.file_name().unwrap_or_default().to_string_lossy(),
suffix
));
let _ = fs::rename(path, &rotated);
self.prune_old_logs(path);
}
fn prune_old_logs(&self, path: &PathBuf) {
let Some(dir) = path.parent() else {
return;
};
let Ok(entries) = fs::read_dir(dir) else {
return;
};
let cutoff = Utc::now() - Duration::days(self.config.retention_days as i64);
let prefix = format!(
"{}.",
path.file_name().unwrap_or_default().to_string_lossy()
);
for entry in entries.flatten() {
let file_name = entry.file_name();
let file_name = file_name.to_string_lossy();
if !file_name.starts_with(&prefix) {
continue;
}
let Ok(metadata) = entry.metadata() else {
continue;
};
let Ok(modified) = metadata.modified() else {
continue;
};
let modified = chrono::DateTime::<Utc>::from(modified);
if modified < cutoff {
let _ = fs::remove_file(entry.path());
}
}
}
}
/// 简化的共享日志存储类型(使用 parking_lot)
pub type SharedLogStore = Arc<parking_lot::RwLock<LogStore>>;
/// P2 安全修复:扩展日志脱敏规则,覆盖更多敏感字段
pub fn sanitize_log_message(message: &str) -> String {
// 简化版本:使用字符串替换而不是正则表达式
let mut sanitized = message.to_string();
// Bearer token
if let Some(pos) = sanitized.find("Bearer ") {
let start = pos + 7;
if let Some(end) =
sanitized[start..].find(|c: char| c.is_whitespace() || c == '"' || c == '\'')
{
sanitized.replace_range(start..start + end, "***");
}
}
sanitized
}
#[cfg(test)]
mod tests {
use super::sanitize_log_message;
#[test]
fn test_sanitize_bearer_token() {
let input = "Authorization: Bearer abcDEF123 end";
let output = sanitize_log_message(input);
assert!(output.contains("***"));
}
#[test]
fn test_plain_text_unchanged() {
let input = "这是一段普通日志,不包含任何敏感字段。";
let output = sanitize_log_message(input);
assert_eq!(output, input);
}
}
@@ -0,0 +1,142 @@
//! Anthropic/Claude API 数据模型
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum AnthropicContentBlock {
#[serde(rename = "text")]
Text { text: String },
#[serde(rename = "tool_use")]
ToolUse {
id: String,
name: String,
input: serde_json::Value,
},
#[serde(rename = "tool_result")]
ToolResult {
tool_use_id: String,
content: serde_json::Value,
},
#[serde(rename = "image")]
Image { source: ImageSource },
/// Extended Thinking 块
#[serde(rename = "thinking")]
Thinking {
thinking: String,
/// 签名字段,用于验证思维内容的完整性
signature: String,
},
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ImageSource {
#[serde(rename = "type")]
pub source_type: String,
pub media_type: String,
pub data: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AnthropicMessage {
pub role: String,
pub content: serde_json::Value,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AnthropicTool {
pub name: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub input_schema: Option<serde_json::Value>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AnthropicMessagesRequest {
pub model: String,
pub messages: Vec<AnthropicMessage>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_tokens: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub system: Option<serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub temperature: Option<f32>,
#[serde(default)]
pub stream: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub tools: Option<Vec<AnthropicTool>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_choice: Option<serde_json::Value>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AnthropicUsage {
pub input_tokens: u32,
pub output_tokens: u32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[allow(dead_code)]
pub struct AnthropicMessagesResponse {
pub id: String,
#[serde(rename = "type")]
pub response_type: String,
pub role: String,
pub content: Vec<AnthropicContentBlock>,
pub model: String,
pub stop_reason: Option<String>,
pub usage: AnthropicUsage,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum AnthropicStreamEvent {
#[serde(rename = "message_start")]
MessageStart { message: AnthropicMessageStart },
#[serde(rename = "content_block_start")]
ContentBlockStart {
index: u32,
content_block: AnthropicContentBlock,
},
#[serde(rename = "content_block_delta")]
ContentBlockDelta { index: u32, delta: AnthropicDelta },
#[serde(rename = "content_block_stop")]
ContentBlockStop { index: u32 },
#[serde(rename = "message_delta")]
MessageDelta {
delta: AnthropicMessageDelta,
usage: AnthropicUsage,
},
#[serde(rename = "message_stop")]
MessageStop,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AnthropicMessageStart {
pub id: String,
#[serde(rename = "type")]
pub msg_type: String,
pub role: String,
pub model: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum AnthropicDelta {
#[serde(rename = "text_delta")]
TextDelta { text: String },
#[serde(rename = "input_json_delta")]
InputJsonDelta { partial_json: String },
/// Extended Thinking delta
#[serde(rename = "thinking_delta")]
ThinkingDelta { thinking: String },
/// Signature delta for thinking blocks
#[serde(rename = "signature_delta")]
SignatureDelta { signature: String },
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AnthropicMessageDelta {
pub stop_reason: Option<String>,
}
@@ -0,0 +1,43 @@
//! 应用类型定义
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum AppType {
ProxyCast,
Claude,
Codex,
Gemini,
}
impl AppType {
pub fn as_str(&self) -> &'static str {
match self {
AppType::ProxyCast => "proxycast",
AppType::Claude => "claude",
AppType::Codex => "codex",
AppType::Gemini => "gemini",
}
}
}
impl std::str::FromStr for AppType {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.to_lowercase().as_str() {
"proxycast" => Ok(AppType::ProxyCast),
"claude" => Ok(AppType::Claude),
"codex" => Ok(AppType::Codex),
"gemini" => Ok(AppType::Gemini),
_ => Err(format!("Invalid app type: {s}")),
}
}
}
impl std::fmt::Display for AppType {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.as_str())
}
}
@@ -0,0 +1,162 @@
//! CodeWhisperer/Kiro API 数据模型
//!
//! 支持标准工具和特殊工具类型(如 web_search)。
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct CodeWhispererRequest {
pub conversation_state: ConversationState,
#[serde(skip_serializing_if = "Option::is_none")]
pub profile_arn: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ConversationState {
pub chat_trigger_type: String,
pub conversation_id: String,
pub current_message: CurrentMessage,
#[serde(skip_serializing_if = "Option::is_none")]
pub history: Option<Vec<HistoryItem>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct CurrentMessage {
pub user_input_message: UserInputMessage,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct UserInputMessage {
pub content: String,
pub model_id: String,
pub origin: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub images: Option<Vec<CWImage>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub user_input_message_context: Option<UserInputMessageContext>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct UserInputMessageContext {
#[serde(skip_serializing_if = "Option::is_none")]
pub tools: Option<Vec<CWToolItem>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_results: Option<Vec<CWToolResult>>,
}
/// CodeWhisperer 工具项
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(untagged)]
pub enum CWToolItem {
/// 标准工具定义
Standard(CWTool),
/// 联网搜索工具
WebSearch(CWWebSearchTool),
}
/// 标准工具定义
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct CWTool {
pub tool_specification: ToolSpecification,
}
/// 联网搜索工具
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CWWebSearchTool {
#[serde(rename = "type")]
pub tool_type: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct ToolSpecification {
pub name: String,
pub description: String,
pub input_schema: InputSchema,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct InputSchema {
pub json: serde_json::Value,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct CWToolResult {
pub content: Vec<CWTextContent>,
pub status: String,
pub tool_use_id: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CWTextContent {
pub text: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CWImage {
pub format: String,
pub source: CWImageSource,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CWImageSource {
pub bytes: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(untagged)]
pub enum HistoryItem {
User(UserHistoryItem),
Assistant(AssistantHistoryItem),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct UserHistoryItem {
pub user_input_message: UserInputMessage,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct AssistantHistoryItem {
pub assistant_response_message: AssistantResponseMessage,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct AssistantResponseMessage {
pub content: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_uses: Option<Vec<CWToolUse>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct CWToolUse {
pub input: serde_json::Value,
pub name: String,
pub tool_use_id: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct CWStreamEvent {
#[serde(skip_serializing_if = "Option::is_none")]
pub assistant_response_event: Option<AssistantResponseEvent>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct AssistantResponseEvent {
#[serde(skip_serializing_if = "Option::is_none")]
pub content: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_use: Option<CWToolUse>,
}
@@ -0,0 +1,95 @@
//! 参数注入类型定义
//!
//! 定义注入规则和注入模式的基础类型
use serde::{Deserialize, Serialize};
/// 注入模式
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(rename_all = "lowercase")]
pub enum InjectionMode {
/// 合并模式:不覆盖已有参数
#[default]
Merge,
/// 覆盖模式:覆盖已有参数
Override,
}
/// 注入规则
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct InjectionRule {
/// 规则 ID
pub id: String,
/// 模型匹配模式(支持通配符)
pub pattern: String,
/// 要注入的参数
pub parameters: serde_json::Value,
/// 注入模式
#[serde(default)]
pub mode: InjectionMode,
/// 优先级(数字越小优先级越高)
#[serde(default = "default_priority")]
pub priority: i32,
/// 是否启用
#[serde(default = "default_enabled")]
pub enabled: bool,
}
fn default_priority() -> i32 {
100
}
fn default_enabled() -> bool {
true
}
impl InjectionRule {
/// 创建新的注入规则
pub fn new(id: &str, pattern: &str, parameters: serde_json::Value) -> Self {
Self {
id: id.to_string(),
pattern: pattern.to_string(),
parameters,
mode: InjectionMode::Merge,
priority: default_priority(),
enabled: true,
}
}
/// 设置注入模式
pub fn with_mode(mut self, mode: InjectionMode) -> Self {
self.mode = mode;
self
}
/// 设置优先级
pub fn with_priority(mut self, priority: i32) -> Self {
self.priority = priority;
self
}
/// 检查是否为精确匹配规则
pub fn is_exact(&self) -> bool {
!self.pattern.contains('*')
}
}
/// 规则排序:精确匹配优先,然后按优先级
impl Ord for InjectionRule {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
match (self.is_exact(), other.is_exact()) {
(true, false) => return std::cmp::Ordering::Less,
(false, true) => return std::cmp::Ordering::Greater,
_ => {}
}
self.priority.cmp(&other.priority)
}
}
impl PartialOrd for InjectionRule {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl Eq for InjectionRule {}
@@ -0,0 +1,186 @@
//! Kiro 凭证指纹绑定模型
//!
//! 为每个 Kiro 凭证存储独立的 Machine ID,实现多账号指纹隔离。
#![allow(dead_code)]
use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::fs;
use std::path::PathBuf;
/// Kiro 凭证指纹绑定
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct KiroFingerprintBinding {
/// 凭证 UUID
pub credential_uuid: String,
/// 绑定的 Machine ID
pub machine_id: String,
/// 创建时间
pub created_at: DateTime<Utc>,
/// 最后切换时间
pub last_switched_at: Option<DateTime<Utc>>,
}
/// 指纹绑定存储
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct KiroFingerprintStore {
/// 凭证 UUID -> 指纹绑定
pub bindings: HashMap<String, KiroFingerprintBinding>,
}
impl KiroFingerprintStore {
/// 获取存储文件路径
pub fn get_storage_path() -> Result<PathBuf, String> {
let app_data_dir = dirs::data_dir()
.ok_or_else(|| "无法获取应用数据目录".to_string())?
.join("proxycast");
if !app_data_dir.exists() {
fs::create_dir_all(&app_data_dir)
.map_err(|e| format!("创建应用数据目录失败: {}", e))?;
}
Ok(app_data_dir.join("kiro_fingerprints.json"))
}
/// 从文件加载
pub fn load() -> Result<Self, String> {
let path = Self::get_storage_path()?;
if !path.exists() {
return Ok(Self::default());
}
let content =
fs::read_to_string(&path).map_err(|e| format!("读取指纹存储文件失败: {}", e))?;
serde_json::from_str(&content).map_err(|e| format!("解析指纹存储文件失败: {}", e))
}
/// 保存到文件
pub fn save(&self) -> Result<(), String> {
let path = Self::get_storage_path()?;
let content =
serde_json::to_string_pretty(self).map_err(|e| format!("序列化指纹存储失败: {}", e))?;
fs::write(&path, content).map_err(|e| format!("写入指纹存储文件失败: {}", e))
}
/// 获取凭证的指纹绑定
pub fn get_binding(&self, credential_uuid: &str) -> Option<&KiroFingerprintBinding> {
self.bindings.get(credential_uuid)
}
/// 获取或创建凭证的指纹绑定
pub fn get_or_create_binding(
&mut self,
credential_uuid: &str,
profile_arn: Option<&str>,
client_id: Option<&str>,
) -> Result<&KiroFingerprintBinding, String> {
if !self.bindings.contains_key(credential_uuid) {
let machine_id = generate_stable_machine_id(credential_uuid, profile_arn, client_id);
let binding = KiroFingerprintBinding {
credential_uuid: credential_uuid.to_string(),
machine_id,
created_at: Utc::now(),
last_switched_at: None,
};
self.bindings.insert(credential_uuid.to_string(), binding);
self.save()?;
}
Ok(self.bindings.get(credential_uuid).unwrap())
}
/// 更新最后切换时间
pub fn update_last_switched(&mut self, credential_uuid: &str) -> Result<(), String> {
if let Some(binding) = self.bindings.get_mut(credential_uuid) {
binding.last_switched_at = Some(Utc::now());
self.save()?;
}
Ok(())
}
/// 删除凭证的指纹绑定
pub fn remove_binding(&mut self, credential_uuid: &str) -> Result<(), String> {
self.bindings.remove(credential_uuid);
self.save()
}
}
/// 生成稳定的 Machine ID
fn generate_stable_machine_id(
credential_uuid: &str,
profile_arn: Option<&str>,
client_id: Option<&str>,
) -> String {
use sha2::{Digest, Sha256};
let seed = format!(
"kiro_fingerprint:{}:{}:{}",
credential_uuid,
profile_arn.unwrap_or(""),
client_id.unwrap_or("")
);
let mut hasher = Sha256::new();
hasher.update(seed.as_bytes());
let result = hasher.finalize();
let hex = format!("{:x}", result);
format!(
"{}-{}-{}-{}-{}",
&hex[0..8],
&hex[8..12],
&hex[12..16],
&hex[16..20],
&hex[20..32]
)
}
/// 切换到本地的结果
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SwitchToLocalResult {
pub success: bool,
pub message: String,
pub requires_action: bool,
pub machine_id: Option<String>,
pub requires_kiro_restart: bool,
}
impl SwitchToLocalResult {
pub fn success(message: impl Into<String>, machine_id: String) -> Self {
Self {
success: true,
message: message.into(),
requires_action: false,
machine_id: Some(machine_id),
requires_kiro_restart: true,
}
}
pub fn error(message: impl Into<String>) -> Self {
Self {
success: false,
message: message.into(),
requires_action: false,
machine_id: None,
requires_kiro_restart: false,
}
}
pub fn requires_admin(message: impl Into<String>) -> Self {
Self {
success: false,
message: message.into(),
requires_action: true,
machine_id: None,
requires_kiro_restart: false,
}
}
}
@@ -0,0 +1,170 @@
//! 机器码相关数据模型
use serde::{Deserialize, Serialize};
/// 机器码信息结构
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MachineIdInfo {
pub current_id: String,
pub original_id: Option<String>,
pub platform: String,
pub can_modify: bool,
pub requires_admin: bool,
pub backup_exists: bool,
pub format_type: MachineIdFormat,
}
/// 机器码操作结果
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MachineIdResult {
pub success: bool,
pub message: String,
pub requires_restart: bool,
pub requires_admin: bool,
pub new_machine_id: Option<String>,
}
/// 管理员权限状态
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AdminStatus {
pub is_admin: bool,
pub platform: String,
pub elevation_method: Option<String>,
pub check_success: bool,
}
/// 机器码格式类型
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum MachineIdFormat {
Uuid,
#[serde(rename = "hex32")]
Hex32,
#[serde(rename = "unknown")]
Unknown,
}
/// 机器码备份信息
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MachineIdBackup {
pub machine_id: String,
pub timestamp: i64,
pub platform: String,
pub format: MachineIdFormat,
pub description: Option<String>,
}
/// 机器码历史记录
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MachineIdHistory {
pub id: String,
pub machine_id: String,
pub timestamp: String,
pub platform: String,
pub backup_path: Option<String>,
}
/// 机器码操作类型
#[derive(Debug, Clone, Serialize, Deserialize)]
#[allow(dead_code)]
pub enum MachineIdOperation {
Get,
Set,
Generate,
Backup,
Restore,
Reset,
}
/// 机器码验证结果
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct MachineIdValidation {
pub is_valid: bool,
pub detected_format: MachineIdFormat,
pub error_message: Option<String>,
pub formatted_id: Option<String>,
}
impl MachineIdFormat {
/// 从字符串检测机器码格式
pub fn detect(machine_id: &str) -> Self {
let cleaned = machine_id.replace("-", "").replace(" ", "").to_lowercase();
if machine_id.contains("-") && machine_id.len() == 36 {
let parts: Vec<&str> = machine_id.split('-').collect();
if parts.len() == 5
&& parts[0].len() == 8
&& parts[1].len() == 4
&& parts[2].len() == 4
&& parts[3].len() == 4
&& parts[4].len() == 12
&& cleaned.chars().all(|c| c.is_ascii_hexdigit())
{
return MachineIdFormat::Uuid;
}
}
if cleaned.len() == 32 && cleaned.chars().all(|c| c.is_ascii_hexdigit()) {
return MachineIdFormat::Hex32;
}
MachineIdFormat::Unknown
}
/// 格式化机器码为标准格式
pub fn format_machine_id(&self, machine_id: &str) -> Result<String, String> {
let cleaned = machine_id.replace("-", "").replace(" ", "").to_lowercase();
match self {
MachineIdFormat::Uuid => {
if cleaned.len() != 32 {
return Err("UUID format requires 32 hex characters".to_string());
}
if !cleaned.chars().all(|c| c.is_ascii_hexdigit()) {
return Err("UUID format requires hex characters only".to_string());
}
Ok(format!(
"{}-{}-{}-{}-{}",
&cleaned[0..8],
&cleaned[8..12],
&cleaned[12..16],
&cleaned[16..20],
&cleaned[20..32]
))
}
MachineIdFormat::Hex32 => {
if cleaned.len() != 32 {
return Err("Hex32 format requires 32 hex characters".to_string());
}
if !cleaned.chars().all(|c| c.is_ascii_hexdigit()) {
return Err("Hex32 format requires hex characters only".to_string());
}
Ok(cleaned)
}
MachineIdFormat::Unknown => Err("Cannot format unknown machine ID format".to_string()),
}
}
}
impl std::fmt::Display for MachineIdFormat {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
MachineIdFormat::Uuid => write!(f, "uuid"),
MachineIdFormat::Hex32 => write!(f, "hex32"),
MachineIdFormat::Unknown => write!(f, "unknown"),
}
}
}
impl std::fmt::Display for MachineIdOperation {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
MachineIdOperation::Get => write!(f, "Get"),
MachineIdOperation::Set => write!(f, "Set"),
MachineIdOperation::Generate => write!(f, "Generate"),
MachineIdOperation::Backup => write!(f, "Backup"),
MachineIdOperation::Restore => write!(f, "Restore"),
MachineIdOperation::Reset => write!(f, "Reset"),
}
}
}
@@ -0,0 +1,40 @@
//! MCP Server 数据模型
use serde::{Deserialize, Serialize};
use serde_json::Value;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct McpServer {
pub id: String,
pub name: String,
pub server_config: Value,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(default)]
pub enabled_proxycast: bool,
#[serde(default)]
pub enabled_claude: bool,
#[serde(default)]
pub enabled_codex: bool,
#[serde(default)]
pub enabled_gemini: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub created_at: Option<i64>,
}
impl McpServer {
#[allow(dead_code)]
pub fn new(id: String, name: String, server_config: Value) -> Self {
Self {
id,
name,
server_config,
description: None,
enabled_proxycast: false,
enabled_claude: false,
enabled_codex: false,
enabled_gemini: false,
created_at: Some(chrono::Utc::now().timestamp()),
}
}
}
+35
View File
@@ -0,0 +1,35 @@
//! 数据模型模块
//!
//! 包含 ProxyCast 的所有核心数据模型定义。
pub mod anthropic;
pub mod app_type;
pub mod codewhisperer;
pub mod injection_types;
pub mod kiro_fingerprint;
pub mod machine_id;
pub mod mcp_model;
pub mod model_registry;
pub mod openai;
pub mod prompt_model;
pub mod provider_model;
pub mod provider_pool_model;
pub mod provider_type;
pub mod route_model;
pub mod skill_model;
#[allow(unused_imports)]
pub use anthropic::*;
pub use app_type::AppType;
#[allow(unused_imports)]
pub use codewhisperer::*;
pub use injection_types::{InjectionMode, InjectionRule};
pub use mcp_model::McpServer;
#[allow(unused_imports)]
pub use openai::*;
pub use prompt_model::Prompt;
pub use provider_model::Provider;
#[allow(unused_imports)]
pub use provider_pool_model::*;
pub use provider_type::ProviderType;
pub use skill_model::{Skill, SkillMetadata, SkillRepo, SkillState, SkillStates};
@@ -0,0 +1,578 @@
//! 模型注册表数据结构
use serde::{Deserialize, Serialize};
/// 模型能力
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct ModelCapabilities {
pub vision: bool,
pub tools: bool,
pub streaming: bool,
pub json_mode: bool,
pub function_calling: bool,
pub reasoning: bool,
}
/// 模型定价
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelPricing {
pub input_per_million: Option<f64>,
pub output_per_million: Option<f64>,
pub cache_read_per_million: Option<f64>,
pub cache_write_per_million: Option<f64>,
pub currency: String,
}
impl Default for ModelPricing {
fn default() -> Self {
Self {
input_per_million: None,
output_per_million: None,
cache_read_per_million: None,
cache_write_per_million: None,
currency: "USD".to_string(),
}
}
}
/// 模型限制
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
pub struct ModelLimits {
pub context_length: Option<u32>,
pub max_output_tokens: Option<u32>,
pub requests_per_minute: Option<u32>,
pub tokens_per_minute: Option<u32>,
}
/// 模型状态
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum ModelStatus {
Active,
Preview,
Alpha,
Beta,
Deprecated,
Legacy,
}
impl Default for ModelStatus {
fn default() -> Self {
Self::Active
}
}
impl std::fmt::Display for ModelStatus {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Active => write!(f, "active"),
Self::Preview => write!(f, "preview"),
Self::Alpha => write!(f, "alpha"),
Self::Beta => write!(f, "beta"),
Self::Deprecated => write!(f, "deprecated"),
Self::Legacy => write!(f, "legacy"),
}
}
}
impl std::str::FromStr for ModelStatus {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.to_lowercase().as_str() {
"active" => Ok(Self::Active),
"preview" => Ok(Self::Preview),
"alpha" => Ok(Self::Alpha),
"beta" => Ok(Self::Beta),
"deprecated" => Ok(Self::Deprecated),
"legacy" => Ok(Self::Legacy),
_ => Err(format!("Unknown model status: {}", s)),
}
}
}
/// 模型服务等级
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum ModelTier {
Mini,
Pro,
Max,
}
impl Default for ModelTier {
fn default() -> Self {
Self::Pro
}
}
impl std::fmt::Display for ModelTier {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Mini => write!(f, "mini"),
Self::Pro => write!(f, "pro"),
Self::Max => write!(f, "max"),
}
}
}
impl std::str::FromStr for ModelTier {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.to_lowercase().as_str() {
"mini" => Ok(Self::Mini),
"pro" => Ok(Self::Pro),
"max" => Ok(Self::Max),
_ => Err(format!("Unknown model tier: {}", s)),
}
}
}
/// 模型数据来源
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum ModelSource {
Embedded,
ModelsDev,
Local,
Custom,
}
impl Default for ModelSource {
fn default() -> Self {
Self::Local
}
}
impl std::fmt::Display for ModelSource {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Embedded => write!(f, "embedded"),
Self::ModelsDev => write!(f, "models.dev"),
Self::Local => write!(f, "local"),
Self::Custom => write!(f, "custom"),
}
}
}
impl std::str::FromStr for ModelSource {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.to_lowercase().as_str() {
"embedded" => Ok(Self::Embedded),
"models.dev" | "modelsdev" => Ok(Self::ModelsDev),
"local" => Ok(Self::Local),
"custom" => Ok(Self::Custom),
_ => Err(format!("Unknown model source: {}", s)),
}
}
}
/// 增强的模型元数据
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct EnhancedModelMetadata {
pub id: String,
pub display_name: String,
pub provider_id: String,
pub provider_name: String,
pub family: Option<String>,
pub tier: ModelTier,
pub capabilities: ModelCapabilities,
pub pricing: Option<ModelPricing>,
pub limits: ModelLimits,
pub status: ModelStatus,
pub release_date: Option<String>,
pub is_latest: bool,
pub description: Option<String>,
pub source: ModelSource,
pub created_at: i64,
pub updated_at: i64,
}
impl EnhancedModelMetadata {
pub fn new(
id: String,
display_name: String,
provider_id: String,
provider_name: String,
) -> Self {
let now = chrono::Utc::now().timestamp();
Self {
id,
display_name,
provider_id,
provider_name,
family: None,
tier: ModelTier::Pro,
capabilities: ModelCapabilities::default(),
pricing: None,
limits: ModelLimits::default(),
status: ModelStatus::Active,
release_date: None,
is_latest: false,
description: None,
source: ModelSource::Local,
created_at: now,
updated_at: now,
}
}
pub fn with_family(mut self, family: impl Into<String>) -> Self {
self.family = Some(family.into());
self
}
pub fn with_tier(mut self, tier: ModelTier) -> Self {
self.tier = tier;
self
}
pub fn with_capabilities(mut self, capabilities: ModelCapabilities) -> Self {
self.capabilities = capabilities;
self
}
pub fn with_pricing(mut self, pricing: ModelPricing) -> Self {
self.pricing = Some(pricing);
self
}
pub fn with_limits(mut self, limits: ModelLimits) -> Self {
self.limits = limits;
self
}
pub fn with_status(mut self, status: ModelStatus) -> Self {
self.status = status;
self
}
pub fn with_release_date(mut self, date: impl Into<String>) -> Self {
self.release_date = Some(date.into());
self
}
pub fn with_is_latest(mut self, is_latest: bool) -> Self {
self.is_latest = is_latest;
self
}
pub fn with_description(mut self, description: impl Into<String>) -> Self {
self.description = Some(description.into());
self
}
pub fn with_source(mut self, source: ModelSource) -> Self {
self.source = source;
self
}
}
/// 用户模型偏好
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct UserModelPreference {
pub model_id: String,
pub is_favorite: bool,
pub is_hidden: bool,
pub custom_alias: Option<String>,
pub usage_count: u32,
pub last_used_at: Option<i64>,
pub created_at: i64,
pub updated_at: i64,
}
impl UserModelPreference {
pub fn new(model_id: String) -> Self {
let now = chrono::Utc::now().timestamp();
Self {
model_id,
is_favorite: false,
is_hidden: false,
custom_alias: None,
usage_count: 0,
last_used_at: None,
created_at: now,
updated_at: now,
}
}
}
/// 模型同步状态
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelSyncState {
pub last_sync_at: Option<i64>,
pub model_count: u32,
pub is_syncing: bool,
pub last_error: Option<String>,
}
impl Default for ModelSyncState {
fn default() -> Self {
Self {
last_sync_at: None,
model_count: 0,
is_syncing: false,
last_error: None,
}
}
}
// Provider Alias 相关类型
/// 单个模型别名映射
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelAlias {
pub actual: String,
pub internal_name: Option<String>,
pub provider: Option<String>,
pub description: Option<String>,
}
/// Provider 的别名配置
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ProviderAliasConfig {
pub provider: String,
pub description: Option<String>,
#[serde(default)]
pub models: Vec<String>,
pub aliases: std::collections::HashMap<String, ModelAlias>,
pub updated_at: Option<String>,
}
impl ProviderAliasConfig {
pub fn supports_model(&self, model: &str) -> bool {
self.models.contains(&model.to_string()) || self.aliases.contains_key(model)
}
pub fn get_internal_name(&self, model: &str) -> Option<&str> {
self.aliases
.get(model)
.and_then(|a| a.internal_name.as_deref())
}
pub fn get_actual_model(&self, model: &str) -> Option<&str> {
self.aliases.get(model).map(|a| a.actual.as_str())
}
}
/// models.dev API 响应中的 Provider 结构
#[derive(Debug, Clone, Serialize, Deserialize)]
#[allow(dead_code)]
pub struct ModelsDevProvider {
pub id: String,
pub name: String,
#[serde(default)]
pub api: Option<String>,
#[serde(default)]
pub npm: Option<String>,
#[serde(default)]
pub models: std::collections::HashMap<String, ModelsDevModel>,
}
/// models.dev API 响应中的 Model 结构
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelsDevModel {
pub id: String,
pub name: String,
#[serde(default)]
pub family: Option<String>,
#[serde(default)]
pub release_date: Option<String>,
#[serde(default)]
pub attachment: bool,
#[serde(default)]
pub reasoning: bool,
#[serde(default)]
pub temperature: bool,
#[serde(default)]
pub tool_call: bool,
#[serde(default)]
pub cost: Option<ModelsDevCost>,
#[serde(default)]
pub limit: Option<ModelsDevLimit>,
#[serde(default)]
pub modalities: Option<ModelsDevModalities>,
#[serde(default)]
pub experimental: Option<bool>,
#[serde(default)]
pub status: Option<String>,
}
/// models.dev API 响应中的 Cost 结构
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelsDevCost {
#[serde(default)]
pub input: Option<f64>,
#[serde(default)]
pub output: Option<f64>,
#[serde(default)]
pub cache_read: Option<f64>,
#[serde(default)]
pub cache_write: Option<f64>,
}
/// models.dev API 响应中的 Limit 结构
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelsDevLimit {
#[serde(default)]
pub context: Option<u32>,
#[serde(default)]
pub output: Option<u32>,
}
/// models.dev API 响应中的 Modalities 结构
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelsDevModalities {
#[serde(default)]
pub input: Vec<String>,
#[serde(default)]
pub output: Vec<String>,
}
impl ModelsDevModel {
#[allow(dead_code)]
pub fn to_enhanced_metadata(
&self,
provider_id: &str,
provider_name: &str,
) -> EnhancedModelMetadata {
let now = chrono::Utc::now().timestamp();
let supports_vision = self
.modalities
.as_ref()
.map(|m| m.input.iter().any(|i| i == "image" || i == "video"))
.unwrap_or(false)
|| self.attachment;
let tier = infer_model_tier(&self.id, &self.name);
let status = self
.status
.as_ref()
.and_then(|s| s.parse().ok())
.unwrap_or(ModelStatus::Active);
let is_latest = self.id.contains("latest");
EnhancedModelMetadata {
id: self.id.clone(),
display_name: self.name.clone(),
provider_id: provider_id.to_string(),
provider_name: provider_name.to_string(),
family: self.family.clone(),
tier,
capabilities: ModelCapabilities {
vision: supports_vision,
tools: self.tool_call,
streaming: true,
json_mode: true,
function_calling: self.tool_call,
reasoning: self.reasoning,
},
pricing: self.cost.as_ref().map(|c| ModelPricing {
input_per_million: c.input,
output_per_million: c.output,
cache_read_per_million: c.cache_read,
cache_write_per_million: c.cache_write,
currency: "USD".to_string(),
}),
limits: ModelLimits {
context_length: self.limit.as_ref().and_then(|l| l.context),
max_output_tokens: self.limit.as_ref().and_then(|l| l.output),
requests_per_minute: None,
tokens_per_minute: None,
},
status,
release_date: self.release_date.clone(),
is_latest,
description: None,
source: ModelSource::ModelsDev,
created_at: now,
updated_at: now,
}
}
}
/// 根据模型 ID 和名称推断服务等级
#[allow(dead_code)]
fn infer_model_tier(model_id: &str, model_name: &str) -> ModelTier {
let id_lower = model_id.to_lowercase();
let name_lower = model_name.to_lowercase();
let max_patterns = [
"opus",
"gpt-4o",
"gpt-4-turbo",
"gemini-2.5-pro",
"gemini-ultra",
"claude-3-opus",
"qwen-max",
"glm-4-plus",
"deepseek-v3",
];
for pattern in max_patterns {
if id_lower.contains(pattern) || name_lower.contains(pattern) {
return ModelTier::Max;
}
}
let mini_patterns = [
"mini",
"nano",
"lite",
"flash",
"haiku",
"gpt-4o-mini",
"gemini-flash",
"qwen-turbo",
"glm-4-flash",
];
for pattern in mini_patterns {
if id_lower.contains(pattern) || name_lower.contains(pattern) {
return ModelTier::Mini;
}
}
ModelTier::Pro
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_model_tier_inference() {
assert_eq!(
infer_model_tier("claude-opus-4-5-20250514", "Claude Opus 4.5"),
ModelTier::Max
);
assert_eq!(
infer_model_tier("gpt-4o-mini", "GPT-4o Mini"),
ModelTier::Mini
);
assert_eq!(
infer_model_tier("claude-sonnet-4-5", "Claude Sonnet 4.5"),
ModelTier::Pro
);
assert_eq!(
infer_model_tier("gemini-2.5-flash", "Gemini 2.5 Flash"),
ModelTier::Mini
);
}
#[test]
fn test_model_status_parsing() {
assert_eq!(
"active".parse::<ModelStatus>().unwrap(),
ModelStatus::Active
);
assert_eq!(
"deprecated".parse::<ModelStatus>().unwrap(),
ModelStatus::Deprecated
);
assert_eq!("beta".parse::<ModelStatus>().unwrap(), ModelStatus::Beta);
}
}
+259
View File
@@ -0,0 +1,259 @@
//! OpenAI API 数据模型
//!
//! 支持标准 OpenAI 格式以及扩展的工具类型(如 web_search)。
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ImageUrl {
pub url: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub detail: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum ContentPart {
#[serde(rename = "text")]
Text { text: String },
#[serde(rename = "image_url")]
ImageUrl { image_url: ImageUrl },
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ToolCall {
pub id: String,
#[serde(rename = "type")]
pub call_type: String,
pub function: FunctionCall,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FunctionCall {
pub name: String,
pub arguments: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(untagged)]
pub enum MessageContent {
Text(String),
Parts(Vec<ContentPart>),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ChatMessage {
pub role: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub content: Option<MessageContent>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_calls: Option<Vec<ToolCall>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_call_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub reasoning_content: Option<String>,
}
impl ChatMessage {
pub fn get_content_text(&self) -> String {
match &self.content {
Some(MessageContent::Text(s)) => s.clone(),
Some(MessageContent::Parts(parts)) => parts
.iter()
.filter_map(|p| {
if let ContentPart::Text { text } = p {
Some(text.clone())
} else {
None
}
})
.collect::<Vec<_>>()
.join(""),
None => String::new(),
}
}
/// 提取消息中的图片 URL 列表
pub fn get_images(&self) -> Vec<(String, String)> {
match &self.content {
Some(MessageContent::Parts(parts)) => parts
.iter()
.filter_map(|p| {
if let ContentPart::ImageUrl { image_url } = p {
if image_url.url.starts_with("data:") {
let parts: Vec<&str> = image_url.url.splitn(2, ',').collect();
if parts.len() == 2 {
let header = parts[0];
let data = parts[1];
let media_type = header
.strip_prefix("data:")
.and_then(|s| s.split(';').next())
.unwrap_or("image/jpeg");
let format =
media_type.split('/').nth(1).unwrap_or("jpeg").to_string();
return Some((format, data.to_string()));
}
}
None
} else {
None
}
})
.collect(),
_ => Vec::new(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FunctionDef {
pub name: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub parameters: Option<serde_json::Value>,
}
/// 工具定义
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum Tool {
#[serde(rename = "function")]
Function { function: FunctionDef },
#[serde(rename = "web_search")]
WebSearch,
#[serde(rename = "web_search_20250305")]
WebSearch20250305,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ChatCompletionRequest {
pub model: String,
pub messages: Vec<ChatMessage>,
#[serde(skip_serializing_if = "Option::is_none")]
pub temperature: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_tokens: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub top_p: Option<f32>,
#[serde(default)]
pub stream: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub tools: Option<Vec<Tool>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_choice: Option<serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub reasoning_effort: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Usage {
pub prompt_tokens: u32,
pub completion_tokens: u32,
pub total_tokens: u32,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ResponseMessage {
pub role: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub content: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_calls: Option<Vec<ToolCall>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Choice {
pub index: u32,
pub message: ResponseMessage,
pub finish_reason: String,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ChatCompletionResponse {
pub id: String,
pub object: String,
pub created: u64,
pub model: String,
pub choices: Vec<Choice>,
pub usage: Usage,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StreamDelta {
#[serde(skip_serializing_if = "Option::is_none")]
pub role: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub content: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_calls: Option<Vec<ToolCall>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StreamChoice {
pub index: u32,
pub delta: StreamDelta,
#[serde(skip_serializing_if = "Option::is_none")]
pub finish_reason: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ChatCompletionChunk {
pub id: String,
pub object: String,
pub created: u64,
pub model: String,
pub choices: Vec<StreamChoice>,
}
// 图像生成 API 数据模型
/// OpenAI 图像生成请求
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ImageGenerationRequest {
pub prompt: String,
#[serde(default = "default_image_model")]
pub model: String,
#[serde(default = "default_n")]
pub n: u32,
#[serde(skip_serializing_if = "Option::is_none")]
pub size: Option<String>,
#[serde(default = "default_response_format")]
pub response_format: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub quality: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub style: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub user: Option<String>,
}
fn default_image_model() -> String {
"gemini-3-pro-image-preview".to_string()
}
fn default_n() -> u32 {
1
}
fn default_response_format() -> String {
"url".to_string()
}
/// OpenAI 图像生成响应
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ImageGenerationResponse {
pub created: i64,
pub data: Vec<ImageData>,
}
/// 单个图像数据
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ImageData {
#[serde(skip_serializing_if = "Option::is_none")]
pub b64_json: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub url: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub revised_prompt: Option<String>,
}
@@ -0,0 +1,36 @@
//! Prompt 数据模型
use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Prompt {
pub id: String,
pub app_type: String,
pub name: String,
pub content: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(default)]
pub enabled: bool,
#[serde(rename = "createdAt", skip_serializing_if = "Option::is_none")]
pub created_at: Option<i64>,
#[serde(rename = "updatedAt", skip_serializing_if = "Option::is_none")]
pub updated_at: Option<i64>,
}
impl Prompt {
#[allow(dead_code)]
pub fn new(id: String, app_type: String, name: String, content: String) -> Self {
let now = chrono::Utc::now().timestamp();
Self {
id,
app_type,
name,
content,
description: None,
enabled: false,
created_at: Some(now),
updated_at: Some(now),
}
}
}
@@ -0,0 +1,45 @@
//! Provider 数据模型
use serde::{Deserialize, Serialize};
use serde_json::Value;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Provider {
pub id: String,
pub app_type: String,
pub name: String,
pub settings_config: Value,
#[serde(skip_serializing_if = "Option::is_none")]
pub category: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub icon: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub icon_color: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub notes: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub created_at: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub sort_index: Option<i32>,
#[serde(default)]
pub is_current: bool,
}
impl Provider {
#[allow(dead_code)]
pub fn new(id: String, app_type: String, name: String, settings_config: Value) -> Self {
Self {
id,
app_type,
name,
settings_config,
category: None,
icon: None,
icon_color: None,
notes: None,
created_at: Some(chrono::Utc::now().timestamp()),
sort_index: None,
is_current: false,
}
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,163 @@
//! Provider 类型定义
//!
//! 包含 Provider 类型枚举和相关实现。
use serde::{Deserialize, Serialize};
/// Provider 类型枚举
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum ProviderType {
Kiro,
Gemini,
#[serde(rename = "openai")]
OpenAI,
Claude,
Antigravity,
Vertex,
#[serde(rename = "gemini_api_key")]
GeminiApiKey,
Codex,
#[serde(rename = "claude_oauth")]
ClaudeOAuth,
// API Key Provider 类型
Anthropic,
#[serde(rename = "azure_openai")]
AzureOpenai,
#[serde(rename = "aws_bedrock")]
AwsBedrock,
Ollama,
}
impl std::fmt::Display for ProviderType {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ProviderType::Kiro => write!(f, "kiro"),
ProviderType::Gemini => write!(f, "gemini"),
ProviderType::OpenAI => write!(f, "openai"),
ProviderType::Claude => write!(f, "claude"),
ProviderType::Antigravity => write!(f, "antigravity"),
ProviderType::Vertex => write!(f, "vertex"),
ProviderType::GeminiApiKey => write!(f, "gemini_api_key"),
ProviderType::Codex => write!(f, "codex"),
ProviderType::ClaudeOAuth => write!(f, "claude_oauth"),
ProviderType::Anthropic => write!(f, "anthropic"),
ProviderType::AzureOpenai => write!(f, "azure_openai"),
ProviderType::AwsBedrock => write!(f, "aws_bedrock"),
ProviderType::Ollama => write!(f, "ollama"),
}
}
}
impl std::str::FromStr for ProviderType {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.to_lowercase().as_str() {
"kiro" => Ok(ProviderType::Kiro),
"gemini" => Ok(ProviderType::Gemini),
"openai" => Ok(ProviderType::OpenAI),
"claude" => Ok(ProviderType::Claude),
"antigravity" => Ok(ProviderType::Antigravity),
"vertex" => Ok(ProviderType::Vertex),
"gemini_api_key" => Ok(ProviderType::GeminiApiKey),
"codex" => Ok(ProviderType::Codex),
"claude_oauth" => Ok(ProviderType::ClaudeOAuth),
"anthropic" => Ok(ProviderType::Anthropic),
"azure_openai" | "azure-openai" => Ok(ProviderType::AzureOpenai),
"aws_bedrock" | "aws-bedrock" => Ok(ProviderType::AwsBedrock),
"ollama" => Ok(ProviderType::Ollama),
_ => Err(format!("Invalid provider: {s}")),
}
}
}
/// Antigravity 支持的模型列表(fallback,当无法从 models 仓库获取时使用)
pub const ANTIGRAVITY_MODELS_FALLBACK: &[&str] = &[
"gemini-2.5-computer-use-preview-10-2025",
"gemini-3-pro-image-preview",
"gemini-3-pro-preview",
"gemini-3-flash-preview",
"gemini-2.5-flash-preview",
"gemini-2.5-flash",
"gemini-2.5-pro",
"gemini-3-flash",
"gemini-3-pro-high",
"gemini-3-pro-low",
"gemini-claude-sonnet-4-5",
"gemini-claude-sonnet-4-5-thinking",
"gemini-claude-opus-4-5-thinking",
"claude-sonnet-4-5",
"claude-sonnet-4-5-thinking",
"claude-opus-4-5-thinking",
];
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_provider_type_from_str() {
assert_eq!("kiro".parse::<ProviderType>().unwrap(), ProviderType::Kiro);
assert_eq!(
"gemini".parse::<ProviderType>().unwrap(),
ProviderType::Gemini
);
assert_eq!(
"openai".parse::<ProviderType>().unwrap(),
ProviderType::OpenAI
);
assert_eq!(
"claude".parse::<ProviderType>().unwrap(),
ProviderType::Claude
);
assert_eq!(
"vertex".parse::<ProviderType>().unwrap(),
ProviderType::Vertex
);
assert_eq!(
"gemini_api_key".parse::<ProviderType>().unwrap(),
ProviderType::GeminiApiKey
);
assert_eq!("KIRO".parse::<ProviderType>().unwrap(), ProviderType::Kiro);
assert_eq!(
"Gemini".parse::<ProviderType>().unwrap(),
ProviderType::Gemini
);
assert_eq!(
"VERTEX".parse::<ProviderType>().unwrap(),
ProviderType::Vertex
);
assert!("invalid".parse::<ProviderType>().is_err());
}
#[test]
fn test_provider_type_display() {
assert_eq!(ProviderType::Kiro.to_string(), "kiro");
assert_eq!(ProviderType::Gemini.to_string(), "gemini");
assert_eq!(ProviderType::OpenAI.to_string(), "openai");
assert_eq!(ProviderType::Claude.to_string(), "claude");
assert_eq!(ProviderType::Vertex.to_string(), "vertex");
assert_eq!(ProviderType::GeminiApiKey.to_string(), "gemini_api_key");
}
#[test]
fn test_provider_type_serde() {
assert_eq!(
serde_json::to_string(&ProviderType::Kiro).unwrap(),
"\"kiro\""
);
assert_eq!(
serde_json::to_string(&ProviderType::OpenAI).unwrap(),
"\"openai\""
);
assert_eq!(
serde_json::from_str::<ProviderType>("\"kiro\"").unwrap(),
ProviderType::Kiro
);
assert_eq!(
serde_json::from_str::<ProviderType>("\"openai\"").unwrap(),
ProviderType::OpenAI
);
}
}
@@ -0,0 +1,130 @@
//! 路由模型
//!
//! 用于多供应商路由功能的数据结构定义。
use serde::{Deserialize, Serialize};
/// 单个路由信息
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RouteInfo {
pub selector: String,
pub provider_type: String,
pub credential_count: usize,
pub endpoints: Vec<RouteEndpoint>,
pub tags: Vec<String>,
pub enabled: bool,
}
/// 路由端点
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RouteEndpoint {
pub path: String,
pub protocol: String,
pub url: String,
}
/// 路由列表响应
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RouteListResponse {
pub base_url: String,
pub default_provider: String,
pub routes: Vec<RouteInfo>,
}
/// curl 示例
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CurlExample {
pub description: String,
pub command: String,
}
impl RouteInfo {
pub fn new(selector: String, provider_type: String) -> Self {
Self {
selector,
provider_type,
credential_count: 0,
endpoints: Vec::new(),
tags: Vec::new(),
enabled: true,
}
}
pub fn add_endpoint(&mut self, base_url: &str, protocol: &str) {
let path = match protocol {
"claude" => format!("/{}/v1/messages", self.selector),
"openai" => format!("/{}/v1/chat/completions", self.selector),
_ => return,
};
let url = format!("{}{}", base_url, path);
self.endpoints.push(RouteEndpoint {
path,
protocol: protocol.to_string(),
url,
});
}
pub fn generate_curl_examples(&self, api_key: &str) -> Vec<CurlExample> {
let mut examples = Vec::new();
for endpoint in &self.endpoints {
let (_model, body) = match endpoint.protocol.as_str() {
"claude" => {
let model = match self.provider_type.as_str() {
"kiro" | "claude" => "claude-sonnet-4-5",
"gemini" => "gemini-2.5-flash",
"qwen" => "qwen3-coder-plus",
"openai" => "gpt-4",
_ => "claude-sonnet-4-5",
};
(
model,
format!(
r#"{{
"model": "{}",
"max_tokens": 1024,
"messages": [{{"role": "user", "content": "Hello!"}}]
}}"#,
model
),
)
}
"openai" => {
let model = match self.provider_type.as_str() {
"kiro" | "claude" => "claude-sonnet-4-5",
"gemini" => "gemini-2.5-flash",
"qwen" => "qwen3-coder-plus",
"openai" => "gpt-4",
_ => "claude-sonnet-4-5",
};
(
model,
format!(
r#"{{
"model": "{}",
"messages": [{{"role": "user", "content": "Hello!"}}]
}}"#,
model
),
)
}
_ => continue,
};
let command = format!(
r#"curl {} \
-H "Content-Type: application/json" \
-H "Authorization: Bearer {}" \
-d '{}'"#,
endpoint.url, api_key, body
);
examples.push(CurlExample {
description: format!("{} 协议", endpoint.protocol.to_uppercase()),
command,
});
}
examples
}
}
@@ -0,0 +1,145 @@
//! Skill 数据模型
use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Skill {
pub key: String,
pub name: String,
pub description: String,
pub directory: String,
#[serde(rename = "readmeUrl", skip_serializing_if = "Option::is_none")]
pub readme_url: Option<String>,
pub installed: bool,
#[serde(rename = "repoOwner", skip_serializing_if = "Option::is_none")]
pub repo_owner: Option<String>,
#[serde(rename = "repoName", skip_serializing_if = "Option::is_none")]
pub repo_name: Option<String>,
#[serde(rename = "repoBranch", skip_serializing_if = "Option::is_none")]
pub repo_branch: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SkillRepo {
pub owner: String,
pub name: String,
pub branch: String,
pub enabled: bool,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SkillState {
pub installed: bool,
pub installed_at: DateTime<Utc>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SkillMetadata {
pub name: Option<String>,
pub description: Option<String>,
}
impl Default for SkillRepo {
fn default() -> Self {
Self {
owner: String::new(),
name: String::new(),
branch: "main".to_string(),
enabled: true,
}
}
}
#[allow(dead_code)]
impl SkillRepo {
pub fn new(owner: String, name: String, branch: String) -> Self {
Self {
owner,
name,
branch,
enabled: true,
}
}
pub fn github_url(&self) -> String {
format!("https://github.com/{}/{}", self.owner, self.name)
}
pub fn zip_url(&self) -> String {
format!(
"https://github.com/{}/{}/archive/refs/heads/{}.zip",
self.owner, self.name, self.branch
)
}
}
pub fn get_default_skill_repos() -> Vec<SkillRepo> {
vec![
SkillRepo {
owner: "proxycast".to_string(),
name: "skills".to_string(),
branch: "main".to_string(),
enabled: true,
},
SkillRepo {
owner: "ComposioHQ".to_string(),
name: "awesome-claude-skills".to_string(),
branch: "main".to_string(),
enabled: true,
},
SkillRepo {
owner: "anthropics".to_string(),
name: "skills".to_string(),
branch: "main".to_string(),
enabled: true,
},
SkillRepo {
owner: "cexll".to_string(),
name: "myclaude".to_string(),
branch: "master".to_string(),
enabled: true,
},
]
}
pub type SkillStates = HashMap<String, SkillState>;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_default_repos_include_proxycast_official() {
let repos = get_default_skill_repos();
assert!(!repos.is_empty(), "默认仓库列表不应为空");
let first_repo = &repos[0];
assert_eq!(
first_repo.owner, "proxycast",
"第一个仓库的 owner 应为 proxycast"
);
assert_eq!(first_repo.name, "skills", "第一个仓库的 name 应为 skills");
assert_eq!(first_repo.branch, "main", "第一个仓库的 branch 应为 main");
assert!(first_repo.enabled, "ProxyCast 官方仓库应默认启用");
}
#[test]
fn test_proxycast_repo_exists_in_list() {
let repos = get_default_skill_repos();
let proxycast_repo = repos
.iter()
.find(|r| r.owner == "proxycast" && r.name == "skills");
assert!(
proxycast_repo.is_some(),
"ProxyCast 官方仓库应存在于默认列表中"
);
let repo = proxycast_repo.unwrap();
assert_eq!(repo.branch, "main");
assert!(repo.enabled);
}
}
+39
View File
@@ -0,0 +1,39 @@
[package]
name = "proxycast-infra"
version.workspace = true
edition.workspace = true
authors.workspace = true
repository.workspace = true
[dependencies]
# 项目内 crate
proxycast-core.workspace = true
# 序列化
serde.workspace = true
serde_json.workspace = true
# 异步运行时
tokio.workspace = true
# 错误处理
thiserror.workspace = true
# 日志
tracing.workspace = true
# 时间和 UUID
chrono.workspace = true
uuid.workspace = true
# HTTP 客户端
reqwest.workspace = true
# 工具库
parking_lot.workspace = true
dashmap.workspace = true
dirs.workspace = true
tiktoken-rs.workspace = true
[dev-dependencies]
proptest.workspace = true
+30
View File
@@ -0,0 +1,30 @@
//! 基础设施模块
//!
//! 包含独立的基础设施组件,不依赖业务逻辑:
//! - proxy: HTTP 代理客户端
//! - resilience: 重试、熔断、故障转移
//! - injection: 请求参数注入
//! - telemetry: 遥测统计
//!
//! 注意:plugin 模块因依赖 Tauri 无法迁移,保留在主 crate
pub mod injection;
pub mod proxy;
pub mod resilience;
pub mod telemetry;
// 重新导出常用类型
pub use injection::{InjectionConfig, InjectionMode, InjectionResult, InjectionRule, Injector};
pub use proxy::{ProxyClientFactory, ProxyError, ProxyProtocol};
pub use resilience::{
Failover, FailoverConfig, Retrier, RetryConfig, TimeoutConfig, TimeoutController,
};
pub use telemetry::{
LogRotationConfig, LoggerError, ModelStats, ModelTokenStats, PeriodTokenStats, ProviderStats,
ProviderTokenStats, RequestLog, RequestLogger, RequestStatus, StatsAggregator, StatsSummary,
TimeRange, TokenSource, TokenStatsSummary, TokenTracker, TokenUsageRecord,
};
pub fn version() -> &'static str {
env!("CARGO_PKG_VERSION")
}
@@ -2,7 +2,7 @@
//!
//! 提供 Provider 故障转移和自动切换功能
use crate::ProviderType;
use proxycast_core::ProviderType;
use serde::{Deserialize, Serialize};
use std::collections::HashSet;
@@ -2,12 +2,10 @@
//!
//! 提供请求日志记录、查询和轮转功能
use crate::telemetry::types::{
ModelStats, ProviderStats, RequestLog, RequestStatus, StatsSummary, TimeRange,
};
use crate::ProviderType;
use super::types::{ModelStats, ProviderStats, RequestLog, RequestStatus, StatsSummary, TimeRange};
use chrono::{DateTime, Duration, Utc};
use parking_lot::RwLock;
use proxycast_core::ProviderType;
use serde::{Deserialize, Serialize};
use std::collections::{HashMap, VecDeque};
use std::fs::{self, File, OpenOptions};
@@ -2,12 +2,10 @@
//!
//! 提供请求统计的聚合、分组和查询功能
use crate::telemetry::types::{
ModelStats, ProviderStats, RequestLog, RequestStatus, StatsSummary, TimeRange,
};
use crate::ProviderType;
use super::types::{ModelStats, ProviderStats, RequestLog, RequestStatus, StatsSummary, TimeRange};
use chrono::{Duration, Utc};
use parking_lot::RwLock;
use proxycast_core::ProviderType;
use std::collections::{HashMap, VecDeque};
/// 统计聚合器
@@ -2,12 +2,12 @@
//!
//! 使用 proptest 进行属性测试
use crate::telemetry::{
use super::{
LogRotationConfig, RequestLog, RequestLogger, RequestStatus, StatsAggregator, TimeRange,
};
use crate::ProviderType;
use chrono::{Duration, Utc};
use proptest::prelude::*;
use proxycast_core::ProviderType;
use std::collections::HashSet;
/// 生成随机的 ProviderType
@@ -4,9 +4,9 @@
#![allow(dead_code)]
use crate::ProviderType;
use chrono::{DateTime, Duration, Utc};
use parking_lot::RwLock;
use proxycast_core::ProviderType;
use serde::{Deserialize, Serialize};
use std::collections::{HashMap, VecDeque};
@@ -2,8 +2,8 @@
//!
//! 定义请求日志、统计数据等核心类型
use crate::ProviderType;
use chrono::{DateTime, Utc};
use proxycast_core::ProviderType;
use serde::{Deserialize, Serialize};
/// 请求状态
-15
View File
@@ -1,15 +0,0 @@
[package]
name = "providers"
version = "0.1.0"
edition = "2021"
[dependencies]
core = { path = "../core" }
serde = { version = "1", features = ["derive"] }
serde_json = "1"
anyhow = "1"
thiserror = "1"
tracing = "0.1"
async-trait = "0.1"
tokio = { version = "1", features = ["full"] }
reqwest = { version = "0.12", features = ["json", "stream"] }
-7
View File
@@ -1,7 +0,0 @@
//! Provider 系统模块
//!
//! 包含 providers, credential, converter 等功能
pub fn version() -> &'static str {
env!("CARGO_PKG_VERSION")
}
-17
View File
@@ -1,17 +0,0 @@
[package]
name = "server"
version = "0.1.0"
edition = "2021"
[dependencies]
core = { path = "../core" }
providers = { path = "../providers" }
serde = { version = "1", features = ["derive"] }
serde_json = "1"
anyhow = "1"
thiserror = "1"
tracing = "0.1"
tokio = { version = "1", features = ["full"] }
axum = { version = "0.7", features = ["ws"] }
tower = "0.4"
tower-http = { version = "0.5", features = ["limit", "cors"] }
-7
View File
@@ -1,7 +0,0 @@
//! API 服务器模块
//!
//! 包含 server, streaming, middleware, router 等功能
pub fn version() -> &'static str {
env!("CARGO_PKG_VERSION")
}
+1 -5
View File
@@ -3,14 +3,10 @@
"provider": "kiro",
"description": "Kiro/CodeWhisperer 服务的模型别名映射(基于 AWS Bedrock)",
"models": [
"claude-opus-4-5",
"claude-opus-4-5-20251101",
"claude-haiku-4-5",
"claude-haiku-4-5-20251001",
"claude-sonnet-4-5",
"claude-sonnet-4-5-20250929",
"claude-sonnet-4-20250514",
"claude-3-7-sonnet-20250219"
"claude-sonnet-4-20250514"
],
"aliases": {
"claude-opus-4-5": {
+15 -5
View File
@@ -170,14 +170,24 @@ impl OpenAIProtocol {
match chunk {
Ok(bytes) => {
let text = String::from_utf8_lossy(&bytes);
// 安全截断:使用 char_indices 找到有效的 UTF-8 字符边界
let truncated = if text.len() > 200 {
let mut end = 200;
for (i, _) in text.char_indices() {
if i <= 200 {
end = i;
} else {
break;
}
}
format!("{}...", &text[..end])
} else {
text.to_string()
};
eprintln!(
"[OpenAIProtocol] 收到 chunk: {} bytes, 内容: {}",
bytes.len(),
if text.len() > 200 {
format!("{}...", &text[..200])
} else {
text.to_string()
}
truncated
);
buffer.push_str(&text);
+2
View File
@@ -1160,6 +1160,8 @@ pub fn run() {
commands::model_registry_cmd::get_models_by_tier,
commands::model_registry_cmd::get_provider_alias_config,
commands::model_registry_cmd::get_all_alias_configs,
commands::model_registry_cmd::fetch_provider_models_from_api,
commands::model_registry_cmd::fetch_provider_models_auto,
// Model Management commands (动态模型列表)
commands::model_cmd::get_credential_models,
commands::model_cmd::refresh_credential_models,
+3 -69
View File
@@ -1,8 +1,7 @@
//! 核心类型定义
//!
//! 包含 Provider 类型枚举和相关实现。
//! 包含应用状态类型和相关实现。
use serde::{Deserialize, Serialize};
use std::sync::Arc;
use tauri::Runtime;
use tokio::sync::RwLock;
@@ -12,73 +11,8 @@ use crate::server;
use crate::services::token_cache_service::TokenCacheService;
use crate::tray::TrayManager;
/// Provider 类型枚举
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum ProviderType {
Kiro,
Gemini,
#[serde(rename = "openai")]
OpenAI,
Claude,
Antigravity,
Vertex,
#[serde(rename = "gemini_api_key")]
GeminiApiKey,
Codex,
#[serde(rename = "claude_oauth")]
ClaudeOAuth,
// API Key Provider 类型
Anthropic,
#[serde(rename = "azure_openai")]
AzureOpenai,
#[serde(rename = "aws_bedrock")]
AwsBedrock,
Ollama,
}
impl std::fmt::Display for ProviderType {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
ProviderType::Kiro => write!(f, "kiro"),
ProviderType::Gemini => write!(f, "gemini"),
ProviderType::OpenAI => write!(f, "openai"),
ProviderType::Claude => write!(f, "claude"),
ProviderType::Antigravity => write!(f, "antigravity"),
ProviderType::Vertex => write!(f, "vertex"),
ProviderType::GeminiApiKey => write!(f, "gemini_api_key"),
ProviderType::Codex => write!(f, "codex"),
ProviderType::ClaudeOAuth => write!(f, "claude_oauth"),
ProviderType::Anthropic => write!(f, "anthropic"),
ProviderType::AzureOpenai => write!(f, "azure_openai"),
ProviderType::AwsBedrock => write!(f, "aws_bedrock"),
ProviderType::Ollama => write!(f, "ollama"),
}
}
}
impl std::str::FromStr for ProviderType {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.to_lowercase().as_str() {
"kiro" => Ok(ProviderType::Kiro),
"gemini" => Ok(ProviderType::Gemini),
"openai" => Ok(ProviderType::OpenAI),
"claude" => Ok(ProviderType::Claude),
"antigravity" => Ok(ProviderType::Antigravity),
"vertex" => Ok(ProviderType::Vertex),
"gemini_api_key" => Ok(ProviderType::GeminiApiKey),
"codex" => Ok(ProviderType::Codex),
"claude_oauth" => Ok(ProviderType::ClaudeOAuth),
"anthropic" => Ok(ProviderType::Anthropic),
"azure_openai" | "azure-openai" => Ok(ProviderType::AzureOpenai),
"aws_bedrock" | "aws-bedrock" => Ok(ProviderType::AwsBedrock),
"ollama" => Ok(ProviderType::Ollama),
_ => Err(format!("Invalid provider: {s}")),
}
}
}
// 重新导出 core crate 的 ProviderType
pub use proxycast_core::ProviderType;
/// 应用状态类型别名
pub type AppState = Arc<RwLock<server::ServerState>>;
+72 -1
View File
@@ -5,7 +5,7 @@
use crate::models::model_registry::{
EnhancedModelMetadata, ModelSyncState, ModelTier, ProviderAliasConfig, UserModelPreference,
};
use crate::services::model_registry_service::ModelRegistryService;
use crate::services::model_registry_service::{FetchModelsResult, ModelRegistryService};
use std::sync::Arc;
use tauri::State;
use tokio::sync::RwLock;
@@ -180,3 +180,74 @@ pub async fn refresh_model_registry(state: State<'_, ModelRegistryState>) -> Res
service.force_reload().await
}
/// 从 Provider API 获取模型列表
///
/// 调用 Provider 的 /v1/models 端点获取模型列表,
/// 如果失败则回退到本地 JSON 文件
///
/// # 参数
/// - `provider_id`: Provider ID(如 "siliconflow", "openai")
/// - `api_host`: API 主机地址
/// - `api_key`: API Key
#[tauri::command]
pub async fn fetch_provider_models_from_api(
state: State<'_, ModelRegistryState>,
provider_id: String,
api_host: String,
api_key: String,
) -> Result<FetchModelsResult, String> {
let guard = state.read().await;
let service = guard
.as_ref()
.ok_or_else(|| "模型注册服务未初始化".to_string())?;
service
.fetch_models_from_api(&provider_id, &api_host, &api_key)
.await
}
/// 从 Provider API 获取模型列表(自动获取 API Key)
///
/// 自动从数据库获取 Provider 的 API Key,然后调用 /v1/models 端点
///
/// # 参数
/// - `provider_id`: Provider ID(如 "siliconflow", "openai")
#[tauri::command]
pub async fn fetch_provider_models_auto(
state: State<'_, ModelRegistryState>,
db: tauri::State<'_, crate::database::DbConnection>,
api_key_service: tauri::State<
'_,
crate::commands::api_key_provider_cmd::ApiKeyProviderServiceState,
>,
provider_id: String,
) -> Result<FetchModelsResult, String> {
// 获取 Provider 信息
let provider = api_key_service
.0
.get_provider(&db, &provider_id)?
.ok_or_else(|| format!("Provider 不存在: {}", provider_id))?;
// 获取 API Key
let api_key = api_key_service
.0
.get_next_api_key(&db, &provider_id)?
.ok_or_else(|| format!("Provider {} 没有可用的 API Key", provider_id))?;
// 获取 API Host
let api_host = provider.provider.api_host.clone();
if api_host.is_empty() {
return Err("Provider 没有配置 API Host".to_string());
}
// 调用模型注册服务
let guard = state.read().await;
let service = guard
.as_ref()
.ok_or_else(|| "模型注册服务未初始化".to_string())?;
service
.fetch_models_from_api(&provider_id, &api_host, &api_key)
.await
}
+190
View File
@@ -268,6 +268,138 @@ struct ApiKeyMigrationRow {
provider_name: String,
}
/// 迁移旧的 Provider ID 到新的 ID
///
/// 修复 system_providers.rs 中 Provider ID 与模型注册表 JSON 文件名不匹配的问题。
/// 例如:silicon -> siliconflow, gemini -> google 等
pub fn migrate_provider_ids(conn: &Connection) -> Result<usize, String> {
// 检查是否已经迁移过
let migrated: bool = conn
.query_row(
"SELECT value FROM settings WHERE key = 'migrated_provider_ids_v1'",
[],
|row| row.get::<_, String>(0),
)
.map(|v| v == "true")
.unwrap_or(false);
if migrated {
tracing::debug!("[迁移] Provider ID 已迁移过,跳过");
return Ok(0);
}
tracing::info!("[迁移] 开始迁移旧的 Provider ID");
// 定义需要迁移的 ID 映射(旧 ID -> 新 ID)
let id_mappings = [
("silicon", "siliconflow"),
("gemini", "google"),
("zhipu", "zhipuai"),
("dashscope", "alibaba"),
("moonshot", "moonshotai"),
("grok", "xai"),
("github", "github-models"),
("copilot", "github-copilot"),
("vertexai", "google-vertex"),
("aws-bedrock", "amazon-bedrock"),
("together", "togetherai"),
("fireworks", "fireworks-ai"),
("mimo", "xiaomi"),
];
let mut migrated_count = 0;
for (old_id, new_id) in &id_mappings {
// 检查旧 ID 是否存在
let old_exists: bool = conn
.query_row(
"SELECT COUNT(*) > 0 FROM api_key_providers WHERE id = ?1",
params![old_id],
|r| r.get(0),
)
.unwrap_or(false);
if !old_exists {
continue;
}
// 检查新 ID 是否存在
let new_exists: bool = conn
.query_row(
"SELECT COUNT(*) > 0 FROM api_key_providers WHERE id = ?1",
params![new_id],
|r| r.get(0),
)
.unwrap_or(false);
// 检查旧 ID 是否有 API Keys
let has_keys: bool = conn
.query_row(
"SELECT COUNT(*) > 0 FROM api_keys WHERE provider_id = ?1",
params![old_id],
|r| r.get(0),
)
.unwrap_or(false);
if has_keys {
// 如果旧 ID 有 API Keys,需要迁移到新 ID
if new_exists {
// 新 ID 已存在,将 API Keys 迁移过去
conn.execute(
"UPDATE api_keys SET provider_id = ?1 WHERE provider_id = ?2",
params![new_id, old_id],
)
.map_err(|e| format!("迁移 API Keys 失败: {}", e))?;
tracing::info!("[迁移] 已将 {} 的 API Keys 迁移到 {}", old_id, new_id);
} else {
// 新 ID 不存在,直接更新旧 ID
conn.execute(
"UPDATE api_key_providers SET id = ?1 WHERE id = ?2",
params![new_id, old_id],
)
.map_err(|e| format!("更新 Provider ID 失败: {}", e))?;
conn.execute(
"UPDATE api_keys SET provider_id = ?1 WHERE provider_id = ?2",
params![new_id, old_id],
)
.map_err(|e| format!("更新 API Keys provider_id 失败: {}", e))?;
tracing::info!("[迁移] 已将 Provider {} 重命名为 {}", old_id, new_id);
migrated_count += 1;
continue;
}
}
// 删除旧的 Provider(无论是否有 API Keys,因为 Keys 已迁移)
conn.execute(
"DELETE FROM api_key_providers WHERE id = ?1",
params![old_id],
)
.map_err(|e| format!("删除旧 Provider 失败: {}", e))?;
tracing::info!("[迁移] 已删除旧 Provider: {}", old_id);
migrated_count += 1;
}
// 标记迁移完成
conn.execute(
"INSERT OR REPLACE INTO settings (key, value) VALUES ('migrated_provider_ids_v1', 'true')",
[],
)
.map_err(|e| format!("标记迁移完成失败: {}", e))?;
if migrated_count > 0 {
tracing::info!(
"[迁移] Provider ID 迁移完成,共处理 {} 个 Provider",
migrated_count
);
}
Ok(migrated_count)
}
/// 清理旧的 API Key 凭证(OpenAIKey 和 ClaudeKey 类型)
///
/// 这些凭证是通过旧的 UI 添加的,现在已经被新的 API Key Provider 系统取代。
@@ -364,3 +496,61 @@ pub fn cleanup_legacy_api_key_credentials(conn: &Connection) -> Result<usize, St
Ok(deleted)
}
/// 当前模型注册表版本
/// 每次更新模型数据结构或添加新 Provider 时,增加此版本号
const MODEL_REGISTRY_VERSION: &str = "2026.01.16.1";
/// 标记需要刷新模型注册表
pub fn mark_model_registry_refresh_needed(conn: &Connection) {
let _ = conn.execute(
"INSERT OR REPLACE INTO settings (key, value) VALUES ('model_registry_refresh_needed', 'true')",
[],
);
tracing::info!("[迁移] 已标记需要刷新模型注册表");
}
/// 检查模型注册表版本,如果版本不匹配则标记需要刷新
pub fn check_model_registry_version(conn: &Connection) {
let current_version: Option<String> = conn
.query_row(
"SELECT value FROM settings WHERE key = 'model_registry_version'",
[],
|row| row.get(0),
)
.ok();
if current_version.as_deref() != Some(MODEL_REGISTRY_VERSION) {
tracing::info!(
"[迁移] 模型注册表版本不匹配: {:?} -> {},标记需要刷新",
current_version,
MODEL_REGISTRY_VERSION
);
mark_model_registry_refresh_needed(conn);
// 更新版本号
let _ = conn.execute(
"INSERT OR REPLACE INTO settings (key, value) VALUES ('model_registry_version', ?1)",
params![MODEL_REGISTRY_VERSION],
);
}
}
/// 检查是否需要刷新模型注册表
pub fn is_model_registry_refresh_needed(conn: &Connection) -> bool {
conn.query_row(
"SELECT value FROM settings WHERE key = 'model_registry_refresh_needed'",
[],
|row| row.get::<_, String>(0),
)
.map(|v| v == "true")
.unwrap_or(false)
}
/// 清除模型注册表刷新标记
pub fn clear_model_registry_refresh_flag(conn: &Connection) {
let _ = conn.execute(
"DELETE FROM settings WHERE key = 'model_registry_refresh_needed'",
[],
);
}
+17
View File
@@ -31,6 +31,23 @@ pub fn init_database() -> Result<DbConnection, String> {
schema::create_tables(&conn).map_err(|e| e.to_string())?;
migration::migrate_from_json(&conn)?;
// 执行 Provider ID 迁移(修复旧 ID 与模型注册表不匹配的问题)
match migration::migrate_provider_ids(&conn) {
Ok(count) => {
if count > 0 {
tracing::info!("[数据库] 已迁移 {} 个 Provider ID", count);
// 标记需要刷新模型注册表
migration::mark_model_registry_refresh_needed(&conn);
}
}
Err(e) => {
tracing::warn!("[数据库] Provider ID 迁移失败(非致命): {}", e);
}
}
// 检查是否需要刷新模型注册表(版本升级时)
migration::check_model_registry_version(&conn);
// 执行 API Keys 到 Provider Pool 的迁移
match migration::migrate_api_keys_to_pool(&conn) {
Ok(count) => {
+22 -13
View File
@@ -44,7 +44,7 @@ pub fn get_system_providers() -> Vec<SystemProviderDef> {
api_version: None,
},
SystemProviderDef {
id: "gemini",
id: "google",
name: "Gemini",
provider_type: ApiProviderType::Gemini,
api_host: "https://generativelanguage.googleapis.com",
@@ -62,7 +62,7 @@ pub fn get_system_providers() -> Vec<SystemProviderDef> {
api_version: None,
},
SystemProviderDef {
id: "moonshot",
id: "moonshotai",
name: "Moonshot",
provider_type: ApiProviderType::Openai,
api_host: "https://api.moonshot.cn",
@@ -80,7 +80,7 @@ pub fn get_system_providers() -> Vec<SystemProviderDef> {
api_version: None,
},
SystemProviderDef {
id: "grok",
id: "xai",
name: "Grok (xAI)",
provider_type: ApiProviderType::Openai,
api_host: "https://api.x.ai",
@@ -119,7 +119,7 @@ pub fn get_system_providers() -> Vec<SystemProviderDef> {
// 国内 AI (15个) - Requirements 3.2
// =========================================================================
SystemProviderDef {
id: "zhipu",
id: "zhipuai",
name: "智谱 (ZhiPu)",
provider_type: ApiProviderType::Openai,
api_host: "https://open.bigmodel.cn/api/paas/v4/",
@@ -137,7 +137,7 @@ pub fn get_system_providers() -> Vec<SystemProviderDef> {
api_version: None,
},
SystemProviderDef {
id: "dashscope",
id: "alibaba",
name: "百炼/通义千问 (Dashscope)",
provider_type: ApiProviderType::Openai,
api_host: "https://dashscope.aliyuncs.com/compatible-mode/v1/",
@@ -236,7 +236,7 @@ pub fn get_system_providers() -> Vec<SystemProviderDef> {
api_version: None,
},
SystemProviderDef {
id: "mimo",
id: "xiaomi",
name: "小米 MiMo",
provider_type: ApiProviderType::Openai,
api_host: "https://api.xiaomimimo.com",
@@ -266,7 +266,7 @@ pub fn get_system_providers() -> Vec<SystemProviderDef> {
api_version: Some("2024-02-15-preview"),
},
SystemProviderDef {
id: "vertexai",
id: "google-vertex",
name: "VertexAI",
provider_type: ApiProviderType::Vertexai,
api_host: "",
@@ -275,7 +275,7 @@ pub fn get_system_providers() -> Vec<SystemProviderDef> {
api_version: None,
},
SystemProviderDef {
id: "aws-bedrock",
id: "amazon-bedrock",
name: "AWS Bedrock",
provider_type: ApiProviderType::AwsBedrock,
api_host: "",
@@ -284,7 +284,7 @@ pub fn get_system_providers() -> Vec<SystemProviderDef> {
api_version: None,
},
SystemProviderDef {
id: "github",
id: "github-models",
name: "Github Models",
provider_type: ApiProviderType::Openai,
api_host: "https://models.github.ai/inference",
@@ -293,7 +293,7 @@ pub fn get_system_providers() -> Vec<SystemProviderDef> {
api_version: None,
},
SystemProviderDef {
id: "copilot",
id: "github-copilot",
name: "Github Copilot",
provider_type: ApiProviderType::Openai,
api_host: "https://api.githubcopilot.com/",
@@ -305,7 +305,7 @@ pub fn get_system_providers() -> Vec<SystemProviderDef> {
// API 聚合/中转服务 (25个) - Requirements 3.4
// =========================================================================
SystemProviderDef {
id: "silicon",
id: "siliconflow",
name: "Silicon Flow",
provider_type: ApiProviderType::Openai,
api_host: "https://api.siliconflow.cn",
@@ -313,6 +313,15 @@ pub fn get_system_providers() -> Vec<SystemProviderDef> {
sort_order: 31,
api_version: None,
},
SystemProviderDef {
id: "siliconflow-cn",
name: "Silicon Flow (国内)",
provider_type: ApiProviderType::Openai,
api_host: "https://api.siliconflow.cn",
group: ProviderGroup::Aggregator,
sort_order: 32,
api_version: None,
},
SystemProviderDef {
id: "openrouter",
name: "OpenRouter",
@@ -341,7 +350,7 @@ pub fn get_system_providers() -> Vec<SystemProviderDef> {
api_version: None,
},
SystemProviderDef {
id: "together",
id: "togetherai",
name: "Together",
provider_type: ApiProviderType::Openai,
api_host: "https://api.together.xyz",
@@ -350,7 +359,7 @@ pub fn get_system_providers() -> Vec<SystemProviderDef> {
api_version: None,
},
SystemProviderDef {
id: "fireworks",
id: "fireworks-ai",
name: "Fireworks",
provider_type: ApiProviderType::Openai,
api_host: "https://api.fireworks.ai/inference",
+28 -10
View File
@@ -1,11 +1,31 @@
//! ProxyCast - AI API 代理服务
//!
//! 这是一个 Tauri 应用,提供 AI API 的代理和管理功能。
//!
//! ## Workspace 结构(方案 A - 最小化拆分)
//!
//! 采用最小化拆分策略,只迁移真正独立的模块:
//! - ✅ proxycast-core crate(models, data, logger)
//! - ✅ proxycast-infra crate(proxy, resilience, injection, telemetry)
//! - 主 crate 保留所有业务逻辑模块(包括 plugin,因依赖 Tauri)
// 抑制 objc crate 宏内部的 unexpected_cfgs 警告
// 该警告来自 cocoa/objc 依赖的 msg_send! 宏,是已知的 issue
#![allow(unexpected_cfgs)]
// 重新导出子 crate 的类型
// 注意:主 crate 保留了自己的 data, logger, models 模块,所以只导出 core 的具体类型
pub use proxycast_core::{LogEntry, LogStore, LogStoreConfig, SharedLogStore};
// infra crate 的类型通过 proxycast_infra 前缀访问,避免与 core 的 InjectionMode/InjectionRule 冲突
pub use proxycast_infra::{
injection, proxy, resilience, telemetry, Failover, FailoverConfig, InjectionConfig,
InjectionMode, InjectionResult, InjectionRule, Injector, LogRotationConfig, LoggerError,
ModelStats, ModelTokenStats, PeriodTokenStats, ProviderStats, ProviderTokenStats,
ProxyClientFactory, ProxyError, ProxyProtocol, RequestLog, RequestLogger, RequestStatus,
Retrier, RetryConfig, StatsAggregator, StatsSummary, TimeRange, TimeoutConfig,
TimeoutController, TokenSource, TokenStatsSummary, TokenTracker, TokenUsageRecord,
};
// 核心模块
pub mod agent;
pub mod app;
@@ -15,25 +35,16 @@ pub mod connect;
pub mod credential;
pub mod database;
pub mod flow_monitor;
pub mod injection;
pub mod middleware;
pub mod orchestrator;
pub mod plugin;
pub mod processor;
pub mod proxy;
pub mod resilience;
pub mod router;
pub mod screenshot;
pub mod services;
pub mod session;
pub mod session_files;
pub mod stream;
pub mod streaming;
pub mod telemetry;
pub mod terminal;
pub mod translator;
pub mod tray;
pub mod websocket;
// 内部模块
mod commands;
@@ -45,9 +56,16 @@ mod dev_bridge;
mod logger;
mod models;
mod providers;
mod server;
mod server_utils;
// 服务器相关模块
mod middleware;
mod processor;
mod router;
mod server;
mod streaming;
mod websocket;
// 重新导出核心类型以保持向后兼容
pub use app::{AppState, LogState, ProviderType, TokenCacheServiceState, TrayManagerState};
pub use services::provider_pool_service::ProviderPoolService;
+4
View File
@@ -167,6 +167,8 @@ pub enum ModelSource {
Local,
/// 用户自定义
Custom,
/// 从 Provider API 获取
Api,
}
impl Default for ModelSource {
@@ -182,6 +184,7 @@ impl std::fmt::Display for ModelSource {
Self::ModelsDev => write!(f, "models.dev"),
Self::Local => write!(f, "local"),
Self::Custom => write!(f, "custom"),
Self::Api => write!(f, "api"),
}
}
}
@@ -195,6 +198,7 @@ impl std::str::FromStr for ModelSource {
"models.dev" | "modelsdev" => Ok(Self::ModelsDev),
"local" => Ok(Self::Local),
"custom" => Ok(Self::Custom),
"api" => Ok(Self::Api),
_ => Err(format!("Unknown model source: {}", s)),
}
}
@@ -9,7 +9,7 @@ use crate::models::model_registry::{
ModelSyncState, ModelTier, ProviderAliasConfig, UserModelPreference,
};
use rusqlite::params;
use serde::Deserialize;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
@@ -218,6 +218,11 @@ impl ModelRegistryService {
match std::fs::read_to_string(&provider_file) {
Ok(content) => match serde_json::from_str::<RepoProviderData>(&content) {
Ok(provider_data) => {
tracing::info!(
"[ModelRegistry] 加载 Provider: {} ({} 个模型)",
provider_id,
provider_data.models.len()
);
for model in provider_data.models {
let enhanced = self.convert_repo_model(
model,
@@ -810,4 +815,215 @@ impl ModelRegistryService {
pub async fn get_all_alias_configs(&self) -> HashMap<String, ProviderAliasConfig> {
self.aliases_cache.read().await.clone()
}
// ========== 从 Provider API 获取模型 ==========
/// 从 Provider API 获取模型列表
///
/// 调用 Provider 的 /v1/models 端点获取模型列表,
/// 如果失败则回退到本地 JSON 文件
///
/// # 参数
/// - `provider_id`: Provider ID(如 "siliconflow", "openai")
/// - `api_host`: API 主机地址
/// - `api_key`: API Key
///
/// # 返回
/// - `Ok(FetchModelsResult)`: 获取结果,包含模型列表和来源
pub async fn fetch_models_from_api(
&self,
provider_id: &str,
api_host: &str,
api_key: &str,
) -> Result<FetchModelsResult, String> {
tracing::info!(
"[ModelRegistry] 从 API 获取模型: provider={}, host={}",
provider_id,
api_host
);
// 构建 API URL
let api_url = Self::build_models_api_url(api_host);
tracing::info!("[ModelRegistry] API URL: {}", api_url);
// 尝试从 API 获取
match self.call_models_api(&api_url, api_key).await {
Ok(api_models) => {
tracing::info!("[ModelRegistry] 从 API 获取到 {} 个模型", api_models.len());
// 转换为内部格式
let now = chrono::Utc::now().timestamp();
let models: Vec<EnhancedModelMetadata> = api_models
.into_iter()
.map(|m| self.convert_api_model(m, provider_id, now))
.collect();
Ok(FetchModelsResult {
models,
source: ModelFetchSource::Api,
error: None,
})
}
Err(api_error) => {
tracing::warn!(
"[ModelRegistry] API 获取失败: {}, 回退到本地文件",
api_error
);
// 回退到本地 JSON 文件
let local_models = self.get_models_by_provider(provider_id).await;
if local_models.is_empty() {
Ok(FetchModelsResult {
models: vec![],
source: ModelFetchSource::LocalFallback,
error: Some(format!("API 获取失败: {}, 本地也无数据", api_error)),
})
} else {
Ok(FetchModelsResult {
models: local_models,
source: ModelFetchSource::LocalFallback,
error: Some(format!("API 获取失败: {}, 已使用本地数据", api_error)),
})
}
}
}
}
/// 构建 /v1/models API URL
fn build_models_api_url(api_host: &str) -> String {
let host = api_host.trim_end_matches('/');
// 检查是否已经包含 /v1 路径
if host.ends_with("/v1") || host.ends_with("/v1/") {
format!("{}/models", host.trim_end_matches('/'))
} else if host.contains("/v1/") {
// 如果路径中间有 /v1/,直接追加 models
format!("{}models", host.trim_end_matches('/').to_string() + "/")
} else {
format!("{}/v1/models", host)
}
}
/// 调用 /v1/models API
async fn call_models_api(
&self,
url: &str,
api_key: &str,
) -> Result<Vec<ApiModelResponse>, String> {
let client = reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(30))
.build()
.map_err(|e| format!("创建 HTTP 客户端失败: {}", e))?;
let response = client
.get(url)
.header("Authorization", format!("Bearer {}", api_key))
.header("Content-Type", "application/json")
.send()
.await
.map_err(|e| format!("请求失败: {}", e))?;
if !response.status().is_success() {
let status = response.status();
let body = response
.text()
.await
.unwrap_or_else(|_| "无法读取响应体".to_string());
return Err(format!("API 返回错误 {}: {}", status, body));
}
let body = response
.text()
.await
.map_err(|e| format!("读取响应失败: {}", e))?;
// 解析 OpenAI 格式的响应
let api_response: ApiModelsResponse =
serde_json::from_str(&body).map_err(|e| format!("解析响应失败: {}", e))?;
Ok(api_response.data)
}
/// 转换 API 模型格式为内部格式
fn convert_api_model(
&self,
model: ApiModelResponse,
provider_id: &str,
now: i64,
) -> EnhancedModelMetadata {
// 从 model id 推断显示名称
let display_name = model.id.split('/').last().unwrap_or(&model.id).to_string();
EnhancedModelMetadata {
id: model.id.clone(),
display_name,
provider_id: provider_id.to_string(),
provider_name: model.owned_by.unwrap_or_else(|| provider_id.to_string()),
family: None,
tier: ModelTier::Pro,
capabilities: ModelCapabilities {
vision: false,
tools: false,
streaming: true,
json_mode: false,
function_calling: false,
reasoning: false,
},
pricing: None,
limits: ModelLimits {
context_length: model.context_length,
max_output_tokens: None,
requests_per_minute: None,
tokens_per_minute: None,
},
status: ModelStatus::Active,
release_date: None,
is_latest: false,
description: None,
source: ModelSource::Api,
created_at: now,
updated_at: now,
}
}
}
// ============================================================================
// API 响应类型
// ============================================================================
/// OpenAI /v1/models API 响应格式
#[derive(Debug, Deserialize)]
struct ApiModelsResponse {
data: Vec<ApiModelResponse>,
}
/// 单个模型的 API 响应
#[derive(Debug, Deserialize)]
struct ApiModelResponse {
id: String,
#[serde(default)]
owned_by: Option<String>,
#[serde(default)]
context_length: Option<u32>,
}
/// 模型获取来源
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub enum ModelFetchSource {
/// 从 API 获取
Api,
/// 从本地文件回退
LocalFallback,
}
/// 从 API 获取模型的结果
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FetchModelsResult {
/// 模型列表
pub models: Vec<EnhancedModelMetadata>,
/// 数据来源
pub source: ModelFetchSource,
/// 错误信息(如果有)
pub error: Option<String>,
}
+1 -1
View File
@@ -1,7 +1,7 @@
{
"$schema": "https://schema.tauri.app/config/2",
"productName": "ProxyCast",
"version": "0.47.3",
"version": "0.47.4",
"identifier": "com.proxycast.app",
"build": {
"beforeDevCommand": "npm run dev",
+26 -1
View File
@@ -8,6 +8,7 @@ import {
RefreshCw,
} from "lucide-react";
import * as Select from "@radix-ui/react-select";
import { invoke } from "@tauri-apps/api/core";
import { LogsTab } from "./LogsTab";
import { RoutesTab } from "./RoutesTab";
import { ProviderIcon } from "@/icons/providers";
@@ -40,6 +41,13 @@ import {
} from "@/lib/api/modelRegistry";
import type { EnhancedModelMetadata } from "@/lib/types/modelRegistry";
// API 获取模型结果类型
interface FetchModelsResult {
models: EnhancedModelMetadata[];
source: "Api" | "LocalFallback";
error: string | null;
}
interface TestState {
endpoint: string;
status: "idle" | "loading" | "success" | "error";
@@ -348,8 +356,25 @@ export function ApiServerPage() {
}
// 2. 添加模型注册表中的模型
// 优先使用 registryId,如果没有模型则回退到 fallbackRegistryId
// 优先使用 registryId,如果没有模型则尝试从 API 获取,最后才回退到 fallbackRegistryId
let registryModels = await getModelsForProvider(registryId);
// 3. 如果本地模型注册表没有模型,优先尝试从 Provider API 获取
if (registryModels.length === 0) {
try {
const result = await invoke<FetchModelsResult>(
"fetch_provider_models_auto",
{ providerId: provider },
);
if (result && result.models && result.models.length > 0) {
registryModels = result.models;
}
} catch {
// API 获取失败,继续尝试 fallback
}
}
// 4. 如果 API 也没有获取到模型,回退到 fallbackRegistryId
if (
registryModels.length === 0 &&
fallbackRegistryId &&
@@ -1,15 +1,32 @@
/**
* @file ProviderModelList 组件
* @description 显示 Provider 支持的模型列表
* @description 显示 Provider 支持的模型列表,支持从 API 刷新
* @module components/provider-pool/api-key/ProviderModelList
*/
import React, { useMemo } from "react";
import React, { useMemo, useState, useCallback } from "react";
import { cn } from "@/lib/utils";
import { useModelRegistry } from "@/hooks/useModelRegistry";
import { Eye, Wrench, Brain, Sparkles, Loader2 } from "lucide-react";
import {
Eye,
Wrench,
Brain,
Sparkles,
Loader2,
RefreshCw,
Cloud,
HardDrive,
} from "lucide-react";
import { Button } from "@/components/ui/button";
import {
Tooltip,
TooltipContent,
TooltipProvider,
TooltipTrigger,
} from "@/components/ui/tooltip";
import type { EnhancedModelMetadata } from "@/lib/types/modelRegistry";
import { mapProviderIdToRegistryId } from "./providerTypeMapping";
import { invoke } from "@tauri-apps/api/core";
// ============================================================================
// 类型定义
@@ -20,12 +37,24 @@ export interface ProviderModelListProps {
providerId: string;
/** Provider 类型(API 协议),如 "anthropic", "openai", "gemini" */
providerType: string;
/** 是否有可用的 API Key(用于显示刷新按钮) */
hasApiKey?: boolean;
/** 额外的 CSS 类名 */
className?: string;
/** 最大显示数量,默认显示全部 */
maxItems?: number;
}
// ============================================================================
// API 响应类型
// ============================================================================
interface FetchModelsResult {
models: EnhancedModelMetadata[];
source: "Api" | "LocalFallback";
error: string | null;
}
// ============================================================================
// 子组件
// ============================================================================
@@ -103,6 +132,7 @@ const ModelItem: React.FC<ModelItemProps> = ({ model }) => {
export const ProviderModelList: React.FC<ProviderModelListProps> = ({
providerId,
providerType,
hasApiKey = false,
className,
maxItems,
}) => {
@@ -118,18 +148,60 @@ export const ProviderModelList: React.FC<ProviderModelListProps> = ({
providerFilter: [registryProviderId],
});
// 从 API 刷新状态
const [refreshing, setRefreshing] = useState(false);
const [apiModels, setApiModels] = useState<EnhancedModelMetadata[] | null>(
null,
);
const [apiSource, setApiSource] = useState<"Api" | "LocalFallback" | null>(
null,
);
const [apiError, setApiError] = useState<string | null>(null);
// 从 API 获取模型列表(自动获取 API Key)
const handleRefreshFromApi = useCallback(async () => {
setRefreshing(true);
setApiError(null);
try {
const result = await invoke<FetchModelsResult>(
"fetch_provider_models_auto",
{
providerId,
},
);
if (result && result.models) {
setApiModels(result.models);
setApiSource(result.source);
if (result.error) {
setApiError(result.error);
}
} else {
setApiError("返回结果格式错误");
}
} catch (err) {
setApiError(err instanceof Error ? err.message : String(err));
} finally {
setRefreshing(false);
}
}, [providerId]);
// 使用 API 模型或本地模型
const displayModelsSource = apiModels ?? models;
// 限制显示数量
const displayModels = useMemo(() => {
if (maxItems && maxItems > 0) {
return models.slice(0, maxItems);
return displayModelsSource.slice(0, maxItems);
}
return models;
}, [models, maxItems]);
return displayModelsSource;
}, [displayModelsSource, maxItems]);
const hasMore = maxItems && models.length > maxItems;
const hasMore = maxItems && displayModelsSource.length > maxItems;
// 加载状态
if (loading) {
if (loading && !apiModels) {
return (
<div
className={cn(
@@ -145,7 +217,7 @@ export const ProviderModelList: React.FC<ProviderModelListProps> = ({
}
// 错误状态
if (error) {
if (error && !apiModels) {
return (
<div
className={cn("py-4 text-center text-sm text-red-500", className)}
@@ -157,16 +229,57 @@ export const ProviderModelList: React.FC<ProviderModelListProps> = ({
}
// 空状态
if (models.length === 0) {
if (displayModelsSource.length === 0) {
return (
<div
className={cn(
"py-4 text-center text-sm text-muted-foreground",
className,
<div className={cn("space-y-2", className)}>
<div className="flex items-center justify-between mb-2">
<h4 className="text-sm font-medium text-foreground flex items-center gap-2">
<Sparkles className="h-4 w-4 text-muted-foreground" />
支持的模型
</h4>
{hasApiKey && (
<TooltipProvider>
<Tooltip>
<TooltipTrigger asChild>
<Button
variant="ghost"
size="sm"
onClick={handleRefreshFromApi}
disabled={refreshing}
className="h-7 px-2"
>
{refreshing ? (
<Loader2 className="h-3.5 w-3.5 animate-spin" />
) : (
<RefreshCw className="h-3.5 w-3.5" />
)}
</Button>
</TooltipTrigger>
<TooltipContent>从 API 获取模型列表</TooltipContent>
</Tooltip>
</TooltipProvider>
)}
</div>
<div
className="py-4 text-center text-sm text-muted-foreground"
data-testid="provider-model-list-empty"
>
暂无模型数据
{hasApiKey && (
<Button
variant="ghost"
size="sm"
onClick={handleRefreshFromApi}
disabled={refreshing}
className="ml-1 h-auto p-0 text-primary underline-offset-4 hover:underline"
>
点击从 API 获取
</Button>
)}
</div>
{apiError && (
<div className="text-xs text-amber-500 text-center">{apiError}</div>
)}
data-testid="provider-model-list-empty"
>
暂无模型数据
</div>
);
}
@@ -182,11 +295,73 @@ export const ProviderModelList: React.FC<ProviderModelListProps> = ({
<Sparkles className="h-4 w-4 text-muted-foreground" />
支持的模型
<span className="text-xs text-muted-foreground font-normal">
({models.length})
({displayModelsSource.length})
</span>
{/* 数据来源标识 */}
{apiSource && (
<TooltipProvider>
<Tooltip>
<TooltipTrigger asChild>
<span
className={cn(
"inline-flex items-center gap-1 text-xs px-1.5 py-0.5 rounded",
apiSource === "Api"
? "bg-green-100 text-green-700 dark:bg-green-900/30 dark:text-green-400"
: "bg-amber-100 text-amber-700 dark:bg-amber-900/30 dark:text-amber-400",
)}
>
{apiSource === "Api" ? (
<>
<Cloud className="h-3 w-3" />
API
</>
) : (
<>
<HardDrive className="h-3 w-3" />
本地
</>
)}
</span>
</TooltipTrigger>
<TooltipContent>
{apiSource === "Api"
? "数据来自 Provider API"
: "API 获取失败,使用本地数据"}
</TooltipContent>
</Tooltip>
</TooltipProvider>
)}
</h4>
{/* 刷新按钮 */}
{hasApiKey && (
<TooltipProvider>
<Tooltip>
<TooltipTrigger asChild>
<Button
variant="ghost"
size="sm"
onClick={handleRefreshFromApi}
disabled={refreshing}
className="h-7 px-2"
>
{refreshing ? (
<Loader2 className="h-3.5 w-3.5 animate-spin" />
) : (
<RefreshCw className="h-3.5 w-3.5" />
)}
</Button>
</TooltipTrigger>
<TooltipContent>从 API 获取最新模型列表</TooltipContent>
</Tooltip>
</TooltipProvider>
)}
</div>
{/* API 错误提示 */}
{apiError && (
<div className="text-xs text-amber-500 mb-2 px-1">{apiError}</div>
)}
{/* 模型列表 */}
<div className="border rounded-md divide-y divide-border">
{displayModels.map((model) => (
@@ -197,7 +372,7 @@ export const ProviderModelList: React.FC<ProviderModelListProps> = ({
{/* 显示更多提示 */}
{hasMore && (
<p className="text-xs text-muted-foreground text-center pt-2">
还有 {models.length - maxItems!} 个模型未显示
还有 {displayModelsSource.length - maxItems!} 个模型未显示
</p>
)}
</div>
@@ -242,6 +242,9 @@ export const ProviderSetting: React.FC<ProviderSettingProps> = ({
<ProviderModelList
providerId={provider.id}
providerType={provider.type}
hasApiKey={
(provider.api_keys?.filter((k) => k.enabled).length ?? 0) > 0
}
/>
</section>
</div>
@@ -20,38 +20,60 @@ const PROVIDER_ID_TO_REGISTRY_ID: Record<string, string> = {
// 主流 AI
openai: "openai",
anthropic: "anthropic",
gemini: "gemini",
google: "google", // Gemini
deepseek: "deepseek",
moonshot: "moonshot",
moonshotai: "moonshotai",
groq: "groq",
grok: "grok",
xai: "xai", // Grok
mistral: "mistral",
perplexity: "perplexity",
cohere: "cohere",
// 国内 AI
zhipu: "zhipu",
zhipuai: "zhipuai",
baichuan: "baichuan",
dashscope: "dashscope",
alibaba: "alibaba", // 百炼/通义千问
doubao: "doubao",
minimax: "minimax",
stepfun: "stepfun",
lingyi: "lingyi",
baidu: "baidu",
yi: "yi", // 零一万物
"baidu-cloud": "baidu-cloud",
hunyuan: "hunyuan",
spark: "spark",
xiaomi: "xiaomi", // 小米 MiMo
// 云服务
"azure-openai": "openai",
vertexai: "google",
"aws-bedrock": "anthropic",
"google-vertex": "google-vertex",
"amazon-bedrock": "amazon-bedrock",
"github-models": "github-models",
"github-copilot": "github-copilot",
// API 聚合服务
siliconflow: "siliconflow",
"siliconflow-cn": "siliconflow-cn",
openrouter: "openrouter",
togetherai: "togetherai",
"fireworks-ai": "fireworks-ai",
aihubmix: "aihubmix",
"302ai": "302ai",
// 代理服务
iflow: "deepseek", // iFlow 是 DeepSeek 的代理
antigravity: "antigravity", // Antigravity 使用自己的模型列表
codex: "codex", // Codex 使用自己的模型列表
// 其他
iflow: "deepseek",
antigravity: "antigravity",
codex: "codex",
// 本地服务
ollama: "ollama",
together: "together",
fireworks: "fireworks",
replicate: "replicate",
lmstudio: "lmstudio",
// 兼容旧 ID(向后兼容)
gemini: "google",
zhipu: "zhipuai",
dashscope: "alibaba",
moonshot: "moonshotai",
grok: "xai",
github: "github-models",
copilot: "github-copilot",
vertexai: "google-vertex",
"aws-bedrock": "amazon-bedrock",
together: "togetherai",
fireworks: "fireworks-ai",
mimo: "xiaomi",
silicon: "siliconflow",
};
/**
@@ -196,6 +196,24 @@ export const TerminalAIModeSelector: React.FC<TerminalAIModeSelectorProps> = ({
return hookModels;
}
// 自定义 API Key Provider(非系统预设)直接使用 hook 结果
// 判断依据:providerId 不在系统预设列表中
const systemProviders = [
"kiro",
"codex",
"gemini",
"gemini_api_key",
"antigravity",
"qwen",
"claude",
"claude_oauth",
"openai",
"iflow",
];
if (!systemProviders.includes(selectedProvider.key.toLowerCase())) {
return hookModels;
}
// 优先使用凭证池中的模型列表(从后端 extract_supported_models 逻辑)
const credentialModels = extractSupportedModels(
selectedProvider.type,
+142 -19
View File
@@ -4,7 +4,8 @@
* @module hooks/useProviderModels
*/
import { useMemo } from "react";
import { useMemo, useState, useEffect } from "react";
import { invoke } from "@tauri-apps/api/core";
import { useModelRegistry } from "./useModelRegistry";
import { useAliasConfig } from "./useAliasConfig";
import { isAliasProvider } from "@/lib/constants/providerMappings";
@@ -36,6 +37,13 @@ export interface UseProviderModelsResult {
error: string | null;
}
// API 获取模型结果类型
interface FetchModelsResult {
models: EnhancedModelMetadata[];
source: "Api" | "LocalFallback";
error: string | null;
}
// ============================================================================
// 工具函数
// ============================================================================
@@ -157,6 +165,7 @@ function convertAliasModelsToMetadata(
* 获取 Provider 的模型列表
*
* 根据 Provider 类型,从别名配置或模型注册表获取模型列表。
* 如果本地没有模型,会尝试从 Provider API 获取。
* 支持返回模型 ID 列表或完整的模型元数据。
*
* @param selectedProvider 当前选中的 Provider
@@ -191,10 +200,15 @@ export function useProviderModels(
const { aliasConfig, loading: aliasLoading } =
useAliasConfig(selectedProvider);
// 计算模型列表
const result = useMemo(() => {
// API 获取的模型缓存
const [apiModels, setApiModels] = useState<EnhancedModelMetadata[]>([]);
const [apiLoading, setApiLoading] = useState(false);
const [apiError, setApiError] = useState<string | null>(null);
// 计算本地模型列表
const localResult = useMemo(() => {
if (!selectedProvider) {
return { modelIds: [], models: [] };
return { modelIds: [], models: [], hasLocalModels: false };
}
// 收集所有模型
@@ -236,16 +250,6 @@ export function useProviderModels(
(m) => m.provider_id === selectedProvider.registryId,
);
// 如果没有找到模型,尝试使用 fallbackRegistryId
if (
registryFilteredModels.length === 0 &&
selectedProvider.fallbackRegistryId
) {
registryFilteredModels = registryModels.filter(
(m) => m.provider_id === selectedProvider.fallbackRegistryId,
);
}
// 过滤掉已存在的模型(避免重复)
const newRegistryModels = registryFilteredModels.filter(
(m) => !allModelIds.includes(m.id),
@@ -257,20 +261,139 @@ export function useProviderModels(
allModels = [...allModels, ...sortedRegistryModels];
allModelIds = [...allModelIds, ...sortedRegistryModels.map((m) => m.id)];
// 判断是否有本地模型(不包括自定义模型)
const hasLocalModels =
sortedRegistryModels.length > 0 ||
(isAliasProvider(selectedProvider.key) &&
aliasConfig &&
aliasConfig.models.length > 0);
return {
modelIds: allModelIds,
models: returnFullMetadata ? allModels : [],
models: allModels,
hasLocalModels,
};
}, [selectedProvider, registryModels, aliasConfig, returnFullMetadata]);
}, [selectedProvider, registryModels, aliasConfig]);
// 当本地没有模型时,从 API 获取
useEffect(() => {
if (!selectedProvider) {
setApiModels([]);
return;
}
// 如果是别名 Provider,不从 API 获取
if (isAliasProvider(selectedProvider.key)) {
return;
}
// 如果本地有模型,不需要从 API 获取
if (localResult.hasLocalModels) {
setApiModels([]);
return;
}
// 如果还在加载本地数据,等待
if (registryLoading || aliasLoading) {
return;
}
// 从 API 获取模型
const fetchFromApi = async () => {
setApiLoading(true);
setApiError(null);
try {
const result = await invoke<FetchModelsResult>(
"fetch_provider_models_auto",
{ providerId: selectedProvider.key },
);
if (result && result.models && result.models.length > 0) {
setApiModels(result.models);
} else {
// API 没有返回模型,尝试 fallback
if (selectedProvider.fallbackRegistryId) {
const fallbackModels = registryModels.filter(
(m) => m.provider_id === selectedProvider.fallbackRegistryId,
);
if (fallbackModels.length > 0) {
setApiModels(sortModels(fallbackModels));
}
}
}
} catch (err) {
setApiError(err instanceof Error ? err.message : String(err));
// API 失败,尝试 fallback
if (selectedProvider.fallbackRegistryId) {
const fallbackModels = registryModels.filter(
(m) => m.provider_id === selectedProvider.fallbackRegistryId,
);
if (fallbackModels.length > 0) {
setApiModels(sortModels(fallbackModels));
}
}
} finally {
setApiLoading(false);
}
};
fetchFromApi();
}, [
selectedProvider,
localResult.hasLocalModels,
registryLoading,
aliasLoading,
registryModels,
]);
// 合并本地模型和 API 模型
const finalResult = useMemo(() => {
// 如果有本地模型,使用本地模型
if (localResult.hasLocalModels || localResult.models.length > 0) {
return {
modelIds: localResult.modelIds,
models: returnFullMetadata ? localResult.models : [],
};
}
// 否则使用 API 模型
if (apiModels.length > 0) {
// 合并自定义模型和 API 模型
const customModels = selectedProvider?.customModels || [];
const customModelMetadata =
customModels.length > 0
? convertCustomModelsToMetadata(
customModels,
selectedProvider!.key,
selectedProvider!.label,
)
: [];
const allModels = [...customModelMetadata, ...apiModels];
const allModelIds = allModels.map((m) => m.id);
return {
modelIds: allModelIds,
models: returnFullMetadata ? allModels : [],
};
}
return {
modelIds: localResult.modelIds,
models: returnFullMetadata ? localResult.models : [],
};
}, [localResult, apiModels, returnFullMetadata, selectedProvider]);
// 计算加载状态
const loading = registryLoading || aliasLoading;
const loading = registryLoading || aliasLoading || apiLoading;
// 计算错误状态
const error = registryError || null;
const error = registryError || apiError || null;
return {
...result,
...finalResult,
loading,
error,
};
+24 -17
View File
@@ -11,10 +11,26 @@ const __dirname = path.dirname(fileURLToPath(import.meta.url));
// 获取 Tauri mock 目录路径
const tauriMockDir = path.resolve(__dirname, "./src/lib/tauri-mock");
export default defineConfig(({ mode }) => ({
export default defineConfig(({ mode }) => {
// 检查是否在 Tauri 环境中运行(通过环境变量判断)
const isTauri = process.env.TAURI_ENV_PLATFORM !== undefined;
// 只在非 Tauri 环境(纯浏览器开发)下使用 mock
const tauriAliases = isTauri ? {} : {
"@tauri-apps/api/core": path.resolve(tauriMockDir, "core.ts"),
"@tauri-apps/api/event": path.resolve(tauriMockDir, "event.ts"),
"@tauri-apps/api/window": path.resolve(tauriMockDir, "window.ts"),
"@tauri-apps/api/app": path.resolve(tauriMockDir, "window.ts"),
"@tauri-apps/api/path": path.resolve(tauriMockDir, "window.ts"),
"@tauri-apps/plugin-dialog": path.resolve(tauriMockDir, "plugin-dialog.ts"),
"@tauri-apps/plugin-shell": path.resolve(tauriMockDir, "plugin-shell.ts"),
"@tauri-apps/plugin-deep-link": path.resolve(tauriMockDir, "plugin-deep-link.ts"),
"@tauri-apps/plugin-global-shortcut": path.resolve(tauriMockDir, "plugin-global-shortcut.ts"),
};
return {
plugins: [
react({
// 开发模式下启用 jsxDev 以获取组件源码位置
jsxRuntime: mode === "development" ? "automatic" : "automatic",
jsxImportSource: "react",
}),
@@ -23,23 +39,13 @@ export default defineConfig(({ mode }) => ({
resolve: {
alias: {
"@": path.resolve(__dirname, "./src"),
// 拦截所有 @tauri-apps/* 导入,重定向到 mock 模块
// 这样在浏览器开发模式下可以使用 mock 实现
"@tauri-apps/api/core": path.resolve(tauriMockDir, "core.ts"),
"@tauri-apps/api/event": path.resolve(tauriMockDir, "event.ts"),
"@tauri-apps/api/window": path.resolve(tauriMockDir, "window.ts"),
"@tauri-apps/api/app": path.resolve(tauriMockDir, "window.ts"),
"@tauri-apps/api/path": path.resolve(tauriMockDir, "window.ts"),
// 拦截 Tauri 插件
"@tauri-apps/plugin-dialog": path.resolve(tauriMockDir, "plugin-dialog.ts"),
"@tauri-apps/plugin-shell": path.resolve(tauriMockDir, "plugin-shell.ts"),
"@tauri-apps/plugin-deep-link": path.resolve(tauriMockDir, "plugin-deep-link.ts"),
"@tauri-apps/plugin-global-shortcut": path.resolve(tauriMockDir, "plugin-global-shortcut.ts"),
// 只在非 Tauri 环境下拦截 @tauri-apps/* 导入
...tauriAliases,
},
},
optimizeDeps: {
// 排除 Tauri 包的预构建,确保 alias 生效
exclude: [
// 只在非 Tauri 环境下排除 Tauri 包的预构建
exclude: isTauri ? [] : [
"@tauri-apps/api",
"@tauri-apps/plugin-dialog",
"@tauri-apps/plugin-shell",
@@ -65,4 +71,5 @@ export default defineConfig(({ mode }) => ({
"**/src-tauri/**",
],
},
}));
};
});