diff --git a/package.json b/package.json index c99cb0260..679d5236a 100644 --- a/package.json +++ b/package.json @@ -1,7 +1,7 @@ { "name": "proxycast", "private": true, - "version": "0.47.3", + "version": "0.47.4", "type": "module", "repository": { "type": "git", diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 3b4918439..ab0540b0a 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -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" diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index f8697a777..7c535c2e5 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -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"] diff --git a/src-tauri/crates/agent/Cargo.toml b/src-tauri/crates/agent/Cargo.toml deleted file mode 100644 index c6c8b493f..000000000 --- a/src-tauri/crates/agent/Cargo.toml +++ /dev/null @@ -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" } diff --git a/src-tauri/crates/agent/src/lib.rs b/src-tauri/crates/agent/src/lib.rs deleted file mode 100644 index ca73cd0aa..000000000 --- a/src-tauri/crates/agent/src/lib.rs +++ /dev/null @@ -1,7 +0,0 @@ -//! Aster Agent 集成模块 -//! -//! 包含 agent 相关功能,依赖 aster 框架 - -pub fn version() -> &'static str { - env!("CARGO_PKG_VERSION") -} diff --git a/src-tauri/crates/app/Cargo.toml b/src-tauri/crates/app/Cargo.toml deleted file mode 100644 index 2cf3d1548..000000000 --- a/src-tauri/crates/app/Cargo.toml +++ /dev/null @@ -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"] } diff --git a/src-tauri/crates/app/src/lib.rs b/src-tauri/crates/app/src/lib.rs deleted file mode 100644 index 116cc1858..000000000 --- a/src-tauri/crates/app/src/lib.rs +++ /dev/null @@ -1,7 +0,0 @@ -//! Tauri 应用入口模块 -//! -//! 包含 app, commands, tray, services 等功能 - -pub fn version() -> &'static str { - env!("CARGO_PKG_VERSION") -} diff --git a/src-tauri/crates/core/Cargo.toml b/src-tauri/crates/core/Cargo.toml index 0a74f42d2..dbf317f36 100644 --- a/src-tauri/crates/core/Cargo.toml +++ b/src-tauri/crates/core/Cargo.toml @@ -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 \ No newline at end of file diff --git a/src-tauri/crates/core/src/data/mod.rs b/src-tauri/crates/core/src/data/mod.rs new file mode 100644 index 000000000..60a799247 --- /dev/null +++ b/src-tauri/crates/core/src/data/mod.rs @@ -0,0 +1,4 @@ +//! 静态数据模块 +//! +//! 模型数据现在从 aiclientproxy/models 仓库获取 +//! 本地硬编码数据已迁移到独立仓库: https://github.com/aiclientproxy/models diff --git a/src-tauri/crates/core/src/lib.rs b/src-tauri/crates/core/src/lib.rs index df932d6f7..089201fd2 100644 --- a/src-tauri/crates/core/src/lib.rs +++ b/src-tauri/crates/core/src/lib.rs @@ -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") diff --git a/src-tauri/crates/core/src/logger.rs b/src-tauri/crates/core/src/logger.rs new file mode 100644 index 000000000..6f3ce21ad --- /dev/null +++ b/src-tauri/crates/core/src/logger.rs @@ -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, + max_logs: usize, + config: LogStoreConfig, + log_file_path: Option, +} + +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 { + self.logs.iter().cloned().collect() + } + + pub fn clear(&mut self) { + self.logs.clear(); + } + + pub fn get_log_file_path(&self) -> Option { + 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::::from(modified); + if modified < cutoff { + let _ = fs::remove_file(entry.path()); + } + } + } +} + +/// 简化的共享日志存储类型(使用 parking_lot) +pub type SharedLogStore = Arc>; + +/// 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); + } +} diff --git a/src-tauri/crates/core/src/models/anthropic.rs b/src-tauri/crates/core/src/models/anthropic.rs new file mode 100644 index 000000000..34a956162 --- /dev/null +++ b/src-tauri/crates/core/src/models/anthropic.rs @@ -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, + #[serde(skip_serializing_if = "Option::is_none")] + pub input_schema: Option, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct AnthropicMessagesRequest { + pub model: String, + pub messages: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + pub max_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub system: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub temperature: Option, + #[serde(default)] + pub stream: bool, + #[serde(skip_serializing_if = "Option::is_none")] + pub tools: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub tool_choice: Option, +} + +#[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, + pub model: String, + pub stop_reason: Option, + 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, +} diff --git a/src-tauri/crates/core/src/models/app_type.rs b/src-tauri/crates/core/src/models/app_type.rs new file mode 100644 index 000000000..dd48833fb --- /dev/null +++ b/src-tauri/crates/core/src/models/app_type.rs @@ -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 { + 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()) + } +} diff --git a/src-tauri/crates/core/src/models/codewhisperer.rs b/src-tauri/crates/core/src/models/codewhisperer.rs new file mode 100644 index 000000000..948635247 --- /dev/null +++ b/src-tauri/crates/core/src/models/codewhisperer.rs @@ -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, +} + +#[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>, +} + +#[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>, + #[serde(skip_serializing_if = "Option::is_none")] + pub user_input_message_context: Option, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct UserInputMessageContext { + #[serde(skip_serializing_if = "Option::is_none")] + pub tools: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub tool_results: Option>, +} + +/// 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, + 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>, +} + +#[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, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(rename_all = "camelCase")] +pub struct AssistantResponseEvent { + #[serde(skip_serializing_if = "Option::is_none")] + pub content: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub tool_use: Option, +} diff --git a/src-tauri/crates/core/src/models/injection_types.rs b/src-tauri/crates/core/src/models/injection_types.rs new file mode 100644 index 000000000..446871a08 --- /dev/null +++ b/src-tauri/crates/core/src/models/injection_types.rs @@ -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 { + Some(self.cmp(other)) + } +} + +impl Eq for InjectionRule {} diff --git a/src-tauri/crates/core/src/models/kiro_fingerprint.rs b/src-tauri/crates/core/src/models/kiro_fingerprint.rs new file mode 100644 index 000000000..c3b25fba5 --- /dev/null +++ b/src-tauri/crates/core/src/models/kiro_fingerprint.rs @@ -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, + /// 最后切换时间 + pub last_switched_at: Option>, +} + +/// 指纹绑定存储 +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +pub struct KiroFingerprintStore { + /// 凭证 UUID -> 指纹绑定 + pub bindings: HashMap, +} + +impl KiroFingerprintStore { + /// 获取存储文件路径 + pub fn get_storage_path() -> Result { + 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 { + 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, + pub requires_kiro_restart: bool, +} + +impl SwitchToLocalResult { + pub fn success(message: impl Into, 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) -> Self { + Self { + success: false, + message: message.into(), + requires_action: false, + machine_id: None, + requires_kiro_restart: false, + } + } + + pub fn requires_admin(message: impl Into) -> Self { + Self { + success: false, + message: message.into(), + requires_action: true, + machine_id: None, + requires_kiro_restart: false, + } + } +} diff --git a/src-tauri/crates/core/src/models/machine_id.rs b/src-tauri/crates/core/src/models/machine_id.rs new file mode 100644 index 000000000..edc7b0928 --- /dev/null +++ b/src-tauri/crates/core/src/models/machine_id.rs @@ -0,0 +1,170 @@ +//! 机器码相关数据模型 + +use serde::{Deserialize, Serialize}; + +/// 机器码信息结构 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct MachineIdInfo { + pub current_id: String, + pub original_id: Option, + 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, +} + +/// 管理员权限状态 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct AdminStatus { + pub is_admin: bool, + pub platform: String, + pub elevation_method: Option, + 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, +} + +/// 机器码历史记录 +#[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, +} + +/// 机器码操作类型 +#[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, + pub formatted_id: Option, +} + +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 { + 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"), + } + } +} diff --git a/src-tauri/crates/core/src/models/mcp_model.rs b/src-tauri/crates/core/src/models/mcp_model.rs new file mode 100644 index 000000000..e38dd22b8 --- /dev/null +++ b/src-tauri/crates/core/src/models/mcp_model.rs @@ -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, + #[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, +} + +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()), + } + } +} diff --git a/src-tauri/crates/core/src/models/mod.rs b/src-tauri/crates/core/src/models/mod.rs new file mode 100644 index 000000000..7a833854c --- /dev/null +++ b/src-tauri/crates/core/src/models/mod.rs @@ -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}; diff --git a/src-tauri/crates/core/src/models/model_registry.rs b/src-tauri/crates/core/src/models/model_registry.rs new file mode 100644 index 000000000..e41bb0cc0 --- /dev/null +++ b/src-tauri/crates/core/src/models/model_registry.rs @@ -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, + pub output_per_million: Option, + pub cache_read_per_million: Option, + pub cache_write_per_million: Option, + 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, + pub max_output_tokens: Option, + pub requests_per_minute: Option, + pub tokens_per_minute: Option, +} + +/// 模型状态 +#[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 { + 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 { + 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 { + 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, + pub tier: ModelTier, + pub capabilities: ModelCapabilities, + pub pricing: Option, + pub limits: ModelLimits, + pub status: ModelStatus, + pub release_date: Option, + pub is_latest: bool, + pub description: Option, + 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) -> 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) -> 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) -> 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, + pub usage_count: u32, + pub last_used_at: Option, + 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, + pub model_count: u32, + pub is_syncing: bool, + pub last_error: Option, +} + +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, + pub provider: Option, + pub description: Option, +} + +/// Provider 的别名配置 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ProviderAliasConfig { + pub provider: String, + pub description: Option, + #[serde(default)] + pub models: Vec, + pub aliases: std::collections::HashMap, + pub updated_at: Option, +} + +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, + #[serde(default)] + pub npm: Option, + #[serde(default)] + pub models: std::collections::HashMap, +} + +/// models.dev API 响应中的 Model 结构 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ModelsDevModel { + pub id: String, + pub name: String, + #[serde(default)] + pub family: Option, + #[serde(default)] + pub release_date: Option, + #[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, + #[serde(default)] + pub limit: Option, + #[serde(default)] + pub modalities: Option, + #[serde(default)] + pub experimental: Option, + #[serde(default)] + pub status: Option, +} + +/// models.dev API 响应中的 Cost 结构 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ModelsDevCost { + #[serde(default)] + pub input: Option, + #[serde(default)] + pub output: Option, + #[serde(default)] + pub cache_read: Option, + #[serde(default)] + pub cache_write: Option, +} + +/// models.dev API 响应中的 Limit 结构 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ModelsDevLimit { + #[serde(default)] + pub context: Option, + #[serde(default)] + pub output: Option, +} + +/// models.dev API 响应中的 Modalities 结构 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ModelsDevModalities { + #[serde(default)] + pub input: Vec, + #[serde(default)] + pub output: Vec, +} + +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::().unwrap(), + ModelStatus::Active + ); + assert_eq!( + "deprecated".parse::().unwrap(), + ModelStatus::Deprecated + ); + assert_eq!("beta".parse::().unwrap(), ModelStatus::Beta); + } +} diff --git a/src-tauri/crates/core/src/models/openai.rs b/src-tauri/crates/core/src/models/openai.rs new file mode 100644 index 000000000..0f4927748 --- /dev/null +++ b/src-tauri/crates/core/src/models/openai.rs @@ -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, +} + +#[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), +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ChatMessage { + pub role: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub content: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub tool_calls: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub tool_call_id: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub reasoning_content: Option, +} + +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::>() + .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, + #[serde(skip_serializing_if = "Option::is_none")] + pub parameters: Option, +} + +/// 工具定义 +#[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, + #[serde(skip_serializing_if = "Option::is_none")] + pub temperature: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub max_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub top_p: Option, + #[serde(default)] + pub stream: bool, + #[serde(skip_serializing_if = "Option::is_none")] + pub tools: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub tool_choice: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub reasoning_effort: Option, +} + +#[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, + #[serde(skip_serializing_if = "Option::is_none")] + pub tool_calls: Option>, +} + +#[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, + pub usage: Usage, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct StreamDelta { + #[serde(skip_serializing_if = "Option::is_none")] + pub role: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub content: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub tool_calls: Option>, +} + +#[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, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ChatCompletionChunk { + pub id: String, + pub object: String, + pub created: u64, + pub model: String, + pub choices: Vec, +} + +// 图像生成 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, + #[serde(default = "default_response_format")] + pub response_format: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub quality: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub style: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub user: Option, +} + +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, +} + +/// 单个图像数据 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ImageData { + #[serde(skip_serializing_if = "Option::is_none")] + pub b64_json: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub url: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub revised_prompt: Option, +} diff --git a/src-tauri/crates/core/src/models/prompt_model.rs b/src-tauri/crates/core/src/models/prompt_model.rs new file mode 100644 index 000000000..c2313533d --- /dev/null +++ b/src-tauri/crates/core/src/models/prompt_model.rs @@ -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, + #[serde(default)] + pub enabled: bool, + #[serde(rename = "createdAt", skip_serializing_if = "Option::is_none")] + pub created_at: Option, + #[serde(rename = "updatedAt", skip_serializing_if = "Option::is_none")] + pub updated_at: Option, +} + +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), + } + } +} diff --git a/src-tauri/crates/core/src/models/provider_model.rs b/src-tauri/crates/core/src/models/provider_model.rs new file mode 100644 index 000000000..a0176d8bf --- /dev/null +++ b/src-tauri/crates/core/src/models/provider_model.rs @@ -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, + #[serde(skip_serializing_if = "Option::is_none")] + pub icon: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub icon_color: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub notes: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub created_at: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub sort_index: Option, + #[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, + } + } +} diff --git a/src-tauri/crates/core/src/models/provider_pool_model.rs b/src-tauri/crates/core/src/models/provider_pool_model.rs new file mode 100644 index 000000000..1c2ad32e6 --- /dev/null +++ b/src-tauri/crates/core/src/models/provider_pool_model.rs @@ -0,0 +1,1116 @@ +//! Provider Pool 数据模型 +//! +//! 支持多凭证池管理,包括健康检测、负载均衡、故障转移等功能。 + +use chrono::{DateTime, Utc}; +use serde::{Deserialize, Serialize}; +use std::collections::HashMap; +use uuid::Uuid; + +use super::provider_type::ANTIGRAVITY_MODELS_FALLBACK; + +/// 凭证来源枚举 +/// 用于标识凭证是如何添加到凭证池的 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize, Default)] +#[serde(rename_all = "snake_case")] +pub enum CredentialSource { + /// 手动添加(通过 UI 添加) + #[default] + Manual, + /// 导入(从文件导入) + Imported, + /// 私有凭证(从高级设置迁移) + Private, +} + +/// Provider 类型别名 +/// +/// 为了向后兼容,PoolProviderType 是 crate::ProviderType 的类型别名。 +/// 所有 Provider 类型定义已统一到 lib.rs 中的 ProviderType。 +pub type PoolProviderType = crate::ProviderType; + +/// 凭证数据,根据 Provider 类型不同而不同 +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum CredentialData { + /// Kiro OAuth 凭证(文件路径) + KiroOAuth { creds_file_path: String }, + /// Gemini OAuth 凭证(文件路径) + GeminiOAuth { + creds_file_path: String, + project_id: Option, + }, + + /// Antigravity OAuth 凭证(文件路径)- Google 内部 Gemini 3 Pro + AntigravityOAuth { + creds_file_path: String, + project_id: Option, + }, + /// OpenAI API Key 凭证 + OpenAIKey { + api_key: String, + base_url: Option, + }, + /// Claude API Key 凭证 + ClaudeKey { + api_key: String, + base_url: Option, + }, + /// Vertex AI API Key 凭证 + VertexKey { + api_key: String, + base_url: Option, + /// Model alias mappings (alias -> upstream model name) + #[serde(default)] + model_aliases: std::collections::HashMap, + }, + /// Gemini API Key 凭证(多账号负载均衡) + GeminiApiKey { + api_key: String, + base_url: Option, + /// 排除的模型列表(支持通配符) + #[serde(default)] + excluded_models: Vec, + }, + /// Codex OAuth 凭证(OpenAI Codex) + CodexOAuth { + creds_file_path: String, + /// API Base URL(可选,默认使用凭证文件中的配置) + #[serde(default)] + api_base_url: Option, + }, + /// Claude OAuth 凭证(Anthropic OAuth) + ClaudeOAuth { creds_file_path: String }, + + /// Anthropic API Key 凭证(直接使用 Anthropic API) + AnthropicKey { + api_key: String, + base_url: Option, + }, +} + +impl CredentialData { + /// 获取凭证的显示名称(隐藏敏感信息) + pub fn display_name(&self) -> String { + match self { + CredentialData::KiroOAuth { creds_file_path } => { + format!("Kiro OAuth: {}", mask_path(creds_file_path)) + } + CredentialData::GeminiOAuth { + creds_file_path, .. + } => { + format!("Gemini OAuth: {}", mask_path(creds_file_path)) + } + + CredentialData::AntigravityOAuth { + creds_file_path, .. + } => { + format!("Antigravity OAuth: {}", mask_path(creds_file_path)) + } + CredentialData::OpenAIKey { api_key, .. } => { + format!("OpenAI: {}", mask_key(api_key)) + } + CredentialData::ClaudeKey { api_key, .. } => { + format!("Claude: {}", mask_key(api_key)) + } + CredentialData::VertexKey { api_key, .. } => { + format!("Vertex AI: {}", mask_key(api_key)) + } + CredentialData::GeminiApiKey { api_key, .. } => { + format!("Gemini API Key: {}", mask_key(api_key)) + } + CredentialData::CodexOAuth { + creds_file_path, .. + } => { + format!("Codex OAuth: {}", mask_path(creds_file_path)) + } + CredentialData::ClaudeOAuth { creds_file_path } => { + format!("Claude OAuth: {}", mask_path(creds_file_path)) + } + + CredentialData::AnthropicKey { api_key, .. } => { + format!("Anthropic: {}", mask_key(api_key)) + } + } + } + + /// 获取 Provider 类型 + pub fn provider_type(&self) -> PoolProviderType { + match self { + CredentialData::KiroOAuth { .. } => PoolProviderType::Kiro, + CredentialData::GeminiOAuth { .. } => PoolProviderType::Gemini, + + CredentialData::AntigravityOAuth { .. } => PoolProviderType::Antigravity, + CredentialData::OpenAIKey { .. } => PoolProviderType::OpenAI, + CredentialData::ClaudeKey { .. } => PoolProviderType::Claude, + CredentialData::VertexKey { .. } => PoolProviderType::Vertex, + CredentialData::GeminiApiKey { .. } => PoolProviderType::GeminiApiKey, + CredentialData::CodexOAuth { .. } => PoolProviderType::Codex, + CredentialData::ClaudeOAuth { .. } => PoolProviderType::ClaudeOAuth, + + CredentialData::AnthropicKey { .. } => PoolProviderType::Anthropic, + } + } +} + +/// 通配符模式匹配 +/// +/// 支持的通配符模式: +/// - 精确匹配: `claude-sonnet-4-5` +/// - 前缀匹配: `claude-*` +/// - 后缀匹配: `*-preview` +/// - 包含匹配: `*flash*` +pub fn pattern_matches(pattern: &str, model: &str) -> bool { + if !pattern.contains('*') { + return pattern == model; + } + + let parts: Vec<&str> = pattern.split('*').collect(); + + match parts.as_slice() { + [prefix, ""] => model.starts_with(prefix), + ["", suffix] => model.ends_with(suffix), + ["", middle, ""] => model.contains(middle), + [prefix, suffix] => model.starts_with(prefix) && model.ends_with(suffix), + _ => false, + } +} + +/// 单个凭证 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ProviderCredential { + /// 唯一标识符 + pub uuid: String, + /// Provider 类型 + pub provider_type: PoolProviderType, + /// 凭证数据 + pub credential: CredentialData, + /// 备注/名称 + pub name: Option, + /// 是否健康 + #[serde(default = "default_true")] + pub is_healthy: bool, + /// 是否禁用(手动禁用) + #[serde(default)] + pub is_disabled: bool, + /// 是否启用自动健康检查 + #[serde(default = "default_true")] + pub check_health: bool, + /// 自定义健康检查模型 + pub check_model_name: Option, + /// 不支持的模型列表(黑名单) + #[serde(default)] + pub not_supported_models: Vec, + /// 支持的模型列表(从 /v1/models 接口获取) + #[serde(default)] + pub supported_models: Vec, + /// 使用次数 + #[serde(default)] + pub usage_count: u64, + /// 错误次数 + #[serde(default)] + pub error_count: u32, + /// 最后使用时间 + pub last_used: Option>, + /// 最后错误时间 + pub last_error_time: Option>, + /// 最后错误消息 + pub last_error_message: Option, + /// 最后健康检查时间 + pub last_health_check_time: Option>, + /// 最后健康检查使用的模型 + pub last_health_check_model: Option, + /// 创建时间 + pub created_at: DateTime, + /// 更新时间 + pub updated_at: DateTime, + /// Token 缓存信息 + #[serde(default)] + pub cached_token: Option, + /// 凭证来源(手动添加/导入/私有) + #[serde(default)] + pub source: CredentialSource, + /// 代理 URL(可覆盖全局代理设置) + pub proxy_url: Option, +} + +fn default_true() -> bool { + true +} + +impl ProviderCredential { + /// 创建新凭证 + pub fn new(provider_type: PoolProviderType, credential: CredentialData) -> Self { + let now = Utc::now(); + Self { + uuid: Uuid::new_v4().to_string(), + provider_type, + credential, + name: None, + is_healthy: true, + is_disabled: false, + check_health: true, + check_model_name: None, + not_supported_models: Vec::new(), + supported_models: Vec::new(), + usage_count: 0, + error_count: 0, + last_used: None, + last_error_time: None, + last_error_message: None, + last_health_check_time: None, + last_health_check_model: None, + created_at: now, + updated_at: now, + cached_token: None, + source: CredentialSource::Manual, + proxy_url: None, + } + } + + /// 创建带来源的新凭证 + pub fn new_with_source( + provider_type: PoolProviderType, + credential: CredentialData, + source: CredentialSource, + ) -> Self { + let mut cred = Self::new(provider_type, credential); + cred.source = source; + cred + } + + /// 是否可用(健康且未禁用) + pub fn is_available(&self) -> bool { + self.is_healthy && !self.is_disabled + } + + /// 是否支持指定模型 + /// + /// 检查两个来源的排除列表: + /// 1. `not_supported_models` - 通用的不支持模型列表(精确匹配) + /// 2. `excluded_models` - 来自 CredentialData::GeminiApiKey 的排除列表(支持通配符) + /// 3. Antigravity 凭证只支持特定的模型列表 + pub fn supports_model(&self, model: &str) -> bool { + // 检查通用的不支持模型列表(精确匹配) + if self.not_supported_models.contains(&model.to_string()) { + return false; + } + + // 检查 GeminiApiKey 的 excluded_models(支持通配符) + if let CredentialData::GeminiApiKey { + excluded_models, .. + } = &self.credential + { + for pattern in excluded_models { + if pattern_matches(pattern, model) { + return false; + } + } + } + + // Antigravity 凭证只支持特定的模型 + // 使用 providers::antigravity 中定义的模型列表(fallback) + // 实际模型列表由 models/aliases/antigravity.json 定义 + if let CredentialData::AntigravityOAuth { .. } = &self.credential { + return ANTIGRAVITY_MODELS_FALLBACK.contains(&model); + } + + true + } + + /// 标记为健康 + pub fn mark_healthy(&mut self, check_model: Option) { + self.is_healthy = true; + self.error_count = 0; + self.last_health_check_time = Some(Utc::now()); + self.last_health_check_model = check_model; + self.updated_at = Utc::now(); + } + + /// 标记为不健康 + pub fn mark_unhealthy(&mut self, error_message: Option) { + self.error_count += 1; + self.last_error_time = Some(Utc::now()); + self.last_error_message = error_message; + self.updated_at = Utc::now(); + // 错误次数达到阈值则标记为不健康 + if self.error_count >= 3 { + self.is_healthy = false; + } + } + + /// 记录使用 + pub fn record_usage(&mut self) { + self.usage_count += 1; + self.last_used = Some(Utc::now()); + self.updated_at = Utc::now(); + } + + /// 重置计数器 + pub fn reset_counters(&mut self) { + self.usage_count = 0; + self.error_count = 0; + self.is_healthy = true; + self.last_error_time = None; + self.last_error_message = None; + self.updated_at = Utc::now(); + } +} + +/// 凭证池统计信息 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct PoolStats { + /// 总凭证数 + pub total_count: usize, + /// 健康凭证数 + pub healthy_count: usize, + /// 禁用凭证数 + pub disabled_count: usize, + /// 总使用次数 + pub total_usage: u64, + /// 总错误次数 + pub total_errors: u64, + /// 最后更新时间 + pub last_update: DateTime, +} + +impl PoolStats { + pub fn from_credentials(credentials: &[ProviderCredential]) -> Self { + Self { + total_count: credentials.len(), + healthy_count: credentials.iter().filter(|c| c.is_healthy).count(), + disabled_count: credentials.iter().filter(|c| c.is_disabled).count(), + total_usage: credentials.iter().map(|c| c.usage_count).sum(), + total_errors: credentials.iter().map(|c| c.error_count as u64).sum(), + last_update: Utc::now(), + } + } +} + +/// 健康检查结果 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct HealthCheckResult { + pub uuid: String, + pub success: bool, + pub model: Option, + pub message: Option, + pub duration_ms: u64, +} + +/// OAuth 凭证状态 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct OAuthStatus { + /// 是否有 access_token + pub has_access_token: bool, + /// 是否有 refresh_token + pub has_refresh_token: bool, + /// token 是否有效 + pub is_token_valid: bool, + /// 过期信息 + pub expiry_info: Option, + /// 凭证文件路径 + pub creds_path: String, +} + +/// Token 缓存状态(用于前端展示) +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct TokenCacheStatus { + /// 是否有缓存的 token + pub has_cached_token: bool, + /// Token 是否有效 + pub is_valid: bool, + /// Token 是否即将过期(5分钟内) + pub is_expiring_soon: bool, + /// 过期时间 + pub expiry_time: Option, + /// 最后刷新时间 + pub last_refresh: Option, + /// 连续刷新失败次数 + pub refresh_error_count: u32, + /// 最后刷新错误信息 + pub last_refresh_error: Option, +} + +/// Token 缓存信息 +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub struct CachedTokenInfo { + /// 缓存的 access_token + pub access_token: Option, + /// 缓存的 refresh_token(刷新后可能变化) + pub refresh_token: Option, + /// Token 过期时间 + pub expiry_time: Option>, + /// 最后刷新时间 + pub last_refresh: Option>, + /// 连续刷新失败次数 + #[serde(default)] + pub refresh_error_count: u32, + /// 最后刷新错误信息 + pub last_refresh_error: Option, +} + +impl CachedTokenInfo { + /// 检查 token 是否有效(存在且未过期) + pub fn is_valid(&self) -> bool { + if self.access_token.is_none() { + return false; + } + match &self.expiry_time { + Some(expiry) => *expiry > Utc::now(), + None => true, // 没有过期时间,假设有效 + } + } + + /// 检查 token 是否即将过期(5分钟内) + pub fn is_expiring_soon(&self) -> bool { + self.is_expiring_within_minutes(5) + } + + /// 检查 token 是否在指定分钟数内过期 + /// + /// # 参数 + /// - `minutes`: 检查的时间阈值(分钟) + /// + /// # 返回 + /// - `true`: Token 将在指定分钟数内过期 + /// - `false`: Token 不会在指定分钟数内过期,或没有过期时间 + pub fn is_expiring_within_minutes(&self, minutes: i64) -> bool { + match &self.expiry_time { + Some(expiry) => { + let threshold = Utc::now() + chrono::Duration::minutes(minutes); + *expiry <= threshold + } + None => false, // 没有过期时间,假设不会过期 + } + } + + /// 检查 token 是否需要刷新(无效或即将过期) + pub fn needs_refresh(&self) -> bool { + !self.is_valid() || self.is_expiring_soon() + } +} + +/// 默认健康检查模型 +pub fn get_default_check_model(provider_type: PoolProviderType) -> &'static str { + match provider_type { + PoolProviderType::Kiro => "claude-haiku-4-5", + PoolProviderType::Gemini => "gemini-2.5-flash", + PoolProviderType::OpenAI => "gpt-3.5-turbo", + // 使用 claude-sonnet-4-5-20250929,兼容更多代理服务器 + PoolProviderType::Claude => "claude-sonnet-4-5-20250929", + PoolProviderType::Antigravity => "gemini-3-pro-preview", + PoolProviderType::Vertex => "gemini-2.0-flash", + PoolProviderType::GeminiApiKey => "gemini-2.5-flash", + PoolProviderType::Codex => "gpt-4o-mini", + PoolProviderType::ClaudeOAuth => "claude-sonnet-4-5-20250929", + // API Key Provider 类型 + PoolProviderType::Anthropic => "claude-sonnet-4-5-20250929", + PoolProviderType::AzureOpenai => "gpt-4o-mini", + PoolProviderType::AwsBedrock => "claude-sonnet-4-5-20250929", + PoolProviderType::Ollama => "llama3.2", + } +} + +/// 凭证池前端展示数据 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct CredentialDisplay { + pub uuid: String, + pub provider_type: String, + pub credential_type: String, + pub name: Option, + pub display_credential: String, + pub is_healthy: bool, + pub is_disabled: bool, + pub check_health: bool, + pub check_model_name: Option, + pub not_supported_models: Vec, + pub supported_models: Vec, + pub usage_count: u64, + pub error_count: u32, + pub last_used: Option, + pub last_error_time: Option, + pub last_error_message: Option, + pub last_health_check_time: Option, + pub last_health_check_model: Option, + pub oauth_status: Option, + pub token_cache_status: Option, + pub created_at: String, + pub updated_at: String, + /// 凭证来源(手动添加/导入/私有) + pub source: CredentialSource, + /// API Key 凭证的 base_url(仅用于 OpenAI/Claude API Key 类型) + pub base_url: Option, + /// API Key 凭证的完整 api_key(仅用于 OpenAI/Claude API Key 类型,用于编辑) + pub api_key: Option, + /// 凭证级代理 URL(可覆盖全局代理设置) + pub proxy_url: Option, +} + +/// 获取凭证类型字符串 +fn get_credential_type(cred: &CredentialData) -> String { + match cred { + CredentialData::KiroOAuth { .. } => "kiro_oauth".to_string(), + CredentialData::GeminiOAuth { .. } => "gemini_oauth".to_string(), + CredentialData::AntigravityOAuth { .. } => "antigravity_oauth".to_string(), + CredentialData::OpenAIKey { .. } => "openai_key".to_string(), + CredentialData::ClaudeKey { .. } => "claude_key".to_string(), + CredentialData::VertexKey { .. } => "vertex_key".to_string(), + CredentialData::GeminiApiKey { .. } => "gemini_api_key".to_string(), + CredentialData::CodexOAuth { .. } => "codex_oauth".to_string(), + CredentialData::ClaudeOAuth { .. } => "claude_oauth".to_string(), + CredentialData::AnthropicKey { .. } => "anthropic_key".to_string(), + } +} + +/// 获取 OAuth 凭证的文件路径 +pub fn get_oauth_creds_path(cred: &CredentialData) -> Option { + match cred { + CredentialData::KiroOAuth { creds_file_path } => Some(creds_file_path.clone()), + CredentialData::GeminiOAuth { + creds_file_path, .. + } => Some(creds_file_path.clone()), + CredentialData::AntigravityOAuth { + creds_file_path, .. + } => Some(creds_file_path.clone()), + CredentialData::CodexOAuth { + creds_file_path, .. + } => Some(creds_file_path.clone()), + CredentialData::ClaudeOAuth { creds_file_path } => Some(creds_file_path.clone()), + _ => None, + } +} + +/// 从 CredentialData 中提取 base_url(仅适用于 API Key 类型) +fn get_base_url(cred: &CredentialData) -> Option { + match cred { + CredentialData::OpenAIKey { base_url, .. } => base_url.clone(), + CredentialData::ClaudeKey { base_url, .. } => base_url.clone(), + CredentialData::AnthropicKey { base_url, .. } => base_url.clone(), + _ => None, + } +} + +/// 从 CredentialData 中提取 api_key(仅适用于 API Key 类型) +fn get_api_key(cred: &CredentialData) -> Option { + match cred { + CredentialData::OpenAIKey { api_key, .. } => Some(api_key.clone()), + CredentialData::ClaudeKey { api_key, .. } => Some(api_key.clone()), + CredentialData::AnthropicKey { api_key, .. } => Some(api_key.clone()), + _ => None, + } +} + +impl From<&ProviderCredential> for CredentialDisplay { + fn from(cred: &ProviderCredential) -> Self { + // 构建 token 缓存状态 + let token_cache_status = cred.cached_token.as_ref().map(|cache| TokenCacheStatus { + has_cached_token: cache.access_token.is_some(), + is_valid: cache.is_valid(), + is_expiring_soon: cache.is_expiring_soon(), + expiry_time: cache.expiry_time.map(|t| t.to_rfc3339()), + last_refresh: cache.last_refresh.map(|t| t.to_rfc3339()), + refresh_error_count: cache.refresh_error_count, + last_refresh_error: cache.last_refresh_error.clone(), + }); + + Self { + uuid: cred.uuid.clone(), + provider_type: cred.provider_type.to_string(), + credential_type: get_credential_type(&cred.credential), + name: cred.name.clone(), + display_credential: cred.credential.display_name(), + is_healthy: cred.is_healthy, + is_disabled: cred.is_disabled, + check_health: cred.check_health, + check_model_name: cred.check_model_name.clone(), + not_supported_models: cred.not_supported_models.clone(), + supported_models: cred.supported_models.clone(), + usage_count: cred.usage_count, + error_count: cred.error_count, + last_used: cred.last_used.map(|t| t.to_rfc3339()), + last_error_time: cred.last_error_time.map(|t| t.to_rfc3339()), + last_error_message: cred.last_error_message.clone(), + last_health_check_time: cred.last_health_check_time.map(|t| t.to_rfc3339()), + last_health_check_model: cred.last_health_check_model.clone(), + oauth_status: None, // 需要单独调用获取 + token_cache_status, + created_at: cred.created_at.to_rfc3339(), + updated_at: cred.updated_at.to_rfc3339(), + source: cred.source, + base_url: get_base_url(&cred.credential), + api_key: get_api_key(&cred.credential), + proxy_url: cred.proxy_url.clone(), + } + } +} + +/// Provider 池概览(按类型分组的统计) +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ProviderPoolOverview { + pub provider_type: String, + pub stats: PoolStats, + pub credentials: Vec, +} + +// 辅助函数:隐藏路径中的用户名 +fn mask_path(path: &str) -> String { + if let Some(home) = dirs::home_dir() { + let home_str = home.to_string_lossy(); + path.replace(&*home_str, "~") + } else { + path.to_string() + } +} + +// 辅助函数:隐藏 API Key +fn mask_key(key: &str) -> String { + if key.len() <= 12 { + "****".to_string() + } else { + format!("{}...{}", &key[..6], &key[key.len() - 4..]) + } +} + +/// 添加凭证的请求结构 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct AddCredentialRequest { + pub provider_type: String, + pub credential: CredentialData, + pub name: Option, + pub check_health: Option, + pub check_model_name: Option, +} + +/// 更新凭证请求 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct UpdateCredentialRequest { + pub name: Option, + pub is_disabled: Option, + pub check_health: Option, + pub check_model_name: Option, + pub not_supported_models: Option>, + /// 新的凭证文件路径(仅适用于OAuth凭证,用于重新上传文件) + pub new_creds_file_path: Option, + /// OAuth相关:新的project_id(仅适用于Gemini) + pub new_project_id: Option, + /// API Key 相关:新的 base_url(仅适用于 API Key 凭证) + pub new_base_url: Option, + /// API Key 相关:新的 api_key(仅适用于 API Key 凭证) + pub new_api_key: Option, + /// 新的代理 URL(可覆盖全局代理设置) + pub new_proxy_url: Option, +} + +pub type ProviderPools = HashMap>; + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_pattern_matches_exact() { + assert!(pattern_matches("gemini-2.5-pro", "gemini-2.5-pro")); + assert!(!pattern_matches("gemini-2.5-pro", "gemini-2.5-flash")); + } + + #[test] + fn test_pattern_matches_prefix() { + assert!(pattern_matches("gemini-*", "gemini-2.5-pro")); + assert!(pattern_matches("gemini-*", "gemini-2.5-flash")); + assert!(!pattern_matches("gemini-*", "claude-sonnet")); + } + + #[test] + fn test_pattern_matches_suffix() { + assert!(pattern_matches("*-preview", "gemini-3-pro-preview")); + assert!(pattern_matches("*-preview", "claude-preview")); + assert!(!pattern_matches("*-preview", "gemini-2.5-pro")); + } + + #[test] + fn test_pattern_matches_contains() { + assert!(pattern_matches("*flash*", "gemini-2.5-flash")); + assert!(pattern_matches("*flash*", "gemini-2.5-flash-lite")); + assert!(!pattern_matches("*flash*", "gemini-2.5-pro")); + } + + #[test] + fn test_pattern_matches_prefix_and_suffix() { + assert!(pattern_matches("gemini-*-pro", "gemini-2.5-pro")); + assert!(pattern_matches("gemini-*-pro", "gemini-3-pro")); + assert!(!pattern_matches("gemini-*-pro", "gemini-2.5-flash")); + } + + #[test] + fn test_supports_model_not_supported_models() { + let cred = ProviderCredential { + uuid: "test-uuid".to_string(), + provider_type: PoolProviderType::Kiro, + credential: CredentialData::KiroOAuth { + creds_file_path: "/path/to/creds".to_string(), + }, + name: None, + is_healthy: true, + is_disabled: false, + check_health: true, + check_model_name: None, + not_supported_models: vec!["claude-opus".to_string()], + supported_models: vec![], + usage_count: 0, + error_count: 0, + last_used: None, + last_error_time: None, + last_error_message: None, + last_health_check_time: None, + last_health_check_model: None, + created_at: Utc::now(), + updated_at: Utc::now(), + cached_token: None, + source: CredentialSource::Manual, + proxy_url: None, + }; + + assert!(!cred.supports_model("claude-opus")); + assert!(cred.supports_model("claude-sonnet")); + } + + #[test] + fn test_supports_model_gemini_api_key_excluded_models_exact() { + let cred = ProviderCredential { + uuid: "test-uuid".to_string(), + provider_type: PoolProviderType::GeminiApiKey, + credential: CredentialData::GeminiApiKey { + api_key: "test-key".to_string(), + base_url: None, + excluded_models: vec!["gemini-2.5-pro".to_string()], + }, + name: None, + is_healthy: true, + is_disabled: false, + check_health: true, + check_model_name: None, + not_supported_models: vec![], + supported_models: vec![], + usage_count: 0, + error_count: 0, + last_used: None, + last_error_time: None, + last_error_message: None, + last_health_check_time: None, + last_health_check_model: None, + created_at: Utc::now(), + updated_at: Utc::now(), + cached_token: None, + source: CredentialSource::Manual, + proxy_url: None, + }; + + // Exact match exclusion + assert!(!cred.supports_model("gemini-2.5-pro")); + // Not excluded + assert!(cred.supports_model("gemini-2.5-flash")); + } + + #[test] + fn test_supports_model_gemini_api_key_excluded_models_wildcard() { + let cred = ProviderCredential { + uuid: "test-uuid".to_string(), + provider_type: PoolProviderType::GeminiApiKey, + credential: CredentialData::GeminiApiKey { + api_key: "test-key".to_string(), + base_url: None, + excluded_models: vec!["gemini-2.5-*".to_string(), "*-preview".to_string()], + }, + name: None, + is_healthy: true, + is_disabled: false, + check_health: true, + check_model_name: None, + not_supported_models: vec![], + supported_models: vec![], + usage_count: 0, + error_count: 0, + last_used: None, + last_error_time: None, + last_error_message: None, + last_health_check_time: None, + last_health_check_model: None, + created_at: Utc::now(), + updated_at: Utc::now(), + cached_token: None, + source: CredentialSource::Manual, + proxy_url: None, + }; + + // Prefix wildcard exclusion + assert!(!cred.supports_model("gemini-2.5-pro")); + assert!(!cred.supports_model("gemini-2.5-flash")); + // Suffix wildcard exclusion + assert!(!cred.supports_model("gemini-3-pro-preview")); + // Not excluded + assert!(cred.supports_model("gemini-2.0-flash")); + assert!(cred.supports_model("gemini-3-pro")); + } + + #[test] + fn test_supports_model_gemini_api_key_excluded_models_contains() { + let cred = ProviderCredential { + uuid: "test-uuid".to_string(), + provider_type: PoolProviderType::GeminiApiKey, + credential: CredentialData::GeminiApiKey { + api_key: "test-key".to_string(), + base_url: None, + excluded_models: vec!["*flash*".to_string()], + }, + name: None, + is_healthy: true, + is_disabled: false, + check_health: true, + check_model_name: None, + not_supported_models: vec![], + supported_models: vec![], + usage_count: 0, + error_count: 0, + last_used: None, + last_error_time: None, + last_error_message: None, + last_health_check_time: None, + last_health_check_model: None, + created_at: Utc::now(), + updated_at: Utc::now(), + cached_token: None, + source: CredentialSource::Manual, + proxy_url: None, + }; + + // Contains wildcard exclusion + assert!(!cred.supports_model("gemini-2.5-flash")); + assert!(!cred.supports_model("gemini-2.5-flash-lite")); + // Not excluded + assert!(cred.supports_model("gemini-2.5-pro")); + } + + #[test] + fn test_supports_model_combined_exclusions() { + let cred = ProviderCredential { + uuid: "test-uuid".to_string(), + provider_type: PoolProviderType::GeminiApiKey, + credential: CredentialData::GeminiApiKey { + api_key: "test-key".to_string(), + base_url: None, + excluded_models: vec!["gemini-2.5-*".to_string()], + }, + name: None, + is_healthy: true, + is_disabled: false, + check_health: true, + check_model_name: None, + not_supported_models: vec!["gemini-3-pro".to_string()], + supported_models: vec![], + usage_count: 0, + error_count: 0, + last_used: None, + last_error_time: None, + last_error_message: None, + last_health_check_time: None, + last_health_check_model: None, + created_at: Utc::now(), + updated_at: Utc::now(), + cached_token: None, + source: CredentialSource::Manual, + proxy_url: None, + }; + + // Excluded by not_supported_models (exact match) + assert!(!cred.supports_model("gemini-3-pro")); + // Excluded by excluded_models (wildcard) + assert!(!cred.supports_model("gemini-2.5-pro")); + assert!(!cred.supports_model("gemini-2.5-flash")); + // Not excluded + assert!(cred.supports_model("gemini-2.0-flash")); + } + + #[test] + fn test_supports_model_non_gemini_api_key_ignores_excluded_models() { + // For non-GeminiApiKey credentials, excluded_models in CredentialData is not checked + let cred = ProviderCredential { + uuid: "test-uuid".to_string(), + provider_type: PoolProviderType::Kiro, + credential: CredentialData::KiroOAuth { + creds_file_path: "/path/to/creds".to_string(), + }, + name: None, + is_healthy: true, + is_disabled: false, + check_health: true, + check_model_name: None, + not_supported_models: vec![], + supported_models: vec![], + usage_count: 0, + error_count: 0, + last_used: None, + last_error_time: None, + last_error_message: None, + last_health_check_time: None, + last_health_check_model: None, + created_at: Utc::now(), + updated_at: Utc::now(), + cached_token: None, + source: CredentialSource::Manual, + proxy_url: None, + }; + + // All models should be supported since not_supported_models is empty + assert!(cred.supports_model("claude-sonnet")); + assert!(cred.supports_model("claude-opus")); + } + + // ======================================================================== + // Property-Based Tests for Token Expiration Check + // ======================================================================== + + use proptest::prelude::*; + + /// 生成随机的过期时间偏移量(分钟) + fn expiry_offset_strategy() -> impl Strategy { + // 生成 -60 到 +120 分钟的偏移量 + -60i64..=120i64 + } + + /// 生成随机的检查阈值(分钟) + fn threshold_strategy() -> impl Strategy { + // 生成 1 到 30 分钟的阈值 + 1i64..=30i64 + } + + proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// **Feature: kiro-streaming-fix, Property 7: Token 过期检查** + /// + /// *对于任意* 即将过期的 Token(指定分钟数内),`is_expiring_within_minutes` + /// 方法应该正确返回 true;对于不会在指定时间内过期的 Token,应该返回 false。 + /// + /// **Validates: Requirements 4.4** + #[test] + fn property_token_expiration_check( + offset_minutes in expiry_offset_strategy(), + threshold_minutes in threshold_strategy() + ) { + let now = Utc::now(); + let expiry_time = now + chrono::Duration::minutes(offset_minutes); + + let cache_info = CachedTokenInfo { + access_token: Some("test_token".to_string()), + refresh_token: None, + expiry_time: Some(expiry_time), + last_refresh: None, + refresh_error_count: 0, + last_refresh_error: None, + }; + + let is_expiring = cache_info.is_expiring_within_minutes(threshold_minutes); + + // Token 应该在 offset_minutes <= threshold_minutes 时被认为即将过期 + // 注意:由于时间精度问题,我们允许 1 秒的误差 + if offset_minutes <= threshold_minutes { + prop_assert!( + is_expiring, + "Token with {}min until expiry should be considered expiring within {}min", + offset_minutes, + threshold_minutes + ); + } else { + prop_assert!( + !is_expiring, + "Token with {}min until expiry should NOT be considered expiring within {}min", + offset_minutes, + threshold_minutes + ); + } + } + + /// **Feature: kiro-streaming-fix, Property 7.1: 无过期时间的 Token 不会被认为即将过期** + /// + /// *对于任意* 没有过期时间的 Token,`is_expiring_within_minutes` 应该返回 false。 + /// + /// **Validates: Requirements 4.4** + #[test] + fn property_no_expiry_time_not_expiring(threshold_minutes in threshold_strategy()) { + let cache_info = CachedTokenInfo { + access_token: Some("test_token".to_string()), + refresh_token: None, + expiry_time: None, // 没有过期时间 + last_refresh: None, + refresh_error_count: 0, + last_refresh_error: None, + }; + + let is_expiring = cache_info.is_expiring_within_minutes(threshold_minutes); + + prop_assert!( + !is_expiring, + "Token without expiry time should NOT be considered expiring within {}min", + threshold_minutes + ); + } + + /// **Feature: kiro-streaming-fix, Property 7.2: is_expiring_soon 等价于 is_expiring_within_minutes(5)** + /// + /// *对于任意* Token,`is_expiring_soon()` 应该等价于 `is_expiring_within_minutes(5)`。 + /// + /// **Validates: Requirements 4.4** + #[test] + fn property_expiring_soon_equivalence(offset_minutes in expiry_offset_strategy()) { + let now = Utc::now(); + let expiry_time = now + chrono::Duration::minutes(offset_minutes); + + let cache_info = CachedTokenInfo { + access_token: Some("test_token".to_string()), + refresh_token: None, + expiry_time: Some(expiry_time), + last_refresh: None, + refresh_error_count: 0, + last_refresh_error: None, + }; + + let is_expiring_soon = cache_info.is_expiring_soon(); + let is_expiring_within_5 = cache_info.is_expiring_within_minutes(5); + + prop_assert_eq!( + is_expiring_soon, + is_expiring_within_5, + "is_expiring_soon() should be equivalent to is_expiring_within_minutes(5)" + ); + } + + /// **Feature: kiro-streaming-fix, Property 7.3: 10分钟阈值检查** + /// + /// *对于任意* 在 10 分钟内过期的 Token,`is_expiring_within_minutes(10)` 应该返回 true。 + /// 这是流式请求前的预检查阈值。 + /// + /// **Validates: Requirements 4.4** + #[test] + fn property_streaming_threshold_check(offset_minutes in 0i64..=10i64) { + let now = Utc::now(); + let expiry_time = now + chrono::Duration::minutes(offset_minutes); + + let cache_info = CachedTokenInfo { + access_token: Some("test_token".to_string()), + refresh_token: None, + expiry_time: Some(expiry_time), + last_refresh: None, + refresh_error_count: 0, + last_refresh_error: None, + }; + + let is_expiring = cache_info.is_expiring_within_minutes(10); + + prop_assert!( + is_expiring, + "Token expiring in {}min should be considered expiring within 10min (streaming threshold)", + offset_minutes + ); + } + } +} diff --git a/src-tauri/crates/core/src/models/provider_type.rs b/src-tauri/crates/core/src/models/provider_type.rs new file mode 100644 index 000000000..085cf0894 --- /dev/null +++ b/src-tauri/crates/core/src/models/provider_type.rs @@ -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 { + 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::().unwrap(), ProviderType::Kiro); + assert_eq!( + "gemini".parse::().unwrap(), + ProviderType::Gemini + ); + assert_eq!( + "openai".parse::().unwrap(), + ProviderType::OpenAI + ); + assert_eq!( + "claude".parse::().unwrap(), + ProviderType::Claude + ); + assert_eq!( + "vertex".parse::().unwrap(), + ProviderType::Vertex + ); + assert_eq!( + "gemini_api_key".parse::().unwrap(), + ProviderType::GeminiApiKey + ); + assert_eq!("KIRO".parse::().unwrap(), ProviderType::Kiro); + assert_eq!( + "Gemini".parse::().unwrap(), + ProviderType::Gemini + ); + assert_eq!( + "VERTEX".parse::().unwrap(), + ProviderType::Vertex + ); + assert!("invalid".parse::().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::("\"kiro\"").unwrap(), + ProviderType::Kiro + ); + assert_eq!( + serde_json::from_str::("\"openai\"").unwrap(), + ProviderType::OpenAI + ); + } +} diff --git a/src-tauri/crates/core/src/models/route_model.rs b/src-tauri/crates/core/src/models/route_model.rs new file mode 100644 index 000000000..1c3a287ae --- /dev/null +++ b/src-tauri/crates/core/src/models/route_model.rs @@ -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, + pub tags: Vec, + 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, +} + +/// 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 { + 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 + } +} diff --git a/src-tauri/crates/core/src/models/skill_model.rs b/src-tauri/crates/core/src/models/skill_model.rs new file mode 100644 index 000000000..8b9e862b4 --- /dev/null +++ b/src-tauri/crates/core/src/models/skill_model.rs @@ -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, + pub installed: bool, + #[serde(rename = "repoOwner", skip_serializing_if = "Option::is_none")] + pub repo_owner: Option, + #[serde(rename = "repoName", skip_serializing_if = "Option::is_none")] + pub repo_name: Option, + #[serde(rename = "repoBranch", skip_serializing_if = "Option::is_none")] + pub repo_branch: Option, +} + +#[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, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct SkillMetadata { + pub name: Option, + pub description: Option, +} + +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 { + 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; + +#[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); + } +} diff --git a/src-tauri/crates/infra/Cargo.toml b/src-tauri/crates/infra/Cargo.toml new file mode 100644 index 000000000..6488b0fdc --- /dev/null +++ b/src-tauri/crates/infra/Cargo.toml @@ -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 diff --git a/src-tauri/src/injection/mod.rs b/src-tauri/crates/infra/src/injection/mod.rs similarity index 100% rename from src-tauri/src/injection/mod.rs rename to src-tauri/crates/infra/src/injection/mod.rs diff --git a/src-tauri/src/injection/tests.rs b/src-tauri/crates/infra/src/injection/tests.rs similarity index 100% rename from src-tauri/src/injection/tests.rs rename to src-tauri/crates/infra/src/injection/tests.rs diff --git a/src-tauri/src/injection/types.rs b/src-tauri/crates/infra/src/injection/types.rs similarity index 100% rename from src-tauri/src/injection/types.rs rename to src-tauri/crates/infra/src/injection/types.rs diff --git a/src-tauri/crates/infra/src/lib.rs b/src-tauri/crates/infra/src/lib.rs new file mode 100644 index 000000000..02adac8f7 --- /dev/null +++ b/src-tauri/crates/infra/src/lib.rs @@ -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") +} diff --git a/src-tauri/src/proxy/client_factory.rs b/src-tauri/crates/infra/src/proxy/client_factory.rs similarity index 100% rename from src-tauri/src/proxy/client_factory.rs rename to src-tauri/crates/infra/src/proxy/client_factory.rs diff --git a/src-tauri/src/proxy/mod.rs b/src-tauri/crates/infra/src/proxy/mod.rs similarity index 100% rename from src-tauri/src/proxy/mod.rs rename to src-tauri/crates/infra/src/proxy/mod.rs diff --git a/src-tauri/src/proxy/tests.rs b/src-tauri/crates/infra/src/proxy/tests.rs similarity index 100% rename from src-tauri/src/proxy/tests.rs rename to src-tauri/crates/infra/src/proxy/tests.rs diff --git a/src-tauri/src/resilience/failover.rs b/src-tauri/crates/infra/src/resilience/failover.rs similarity index 99% rename from src-tauri/src/resilience/failover.rs rename to src-tauri/crates/infra/src/resilience/failover.rs index 2651cf83b..5a84bfd27 100644 --- a/src-tauri/src/resilience/failover.rs +++ b/src-tauri/crates/infra/src/resilience/failover.rs @@ -2,7 +2,7 @@ //! //! 提供 Provider 故障转移和自动切换功能 -use crate::ProviderType; +use proxycast_core::ProviderType; use serde::{Deserialize, Serialize}; use std::collections::HashSet; diff --git a/src-tauri/src/resilience/mod.rs b/src-tauri/crates/infra/src/resilience/mod.rs similarity index 100% rename from src-tauri/src/resilience/mod.rs rename to src-tauri/crates/infra/src/resilience/mod.rs diff --git a/src-tauri/src/resilience/retry.rs b/src-tauri/crates/infra/src/resilience/retry.rs similarity index 100% rename from src-tauri/src/resilience/retry.rs rename to src-tauri/crates/infra/src/resilience/retry.rs diff --git a/src-tauri/src/resilience/tests.rs b/src-tauri/crates/infra/src/resilience/tests.rs similarity index 100% rename from src-tauri/src/resilience/tests.rs rename to src-tauri/crates/infra/src/resilience/tests.rs diff --git a/src-tauri/src/resilience/timeout.rs b/src-tauri/crates/infra/src/resilience/timeout.rs similarity index 100% rename from src-tauri/src/resilience/timeout.rs rename to src-tauri/crates/infra/src/resilience/timeout.rs diff --git a/src-tauri/src/telemetry/logger.rs b/src-tauri/crates/infra/src/telemetry/logger.rs similarity index 99% rename from src-tauri/src/telemetry/logger.rs rename to src-tauri/crates/infra/src/telemetry/logger.rs index 8284b189d..1815c77d1 100644 --- a/src-tauri/src/telemetry/logger.rs +++ b/src-tauri/crates/infra/src/telemetry/logger.rs @@ -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}; diff --git a/src-tauri/src/telemetry/mod.rs b/src-tauri/crates/infra/src/telemetry/mod.rs similarity index 100% rename from src-tauri/src/telemetry/mod.rs rename to src-tauri/crates/infra/src/telemetry/mod.rs diff --git a/src-tauri/src/telemetry/stats.rs b/src-tauri/crates/infra/src/telemetry/stats.rs similarity index 98% rename from src-tauri/src/telemetry/stats.rs rename to src-tauri/crates/infra/src/telemetry/stats.rs index a39241d25..f65a73394 100644 --- a/src-tauri/src/telemetry/stats.rs +++ b/src-tauri/crates/infra/src/telemetry/stats.rs @@ -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}; /// 统计聚合器 diff --git a/src-tauri/src/telemetry/tests.rs b/src-tauri/crates/infra/src/telemetry/tests.rs similarity index 99% rename from src-tauri/src/telemetry/tests.rs rename to src-tauri/crates/infra/src/telemetry/tests.rs index 38ecaa900..fcf340659 100644 --- a/src-tauri/src/telemetry/tests.rs +++ b/src-tauri/crates/infra/src/telemetry/tests.rs @@ -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 diff --git a/src-tauri/src/telemetry/tokens.rs b/src-tauri/crates/infra/src/telemetry/tokens.rs similarity index 99% rename from src-tauri/src/telemetry/tokens.rs rename to src-tauri/crates/infra/src/telemetry/tokens.rs index b0de7477d..b6220eeb3 100644 --- a/src-tauri/src/telemetry/tokens.rs +++ b/src-tauri/crates/infra/src/telemetry/tokens.rs @@ -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}; diff --git a/src-tauri/src/telemetry/types.rs b/src-tauri/crates/infra/src/telemetry/types.rs similarity index 99% rename from src-tauri/src/telemetry/types.rs rename to src-tauri/crates/infra/src/telemetry/types.rs index 2f2e2b7e0..f5ef4a715 100644 --- a/src-tauri/src/telemetry/types.rs +++ b/src-tauri/crates/infra/src/telemetry/types.rs @@ -2,8 +2,8 @@ //! //! 定义请求日志、统计数据等核心类型 -use crate::ProviderType; use chrono::{DateTime, Utc}; +use proxycast_core::ProviderType; use serde::{Deserialize, Serialize}; /// 请求状态 diff --git a/src-tauri/crates/providers/Cargo.toml b/src-tauri/crates/providers/Cargo.toml deleted file mode 100644 index 80390d44f..000000000 --- a/src-tauri/crates/providers/Cargo.toml +++ /dev/null @@ -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"] } diff --git a/src-tauri/crates/providers/src/lib.rs b/src-tauri/crates/providers/src/lib.rs deleted file mode 100644 index c2d83b196..000000000 --- a/src-tauri/crates/providers/src/lib.rs +++ /dev/null @@ -1,7 +0,0 @@ -//! Provider 系统模块 -//! -//! 包含 providers, credential, converter 等功能 - -pub fn version() -> &'static str { - env!("CARGO_PKG_VERSION") -} diff --git a/src-tauri/crates/server/Cargo.toml b/src-tauri/crates/server/Cargo.toml deleted file mode 100644 index 949ee8b40..000000000 --- a/src-tauri/crates/server/Cargo.toml +++ /dev/null @@ -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"] } diff --git a/src-tauri/crates/server/src/lib.rs b/src-tauri/crates/server/src/lib.rs deleted file mode 100644 index 6071ab71e..000000000 --- a/src-tauri/crates/server/src/lib.rs +++ /dev/null @@ -1,7 +0,0 @@ -//! API 服务器模块 -//! -//! 包含 server, streaming, middleware, router 等功能 - -pub fn version() -> &'static str { - env!("CARGO_PKG_VERSION") -} diff --git a/src-tauri/resources/models/aliases/kiro.json b/src-tauri/resources/models/aliases/kiro.json index 0132f00a8..66dba5fdd 100644 --- a/src-tauri/resources/models/aliases/kiro.json +++ b/src-tauri/resources/models/aliases/kiro.json @@ -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": { diff --git a/src-tauri/src/agent/protocols/openai.rs b/src-tauri/src/agent/protocols/openai.rs index 29f7675d6..3f17e0ca8 100644 --- a/src-tauri/src/agent/protocols/openai.rs +++ b/src-tauri/src/agent/protocols/openai.rs @@ -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); diff --git a/src-tauri/src/app/runner.rs b/src-tauri/src/app/runner.rs index c499baf34..7ae04e998 100644 --- a/src-tauri/src/app/runner.rs +++ b/src-tauri/src/app/runner.rs @@ -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, diff --git a/src-tauri/src/app/types.rs b/src-tauri/src/app/types.rs index 4345f0321..03b14fa70 100644 --- a/src-tauri/src/app/types.rs +++ b/src-tauri/src/app/types.rs @@ -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 { - 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>; diff --git a/src-tauri/src/commands/model_registry_cmd.rs b/src-tauri/src/commands/model_registry_cmd.rs index ed68c2639..7edaccd80 100644 --- a/src-tauri/src/commands/model_registry_cmd.rs +++ b/src-tauri/src/commands/model_registry_cmd.rs @@ -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 { + 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 { + // 获取 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 +} diff --git a/src-tauri/src/database/migration.rs b/src-tauri/src/database/migration.rs index 771b77d71..5c7970c3b 100644 --- a/src-tauri/src/database/migration.rs +++ b/src-tauri/src/database/migration.rs @@ -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 { + // 检查是否已经迁移过 + 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 = 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'", + [], + ); +} diff --git a/src-tauri/src/database/mod.rs b/src-tauri/src/database/mod.rs index 3e9bfa3f5..0bbdc4779 100644 --- a/src-tauri/src/database/mod.rs +++ b/src-tauri/src/database/mod.rs @@ -31,6 +31,23 @@ pub fn init_database() -> Result { 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) => { diff --git a/src-tauri/src/database/system_providers.rs b/src-tauri/src/database/system_providers.rs index 646b89233..5c0dea3be 100644 --- a/src-tauri/src/database/system_providers.rs +++ b/src-tauri/src/database/system_providers.rs @@ -44,7 +44,7 @@ pub fn get_system_providers() -> Vec { 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 { 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 { 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 { // 国内 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 { 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 { 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 { 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 { 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 { 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 { 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 { // 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 { 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 { 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 { api_version: None, }, SystemProviderDef { - id: "fireworks", + id: "fireworks-ai", name: "Fireworks", provider_type: ApiProviderType::Openai, api_host: "https://api.fireworks.ai/inference", diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index faa9edd98..f6cab7053 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -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; diff --git a/src-tauri/src/models/model_registry.rs b/src-tauri/src/models/model_registry.rs index 7ad37f36b..db4d1ffb9 100644 --- a/src-tauri/src/models/model_registry.rs +++ b/src-tauri/src/models/model_registry.rs @@ -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)), } } diff --git a/src-tauri/src/services/model_registry_service.rs b/src-tauri/src/services/model_registry_service.rs index 78cdf9235..d27b3a0ed 100644 --- a/src-tauri/src/services/model_registry_service.rs +++ b/src-tauri/src/services/model_registry_service.rs @@ -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::(&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 { 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 { + 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 = 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, 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, +} + +/// 单个模型的 API 响应 +#[derive(Debug, Deserialize)] +struct ApiModelResponse { + id: String, + #[serde(default)] + owned_by: Option, + #[serde(default)] + context_length: Option, +} + +/// 模型获取来源 +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)] +pub enum ModelFetchSource { + /// 从 API 获取 + Api, + /// 从本地文件回退 + LocalFallback, +} + +/// 从 API 获取模型的结果 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct FetchModelsResult { + /// 模型列表 + pub models: Vec, + /// 数据来源 + pub source: ModelFetchSource, + /// 错误信息(如果有) + pub error: Option, } diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index e20922e0b..7c83b7edb 100644 --- a/src-tauri/tauri.conf.json +++ b/src-tauri/tauri.conf.json @@ -1,7 +1,7 @@ { "$schema": "https://schema.tauri.app/config/2", "productName": "ProxyCast", - "version": "0.47.3", + "version": "0.47.4", "identifier": "com.proxycast.app", "build": { "beforeDevCommand": "npm run dev", diff --git a/src/components/api-server/ApiServerPage.tsx b/src/components/api-server/ApiServerPage.tsx index b852bc55c..42ba672f7 100644 --- a/src/components/api-server/ApiServerPage.tsx +++ b/src/components/api-server/ApiServerPage.tsx @@ -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( + "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 && diff --git a/src/components/provider-pool/api-key/ProviderModelList.tsx b/src/components/provider-pool/api-key/ProviderModelList.tsx index 3138bb69a..aba34180c 100644 --- a/src/components/provider-pool/api-key/ProviderModelList.tsx +++ b/src/components/provider-pool/api-key/ProviderModelList.tsx @@ -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 = ({ model }) => { export const ProviderModelList: React.FC = ({ providerId, providerType, + hasApiKey = false, className, maxItems, }) => { @@ -118,18 +148,60 @@ export const ProviderModelList: React.FC = ({ providerFilter: [registryProviderId], }); + // 从 API 刷新状态 + const [refreshing, setRefreshing] = useState(false); + const [apiModels, setApiModels] = useState( + null, + ); + const [apiSource, setApiSource] = useState<"Api" | "LocalFallback" | null>( + null, + ); + const [apiError, setApiError] = useState(null); + + // 从 API 获取模型列表(自动获取 API Key) + const handleRefreshFromApi = useCallback(async () => { + setRefreshing(true); + setApiError(null); + + try { + const result = await invoke( + "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 (
= ({ } // 错误状态 - if (error) { + if (error && !apiModels) { return (
= ({ } // 空状态 - if (models.length === 0) { + if (displayModelsSource.length === 0) { return ( -
+
+

+ + 支持的模型 +

+ {hasApiKey && ( + + + + + + 从 API 获取模型列表 + + + )} +
+
+ 暂无模型数据 + {hasApiKey && ( + + )} +
+ {apiError && ( +
{apiError}
)} - data-testid="provider-model-list-empty" - > - 暂无模型数据
); } @@ -182,11 +295,73 @@ export const ProviderModelList: React.FC = ({ 支持的模型 - ({models.length}) + ({displayModelsSource.length}) + {/* 数据来源标识 */} + {apiSource && ( + + + + + {apiSource === "Api" ? ( + <> + + API + + ) : ( + <> + + 本地 + + )} + + + + {apiSource === "Api" + ? "数据来自 Provider API" + : "API 获取失败,使用本地数据"} + + + + )} + {/* 刷新按钮 */} + {hasApiKey && ( + + + + + + 从 API 获取最新模型列表 + + + )}
+ {/* API 错误提示 */} + {apiError && ( +
{apiError}
+ )} + {/* 模型列表 */}
{displayModels.map((model) => ( @@ -197,7 +372,7 @@ export const ProviderModelList: React.FC = ({ {/* 显示更多提示 */} {hasMore && (

- 还有 {models.length - maxItems!} 个模型未显示 + 还有 {displayModelsSource.length - maxItems!} 个模型未显示

)}
diff --git a/src/components/provider-pool/api-key/ProviderSetting.tsx b/src/components/provider-pool/api-key/ProviderSetting.tsx index e7eacf7be..3b1f55085 100644 --- a/src/components/provider-pool/api-key/ProviderSetting.tsx +++ b/src/components/provider-pool/api-key/ProviderSetting.tsx @@ -242,6 +242,9 @@ export const ProviderSetting: React.FC = ({ k.enabled).length ?? 0) > 0 + } />
diff --git a/src/components/provider-pool/api-key/providerTypeMapping.ts b/src/components/provider-pool/api-key/providerTypeMapping.ts index 2a9df273d..f9f6724bc 100644 --- a/src/components/provider-pool/api-key/providerTypeMapping.ts +++ b/src/components/provider-pool/api-key/providerTypeMapping.ts @@ -20,38 +20,60 @@ const PROVIDER_ID_TO_REGISTRY_ID: Record = { // 主流 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", }; /** diff --git a/src/components/terminal/ai/TerminalAIModeSelector.tsx b/src/components/terminal/ai/TerminalAIModeSelector.tsx index ae346b798..8918f88d9 100644 --- a/src/components/terminal/ai/TerminalAIModeSelector.tsx +++ b/src/components/terminal/ai/TerminalAIModeSelector.tsx @@ -196,6 +196,24 @@ export const TerminalAIModeSelector: React.FC = ({ 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, diff --git a/src/hooks/useProviderModels.ts b/src/hooks/useProviderModels.ts index 222453c81..d362985a8 100644 --- a/src/hooks/useProviderModels.ts +++ b/src/hooks/useProviderModels.ts @@ -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([]); + const [apiLoading, setApiLoading] = useState(false); + const [apiError, setApiError] = useState(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( + "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, }; diff --git a/vite.config.ts b/vite.config.ts index 51f774b8b..c1c8bcfcf 100644 --- a/vite.config.ts +++ b/vite.config.ts @@ -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/**", ], }, -})); +}; +});