feat: 实现自动下载更新功能

- 添加 check_for_updates 命令检查 GitHub 最新版本
- 添加 download_update 命令自动下载对应平台安装包
- 支持 Windows (.exe/.msi)、macOS (.dmg)、Linux (.deb/.appimage)
- 动态从 GitHub API 获取实际文件列表并匹配平台
- 下载完成后自动运行安装程序
- 安装程序启动成功后自动退出应用避免文件占用
- 失败时提供网页下载备选方案
- 前端 UI 显示下载状态和友好提示
- 添加版本比较和平台检测单元测试
This commit is contained in:
L_SEN
2025-12-21 12:14:39 +08:00
parent 6923f5062a
commit 942fab1f10
42 changed files with 686 additions and 301 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")]
@@ -111,117 +111,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)
}
@@ -731,3 +676,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(())
}
+1 -2
View File
@@ -1,8 +1,7 @@
//! 插件系统相关命令
use crate::plugin::{PluginConfig, PluginInfo, PluginManager, PluginStatus};
use crate::plugin::{PluginConfig, PluginInfo, PluginManager};
use serde::{Deserialize, Serialize};
use std::path::PathBuf;
use std::sync::Arc;
use tokio::sync::RwLock;
+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))?;
}
// 返回新的文件路径
@@ -1179,7 +1178,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(
@@ -1227,7 +1226,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);
let api_key = &config.server.api_key;
+2 -2
View File
@@ -7,10 +7,10 @@
//! - 7.2: 凭证健康状态变化时在 1 秒内更新托盘图标
//! - 7.3: 托盘菜单打开时获取并显示最新信息
use crate::tray::{CredentialHealth, TrayIconStatus, TrayStateSnapshot};
use crate::tray::{TrayIconStatus, TrayStateSnapshot};
use crate::TrayManagerState;
use tauri::State;
use tracing::{debug, error, info};
use tracing::{debug, info};
/// 同步托盘状态
///
+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()
+1 -1
View File
@@ -1,6 +1,6 @@
//! WebSocket 相关的 Tauri 命令
use crate::websocket::{WsConfig, WsConnection, WsStatsSnapshot};
use crate::websocket::{WsConnection, WsStatsSnapshot};
use serde::{Deserialize, Serialize};
use std::sync::Arc;
use tokio::sync::RwLock;
+2 -1
View File
@@ -496,7 +496,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 {
+8 -16
View File
@@ -10,26 +10,18 @@ mod path_utils;
mod types;
mod yaml;
pub use export::{
base64_decode, base64_encode, ExportBundle, ExportError, ExportOptions, ExportService,
REDACTED_PLACEHOLDER,
};
pub use export::{ExportBundle, ExportOptions, ExportService};
pub use hot_reload::{
ConfigChangeEvent, ConfigChangeKind, FileWatcher, HotReloadError, HotReloadManager,
HotReloadStatus, ReloadResult,
ConfigChangeEvent, ConfigChangeKind, FileWatcher, HotReloadManager, ReloadResult,
};
pub use import::{ImportError, ImportOptions, ImportResult, ImportService, ValidationResult};
pub use path_utils::{collapse_tilde, contains_tilde, expand_tilde};
pub use import::{ImportOptions, ImportService, ValidationResult};
pub use path_utils::expand_tilde;
pub use types::{
AmpConfig, AmpModelMapping, ApiKeyEntry, Config, CredentialEntry, CredentialPoolConfig,
CustomProviderConfig, GeminiApiKeyEntry, IFlowCredentialEntry, InjectionRuleConfig,
InjectionSettings, LoggingConfig, ProviderConfig, ProvidersConfig, QuotaExceededConfig,
RemoteManagementConfig, RetrySettings, RoutingConfig, RoutingRuleConfig, ServerConfig,
TlsConfig, VertexApiKeyEntry, VertexModelAlias,
};
pub use yaml::{
load_config, save_config, save_config_yaml, ConfigError, ConfigManager, YamlService,
AmpConfig, AmpModelMapping, ApiKeyEntry, Config, CredentialEntry, GeminiApiKeyEntry,
InjectionRuleConfig, InjectionSettings, QuotaExceededConfig, RemoteManagementConfig,
VertexApiKeyEntry, VertexModelAlias,
};
pub use yaml::{load_config, save_config, ConfigError, ConfigManager, YamlService};
#[cfg(test)]
mod tests;
+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)
+2 -3
View File
@@ -33,9 +33,7 @@ use commands::skill_cmd::SkillServiceState;
use services::provider_pool_service::ProviderPoolService;
use services::skill_service::SkillService;
use services::token_cache_service::TokenCacheService;
use tray::{
calculate_icon_status, CredentialHealth, TrayIconStatus, TrayManager, TrayStateSnapshot,
};
use tray::{TrayIconStatus, TrayManager, TrayStateSnapshot};
/// TokenCacheService 状态封装
pub struct TokenCacheServiceState(pub Arc<TokenCacheService>);
@@ -1647,6 +1645,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
@@ -77,8 +77,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());
}
}
}
+1 -2
View File
@@ -13,8 +13,7 @@ use tokio::time::timeout;
use super::loader::PluginLoader;
use super::types::{
HookResult, Plugin, PluginConfig, PluginContext, PluginError, PluginInfo, PluginInstance,
PluginStatus,
HookResult, PluginConfig, PluginContext, PluginError, PluginInfo, PluginInstance, PluginStatus,
};
/// 插件管理器配置
+1 -1
View File
@@ -16,4 +16,4 @@ pub use plugin::{PluginPostStep, PluginPreStep};
pub use provider::ProviderStep;
pub use routing::RoutingStep;
pub use telemetry::TelemetryStep;
pub use traits::{PipelineStep, StepError};
pub use traits::PipelineStep;
+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();
+12 -12
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(),
};
@@ -1421,7 +1421,7 @@ pub async fn start_oauth_login(
// 构建凭证
let now = chrono::Utc::now();
let mut credentials = AntigravityCredentials {
let credentials = AntigravityCredentials {
access_token: Some(access_token.to_string()),
refresh_token,
token_type: Some("Bearer".to_string()),
@@ -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
@@ -540,7 +540,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))
}
@@ -1072,11 +1072,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 {
@@ -1108,14 +1106,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<_>>()
});
@@ -1810,8 +1808,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 -9
View File
@@ -3,29 +3,21 @@
//! 处理 OpenAI 和 Anthropic 格式的 API 请求
use axum::{
body::Body,
extract::State,
http::{header, HeaderMap, StatusCode},
http::{HeaderMap, StatusCode},
response::{IntoResponse, Response},
Json,
};
use futures::stream;
use crate::converter::anthropic_to_openai::convert_anthropic_to_openai;
use crate::converter::openai_to_antigravity::{
convert_antigravity_to_openai_response, convert_openai_to_antigravity_with_context,
};
use crate::models::anthropic::AnthropicMessagesRequest;
use crate::models::openai::ChatCompletionRequest;
use crate::processor::RequestContext;
use crate::providers::{AntigravityProvider, GeminiProvider, KiroProvider, QwenProvider};
use crate::server::{record_request_telemetry, record_token_usage, AppState};
use crate::server_utils::{
build_anthropic_response, build_anthropic_stream_response, message_content_len,
parse_cw_response, safe_truncate,
};
use crate::telemetry::RequestStatus;
use crate::ProviderType;
use super::{call_provider_anthropic, call_provider_openai};
@@ -8,7 +8,6 @@ use axum::{
response::{IntoResponse, Response},
Json,
};
use futures::stream;
use crate::converter::anthropic_to_openai::convert_anthropic_to_openai;
use crate::converter::openai_to_antigravity::{
@@ -18,8 +17,7 @@ use crate::models::anthropic::AnthropicMessagesRequest;
use crate::models::openai::ChatCompletionRequest;
use crate::models::provider_pool_model::{CredentialData, ProviderCredential};
use crate::providers::{
AntigravityProvider, ClaudeCustomProvider, GeminiProvider, KiroProvider, OpenAICustomProvider,
QwenProvider, VertexProvider,
AntigravityProvider, ClaudeCustomProvider, KiroProvider, OpenAICustomProvider, VertexProvider,
};
use crate::server::AppState;
use crate::server_utils::{
@@ -650,7 +648,7 @@ pub async fn call_provider_openai(
credential: &ProviderCredential,
request: &ChatCompletionRequest,
) -> Response {
let start_time = std::time::Instant::now();
let _start_time = std::time::Instant::now();
match &credential.credential {
CredentialData::KiroOAuth { creds_file_path } => {
let mut kiro = KiroProvider::new();
+3 -15
View File
@@ -110,11 +110,7 @@ pub async fn handle_websocket(socket: WebSocket, state: AppState, client_info: O
let response = handle_ws_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(WsMessage::Text(resp_text.into()))
.await
.is_err()
{
if sender.send(WsMessage::Text(resp_text)).await.is_err() {
break;
}
}
@@ -126,11 +122,7 @@ pub async fn handle_websocket(socket: WebSocket, state: AppState, client_info: O
e
)));
let error_text = serde_json::to_string(&error).unwrap_or_default();
if sender
.send(WsMessage::Text(error_text.into()))
.await
.is_err()
{
if sender.send(WsMessage::Text(error_text)).await.is_err() {
break;
}
}
@@ -142,11 +134,7 @@ pub async fn handle_websocket(socket: WebSocket, state: AppState, client_info: O
"Binary messages not supported",
));
let error_text = serde_json::to_string(&error).unwrap_or_default();
if sender
.send(WsMessage::Text(error_text.into()))
.await
.is_err()
{
if sender.send(WsMessage::Text(error_text)).await.is_err() {
break;
}
}
+4 -12
View File
@@ -4,9 +4,6 @@ use crate::config::{
ReloadResult,
};
use crate::converter::anthropic_to_openai::convert_anthropic_to_openai;
use crate::converter::openai_to_antigravity::{
convert_antigravity_to_openai_response, convert_openai_to_antigravity_with_context,
};
use crate::credential::CredentialSyncService;
use crate::database::dao::provider_pool::ProviderPoolDao;
use crate::database::DbConnection;
@@ -23,28 +20,25 @@ use crate::providers::gemini::GeminiProvider;
use crate::providers::kiro::KiroProvider;
use crate::providers::openai_custom::OpenAICustomProvider;
use crate::providers::qwen::QwenProvider;
use crate::providers::vertex::VertexProvider;
use crate::server_utils::{
build_anthropic_response, build_anthropic_stream_response, build_gemini_native_request, health,
message_content_len, models, parse_cw_response, safe_truncate, CWParsedResponse,
models, parse_cw_response,
};
use crate::services::provider_pool_service::ProviderPoolService;
use crate::services::token_cache_service::TokenCacheService;
use crate::telemetry::{RequestLog, RequestStatus};
use crate::websocket::{WsConfig, WsConnectionManager, WsStats};
use axum::{
body::Body,
extract::{DefaultBodyLimit, Path, State},
http::{header, HeaderMap, StatusCode},
http::{HeaderMap, StatusCode},
response::{IntoResponse, Response},
routing::{get, post},
Json, Router,
};
use futures::stream;
use serde::{Deserialize, Serialize};
use std::path::PathBuf;
use std::sync::Arc;
use tokio::sync::{mpsc, oneshot, RwLock};
use tokio::sync::{oneshot, RwLock};
/// 记录请求统计到遥测系统
pub fn record_request_telemetry(
@@ -1527,7 +1521,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,
@@ -1797,5 +1791,3 @@ async fn chat_completions_internal(state: &AppState, request: &ChatCompletionReq
.into_response(),
}
}
use crate::models::provider_pool_model::ProviderCredential;
@@ -519,7 +519,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)
@@ -1134,7 +1134,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());
+2 -2
View File
@@ -9,9 +9,9 @@
use super::state::{calculate_icon_status, CredentialHealth, TrayIconStatus, TrayStateSnapshot};
use super::TrayManager;
use std::sync::Arc;
use tauri::{AppHandle, Manager, Runtime};
use tauri::{AppHandle, Runtime};
use tokio::sync::RwLock;
use tracing::{debug, error, info};
use tracing::{debug, info};
/// 托盘状态同步器
///
+3 -3
View File
@@ -134,7 +134,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;
}
}
@@ -146,7 +146,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;
}
}
@@ -158,7 +158,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;
}
}
+1 -3
View File
@@ -2,11 +2,9 @@
//!
//! 提供心跳检测、优雅关闭和资源清理功能
use super::{WsConnection, WsConnectionStatus, WsError, WsMessage};
use super::WsMessage;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::mpsc;
/// 心跳管理器
#[derive(Debug)]
+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/**"],
},