mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
@@ -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(())
|
||||
}
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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}";
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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};
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
|
||||
/// 检查是否启用自动切换项目
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
|
||||
|
||||
@@ -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(),
|
||||
};
|
||||
|
||||
|
||||
@@ -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()),
|
||||
}
|
||||
|
||||
@@ -604,7 +604,7 @@ impl CodexProvider {
|
||||
("codex_cli_simplified_flow", "true"),
|
||||
];
|
||||
|
||||
let query = serde_urlencoded::to_string(¶ms)?;
|
||||
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()),
|
||||
}
|
||||
|
||||
@@ -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::*;
|
||||
|
||||
@@ -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");
|
||||
|
||||
|
||||
@@ -608,7 +608,7 @@ impl IFlowProvider {
|
||||
("code_challenge_method", "S256"),
|
||||
];
|
||||
|
||||
let query = serde_urlencoded::to_string(¶ms)?;
|
||||
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()),
|
||||
}
|
||||
|
||||
@@ -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 } => {
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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());
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
};
|
||||
|
||||
@@ -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
@@ -13,7 +13,7 @@ export default defineConfig({
|
||||
clearScreen: false,
|
||||
server: {
|
||||
port: 1420,
|
||||
strictPort: true,
|
||||
strictPort: false,
|
||||
watch: {
|
||||
ignored: ["**/src-tauri/**"],
|
||||
},
|
||||
|
||||
Reference in New Issue
Block a user