Merge pull request #22 from zss823158062/main

feat: 实现自动下载更新功能
This commit is contained in:
coso
2025-12-22 20:11:32 +08:00
committed by GitHub
35 changed files with 667 additions and 279 deletions
+472 -106
View File
@@ -5,7 +5,7 @@ use crate::config::{
use crate::models::AppType;
use serde::{Deserialize, Serialize};
use std::path::PathBuf;
use tauri::AppHandle;
use tauri::{AppHandle, Manager};
use tauri_plugin_autostart::ManagerExt;
#[cfg(target_os = "windows")]
@@ -120,117 +120,62 @@ pub struct ToolVersion {
pub installed: bool,
}
/// 检测工具版本的辅助函数
fn check_tool_version(command: &str, args: &[&str]) -> Option<String> {
// 在 Windows 上,先尝试直接执行命令
let mut cmd = std::process::Command::new(command);
cmd.args(args);
#[cfg(target_os = "windows")]
cmd.creation_flags(0x08000000); // CREATE_NO_WINDOW
let output = cmd.output().ok();
// 如果直接执行失败,在 Windows 上尝试通过 PowerShell 执行
#[cfg(target_os = "windows")]
let output = output.or_else(|| {
std::process::Command::new("powershell")
.args(["-Command", &format!("{} {}", command, args.join(" "))])
.creation_flags(0x08000000)
.output()
.ok()
});
output
.and_then(|o| {
if o.status.success() {
// 先尝试 stdout,失败则尝试 stderr
String::from_utf8(o.stdout.clone())
.or_else(|_| String::from_utf8(o.stderr))
.ok()
} else {
None
}
})
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
}
#[tauri::command]
pub async fn get_tool_versions() -> Result<Vec<ToolVersion>, String> {
let mut versions = Vec::new();
// Check Claude Code version
#[cfg(target_os = "windows")]
let claude_version = std::process::Command::new("claude")
.arg("--version")
.creation_flags(0x08000000) // CREATE_NO_WINDOW
.output()
.ok()
.and_then(|o| {
if o.status.success() {
String::from_utf8(o.stdout).ok()
} else {
None
}
})
.map(|s| s.trim().to_string());
// 定义要检测的工具列表
let tools = vec![
("Claude Code", "claude", vec!["--version"]),
("Codex", "codex", vec!["--version"]),
("Gemini CLI", "gemini", vec!["--version"]),
];
#[cfg(not(target_os = "windows"))]
let claude_version = std::process::Command::new("claude")
.arg("--version")
.output()
.ok()
.and_then(|o| {
if o.status.success() {
String::from_utf8(o.stdout).ok()
} else {
None
}
})
.map(|s| s.trim().to_string());
for (name, command, args) in tools {
let version = check_tool_version(command, &args);
versions.push(ToolVersion {
name: "Claude Code".to_string(),
version: claude_version.clone(),
installed: claude_version.is_some(),
});
// Check Codex version
#[cfg(target_os = "windows")]
let codex_version = std::process::Command::new("codex")
.arg("--version")
.creation_flags(0x08000000) // CREATE_NO_WINDOW
.output()
.ok()
.and_then(|o| {
if o.status.success() {
String::from_utf8(o.stdout).ok()
} else {
None
}
})
.map(|s| s.trim().to_string());
#[cfg(not(target_os = "windows"))]
let codex_version = std::process::Command::new("codex")
.arg("--version")
.output()
.ok()
.and_then(|o| {
if o.status.success() {
String::from_utf8(o.stdout).ok()
} else {
None
}
})
.map(|s| s.trim().to_string());
versions.push(ToolVersion {
name: "Codex".to_string(),
version: codex_version.clone(),
installed: codex_version.is_some(),
});
// Check Gemini CLI version
#[cfg(target_os = "windows")]
let gemini_version = std::process::Command::new("gemini")
.arg("--version")
.creation_flags(0x08000000) // CREATE_NO_WINDOW
.output()
.ok()
.and_then(|o| {
if o.status.success() {
String::from_utf8(o.stdout).ok()
} else {
None
}
})
.map(|s| s.trim().to_string());
#[cfg(not(target_os = "windows"))]
let gemini_version = std::process::Command::new("gemini")
.arg("--version")
.output()
.ok()
.and_then(|o| {
if o.status.success() {
String::from_utf8(o.stdout).ok()
} else {
None
}
})
.map(|s| s.trim().to_string());
versions.push(ToolVersion {
name: "Gemini CLI".to_string(),
version: gemini_version.clone(),
installed: gemini_version.is_some(),
});
versions.push(ToolVersion {
name: name.to_string(),
version: version.clone(),
installed: version.is_some(),
});
}
Ok(versions)
}
@@ -740,3 +685,424 @@ fn version_compare(current: &str, latest: &str) -> bool {
false
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_version_compare() {
// 测试版本比较逻辑
assert!(version_compare("0.14.0", "0.14.1"));
assert!(version_compare("0.14.0", "0.15.0"));
assert!(version_compare("0.14.0", "1.0.0"));
assert!(!version_compare("0.14.1", "0.14.0"));
assert!(!version_compare("0.14.0", "0.14.0"));
assert!(!version_compare("1.0.0", "0.14.0"));
}
#[test]
fn test_get_platform_patterns() {
let patterns = get_platform_patterns();
// 在支持的平台上应该返回非空的模式列表
#[cfg(any(
all(
target_os = "windows",
any(target_arch = "x86_64", target_arch = "aarch64")
),
all(
target_os = "macos",
any(target_arch = "x86_64", target_arch = "aarch64")
),
all(
target_os = "linux",
any(target_arch = "x86_64", target_arch = "aarch64")
)
))]
{
assert!(!patterns.is_empty());
}
// 在不支持的平台上应该返回空列表
#[cfg(not(any(
all(
target_os = "windows",
any(target_arch = "x86_64", target_arch = "aarch64")
),
all(
target_os = "macos",
any(target_arch = "x86_64", target_arch = "aarch64")
),
all(
target_os = "linux",
any(target_arch = "x86_64", target_arch = "aarch64")
)
)))]
{
assert!(patterns.is_empty());
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DownloadResult {
pub success: bool,
pub message: String,
#[serde(rename = "filePath")]
pub file_path: Option<String>,
}
/// 下载更新安装包
///
/// 从 GitHub Releases 下载对应平台的安装包到下载目录
#[tauri::command]
pub async fn download_update(app_handle: AppHandle) -> Result<DownloadResult, String> {
// 首先检查是否有更新
let version_info = check_for_updates().await?;
if !version_info.has_update {
return Ok(DownloadResult {
success: false,
message: "当前已是最新版本".to_string(),
file_path: None,
});
}
let latest_version = version_info.latest.ok_or("无法获取最新版本信息")?;
// 从 GitHub API 获取实际的文件列表并匹配平台
let (filename, download_url) = get_platform_download_from_github(&latest_version).await?;
// 获取下载目录
let download_dir = get_download_directory(&app_handle)?;
let file_path = download_dir.join(&filename);
// 如果文件已存在,先删除
if file_path.exists() {
if let Err(e) = std::fs::remove_file(&file_path) {
tracing::warn!("删除旧文件失败: {}", e);
}
}
// 下载文件
let client = reqwest::Client::new();
match client
.get(&download_url)
.header("User-Agent", "ProxyCast")
.send()
.await
{
Ok(response) => {
if !response.status().is_success() {
return Ok(DownloadResult {
success: false,
message: format!("下载失败: HTTP {}", response.status()),
file_path: None,
});
}
// 获取文件内容
match response.bytes().await {
Ok(bytes) => {
// 写入文件
match std::fs::write(&file_path, bytes) {
Ok(_) => {
tracing::info!("安装包下载成功: {:?}", file_path);
// 尝试直接运行安装程序
match run_installer(&file_path) {
Ok(_) => {
tracing::info!("已启动安装程序,准备退出当前应用");
// 延迟退出,给安装程序时间启动
tokio::spawn(async {
tokio::time::sleep(tokio::time::Duration::from_secs(2))
.await;
tracing::info!("自动退出应用以便安装程序运行");
std::process::exit(0);
});
}
Err(e) => {
tracing::warn!("启动安装程序失败: {},尝试打开文件位置", e);
// 如果无法运行安装程序,则打开文件所在目录
if let Err(open_err) = open_file_location(&file_path) {
tracing::warn!("打开文件所在目录也失败: {}", open_err);
}
}
}
Ok(DownloadResult {
success: true,
message: format!("下载完成: {}", filename),
file_path: Some(file_path.to_string_lossy().to_string()),
})
}
Err(e) => Ok(DownloadResult {
success: false,
message: format!("保存文件失败: {}", e),
file_path: None,
}),
}
}
Err(e) => Ok(DownloadResult {
success: false,
message: format!("读取下载内容失败: {}", e),
file_path: None,
}),
}
}
Err(e) => Ok(DownloadResult {
success: false,
message: format!("网络请求失败: {}", e),
file_path: None,
}),
}
}
/// 从 GitHub API 获取实际的文件列表并匹配平台
async fn get_platform_download_from_github(version: &str) -> Result<(String, String), String> {
let api_url = format!(
"https://api.github.com/repos/aiclientproxy/proxycast/releases/tags/v{}",
version
);
let client = reqwest::Client::new();
let response = client
.get(&api_url)
.header("User-Agent", "ProxyCast")
.send()
.await
.map_err(|e| format!("请求 GitHub API 失败: {}", e))?;
if !response.status().is_success() {
return Err(format!("GitHub API 请求失败: {}", response.status()));
}
let data: serde_json::Value = response
.json()
.await
.map_err(|e| format!("解析 GitHub API 响应失败: {}", e))?;
let assets = data["assets"]
.as_array()
.ok_or("GitHub API 响应中没有找到 assets")?;
// 根据当前平台匹配文件
let platform_patterns = get_platform_patterns();
for asset in assets {
let name = asset["name"].as_str().unwrap_or("");
let download_url = asset["browser_download_url"].as_str().unwrap_or("");
for pattern in &platform_patterns {
if name.contains(pattern) {
return Ok((name.to_string(), download_url.to_string()));
}
}
}
Err("未找到适合当前平台的安装包".to_string())
}
/// 获取当前平台的文件名匹配模式
fn get_platform_patterns() -> Vec<&'static str> {
#[cfg(all(target_os = "windows", target_arch = "x86_64"))]
{
vec!["x64-setup.exe", "x64_en-US.msi"]
}
#[cfg(all(target_os = "windows", target_arch = "aarch64"))]
{
vec!["arm64-setup.exe", "arm64_en-US.msi"]
}
#[cfg(all(target_os = "macos", target_arch = "x86_64"))]
{
vec!["x64.dmg"]
}
#[cfg(all(target_os = "macos", target_arch = "aarch64"))]
{
vec!["aarch64.dmg"]
}
#[cfg(all(target_os = "linux", target_arch = "x86_64"))]
{
vec!["amd64.deb", "amd64.AppImage"]
}
#[cfg(all(target_os = "linux", target_arch = "aarch64"))]
{
vec!["arm64.deb", "arm64.AppImage"]
}
#[cfg(not(any(
all(
target_os = "windows",
any(target_arch = "x86_64", target_arch = "aarch64")
),
all(
target_os = "macos",
any(target_arch = "x86_64", target_arch = "aarch64")
),
all(
target_os = "linux",
any(target_arch = "x86_64", target_arch = "aarch64")
)
)))]
{
vec![]
}
}
/// 获取下载目录
fn get_download_directory(app_handle: &AppHandle) -> Result<PathBuf, String> {
// 优先使用系统下载目录
if let Some(download_dir) = dirs::download_dir() {
return Ok(download_dir);
}
// 回退到应用数据目录
let app_data_dir = app_handle
.path()
.app_data_dir()
.map_err(|e| format!("无法获取应用数据目录: {}", e))?;
let download_dir = app_data_dir.join("downloads");
// 确保目录存在
std::fs::create_dir_all(&download_dir).map_err(|e| format!("创建下载目录失败: {}", e))?;
Ok(download_dir)
}
/// 运行安装程序
fn run_installer(file_path: &PathBuf) -> Result<(), String> {
let extension = file_path
.extension()
.and_then(|ext| ext.to_str())
.unwrap_or("");
match extension.to_lowercase().as_str() {
"exe" | "msi" => {
#[cfg(target_os = "windows")]
{
tracing::info!("Windows: 启动安装程序: {:?}", file_path);
std::process::Command::new(file_path)
.spawn()
.map_err(|e| format!("启动 Windows 安装程序失败: {}", e))?;
}
#[cfg(not(target_os = "windows"))]
{
return Err("Windows 安装程序只能在 Windows 系统上运行".to_string());
}
}
"dmg" => {
#[cfg(target_os = "macos")]
{
tracing::info!("macOS: 打开 DMG 文件: {:?}", file_path);
std::process::Command::new("open")
.arg(&file_path)
.spawn()
.map_err(|e| format!("打开 macOS DMG 文件失败: {}", e))?;
}
#[cfg(not(target_os = "macos"))]
{
return Err("DMG 文件只能在 macOS 系统上打开".to_string());
}
}
"deb" => {
#[cfg(target_os = "linux")]
{
tracing::info!("Linux: 尝试安装 DEB 包: {:?}", file_path);
// 尝试使用系统默认的包管理器打开
let result = std::process::Command::new("xdg-open")
.arg(&file_path)
.spawn();
if result.is_err() {
// 如果 xdg-open 失败,尝试使用 dpkg
tracing::info!("xdg-open 失败,尝试使用 gdebi 或提示用户手动安装");
return Err("请手动安装 DEB 包,或使用: sudo dpkg -i filename.deb".to_string());
}
}
#[cfg(not(target_os = "linux"))]
{
return Err("DEB 包只能在 Linux 系统上安装".to_string());
}
}
"appimage" => {
#[cfg(target_os = "linux")]
{
tracing::info!("Linux: 设置 AppImage 可执行权限并运行: {:?}", file_path);
// 设置可执行权限
std::process::Command::new("chmod")
.args(&["+x", &file_path.to_string_lossy()])
.output()
.map_err(|e| format!("设置 AppImage 可执行权限失败: {}", e))?;
// 运行 AppImage
std::process::Command::new(&file_path)
.spawn()
.map_err(|e| format!("运行 AppImage 失败: {}", e))?;
}
#[cfg(not(target_os = "linux"))]
{
return Err("AppImage 只能在 Linux 系统上运行".to_string());
}
}
_ => {
return Err(format!("不支持的文件类型: {}", extension));
}
}
Ok(())
}
/// 打开文件所在位置
fn open_file_location(file_path: &PathBuf) -> Result<(), String> {
#[cfg(target_os = "windows")]
{
tracing::info!("Windows: 使用 explorer 打开文件位置: {:?}", file_path);
std::process::Command::new("explorer")
.args(["/select,", &file_path.to_string_lossy()])
.creation_flags(0x08000000) // CREATE_NO_WINDOW
.spawn()
.map_err(|e| format!("Windows explorer 启动失败: {}", e))?;
}
#[cfg(target_os = "macos")]
{
tracing::info!("macOS: 使用 open -R 打开文件位置: {:?}", file_path);
std::process::Command::new("open")
.args(&["-R", &file_path.to_string_lossy()])
.spawn()
.map_err(|e| format!("macOS open 命令失败: {}", e))?;
}
#[cfg(target_os = "linux")]
{
if let Some(parent) = file_path.parent() {
tracing::info!("Linux: 使用 xdg-open 打开目录: {:?}", parent);
std::process::Command::new("xdg-open")
.arg(parent)
.spawn()
.map_err(|e| format!("Linux xdg-open 命令失败: {}", e))?;
} else {
return Err("无法获取文件的父目录".to_string());
}
}
#[cfg(not(any(target_os = "windows", target_os = "macos", target_os = "linux")))]
{
return Err("不支持的操作系统".to_string());
}
Ok(())
}
+9 -10
View File
@@ -22,9 +22,9 @@ pub struct CredentialSyncServiceState(pub Option<Arc<CredentialSyncService>>);
/// 展开路径中的 ~ 为用户主目录
fn expand_tilde(path: &str) -> String {
if path.starts_with("~/") {
if let Some(stripped) = path.strip_prefix("~/") {
if let Some(home) = dirs::home_dir() {
return home.join(&path[2..]).to_string_lossy().to_string();
return home.join(stripped).to_string_lossy().to_string();
}
}
path.to_string()
@@ -82,8 +82,7 @@ fn copy_and_rename_credential_file(
// 对于 Kiro 凭证,需要合并 clientIdHash 文件中的 client_id/client_secret
if provider_type == "kiro" {
let content =
fs::read_to_string(&source).map_err(|e| format!("读取凭证文件失败: {}", e))?;
let content = fs::read_to_string(source).map_err(|e| format!("读取凭证文件失败: {}", e))?;
let mut creds: serde_json::Value =
serde_json::from_str(&content).map_err(|e| format!("解析凭证文件失败: {}", e))?;
@@ -200,9 +199,9 @@ fn copy_and_rename_credential_file(
tracing::error!(
"[KIRO] IdC 认证方式缺少 clientId/clientSecret,无法创建有效的凭证副本"
);
return Err(format!(
"IdC 认证凭证不完整:缺少 clientId/clientSecret。\n\n💡 解决方案:\n1. 确保 ~/.aws/sso/cache/ 目录下有对应的 clientIdHash 文件\n2. 如果使用 AWS IAM Identity Center,请确保已完成完整的 SSO 登录流程\n3. 或者尝试使用 Social 认证方式的凭证"
));
return Err(
"IdC 认证凭证不完整:缺少 clientId/clientSecret。\n\n💡 解决方案:\n1. 确保 ~/.aws/sso/cache/ 目录下有对应的 clientIdHash 文件\n2. 如果使用 AWS IAM Identity Center,请确保已完成完整的 SSO 登录流程\n3. 或者尝试使用 Social 认证方式的凭证".to_string()
);
} else {
tracing::warn!("[KIRO] 未找到 client_id/client_secret,将使用 social 认证方式");
}
@@ -214,7 +213,7 @@ fn copy_and_rename_credential_file(
fs::write(&target_path, merged_content).map_err(|e| format!("写入凭证文件失败: {}", e))?;
} else {
// 其他类型直接复制
fs::copy(&source, &target_path).map_err(|e| format!("复制凭证文件失败: {}", e))?;
fs::copy(source, &target_path).map_err(|e| format!("复制凭证文件失败: {}", e))?;
}
// 返回新的文件路径
@@ -1117,7 +1116,7 @@ pub async fn get_antigravity_auth_url_and_wait(
);
// 从凭证中获取 project_id
let project_id = result.credentials.projectId.clone();
let project_id = result.credentials.project_id.clone();
// 添加到凭证池
let credential = pool_service.0.add_credential(
@@ -1165,7 +1164,7 @@ pub async fn start_antigravity_oauth_login(
);
// 从凭证中获取 project_id
let project_id = result.credentials.projectId.clone();
let project_id = result.credentials.project_id.clone();
// 添加到凭证池
let credential = pool_service.0.add_credential(
+1 -8
View File
@@ -65,14 +65,7 @@ pub async fn get_route_curl_examples(
.map_err(|e| e.to_string())?;
// 查找匹配的路由
let route = routes.iter().find(|r| r.selector == selector).or_else(|| {
// 如果是默认路由
if selector == "default" {
None // 返回 None 让下面的代码生成默认示例
} else {
None
}
});
let route = routes.iter().find(|r| r.selector == selector);
// P0 安全修复:curl 示例使用占位符,不暴露真实 API Key
let api_key = "${PROXYCAST_API_KEY}";
+2 -2
View File
@@ -116,9 +116,9 @@ fn read_kiro_credential_info(creds_file_path: &str) -> Result<(String, Option<St
/// 展开路径中的 ~ 为用户主目录
fn expand_tilde(path: &str) -> String {
if path.starts_with("~/") {
if let Some(stripped) = path.strip_prefix("~/") {
if let Some(home) = dirs::home_dir() {
return home.join(&path[2..]).to_string_lossy().to_string();
return home.join(stripped).to_string_lossy().to_string();
}
}
path.to_string()
+2 -1
View File
@@ -500,7 +500,8 @@ mod base64 {
}
}
pub use self::base64::{decode as base64_decode, encode as base64_encode};
pub use self::base64::decode as base64_decode;
pub use self::base64::encode as base64_encode;
#[cfg(test)]
mod unit_tests {
+5 -6
View File
@@ -17,12 +17,11 @@ pub use hot_reload::{
pub use import::{ImportOptions, ImportService, ValidationResult};
pub use path_utils::{collapse_tilde, contains_tilde, expand_tilde};
pub use types::{
generate_secure_api_key, is_default_api_key, AmpConfig, AmpModelMapping, ApiKeyEntry, Config,
CredentialEntry, CredentialPoolConfig, CustomProviderConfig, GeminiApiKeyEntry,
IFlowCredentialEntry, InjectionRuleConfig, InjectionSettings, LoggingConfig, ProviderConfig,
ProvidersConfig, QuotaExceededConfig, RemoteManagementConfig, RetrySettings, RoutingConfig,
RoutingRuleConfig, ServerConfig, TlsConfig, VertexApiKeyEntry, VertexModelAlias,
DEFAULT_API_KEY,
AmpConfig, AmpModelMapping, ApiKeyEntry, Config, CredentialEntry, CredentialPoolConfig,
CustomProviderConfig, GeminiApiKeyEntry, IFlowCredentialEntry, InjectionRuleConfig,
InjectionSettings, LoggingConfig, ProviderConfig, ProvidersConfig, QuotaExceededConfig,
RemoteManagementConfig, RetrySettings, RoutingConfig, ServerConfig, TlsConfig,
VertexApiKeyEntry, VertexModelAlias,
};
pub use yaml::{load_config, save_config, ConfigError, ConfigManager, YamlService};
+2 -2
View File
@@ -44,9 +44,9 @@ pub fn expand_tilde<P: AsRef<Path>>(path: P) -> PathBuf {
if path_str == "~" {
// 仅 ~
home_dir
} else if path_str.starts_with("~/") {
} else if let Some(rest) = path_str.strip_prefix("~/") {
// ~/path 格式
let rest = &path_str[2..]; // 跳过 "~/"
// 跳过 "~/"
home_dir.join(rest)
} else {
// ~user/path 格式,不支持,返回原路径
@@ -331,7 +331,7 @@ fn find_function_name(contents: &[GeminiContent], tool_id: &str) -> String {
/// 清理参数中不需要的字段
fn clean_parameters(params: Option<serde_json::Value>) -> Option<serde_json::Value> {
params.map(|v| clean_value(v))
params.map(clean_value)
}
fn clean_value(value: serde_json::Value) -> serde_json::Value {
@@ -489,8 +489,7 @@ pub fn convert_antigravity_to_openai_response(
content.push_str(text);
}
if let Some(fc) = part.get("functionCall") {
let call_id =
format!("call_{}", uuid::Uuid::new_v4().to_string()[..8].to_string());
let call_id = format!("call_{}", &uuid::Uuid::new_v4().to_string()[..8]);
tool_calls.push(serde_json::json!({
"id": call_id,
"type": "function",
+10 -9
View File
@@ -149,10 +149,11 @@ impl ProtocolSelector {
/// 获取推荐的中间协议(用于不支持直接转换的情况)
pub fn intermediate_protocol(source: Protocol, target: Protocol) -> Option<Protocol> {
// 大多数情况下,OpenAI 是最好的中间协议
if !Self::supports_direct_conversion(source, target) {
if source != Protocol::OpenAI && target != Protocol::OpenAI {
return Some(Protocol::OpenAI);
}
if !Self::supports_direct_conversion(source, target)
&& source != Protocol::OpenAI
&& target != Protocol::OpenAI
{
return Some(Protocol::OpenAI);
}
None
}
@@ -172,11 +173,11 @@ impl ProtocolSelector {
) -> bool {
// 工具调用在某些转换中需要特殊处理
if has_tools {
match (source, target_provider) {
(Protocol::OpenAI, PoolProviderType::Kiro) => true,
(Protocol::Anthropic, PoolProviderType::Kiro) => true,
_ => false,
}
matches!(
(source, target_provider),
(Protocol::OpenAI, PoolProviderType::Kiro)
| (Protocol::Anthropic, PoolProviderType::Kiro)
)
} else if has_images {
// 图片在某些 Provider 中需要特殊处理
match target_provider {
+4 -4
View File
@@ -225,10 +225,10 @@ impl HealthChecker {
let mut recovered = Vec::new();
for cred in pool.all() {
if matches!(cred.status, CredentialStatus::Unhealthy { .. }) {
if pool.mark_active(&cred.id).is_ok() {
recovered.push(cred.id.clone());
}
if matches!(cred.status, CredentialStatus::Unhealthy { .. })
&& pool.mark_active(&cred.id).is_ok()
{
recovered.push(cred.id.clone());
}
}
+1 -5
View File
@@ -295,11 +295,7 @@ impl QuotaManager {
}
// 移除 -preview 后缀或 -preview-xxx 部分
if let Some(pos) = model.find("-preview") {
Some(model[..pos].to_string())
} else {
None
}
model.find("-preview").map(|pos| model[..pos].to_string())
}
/// 检查是否启用自动切换项目
+6 -13
View File
@@ -23,13 +23,11 @@ impl ProviderPoolDao {
ORDER BY provider_type, created_at ASC",
)?;
let rows = stmt.query_map([], |row| Self::row_to_credential(row))?;
let rows = stmt.query_map([], Self::row_to_credential)?;
let mut credentials = Vec::new();
for row in rows {
if let Ok(cred) = row {
credentials.push(cred);
}
for cred in rows.flatten() {
credentials.push(cred);
}
Ok(credentials)
}
@@ -54,10 +52,8 @@ impl ProviderPoolDao {
})?;
let mut credentials = Vec::new();
for row in rows {
if let Ok(cred) = row {
credentials.push(cred);
}
for cred in rows.flatten() {
credentials.push(cred);
}
Ok(credentials)
}
@@ -112,10 +108,7 @@ impl ProviderPoolDao {
let mut grouped: ProviderPools = std::collections::HashMap::new();
for cred in all {
grouped
.entry(cred.provider_type)
.or_insert_with(Vec::new)
.push(cred);
grouped.entry(cred.provider_type).or_default().push(cred);
}
Ok(grouped)
+1
View File
@@ -1826,6 +1826,7 @@ pub fn run() {
commands::config_cmd::expand_path,
commands::config_cmd::open_auth_dir,
commands::config_cmd::check_for_updates,
commands::config_cmd::download_update,
// MCP commands
commands::mcp_cmd::get_mcp_servers,
commands::mcp_cmd::add_mcp_server,
+2 -2
View File
@@ -106,8 +106,8 @@ impl<S> ManagementAuthService<S> {
// 支持两种方式:Authorization: Bearer <key> 或 X-Management-Key: <key>
if let Some(auth) = req.headers().get("authorization") {
if let Ok(auth_str) = auth.to_str() {
if auth_str.starts_with("Bearer ") {
return Some(auth_str[7..].to_string());
if let Some(stripped) = auth_str.strip_prefix("Bearer ") {
return Some(stripped.to_string());
}
}
}
+2 -2
View File
@@ -201,7 +201,7 @@ impl ProviderStep {
// 检查状态码是否可重试
let should_retry = err
.status_code
.map_or(true, |code| self.retrier.config().is_retryable(code));
.is_none_or(|code| self.retrier.config().is_retryable(code));
let should_failover = err.should_failover || err.is_quota_exceeded();
@@ -387,7 +387,7 @@ impl ProviderStep {
// 检查状态码是否可重试
let should_retry = err
.status_code
.map_or(true, |code| self.retrier.config().is_retryable(code));
.is_none_or(|code| self.retrier.config().is_retryable(code));
let should_failover = err.should_failover || err.is_quota_exceeded();
+11 -11
View File
@@ -134,8 +134,8 @@ pub struct AntigravityCredentials {
#[serde(skip_serializing_if = "Option::is_none")]
pub enable: Option<bool>,
/// 项目 ID
#[serde(skip_serializing_if = "Option::is_none", alias = "project_id")]
pub projectId: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub project_id: Option<String>,
/// 用户邮箱
#[serde(skip_serializing_if = "Option::is_none")]
pub email: Option<String>,
@@ -159,7 +159,7 @@ impl Default for AntigravityCredentials {
expires_in: None,
timestamp: None,
enable: None,
projectId: None,
project_id: None,
email: None,
}
}
@@ -225,8 +225,8 @@ impl AntigravityProvider {
// 尝试解析为单个凭证对象
if let Ok(creds) = serde_json::from_str::<AntigravityCredentials>(&content) {
self.credentials = creds;
// 如果凭证中有 projectId,设置到 provider
if let Some(ref pid) = self.credentials.projectId {
// 如果凭证中有 project_id,设置到 provider
if let Some(ref pid) = self.credentials.project_id {
self.project_id = Some(pid.clone());
}
return Ok(());
@@ -237,8 +237,8 @@ impl AntigravityProvider {
// 找到第一个启用的凭证
if let Some(creds) = creds_array.into_iter().find(|c| c.enable != Some(false)) {
self.credentials = creds;
// 如果凭证中有 projectId,设置到 provider
if let Some(ref pid) = self.credentials.projectId {
// 如果凭证中有 project_id,设置到 provider
if let Some(ref pid) = self.credentials.project_id {
self.project_id = Some(pid.clone());
}
return Ok(());
@@ -765,7 +765,7 @@ pub async fn fetch_project_id_for_oauth(
tracing::info!("[Antigravity OAuth] cloudaicompanionProject 为空字符串,有资格但无 projectId");
Ok(Some(FetchedProjectId::NoProject)) // 空字符串,有资格但无 projectId
} else {
tracing::info!("[Antigravity OAuth] 获取到 projectId: {}", s);
tracing::info!("[Antigravity OAuth] 获取到 project_id: {}", s);
Ok(Some(FetchedProjectId::HasProject(s.to_string()))) // 有 projectId
}
} else {
@@ -986,7 +986,7 @@ pub async fn start_oauth_server_and_get_url(
expires_in,
timestamp: Some(now.timestamp_millis()),
enable: Some(true),
projectId: project_id,
project_id: project_id,
email: email.clone(),
};
@@ -1207,7 +1207,7 @@ pub async fn start_oauth_login_with_port(
expires_in,
timestamp: Some(now.timestamp_millis()),
enable: Some(true),
projectId: project_id,
project_id: project_id,
email: email.clone(),
};
@@ -1433,7 +1433,7 @@ pub async fn start_oauth_login(
expires_in,
timestamp: Some(now.timestamp_millis()),
enable: Some(true),
projectId: project_id,
project_id: project_id,
email: email.clone(),
};
+1 -2
View File
@@ -680,8 +680,7 @@ pub async fn start_claude_oauth_server_and_get_url() -> Result<
// 等待回调结果
match rx.await {
Ok(result) => result.map_err(|e| {
Box::new(std::io::Error::new(std::io::ErrorKind::Other, e))
as Box<dyn Error + Send + Sync>
Box::new(std::io::Error::other(e)) as Box<dyn Error + Send + Sync>
}),
Err(_) => Err("OAuth 回调通道关闭".into()),
}
+8 -11
View File
@@ -604,7 +604,7 @@ impl CodexProvider {
("codex_cli_simplified_flow", "true"),
];
let query = serde_urlencoded::to_string(&params)?;
let query = serde_urlencoded::to_string(params)?;
Ok(format!("{}?{}", OPENAI_AUTH_URL, query))
}
@@ -1252,11 +1252,9 @@ fn transform_to_codex_format(
} else if let Some(arr) = content.as_array() {
arr.iter()
.filter_map(|part| {
if let Some(text) = part["text"].as_str() {
Some(serde_json::json!({"type": "input_text", "text": text}))
} else {
None
}
part["text"].as_str().map(
|text| serde_json::json!({"type": "input_text", "text": text}),
)
})
.collect()
} else {
@@ -1288,14 +1286,14 @@ fn transform_to_codex_format(
let tools = request["tools"].as_array().map(|tools| {
tools
.iter()
.filter_map(|tool| {
.map(|tool| {
let func = &tool["function"];
Some(serde_json::json!({
serde_json::json!({
"type": "function",
"name": func["name"],
"description": func["description"],
"parameters": func["parameters"]
}))
})
})
.collect::<Vec<_>>()
});
@@ -2112,8 +2110,7 @@ pub async fn start_codex_oauth_server_and_get_url() -> Result<
// 等待回调结果
match rx.await {
Ok(result) => result.map_err(|e| {
Box::new(std::io::Error::new(std::io::ErrorKind::Other, e))
as Box<dyn Error + Send + Sync>
Box::new(std::io::Error::other(e)) as Box<dyn Error + Send + Sync>
}),
Err(_) => Err("OAuth 回调通道关闭".into()),
}
-8
View File
@@ -245,9 +245,6 @@ fn truncate_message(msg: &str, max_len: usize) -> String {
}
}
/// Provider 操作结果类型别名
pub type ProviderResult<T> = Result<T, ProviderError>;
/// 从 HTTP 响应创建用户友好的错误
///
/// 用于 Provider 中的 Token 刷新等操作
@@ -304,11 +301,6 @@ pub fn create_auth_error(message: &str) -> Box<dyn Error + Send + Sync> {
Box::new(ProviderError::AuthenticationError(message.to_string()))
}
/// 创建解析错误
pub fn create_parse_error(message: &str) -> Box<dyn Error + Send + Sync> {
Box::new(ProviderError::ParseError(message.to_string()))
}
#[cfg(test)]
mod tests {
use super::*;
+1 -1
View File
@@ -1080,7 +1080,7 @@ async fn save_gemini_credentials_to_file(
// 获取凭证存储目录
let credentials_dir = dirs::data_dir()
.ok_or_else(|| "无法获取应用数据目录")?
.ok_or("无法获取应用数据目录")?
.join("proxycast")
.join("credentials");
+2 -3
View File
@@ -608,7 +608,7 @@ impl IFlowProvider {
("code_challenge_method", "S256"),
];
let query = serde_urlencoded::to_string(&params)?;
let query = serde_urlencoded::to_string(params)?;
Ok(format!("{}?{}", IFLOW_AUTH_URL, query))
}
@@ -1973,8 +1973,7 @@ pub async fn start_iflow_oauth_server_and_get_url() -> Result<
// 等待回调结果
match rx.await {
Ok(result) => result.map_err(|e| {
Box::new(std::io::Error::new(std::io::ErrorKind::Other, e))
as Box<dyn Error + Send + Sync>
Box::new(std::io::Error::other(e)) as Box<dyn Error + Send + Sync>
}),
Err(_) => Err("OAuth 回调通道关闭".into()),
}
+1 -11
View File
@@ -14,21 +14,13 @@
//! - 需求 5.3: 流中发生错误时发送错误事件并优雅关闭流
use axum::{
body::Body,
extract::State,
http::{header, HeaderMap, StatusCode},
http::{HeaderMap, StatusCode},
response::{IntoResponse, Response},
Json,
};
use chrono::Utc;
use std::collections::HashMap;
use crate::converter::anthropic_to_openai::convert_anthropic_to_openai;
use crate::flow_monitor::{
ClientInfo, FlowError, FlowErrorType, FlowMetadata, FlowType, InterceptAction, InterceptType,
LLMFlow, LLMRequest, LLMResponse, Message, MessageContent, MessageRole, RequestParameters,
RoutingInfo, TokenUsage,
};
use crate::models::anthropic::AnthropicMessagesRequest;
use crate::models::openai::ChatCompletionRequest;
use crate::processor::RequestContext;
@@ -37,8 +29,6 @@ use crate::server_utils::{
build_anthropic_response, build_anthropic_stream_response, message_content_len,
parse_cw_response, safe_truncate,
};
use crate::streaming::StreamFormat as StreamingFormat;
use crate::ProviderType;
use super::{call_provider_anthropic, call_provider_openai};
@@ -20,7 +20,6 @@ use axum::{
response::{IntoResponse, Response},
Json,
};
use futures::StreamExt;
use crate::converter::anthropic_to_openai::convert_anthropic_to_openai;
use crate::converter::openai_to_antigravity::{
@@ -696,20 +695,6 @@ pub async fn call_provider_openai(
request: &ChatCompletionRequest,
flow_id: Option<&str>,
) -> Response {
// 如果是流式请求且有 flow_id,设置流式状态
if request.stream {
if let Some(fid) = flow_id {
// 根据凭证类型确定流格式
let format = match &credential.credential {
CredentialData::KiroOAuth { .. } => StreamFormat::OpenAI,
CredentialData::ClaudeKey { .. } => StreamFormat::Anthropic,
CredentialData::AntigravityOAuth { .. } => StreamFormat::Gemini,
_ => StreamFormat::OpenAI,
};
state.flow_monitor.set_streaming(fid, format).await;
}
}
let _start_time = std::time::Instant::now();
match &credential.credential {
CredentialData::KiroOAuth { creds_file_path } => {
+3 -18
View File
@@ -196,12 +196,7 @@ pub async fn handle_websocket(
handle_ws_message(&state, &conn_id, ws_msg, &flow_subscribed).await;
if let Some(resp) = response {
let resp_text = serde_json::to_string(&resp).unwrap_or_default();
let mut sender_guard = sender.lock().await;
if sender_guard
.send(WsMessage::Text(resp_text.into()))
.await
.is_err()
{
if sender.send(WsMessage::Text(resp_text)).await.is_err() {
break;
}
}
@@ -213,12 +208,7 @@ pub async fn handle_websocket(
e
)));
let error_text = serde_json::to_string(&error).unwrap_or_default();
let mut sender_guard = sender.lock().await;
if sender_guard
.send(WsMessage::Text(error_text.into()))
.await
.is_err()
{
if sender.send(WsMessage::Text(error_text)).await.is_err() {
break;
}
}
@@ -230,12 +220,7 @@ pub async fn handle_websocket(
"Binary messages not supported",
));
let error_text = serde_json::to_string(&error).unwrap_or_default();
let mut sender_guard = sender.lock().await;
if sender_guard
.send(WsMessage::Text(error_text.into()))
.await
.is_err()
{
if sender.send(WsMessage::Text(error_text)).await.is_err() {
break;
}
}
+1 -1
View File
@@ -1575,7 +1575,7 @@ async fn amp_management_proxy_internal(
None => {
state.logs.write().await.add(
"warn",
&format!("[AMP] No upstream URL configured for management proxy"),
"[AMP] No upstream URL configured for management proxy",
);
return (
StatusCode::SERVICE_UNAVAILABLE,
@@ -523,7 +523,7 @@ impl ProviderPoolService {
format!("{} 服务暂时不可用。\n💡 解决方案:\n1. 这通常是服务提供方的临时问题\n2. 请稍后重试\n3. 如问题持续,可尝试其他凭证", provider_type)
} else if error.contains("读取凭证文件失败") || error.contains("解析凭证失败")
{
format!("凭证文件损坏或不可读。\n💡 解决方案:\n1. 凭证文件可能已损坏\n2. 建议删除此凭证后重新添加\n3. 确保文件权限正确且格式为有效的 JSON")
"凭证文件损坏或不可读。\n💡 解决方案:\n1. 凭证文件可能已损坏\n2. 建议删除此凭证后重新添加\n3. 确保文件权限正确且格式为有效的 JSON".to_string()
} else {
// 对于其他未识别的错误,提供通用建议
format!("操作失败:{}\n💡 建议:\n1. 检查网络连接和凭证状态\n2. 尝试刷新 Token 或重新添加凭证\n3. 如问题持续,请联系技术支持", error)
@@ -1196,7 +1196,7 @@ impl ProviderPoolService {
}
// 为每个命名凭证创建路由
for (_provider_type, credentials) in &grouped {
for credentials in grouped.values() {
for cred in credentials {
if let Some(name) = &cred.name {
if cred.is_available() {
+2
View File
@@ -204,6 +204,7 @@ pub fn build_request_headers(
}
/// 构造 User-Agent 字符串(用于测试)
#[cfg(test)]
pub fn build_user_agent(kiro_version: &str, machine_id: &str) -> String {
let os_name = std::env::consts::OS;
format!(
@@ -213,6 +214,7 @@ pub fn build_user_agent(kiro_version: &str, machine_id: &str) -> String {
}
/// 构造 x-amz-user-agent 字符串(用于测试)
#[cfg(test)]
pub fn build_x_amz_user_agent(kiro_version: &str, machine_id: &str) -> String {
format!("aws-sdk-js/1.0.0 KiroIDE-{}-{}", kiro_version, machine_id)
}
+4 -3
View File
@@ -12,6 +12,7 @@ use serde::{Deserialize, Serialize};
use std::collections::{HashMap, VecDeque};
use std::fs::{self, File, OpenOptions};
use std::io::{BufRead, BufReader, Write};
use std::path::Path;
use std::path::PathBuf;
/// 日志记录器错误
@@ -282,7 +283,7 @@ impl RequestLogger {
let entry = entry?;
let path = entry.path();
if path.is_file() && path.extension().map_or(false, |ext| ext == "jsonl") {
if path.is_file() && path.extension().is_some_and(|ext| ext == "jsonl") {
// 从文件名解析日期
if let Some(file_date) = self.parse_log_file_date(&path) {
if file_date < cutoff {
@@ -393,7 +394,7 @@ impl RequestLogger {
}
/// 从日志文件名解析日期
fn parse_log_file_date(&self, path: &PathBuf) -> Option<DateTime<Utc>> {
fn parse_log_file_date(&self, path: &Path) -> Option<DateTime<Utc>> {
let file_name = path.file_stem()?.to_str()?;
// 文件名格式: requests_YYYY-MM-DD 或 requests_YYYY-MM-DD_N
let date_part = file_name.strip_prefix("requests_")?;
@@ -441,7 +442,7 @@ impl RequestLogger {
let entry = entry?;
let path = entry.path();
if path.is_file() && path.extension().map_or(false, |ext| ext == "jsonl") {
if path.is_file() && path.extension().is_some_and(|ext| ext == "jsonl") {
if let Some(file_date) = self.parse_log_file_date(&path) {
if file_date >= cutoff {
log_files.push(path);
+1 -1
View File
@@ -226,7 +226,7 @@ impl<R: Runtime> TrayManager<R> {
.tooltip("ProxyCast - AI API 代理")
.on_tray_icon_event(|tray, event| {
let app = tray.app_handle();
handle_tray_icon_event(&app, event);
handle_tray_icon_event(app, event);
})
.on_menu_event(|app, event| {
handle_menu_event(app, event.id().as_ref());
+3 -3
View File
@@ -150,7 +150,7 @@ async fn handle_socket(socket: WebSocket, state: WsHandlerState, client_info: Op
let response = handle_message(&state, &conn_id, ws_msg).await;
if let Some(resp) = response {
let resp_text = serde_json::to_string(&resp).unwrap_or_default();
if sender.send(Message::Text(resp_text.into())).await.is_err() {
if sender.send(Message::Text(resp_text)).await.is_err() {
break;
}
}
@@ -162,7 +162,7 @@ async fn handle_socket(socket: WebSocket, state: WsHandlerState, client_info: Op
e
)));
let error_text = serde_json::to_string(&error).unwrap_or_default();
if sender.send(Message::Text(error_text.into())).await.is_err() {
if sender.send(Message::Text(error_text)).await.is_err() {
break;
}
}
@@ -174,7 +174,7 @@ async fn handle_socket(socket: WebSocket, state: WsHandlerState, client_info: Op
let error =
WsMessage::Error(WsError::invalid_message("Binary messages not supported"));
let error_text = serde_json::to_string(&error).unwrap_or_default();
if sender.send(Message::Text(error_text.into())).await.is_err() {
if sender.send(Message::Text(error_text)).await.is_err() {
break;
}
}
+4 -4
View File
@@ -42,10 +42,10 @@ impl StreamForwarder {
}
// 处理 data: 前缀
let data = if trimmed.starts_with("data: ") {
&trimmed[6..]
} else if trimmed.starts_with("data:") {
&trimmed[5..]
let data = if let Some(stripped) = trimmed.strip_prefix("data: ") {
stripped
} else if let Some(stripped) = trimmed.strip_prefix("data:") {
stripped
} else {
trimmed
};
+101 -11
View File
@@ -15,6 +15,12 @@ interface VersionInfo {
error?: string;
}
interface DownloadResult {
success: boolean;
message: string;
filePath?: string;
}
interface ToolVersion {
name: string;
version: string | null;
@@ -30,6 +36,10 @@ export function AboutSection() {
error: undefined,
});
const [checking, setChecking] = useState(false);
const [downloading, setDownloading] = useState(false);
const [downloadResult, setDownloadResult] = useState<DownloadResult | null>(
null,
);
const [toolVersions, setToolVersions] = useState<ToolVersion[]>([]);
const [loadingTools, setLoadingTools] = useState(true);
@@ -67,6 +77,7 @@ export function AboutSection() {
const handleCheckUpdate = async () => {
setChecking(true);
setDownloadResult(null);
try {
const result = await invoke<VersionInfo>("check_for_updates");
setVersionInfo(result);
@@ -81,6 +92,37 @@ export function AboutSection() {
}
};
const handleDownloadUpdate = async () => {
setDownloading(true);
setDownloadResult(null);
try {
const result = await invoke<DownloadResult>("download_update");
setDownloadResult(result);
if (result.success) {
// 下载成功,显示安装提示
setTimeout(() => {
setDownloadResult({
...result,
message: "安装程序已启动,应用将自动关闭以完成更新",
});
}, 1000);
} else {
// 下载失败,显示错误信息
console.error("Download failed:", result.message);
}
} catch (error) {
console.error("Failed to download update:", error);
setDownloadResult({
success: false,
message: "下载失败,请手动下载",
filePath: undefined,
});
} finally {
setDownloading(false);
}
};
return (
<div className="space-y-6 max-w-2xl">
{/* 应用信息 */}
@@ -116,7 +158,7 @@ export function AboutSection() {
<div className="flex items-center justify-center gap-2">
<button
onClick={handleCheckUpdate}
disabled={checking}
disabled={checking || downloading}
className="inline-flex items-center gap-2 px-4 py-2 rounded-lg border text-sm hover:bg-muted disabled:opacity-50"
>
<RefreshCw
@@ -125,18 +167,66 @@ export function AboutSection() {
检查更新
</button>
{versionInfo.hasUpdate && versionInfo.downloadUrl && (
<a
href={versionInfo.downloadUrl}
target="_blank"
rel="noopener noreferrer"
className="inline-flex items-center gap-2 px-4 py-2 rounded-lg bg-green-600 text-white text-sm hover:bg-green-700"
>
<ExternalLink className="h-4 w-4" />
下载新版本
</a>
{versionInfo.hasUpdate && (
<>
<button
onClick={handleDownloadUpdate}
disabled={downloading}
className="inline-flex items-center gap-2 px-4 py-2 rounded-lg bg-green-600 text-white text-sm hover:bg-green-700 disabled:opacity-50"
>
<RefreshCw
className={`h-4 w-4 ${downloading ? "animate-spin" : ""}`}
/>
{downloading ? "下载中..." : "下载更新"}
</button>
{versionInfo.downloadUrl && (
<a
href={versionInfo.downloadUrl}
target="_blank"
rel="noopener noreferrer"
className="inline-flex items-center gap-2 px-4 py-2 rounded-lg border text-sm hover:bg-muted"
>
<ExternalLink className="h-4 w-4" />
网页下载
</a>
)}
</>
)}
</div>
{/* 下载结果提示 */}
{downloadResult && (
<div
className={`mt-2 p-3 rounded-lg text-sm ${
downloadResult.success
? "bg-green-50 text-green-700 border border-green-200"
: "bg-red-50 text-red-700 border border-red-200"
}`}
>
<div className="flex items-start gap-2">
{downloadResult.success ? (
<CheckCircle2 className="h-4 w-4 mt-0.5 flex-shrink-0" />
) : (
<AlertCircle className="h-4 w-4 mt-0.5 flex-shrink-0" />
)}
<div className="flex-1">
<p>{downloadResult.message}</p>
{!downloadResult.success && versionInfo.downloadUrl && (
<a
href={versionInfo.downloadUrl}
target="_blank"
rel="noopener noreferrer"
className="inline-flex items-center gap-1 mt-2 underline hover:no-underline"
>
<ExternalLink className="h-3 w-3" />
前往网页下载
</a>
)}
</div>
</div>
</div>
)}
</div>
{/* 链接 */}
Binary file not shown.
Binary file not shown.
+1 -1
View File
@@ -13,7 +13,7 @@ export default defineConfig({
clearScreen: false,
server: {
port: 1420,
strictPort: true,
strictPort: false,
watch: {
ignored: ["**/src-tauri/**"],
},