mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
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:
+1
-1
@@ -1,7 +1,7 @@
|
||||
{
|
||||
"name": "proxycast",
|
||||
"private": true,
|
||||
"version": "0.47.3",
|
||||
"version": "0.47.4",
|
||||
"type": "module",
|
||||
"repository": {
|
||||
"type": "git",
|
||||
|
||||
Generated
+39
-77
@@ -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
@@ -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"]
|
||||
|
||||
@@ -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" }
|
||||
@@ -1,7 +0,0 @@
|
||||
//! Aster Agent 集成模块
|
||||
//!
|
||||
//! 包含 agent 相关功能,依赖 aster 框架
|
||||
|
||||
pub fn version() -> &'static str {
|
||||
env!("CARGO_PKG_VERSION")
|
||||
}
|
||||
@@ -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"] }
|
||||
@@ -1,7 +0,0 @@
|
||||
//! Tauri 应用入口模块
|
||||
//!
|
||||
//! 包含 app, commands, tray, services 等功能
|
||||
|
||||
pub fn version() -> &'static str {
|
||||
env!("CARGO_PKG_VERSION")
|
||||
}
|
||||
@@ -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
|
||||
@@ -0,0 +1,4 @@
|
||||
//! 静态数据模块
|
||||
//!
|
||||
//! 模型数据现在从 aiclientproxy/models 仓库获取
|
||||
//! 本地硬编码数据已迁移到独立仓库: https://github.com/aiclientproxy/models
|
||||
@@ -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")
|
||||
|
||||
@@ -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()),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
@@ -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")
|
||||
}
|
||||
+1
-1
@@ -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};
|
||||
|
||||
/// 请求状态
|
||||
@@ -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"] }
|
||||
@@ -1,7 +0,0 @@
|
||||
//! Provider 系统模块
|
||||
//!
|
||||
//! 包含 providers, credential, converter 等功能
|
||||
|
||||
pub fn version() -> &'static str {
|
||||
env!("CARGO_PKG_VERSION")
|
||||
}
|
||||
@@ -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"] }
|
||||
@@ -1,7 +0,0 @@
|
||||
//! API 服务器模块
|
||||
//!
|
||||
//! 包含 server, streaming, middleware, router 等功能
|
||||
|
||||
pub fn version() -> &'static str {
|
||||
env!("CARGO_PKG_VERSION")
|
||||
}
|
||||
@@ -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": {
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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>>;
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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'",
|
||||
[],
|
||||
);
|
||||
}
|
||||
|
||||
@@ -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) => {
|
||||
|
||||
@@ -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
@@ -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;
|
||||
|
||||
@@ -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,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",
|
||||
|
||||
@@ -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
@@ -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
@@ -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/**",
|
||||
],
|
||||
},
|
||||
}));
|
||||
};
|
||||
});
|
||||
|
||||
Reference in New Issue
Block a user