diff --git a/package.json b/package.json index b13f8587b..9de4a769a 100644 --- a/package.json +++ b/package.json @@ -1,7 +1,7 @@ { "name": "proxycast", "private": true, - "version": "0.12.4", + "version": "0.12.5", "type": "module", "scripts": { "dev": "vite", diff --git a/src-tauri/Cargo.lock b/src-tauri/Cargo.lock index 68bc42c37..ebfd84398 100644 --- a/src-tauri/Cargo.lock +++ b/src-tauri/Cargo.lock @@ -3349,7 +3349,7 @@ dependencies = [ [[package]] name = "proxycast" -version = "0.12.4" +version = "0.12.5" dependencies = [ "anyhow", "async-stream", @@ -3390,6 +3390,7 @@ dependencies = [ "tower-http 0.5.2", "tracing", "tracing-subscriber", + "url", "urlencoding", "uuid", "zip", diff --git a/src-tauri/Cargo.toml b/src-tauri/Cargo.toml index 067f231f4..08a7a30ac 100644 --- a/src-tauri/Cargo.toml +++ b/src-tauri/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "proxycast" -version = "0.12.4" +version = "0.12.5" description = "AI API Proxy Desktop App" authors = ["you"] edition = "2021" @@ -52,6 +52,7 @@ rand = "0.8" sha2 = "0.10" serde_urlencoded = "0.7" open = "5" +url = "2" [dev-dependencies] proptest = "1" diff --git a/src-tauri/src/commands/mod.rs b/src-tauri/src/commands/mod.rs index 58e8d2b1e..656d134da 100644 --- a/src-tauri/src/commands/mod.rs +++ b/src-tauri/src/commands/mod.rs @@ -11,4 +11,5 @@ pub mod router_cmd; pub mod skill_cmd; pub mod switch_cmd; pub mod telemetry_cmd; +pub mod usage_cmd; pub mod websocket_cmd; diff --git a/src-tauri/src/commands/provider_pool_cmd.rs b/src-tauri/src/commands/provider_pool_cmd.rs index 7c8364ae6..257169fcf 100644 --- a/src-tauri/src/commands/provider_pool_cmd.rs +++ b/src-tauri/src/commands/provider_pool_cmd.rs @@ -87,6 +87,33 @@ fn copy_and_rename_credential_file( let mut creds: serde_json::Value = serde_json::from_str(&content).map_err(|e| format!("解析凭证文件失败: {}", e))?; + // 检测 refreshToken 是否被截断 + // 正常的 refreshToken 长度应该在 500+ 字符,如果小于 100 字符则可能被截断 + if let Some(refresh_token) = creds.get("refreshToken").and_then(|v| v.as_str()) { + let token_len = refresh_token.len(); + + // 检测常见的截断模式 + let is_truncated = + token_len < 100 || refresh_token.ends_with("...") || refresh_token.contains("..."); + + if is_truncated { + tracing::error!( + "[KIRO] 检测到 refreshToken 被截断!长度: {}, 内容: {}", + token_len, + &refresh_token[..std::cmp::min(50, token_len)] + ); + return Err(format!( + "凭证文件中的 refreshToken 已被截断(长度: {} 字符)。\n\n⚠️ 这通常是 Kiro IDE 为了防止凭证被第三方工具使用而故意截断的。\n\n💡 解决方案:\n1. 使用 Kir-Manager 工具获取完整的凭证\n2. 或者使用其他方式获取未截断的凭证文件\n3. 正常的 refreshToken 长度应该在 500+ 字符\n\n当前 refreshToken: {}...", + token_len, + &refresh_token[..std::cmp::min(30, token_len)] + )); + } + + tracing::info!("[KIRO] refreshToken 长度检查通过: {} 字符", token_len); + } else { + tracing::warn!("[KIRO] 凭证文件中没有 refreshToken 字段"); + } + let aws_sso_cache_dir = dirs::home_dir() .ok_or_else(|| "无法获取用户主目录".to_string())? .join(".aws") @@ -164,9 +191,23 @@ fn copy_and_rename_credential_file( } if !found_credentials { - tracing::warn!( - "[KIRO] 未找到 client_id/client_secret,副本可能无法独立刷新 Token(将使用 social 认证)" - ); + // 检查认证方式 + let auth_method = creds + .get("authMethod") + .and_then(|v| v.as_str()) + .unwrap_or("social"); + + if auth_method.to_lowercase() == "idc" { + // IdC 认证必须有 clientId/clientSecret + 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 认证方式的凭证" + )); + } else { + tracing::warn!("[KIRO] 未找到 client_id/client_secret,将使用 social 认证方式"); + } } // 写入合并后的凭证到副本文件 diff --git a/src-tauri/src/commands/usage_cmd.rs b/src-tauri/src/commands/usage_cmd.rs new file mode 100644 index 000000000..2c6a24a87 --- /dev/null +++ b/src-tauri/src/commands/usage_cmd.rs @@ -0,0 +1,349 @@ +//! Usage Tauri 命令 +//! +//! 提供 Kiro 用量查询的 Tauri 命令接口。 + +use crate::database::dao::provider_pool::ProviderPoolDao; +use crate::database::DbConnection; +use crate::models::provider_pool_model::{CredentialData, PoolProviderType}; +use crate::services::usage_service::{self, UsageInfo}; +use crate::TokenCacheServiceState; +use tauri::State; + +/// 默认 Kiro 版本号 +const DEFAULT_KIRO_VERSION: &str = "1.0.0"; + +/// 获取 Kiro 用量信息 +/// +/// **Validates: Requirements 1.1** +/// +/// # Arguments +/// * `credential_uuid` - 凭证的 UUID +/// * `db` - 数据库连接 +/// * `token_cache` - Token 缓存服务 +/// +/// # Returns +/// * `Ok(UsageInfo)` - 成功时返回用量信息 +/// * `Err(String)` - 失败时返回错误消息 +#[tauri::command] +pub async fn get_kiro_usage( + credential_uuid: String, + db: State<'_, DbConnection>, + token_cache: State<'_, TokenCacheServiceState>, +) -> Result { + // 1. 获取凭证信息 + let credential = { + let conn = db.lock().map_err(|e| e.to_string())?; + ProviderPoolDao::get_by_uuid(&conn, &credential_uuid) + .map_err(|e| e.to_string())? + .ok_or_else(|| format!("凭证不存在: {}", credential_uuid))? + }; + + // 2. 验证是否为 Kiro 凭证 + if credential.provider_type != PoolProviderType::Kiro { + return Err(format!( + "不支持的凭证类型: {:?},仅支持 Kiro 凭证", + credential.provider_type + )); + } + + // 3. 获取凭证文件路径 + let creds_file_path = match &credential.credential { + CredentialData::KiroOAuth { creds_file_path } => creds_file_path.clone(), + _ => return Err("凭证数据类型不匹配".to_string()), + }; + + // 4. 获取有效的 access_token + let access_token = token_cache + .0 + .get_valid_token(&db, &credential_uuid) + .await + .map_err(|e| { + // 提供更友好的错误信息 + if e.contains("401") || e.contains("Bad credentials") || e.contains("过期") || e.contains("无效") { + format!("刷新 Kiro Token 失败: OAuth 凭证已过期或无效,需要重新认证。\n💡 解决方案:\n1. 删除当前 OAuth 凭证\n2. 重新添加 OAuth 凭证\n3. 确保使用最新的凭证文件\n\n技术详情:{}", e) + } else { + e + } + })?; + + // 5. 从凭证文件读取 auth_method 和 profile_arn + let (auth_method, profile_arn) = read_kiro_credential_info(&creds_file_path)?; + + // 6. 获取 machine_id + let machine_id = get_machine_id()?; + + // 7. 调用 Usage API + let usage_info = usage_service::get_usage_limits_safe( + &access_token, + &auth_method, + profile_arn.as_deref(), + &machine_id, + DEFAULT_KIRO_VERSION, + ) + .await; + + Ok(usage_info) +} + +/// 从 Kiro 凭证文件读取 auth_method 和 profile_arn +fn read_kiro_credential_info(creds_file_path: &str) -> Result<(String, Option), String> { + // 展开 ~ 路径 + let expanded_path = expand_tilde(creds_file_path); + + // 读取文件 + let content = + std::fs::read_to_string(&expanded_path).map_err(|e| format!("读取凭证文件失败: {}", e))?; + + // 解析 JSON + let json: serde_json::Value = + serde_json::from_str(&content).map_err(|e| format!("解析凭证文件失败: {}", e))?; + + // 获取 auth_method,默认为 "social" + let auth_method = json + .get("authMethod") + .and_then(|v| v.as_str()) + .unwrap_or("social") + .to_string(); + + // 获取 profile_arn(可选) + let profile_arn = json + .get("profileArn") + .and_then(|v| v.as_str()) + .map(|s| s.to_string()); + + Ok((auth_method, profile_arn)) +} + +/// 展开路径中的 ~ 为用户主目录 +fn expand_tilde(path: &str) -> String { + if path.starts_with("~/") { + if let Some(home) = dirs::home_dir() { + return home.join(&path[2..]).to_string_lossy().to_string(); + } + } + path.to_string() +} + +/// 获取设备 ID(SHA256 哈希) +fn get_machine_id() -> Result { + // 尝试获取系统 machine-id + let raw_id = get_raw_machine_id()?; + + // 计算 SHA256 哈希 + use sha2::{Digest, Sha256}; + let mut hasher = Sha256::new(); + hasher.update(raw_id.as_bytes()); + let result = hasher.finalize(); + + Ok(format!("{:x}", result)) +} + +/// 获取原始设备 ID +fn get_raw_machine_id() -> Result { + #[cfg(target_os = "macos")] + { + // macOS: 使用 IOPlatformUUID + use std::process::Command; + let output = Command::new("ioreg") + .args(["-rd1", "-c", "IOPlatformExpertDevice"]) + .output() + .map_err(|e| format!("执行 ioreg 失败: {}", e))?; + + let stdout = String::from_utf8_lossy(&output.stdout); + for line in stdout.lines() { + if line.contains("IOPlatformUUID") { + if let Some(uuid) = line.split('"').nth(3) { + return Ok(uuid.to_string()); + } + } + } + Err("无法获取 IOPlatformUUID".to_string()) + } + + #[cfg(target_os = "linux")] + { + // Linux: 读取 /etc/machine-id + std::fs::read_to_string("/etc/machine-id") + .map(|s| s.trim().to_string()) + .map_err(|e| format!("读取 /etc/machine-id 失败: {}", e)) + } + + #[cfg(target_os = "windows")] + { + // Windows: 使用注册表中的 MachineGuid + use std::process::Command; + let output = Command::new("reg") + .args([ + "query", + "HKEY_LOCAL_MACHINE\\SOFTWARE\\Microsoft\\Cryptography", + "/v", + "MachineGuid", + ]) + .output() + .map_err(|e| format!("执行 reg query 失败: {}", e))?; + + let stdout = String::from_utf8_lossy(&output.stdout); + for line in stdout.lines() { + if line.contains("MachineGuid") { + if let Some(guid) = line.split_whitespace().last() { + return Ok(guid.to_string()); + } + } + } + Err("无法获取 MachineGuid".to_string()) + } + + #[cfg(not(any(target_os = "macos", target_os = "linux", target_os = "windows")))] + { + Err("不支持的操作系统".to_string()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_expand_tilde() { + let path = "~/test/path"; + let expanded = expand_tilde(path); + assert!(!expanded.starts_with("~/")); + assert!(expanded.ends_with("test/path")); + } + + #[test] + fn test_expand_tilde_no_tilde() { + let path = "/absolute/path"; + let expanded = expand_tilde(path); + assert_eq!(expanded, path); + } + + #[test] + fn test_get_machine_id() { + // 这个测试在不同平台上行为不同 + let result = get_machine_id(); + // 应该能成功获取 machine_id + assert!(result.is_ok(), "Failed to get machine_id: {:?}", result); + // machine_id 应该是 64 字符的十六进制字符串(SHA256) + let id = result.unwrap(); + assert_eq!(id.len(), 64, "Machine ID should be 64 hex chars"); + assert!( + id.chars().all(|c| c.is_ascii_hexdigit()), + "Machine ID should be hex" + ); + } +} + +// ============================================================================ +// 集成测试 +// ============================================================================ + +#[cfg(test)] +mod integration_tests { + use super::*; + + /// 测试 read_kiro_credential_info 函数 + /// 验证能正确解析 Kiro 凭证文件中的 auth_method 和 profile_arn + #[test] + fn test_read_kiro_credential_info_social() { + // 创建临时文件 + let temp_dir = std::env::temp_dir(); + let temp_file = temp_dir.join("test_kiro_creds_social.json"); + + let creds_json = serde_json::json!({ + "accessToken": "test_access_token", + "refreshToken": "test_refresh_token", + "authMethod": "social", + "profileArn": "arn:aws:iam::123456789:profile/test" + }); + + std::fs::write(&temp_file, serde_json::to_string(&creds_json).unwrap()).unwrap(); + + let result = read_kiro_credential_info(temp_file.to_str().unwrap()); + assert!(result.is_ok()); + + let (auth_method, profile_arn) = result.unwrap(); + assert_eq!(auth_method, "social"); + assert_eq!( + profile_arn, + Some("arn:aws:iam::123456789:profile/test".to_string()) + ); + + // 清理 + let _ = std::fs::remove_file(&temp_file); + } + + /// 测试 read_kiro_credential_info 函数 - IdC 认证 + #[test] + fn test_read_kiro_credential_info_idc() { + let temp_dir = std::env::temp_dir(); + let temp_file = temp_dir.join("test_kiro_creds_idc.json"); + + let creds_json = serde_json::json!({ + "accessToken": "test_access_token", + "refreshToken": "test_refresh_token", + "authMethod": "idc" + }); + + std::fs::write(&temp_file, serde_json::to_string(&creds_json).unwrap()).unwrap(); + + let result = read_kiro_credential_info(temp_file.to_str().unwrap()); + assert!(result.is_ok()); + + let (auth_method, profile_arn) = result.unwrap(); + assert_eq!(auth_method, "idc"); + assert_eq!(profile_arn, None); + + // 清理 + let _ = std::fs::remove_file(&temp_file); + } + + /// 测试 read_kiro_credential_info 函数 - 默认 auth_method + #[test] + fn test_read_kiro_credential_info_default_auth_method() { + let temp_dir = std::env::temp_dir(); + let temp_file = temp_dir.join("test_kiro_creds_default.json"); + + // 没有 authMethod 字段,应该默认为 "social" + let creds_json = serde_json::json!({ + "accessToken": "test_access_token", + "refreshToken": "test_refresh_token" + }); + + std::fs::write(&temp_file, serde_json::to_string(&creds_json).unwrap()).unwrap(); + + let result = read_kiro_credential_info(temp_file.to_str().unwrap()); + assert!(result.is_ok()); + + let (auth_method, profile_arn) = result.unwrap(); + assert_eq!(auth_method, "social"); + assert_eq!(profile_arn, None); + + // 清理 + let _ = std::fs::remove_file(&temp_file); + } + + /// 测试 read_kiro_credential_info 函数 - 文件不存在 + #[test] + fn test_read_kiro_credential_info_file_not_found() { + let result = read_kiro_credential_info("/nonexistent/path/to/creds.json"); + assert!(result.is_err()); + assert!(result.unwrap_err().contains("读取凭证文件失败")); + } + + /// 测试 read_kiro_credential_info 函数 - 无效 JSON + #[test] + fn test_read_kiro_credential_info_invalid_json() { + let temp_dir = std::env::temp_dir(); + let temp_file = temp_dir.join("test_kiro_creds_invalid.json"); + + std::fs::write(&temp_file, "not valid json").unwrap(); + + let result = read_kiro_credential_info(temp_file.to_str().unwrap()); + assert!(result.is_err()); + assert!(result.unwrap_err().contains("解析凭证文件失败")); + + // 清理 + let _ = std::fs::remove_file(&temp_file); + } +} diff --git a/src-tauri/src/lib.rs b/src-tauri/src/lib.rs index e937d299a..b738d9f2c 100644 --- a/src-tauri/src/lib.rs +++ b/src-tauri/src/lib.rs @@ -1648,6 +1648,8 @@ pub fn run() { commands::injection_cmd::add_injection_rule, commands::injection_cmd::remove_injection_rule, commands::injection_cmd::update_injection_rule, + // Usage commands + commands::usage_cmd::get_kiro_usage, ]) .run(tauri::generate_context!()) .expect("error while running tauri application"); diff --git a/src-tauri/src/providers/kiro.rs b/src-tauri/src/providers/kiro.rs index 0e1b7016a..27a2e9c62 100644 --- a/src-tauri/src/providers/kiro.rs +++ b/src-tauri/src/providers/kiro.rs @@ -6,34 +6,104 @@ use serde::{Deserialize, Serialize}; use std::error::Error; use std::path::PathBuf; -/// 生成设备指纹 (MAC 地址的 SHA256) +/// 生成设备指纹 (Machine ID 的 SHA256) +/// +/// 与 Kiro IDE 保持一致的指纹生成方式(参考 Kir-Manager): +/// - macOS: 使用 IOPlatformUUID(硬件级别唯一标识) +/// - Linux: 使用 /etc/machine-id +/// - Windows: 使用 WMI 获取系统 UUID +/// +/// 最终返回 SHA256 哈希后的 64 字符十六进制字符串 fn get_device_fingerprint() -> String { + use sha2::{Digest, Sha256}; + + let raw_id = + get_raw_machine_id().unwrap_or_else(|| "00000000-0000-0000-0000-000000000000".to_string()); + + // 使用 SHA256 生成 64 字符的十六进制指纹 + let mut hasher = Sha256::new(); + hasher.update(raw_id.as_bytes()); + let result = hasher.finalize(); + format!("{:x}", result) +} + +/// 获取原始 Machine ID(未哈希) +fn get_raw_machine_id() -> Option { use std::process::Command; - // 尝试获取 MAC 地址 - let mac = if cfg!(target_os = "macos") { - Command::new("ifconfig") + if cfg!(target_os = "macos") { + // macOS: 使用 ioreg 获取 IOPlatformUUID + Command::new("ioreg") + .args(["-rd1", "-c", "IOPlatformExpertDevice"]) .output() .ok() .and_then(|o| String::from_utf8(o.stdout).ok()) .and_then(|s| { s.lines() - .find(|l| l.contains("ether ")) - .and_then(|l| l.split_whitespace().nth(1)) - .map(|s| s.to_string()) + .find(|l| l.contains("IOPlatformUUID")) + .and_then(|l| l.split('=').nth(1)) + .map(|s| s.trim().trim_matches('"').to_lowercase()) + }) + } else if cfg!(target_os = "linux") { + // Linux: 读取 /etc/machine-id 或 /var/lib/dbus/machine-id + std::fs::read_to_string("/etc/machine-id") + .or_else(|_| std::fs::read_to_string("/var/lib/dbus/machine-id")) + .ok() + .map(|s| s.trim().to_lowercase()) + } else if cfg!(target_os = "windows") { + // Windows: 使用 wmic 获取系统 UUID + Command::new("wmic") + .args(["csproduct", "get", "UUID"]) + .output() + .ok() + .and_then(|o| String::from_utf8(o.stdout).ok()) + .and_then(|s| { + s.lines() + .skip(1) // 跳过表头 + .find(|l| !l.trim().is_empty()) + .map(|s| s.trim().to_lowercase()) }) } else { None - }; + } +} - let mac = mac.unwrap_or_else(|| "00:00:00:00:00:00".to_string()); +/// 获取 Kiro IDE 版本号 +/// +/// 尝试从 Kiro.app 的 Info.plist 读取实际版本,失败时使用默认值 +fn get_kiro_version() -> String { + use std::process::Command; - // SHA256 hash - use std::collections::hash_map::DefaultHasher; - use std::hash::{Hash, Hasher}; - let mut hasher = DefaultHasher::new(); - mac.hash(&mut hasher); - format!("{:016x}{:016x}", hasher.finish(), hasher.finish()) + if cfg!(target_os = "macos") { + // 尝试从 Kiro.app 读取版本 + let kiro_paths = [ + "/Applications/Kiro.app/Contents/Info.plist", + // 用户目录下的安装 + &format!( + "{}/Applications/Kiro.app/Contents/Info.plist", + dirs::home_dir() + .map(|p| p.to_string_lossy().to_string()) + .unwrap_or_default() + ), + ]; + + for plist_path in &kiro_paths { + if let Ok(output) = Command::new("defaults") + .args(["read", plist_path, "CFBundleShortVersionString"]) + .output() + { + if let Ok(version) = String::from_utf8(output.stdout) { + let version = version.trim(); + if !version.is_empty() { + return version.to_string(); + } + } + } + } + } + + // 默认版本号 + "0.1.25".to_string() } #[derive(Debug, Clone, Serialize, Deserialize)] @@ -472,6 +542,25 @@ impl KiroProvider { return Err("refresh_token 为空。\n💡 解决方案:\n1. 检查凭证文件是否损坏\n2. 重新生成 OAuth 凭证".to_string()); } + let token_len = refresh_token.len(); + + // 检测 refreshToken 是否被截断 + // 正常的 refreshToken 长度应该在 500+ 字符 + let is_truncated = + token_len < 100 || refresh_token.ends_with("...") || refresh_token.contains("..."); + + if is_truncated { + tracing::error!( + "[KIRO] 检测到 refreshToken 被截断!长度: {}, 内容: {}...", + token_len, + &refresh_token[..std::cmp::min(30, token_len)] + ); + return Err(format!( + "refreshToken 已被截断(长度: {} 字符)。\n\n⚠️ 这通常是 Kiro IDE 为了防止凭证被第三方工具使用而故意截断的。\n\n💡 解决方案:\n1. 使用 Kir-Manager 工具获取完整的凭证\n2. 或者使用其他方式获取未截断的凭证文件\n3. 正常的 refreshToken 长度应该在 500+ 字符", + token_len + )); + } + // 检查是否看起来像有效的 token(简单的长度和格式检查) if refresh_token.len() < 10 { return Err("refresh_token 格式异常(长度过短)。\n💡 解决方案:\n1. 凭证文件可能已损坏\n2. 重新获取 OAuth 凭证".to_string()); @@ -480,28 +569,20 @@ impl KiroProvider { Ok(()) } - /// 检测最佳的认证方式 - /// 优先使用 IdC(如果有完整配置),否则回退到 social ��证 + /// 检测认证方式 + /// + /// 注意:不再自动降级!IdC 和 Social 的 refreshToken 不兼容, + /// 不能将 IdC 的 refreshToken 用于 Social 端点。 pub fn detect_auth_method(&self) -> String { - // 检查当前设置的认证方式 - let current_auth = self.credentials.auth_method.as_deref().unwrap_or("social"); + // 直接返回配置中的认证方式,不做降级 + let auth_method = self.credentials.auth_method.as_deref().unwrap_or("social"); + tracing::debug!("[KIRO] 使用配置的认证方式: {}", auth_method); + auth_method.to_lowercase() + } - // 如果当前是 IdC 方式,检查是否有完整的 IdC 配置 - if current_auth.to_lowercase() == "idc" { - if self.credentials.client_id.is_some() && self.credentials.client_secret.is_some() { - // IdC 配置完整,继续使用 IdC - tracing::debug!("[KIRO] IdC 配置完整,使用 IdC 认证"); - "idc".to_string() - } else { - // IdC 配置不完整,降级到 social - tracing::warn!("[KIRO] IdC 配置不完整(缺少 client_id 或 client_secret),自动降级到 social 认证"); - "social".to_string() - } - } else { - // 默认或已设置为 social - tracing::debug!("[KIRO] 使用 social 认证"); - "social".to_string() - } + /// 检查 IdC 认证配置是否完整 + pub fn is_idc_config_complete(&self) -> bool { + self.credentials.client_id.is_some() && self.credentials.client_secret.is_some() } /// 更新认证方式到凭证中(仅在内存中,需要调用 save_credentials 持久化) @@ -533,22 +614,28 @@ impl KiroProvider { .ok_or("No refresh token")? .clone(); - // 使用智能检测的认证方式,而不是直接使用配置中的方式 - let detected_auth_method = self.detect_auth_method(); - tracing::info!("[KIRO] 检测到的认证方式: {}", detected_auth_method); + // 获取认证方式 + let auth_method = self.detect_auth_method(); + tracing::info!("[KIRO] 使用认证方式: {}", auth_method); - // 如果检测到的方式与配置中的不同,更新配置 - let current_auth = self.credentials.auth_method.as_deref().unwrap_or("social"); - if current_auth != detected_auth_method { - tracing::info!( - "[KIRO] 认证方式从 {} 切换到 {}", - current_auth, - detected_auth_method - ); - self.set_auth_method(&detected_auth_method); + // 检查 IdC 认证是否有完整配置 + if auth_method == "idc" && !self.is_idc_config_complete() { + let has_client_id = self.credentials.client_id.is_some(); + let has_client_secret = self.credentials.client_secret.is_some(); + + // IdC 认证缺少必要凭证,返回明确错误(不能降级到 social,因为 refreshToken 不兼容) + let missing = match (has_client_id, has_client_secret) { + (false, false) => "clientId 和 clientSecret", + (false, true) => "clientId", + (true, false) => "clientSecret", + _ => unreachable!(), + }; + + return Err(format!( + "IdC 认证配置不完整:缺少 {}。\n\n⚠️ 注意:IdC 凭证的 refreshToken 无法用于 Social 认证,必须提供完整的 IdC 配置。\n\n💡 解决方案:\n1. 删除当前凭证\n2. 重新从 Kiro IDE 获取最新的凭证文件(确保完成完整的 SSO 登录流程)\n3. 确保 ~/.aws/sso/cache/ 目录下有对应的 clientIdHash 文件\n4. 重新添加凭证到 ProxyCast", + missing + ).into()); } - - let auth_method = detected_auth_method.to_lowercase(); let refresh_url = self.get_refresh_url(); tracing::debug!( @@ -562,8 +649,12 @@ impl KiroProvider { self.credentials.client_secret.is_some() ); + // 获取设备指纹和版本号(用于 Social 认证的 User-Agent) + let device_fp = get_device_fingerprint(); + let kiro_version = get_kiro_version(); + let resp = if auth_method == "idc" { - // IdC 认证使用 JSON 格式(参考 AIClient-2-API 实现) + // IdC 认证使用 JSON 格式(参考 Kir-Manager 实现) let client_id = self .credentials .client_id @@ -575,7 +666,7 @@ impl KiroProvider { .as_ref() .ok_or("IdC 认证配置错误:缺少 client_secret。建议删除后重新添加 OAuth 凭证")?; - // 使用 JSON 格式发送请求(与 AIClient-2-API 保持一致) + // 使用 JSON 格式发送请求(与 Kir-Manager 保持一致) let body = serde_json::json!({ "refreshToken": &refresh_token, "clientId": client_id, @@ -585,20 +676,37 @@ impl KiroProvider { tracing::debug!("[KIRO] IdC 刷新请求体已构建"); + // IdC 认证的 Headers(参考 Kir-Manager) self.client .post(&refresh_url) .header("Content-Type", "application/json") - .header("Accept", "application/json") + .header("Host", "oidc.us-east-1.amazonaws.com") + .header( + "x-amz-user-agent", + "aws-sdk-js/3.738.0 ua/2.1 os/other lang/js api/sso-oidc#3.738.0 m/E KiroIDE", + ) + .header("User-Agent", "node") + .header("Accept", "*/*") + .header("Connection", "keep-alive") .json(&body) .send() .await? } else { - // Social 认证使用简单的 JSON 格式 + // Social 认证使用简单的 JSON 格式(参考 Kir-Manager) let body = serde_json::json!({ "refreshToken": &refresh_token }); + + // Social 认证的 Headers(参考 Kir-Manager) self.client .post(&refresh_url) + .header( + "User-Agent", + format!("KiroIDE-{}-{}", kiro_version, device_fp), + ) + .header("Accept", "application/json, text/plain, */*") + .header("Accept-Encoding", "br, gzip, deflate") .header("Content-Type", "application/json") - .header("Accept", "application/json") + .header("Accept-Language", "*") + .header("Sec-Fetch-Mode", "cors") .json(&body) .send() .await? @@ -810,7 +918,7 @@ impl KiroProvider { // 生成设备指纹用于伪装 Kiro IDE let device_fp = get_device_fingerprint(); - let kiro_version = "0.1.25"; + let kiro_version = get_kiro_version(); let resp = self .client diff --git a/src-tauri/src/services/mod.rs b/src-tauri/src/services/mod.rs index 90b1adac9..07768fb05 100644 --- a/src-tauri/src/services/mod.rs +++ b/src-tauri/src/services/mod.rs @@ -7,3 +7,4 @@ pub mod provider_pool_service; pub mod skill_service; pub mod switch; pub mod token_cache_service; +pub mod usage_service; diff --git a/src-tauri/src/services/token_cache_service.rs b/src-tauri/src/services/token_cache_service.rs index af311c4b1..1609c967c 100644 --- a/src-tauri/src/services/token_cache_service.rs +++ b/src-tauri/src/services/token_cache_service.rs @@ -43,6 +43,7 @@ impl TokenCacheService { /// 1. 检查数据库缓存是否有效 /// 2. 如果缓存有效且未过期,直接返回 /// 3. 如果缓存无效或即将过期,执行刷新 + /// 4. 如果刷新失败(如 refreshToken 被截断),尝试使用源文件中的 accessToken pub async fn get_valid_token(&self, db: &DbConnection, uuid: &str) -> Result { // 首先检查缓存 let cached = { @@ -65,7 +66,69 @@ impl TokenCacheService { } // 需要刷新(无缓存、已过期或即将过期) - self.refresh_and_cache(db, uuid, false).await + match self.refresh_and_cache(db, uuid, false).await { + Ok(token) => Ok(token), + Err(refresh_error) => { + // 刷新失败时,检查是否是因为 refreshToken 被截断 + // 如果是,尝试直接使用源文件中的 accessToken(可能仍然有效) + if refresh_error.contains("截断") || refresh_error.contains("truncated") { + tracing::warn!( + "[TOKEN_CACHE] refreshToken 被截断,尝试使用源文件中的 accessToken: {}", + &uuid[..8] + ); + + // 获取凭证信息 + let credential = { + let conn = db.lock().map_err(|e| e.to_string())?; + ProviderPoolDao::get_by_uuid(&conn, uuid) + .map_err(|e| e.to_string())? + .ok_or_else(|| format!("Credential not found: {}", uuid))? + }; + + // 尝试从源文件读取 accessToken + match self.read_token_from_source(&credential).await { + Ok(token_info) => { + if let Some(token) = token_info.access_token { + tracing::info!( + "[TOKEN_CACHE] 使用源文件中的 accessToken(可能已过期): {}", + &uuid[..8] + ); + // 注意:这个 token 可能已过期,但至少可以尝试使用 + // 缓存这个 token(但不设置过期时间,因为我们不知道它何时过期) + let cache_info = CachedTokenInfo { + access_token: Some(token.clone()), + refresh_token: token_info.refresh_token, + expiry_time: None, // 不知道过期时间 + last_refresh: Some(Utc::now()), + refresh_error_count: 1, + last_refresh_error: Some(format!( + "refreshToken 被截断,使用源文件 accessToken: {}", + refresh_error + )), + }; + + // 缓存到数据库 + if let Ok(conn) = db.lock() { + let _ = ProviderPoolDao::update_token_cache( + &conn, + uuid, + &cache_info, + ); + } + + return Ok(token); + } + } + Err(e) => { + tracing::error!("[TOKEN_CACHE] 无法从源文件读取 accessToken: {}", e); + } + } + } + + // 返回原始刷新错误 + Err(refresh_error) + } + } } /// 刷新 Token 并缓存到数据库 diff --git a/src-tauri/src/services/usage_service.rs b/src-tauri/src/services/usage_service.rs new file mode 100644 index 000000000..d4270d7ab --- /dev/null +++ b/src-tauri/src/services/usage_service.rs @@ -0,0 +1,915 @@ +//! Usage Service - Kiro 用量查询服务 +//! +//! 通过调用 AWS Q 的 getUsageLimits API 获取用户的用量信息。 +//! 参考 Kir-Manager 项目的 usage/usage.go 实现。 + +use reqwest::header::{HeaderMap, HeaderValue, USER_AGENT}; +use serde::{Deserialize, Serialize}; +use std::error::Error; +use uuid::Uuid; + +// ============================================================================ +// 常量定义 +// ============================================================================ + +/// API 端点 +pub const USAGE_LIMITS_URL: &str = "https://q.us-east-1.amazonaws.com/getUsageLimits"; + +/// Query 参数 +pub const ORIGIN_PARAM: &str = "AI_EDITOR"; +pub const RESOURCE_TYPE_PARAM: &str = "AGENTIC_REQUEST"; + +/// HTTP 请求超时(秒) +pub const HTTP_TIMEOUT_SECS: u64 = 10; + +// ============================================================================ +// API Response 数据模型 +// ============================================================================ + +/// API 响应结构 +#[derive(Debug, Clone, Deserialize, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct UsageLimitsResponse { + pub subscription_info: SubscriptionInfo, + pub usage_breakdown_list: Vec, +} + +/// 订阅信息结构 +#[derive(Debug, Clone, Deserialize, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct SubscriptionInfo { + pub subscription_title: String, + #[serde(rename = "type")] + pub subscription_type: String, +} + +/// 用量明细结构 +#[derive(Debug, Clone, Deserialize, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct UsageBreakdown { + pub usage_limit_with_precision: f64, + pub current_usage_with_precision: f64, + pub display_name: String, + #[serde(default)] + pub free_trial_info: Option, + #[serde(default)] + pub bonuses: Option>, +} + +/// 免费试用信息 +#[derive(Debug, Clone, Deserialize, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct FreeTrialInfo { + pub usage_limit_with_precision: f64, + pub current_usage_with_precision: f64, + pub free_trial_status: String, +} + +/// 奖励额度 +#[derive(Debug, Clone, Deserialize, Serialize)] +#[serde(rename_all = "camelCase")] +pub struct Bonus { + pub bonus_code: String, + pub usage_limit: f64, + pub current_usage: f64, + pub status: String, +} + +// ============================================================================ +// 计算结果数据模型 +// ============================================================================ + +/// 计算后的用量信息 +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +#[serde(rename_all = "camelCase")] +pub struct UsageInfo { + /// 订阅类型名称 + pub subscription_title: String, + /// 总额度 + pub usage_limit: f64, + /// 已使用 + pub current_usage: f64, + /// 余额 = usage_limit - current_usage + pub balance: f64, + /// 余额低于 20% + pub is_low_balance: bool, +} + +impl UsageInfo { + /// 创建空的 UsageInfo + pub fn empty() -> Self { + Self::default() + } +} + +// ============================================================================ +// URL 构造函数 +// ============================================================================ + +/// 构造 API 请求 URL +/// +/// **Property 4: Social Auth URL Construction** +/// **Property 5: IdC Auth URL Construction** +/// **Validates: Requirements 2.1, 2.2** +/// +/// - Social 认证: 包含 profileArn 参数 +/// - IdC 认证: 不包含 profileArn 参数 +pub fn build_usage_api_url( + auth_method: &str, + profile_arn: Option<&str>, +) -> Result> { + let mut url = url::Url::parse(USAGE_LIMITS_URL)?; + + { + let mut query = url.query_pairs_mut(); + query.append_pair("origin", ORIGIN_PARAM); + query.append_pair("resourceType", RESOURCE_TYPE_PARAM); + + // Property 4: Social Auth URL Construction + // 只有 social 类型才加入 profileArn + if auth_method == "social" { + match profile_arn { + Some(arn) if !arn.is_empty() => { + query.append_pair("profileArn", arn); + } + _ => { + return Err("social auth requires profileArn".into()); + } + } + } + // Property 5: IdC Auth URL Construction + // IdC 类型不包含 profileArn + } + + Ok(url.to_string()) +} + +// ============================================================================ +// 请求头构造函数 +// ============================================================================ + +/// 构造 API 请求头 +/// +/// **Property 6: User-Agent Header Format** +/// **Validates: Requirements 4.1, 4.2, 4.3** +/// +/// Headers: +/// - User-Agent: aws-sdk-js/1.0.0 ua/2.1 os/{os} lang/rust api/codewhispererruntime#1.0.0 m/N,E KiroIDE-{version}-{machineId} +/// - x-amz-user-agent: aws-sdk-js/1.0.0 KiroIDE-{version}-{machineId} +/// - amz-sdk-invocation-id: UUID +/// - amz-sdk-request: attempt=1; max=1 +pub fn build_request_headers( + access_token: &str, + kiro_version: &str, + machine_id: &str, +) -> Result> { + let mut headers = HeaderMap::new(); + + // Authorization header + let auth_value = format!("Bearer {}", access_token); + headers.insert("Authorization", HeaderValue::from_str(&auth_value)?); + + // User-Agent header + // 格式: aws-sdk-js/1.0.0 ua/2.1 os/{os} lang/rust api/codewhispererruntime#1.0.0 m/N,E KiroIDE-{version}-{machineId} + let os_name = std::env::consts::OS; + let user_agent = format!( + "aws-sdk-js/1.0.0 ua/2.1 os/{} lang/rust api/codewhispererruntime#1.0.0 m/N,E KiroIDE-{}-{}", + os_name, kiro_version, machine_id + ); + headers.insert(USER_AGENT, HeaderValue::from_str(&user_agent)?); + + // x-amz-user-agent header + let x_amz_user_agent = format!("aws-sdk-js/1.0.0 KiroIDE-{}-{}", kiro_version, machine_id); + headers.insert( + "x-amz-user-agent", + HeaderValue::from_str(&x_amz_user_agent)?, + ); + + // amz-sdk-invocation-id: 每次请求随机生成 UUID + headers.insert( + "amz-sdk-invocation-id", + HeaderValue::from_str(&Uuid::new_v4().to_string())?, + ); + + // amz-sdk-request header + headers.insert( + "amz-sdk-request", + HeaderValue::from_static("attempt=1; max=1"), + ); + + // Connection header + headers.insert("Connection", HeaderValue::from_static("close")); + + Ok(headers) +} + +/// 构造 User-Agent 字符串(用于测试) +pub fn build_user_agent(kiro_version: &str, machine_id: &str) -> String { + let os_name = std::env::consts::OS; + format!( + "aws-sdk-js/1.0.0 ua/2.1 os/{} lang/rust api/codewhispererruntime#1.0.0 m/N,E KiroIDE-{}-{}", + os_name, kiro_version, machine_id + ) +} + +/// 构造 x-amz-user-agent 字符串(用于测试) +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) +} + +// ============================================================================ +// API 调用函数 +// ============================================================================ + +/// 调用 AWS Q getUsageLimits API 获取用量信息 +/// +/// **Validates: Requirements 1.1, 4.4** +/// +/// # Arguments +/// * `access_token` - Bearer token +/// * `auth_method` - 认证方式 ("social" 或 "idc") +/// * `profile_arn` - Social 认证需要的 profileArn +/// * `machine_id` - 设备 ID (SHA256 哈希) +/// * `kiro_version` - Kiro 版本号 +/// +/// # Returns +/// * `Ok(UsageInfo)` - 成功时返回计算后的用量信息 +/// * `Err` - 失败时返回错误 +pub async fn get_usage_limits( + access_token: &str, + auth_method: &str, + profile_arn: Option<&str>, + machine_id: &str, + kiro_version: &str, +) -> Result> { + // 验证参数 + if access_token.is_empty() { + return Err("invalid token: missing accessToken".into()); + } + + if machine_id.is_empty() { + return Err("invalid machineID: empty".into()); + } + + // 构造 URL + let url = build_usage_api_url(auth_method, profile_arn)?; + + // 构造请求头 + let headers = build_request_headers(access_token, kiro_version, machine_id)?; + + // 创建 HTTP 客户端(带超时) + let client = reqwest::Client::builder() + .timeout(std::time::Duration::from_secs(HTTP_TIMEOUT_SECS)) + .build()?; + + // 发送请求 + let response = client.get(&url).headers(headers).send().await?; + + // 检查状态码 + let status = response.status(); + if !status.is_success() { + let body = response.text().await.unwrap_or_default(); + return Err(format!("API request failed with status {}: {}", status, body).into()); + } + + // 解析响应 + let usage_response: UsageLimitsResponse = response.json().await?; + + // 计算余额并返回 + Ok(calculate_balance(&usage_response)) +} + +/// 安全地调用 API 获取用量信息 +/// +/// **Property 3: Error Handling Graceful Degradation** +/// **Validates: Requirements 1.4** +/// +/// 当发生任何错误时,返回空的 UsageInfo 而非 panic +pub async fn get_usage_limits_safe( + access_token: &str, + auth_method: &str, + profile_arn: Option<&str>, + machine_id: &str, + kiro_version: &str, +) -> UsageInfo { + match get_usage_limits( + access_token, + auth_method, + profile_arn, + machine_id, + kiro_version, + ) + .await + { + Ok(info) => info, + Err(e) => { + tracing::warn!("Failed to get usage limits: {}", e); + UsageInfo::empty() + } + } +} + +// ============================================================================ +// 余额计算函数 +// ============================================================================ + +/// 低余额阈值 (20%) +pub const LOW_BALANCE_THRESHOLD: f64 = 0.2; + +/// 从 API 响应计算余额 +/// +/// **Property 1: Balance Calculation Correctness** +/// **Validates: Requirements 1.2** +/// +/// 计算逻辑: +/// - 总额度 = Σ(usage_limit_with_precision + free_trial_info?.usage_limit_with_precision + Σ(bonuses[].usage_limit)) +/// - 总使用 = Σ(current_usage_with_precision + free_trial_info?.current_usage_with_precision + Σ(bonuses[].current_usage)) +/// - 余额 = 总额度 - 总使用 +pub fn calculate_balance(response: &UsageLimitsResponse) -> UsageInfo { + calculate_balance_with_threshold(response, LOW_BALANCE_THRESHOLD) +} + +/// 从 API 响应计算余额(使用指定阈值) +/// +/// threshold: 低余额阈值(0.0 ~ 1.0),例如 0.2 表示余额低于 20% 时为低余额 +pub fn calculate_balance_with_threshold( + response: &UsageLimitsResponse, + threshold: f64, +) -> UsageInfo { + let mut total_usage_limit = 0.0; + let mut total_current_usage = 0.0; + + for breakdown in &response.usage_breakdown_list { + // 基本额度 + total_usage_limit += breakdown.usage_limit_with_precision; + total_current_usage += breakdown.current_usage_with_precision; + + // 免费试用额度(如果存在) + if let Some(ref free_trial) = breakdown.free_trial_info { + total_usage_limit += free_trial.usage_limit_with_precision; + total_current_usage += free_trial.current_usage_with_precision; + } + + // 奖励额度(如果存在) + if let Some(ref bonuses) = breakdown.bonuses { + for bonus in bonuses { + total_usage_limit += bonus.usage_limit; + total_current_usage += bonus.current_usage; + } + } + } + + let balance = total_usage_limit - total_current_usage; + + // Property 2: Low Balance Detection + // Validates: Requirements 1.3 + // is_low_balance = (balance / total_usage_limit) < threshold + let is_low_balance = if total_usage_limit > 0.0 { + (balance / total_usage_limit) < threshold + } else { + false + }; + + UsageInfo { + subscription_title: response.subscription_info.subscription_title.clone(), + usage_limit: total_usage_limit, + current_usage: total_current_usage, + balance, + is_low_balance, + } +} + +// ============================================================================ +// 测试模块 +// ============================================================================ + +#[cfg(test)] +mod tests { + use super::*; + use proptest::prelude::*; + use urlencoding; + + // ======================================================================== + // Arbitrary 生成器 + // ======================================================================== + + /// 生成有效的 Bonus + fn arb_bonus() -> impl Strategy { + ( + "[a-zA-Z0-9]{4,10}", // bonus_code + 0.0..1000.0f64, // usage_limit + 0.0..1000.0f64, // current_usage + prop_oneof!["ACTIVE", "EXPIRED", "PENDING"], + ) + .prop_map(|(bonus_code, usage_limit, current_usage, status)| Bonus { + bonus_code, + usage_limit, + current_usage, + status: status.to_string(), + }) + } + + /// 生成有效的 FreeTrialInfo + fn arb_free_trial_info() -> impl Strategy { + ( + 0.0..1000.0f64, // usage_limit_with_precision + 0.0..1000.0f64, // current_usage_with_precision + prop_oneof!["ACTIVE", "EXPIRED"], + ) + .prop_map(|(usage_limit, current_usage, status)| FreeTrialInfo { + usage_limit_with_precision: usage_limit, + current_usage_with_precision: current_usage, + free_trial_status: status.to_string(), + }) + } + + /// 生成有效的 UsageBreakdown + fn arb_usage_breakdown() -> impl Strategy { + ( + 0.0..1000.0f64, // usage_limit_with_precision + 0.0..1000.0f64, // current_usage_with_precision + "[a-zA-Z ]{5,20}", // display_name + prop::option::of(arb_free_trial_info()), // free_trial_info + prop::option::of(prop::collection::vec(arb_bonus(), 0..3)), // bonuses + ) + .prop_map( + |(usage_limit, current_usage, display_name, free_trial_info, bonuses)| { + UsageBreakdown { + usage_limit_with_precision: usage_limit, + current_usage_with_precision: current_usage, + display_name, + free_trial_info, + bonuses, + } + }, + ) + } + + /// 生成有效的 SubscriptionInfo + fn arb_subscription_info() -> impl Strategy { + ( + prop_oneof!["Free Tier", "Pro", "Enterprise"], + prop_oneof!["FREE", "PAID", "TRIAL"], + ) + .prop_map(|(title, sub_type)| SubscriptionInfo { + subscription_title: title.to_string(), + subscription_type: sub_type.to_string(), + }) + } + + /// 生成有效的 UsageLimitsResponse + fn arb_usage_limits_response() -> impl Strategy { + ( + arb_subscription_info(), + prop::collection::vec(arb_usage_breakdown(), 1..5), + ) + .prop_map( + |(subscription_info, usage_breakdown_list)| UsageLimitsResponse { + subscription_info, + usage_breakdown_list, + }, + ) + } + + // ======================================================================== + // Property 1: Balance Calculation Correctness + // **Feature: kiro-usage-api, Property 1: Balance Calculation Correctness** + // **Validates: Requirements 1.2** + // ======================================================================== + + proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// Property 1: 余额计算正确性 + /// + /// *For any* UsageLimitsResponse with valid usage breakdown data, + /// the calculated balance SHALL equal (total_usage_limit - total_current_usage), + /// where totals include base amounts, free trial amounts, and bonus amounts. + #[test] + fn prop_balance_calculation_correctness(response in arb_usage_limits_response()) { + let result = calculate_balance(&response); + + // 手动计算期望值 + let mut expected_limit = 0.0; + let mut expected_usage = 0.0; + + for breakdown in &response.usage_breakdown_list { + expected_limit += breakdown.usage_limit_with_precision; + expected_usage += breakdown.current_usage_with_precision; + + if let Some(ref ft) = breakdown.free_trial_info { + expected_limit += ft.usage_limit_with_precision; + expected_usage += ft.current_usage_with_precision; + } + + if let Some(ref bonuses) = breakdown.bonuses { + for bonus in bonuses { + expected_limit += bonus.usage_limit; + expected_usage += bonus.current_usage; + } + } + } + + let expected_balance = expected_limit - expected_usage; + + // 使用近似比较(浮点数精度问题) + let epsilon = 1e-10; + prop_assert!((result.usage_limit - expected_limit).abs() < epsilon, + "usage_limit mismatch: got {}, expected {}", result.usage_limit, expected_limit); + prop_assert!((result.current_usage - expected_usage).abs() < epsilon, + "current_usage mismatch: got {}, expected {}", result.current_usage, expected_usage); + prop_assert!((result.balance - expected_balance).abs() < epsilon, + "balance mismatch: got {}, expected {}", result.balance, expected_balance); + } + } + + // ======================================================================== + // Property 2: Low Balance Detection + // **Feature: kiro-usage-api, Property 2: Low Balance Detection** + // **Validates: Requirements 1.3** + // ======================================================================== + + /// 生成有效的 UsageInfo(直接生成,用于测试低余额检测) + fn arb_usage_info() -> impl Strategy { + // 生成 usage_limit 和 balance,确保 balance <= usage_limit + (0.01..1000.0f64).prop_flat_map(|usage_limit| { + // balance 可以是 0 到 usage_limit 之间的任意值 + (Just(usage_limit), 0.0..=usage_limit) + }) + } + + proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// Property 2: 低余额检测 + /// + /// *For any* UsageInfo where balance/usage_limit < 0.2 (and usage_limit > 0), + /// is_low_balance SHALL be true; otherwise it SHALL be false. + #[test] + fn prop_low_balance_detection((usage_limit, balance) in arb_usage_info()) { + // 构造一个简单的响应来测试低余额检测 + let current_usage = usage_limit - balance; + let response = UsageLimitsResponse { + subscription_info: SubscriptionInfo { + subscription_title: "Test".to_string(), + subscription_type: "FREE".to_string(), + }, + usage_breakdown_list: vec![UsageBreakdown { + usage_limit_with_precision: usage_limit, + current_usage_with_precision: current_usage, + display_name: "Test".to_string(), + free_trial_info: None, + bonuses: None, + }], + }; + + let result = calculate_balance(&response); + + // 计算期望的 is_low_balance + let ratio = balance / usage_limit; + let expected_low_balance = ratio < LOW_BALANCE_THRESHOLD; + + prop_assert_eq!( + result.is_low_balance, + expected_low_balance, + "is_low_balance mismatch: got {}, expected {} (ratio: {}, threshold: {})", + result.is_low_balance, + expected_low_balance, + ratio, + LOW_BALANCE_THRESHOLD + ); + } + + /// Property 2 边界情况: 当 usage_limit 为 0 时,is_low_balance 应为 false + #[test] + fn prop_low_balance_zero_limit(current_usage in 0.0..100.0f64) { + let response = UsageLimitsResponse { + subscription_info: SubscriptionInfo { + subscription_title: "Test".to_string(), + subscription_type: "FREE".to_string(), + }, + usage_breakdown_list: vec![UsageBreakdown { + usage_limit_with_precision: 0.0, + current_usage_with_precision: current_usage, + display_name: "Test".to_string(), + free_trial_info: None, + bonuses: None, + }], + }; + + let result = calculate_balance(&response); + + // 当 usage_limit 为 0 时,is_low_balance 应为 false(避免除零) + prop_assert!(!result.is_low_balance, + "is_low_balance should be false when usage_limit is 0"); + } + } + + // ======================================================================== + // Property 4: Social Auth URL Construction + // **Feature: kiro-usage-api, Property 4: Social Auth URL Construction** + // **Validates: Requirements 2.1** + // ======================================================================== + + /// 生成有效的 profileArn + fn arb_profile_arn() -> impl Strategy { + "[a-zA-Z0-9:/-]{10,50}".prop_map(|s| format!("arn:aws:iam::{}", s)) + } + + proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// Property 4: Social Auth URL 构造 + /// + /// *For any* request with auth_method="social" and a valid profile_arn, + /// the request URL SHALL contain `profileArn={profile_arn}` as a query parameter. + #[test] + fn prop_social_auth_url_contains_profile_arn(profile_arn in arb_profile_arn()) { + let url = build_usage_api_url("social", Some(&profile_arn)).unwrap(); + + // URL 应该包含 profileArn 参数 + prop_assert!(url.contains("profileArn="), + "Social auth URL should contain profileArn parameter, got: {}", url); + + // URL 应该包含编码后的 profile_arn 值 + let encoded_arn = urlencoding::encode(&profile_arn); + prop_assert!(url.contains(&encoded_arn.to_string()), + "Social auth URL should contain encoded profileArn value '{}', got: {}", encoded_arn, url); + + // URL 应该包含基本参数 + prop_assert!(url.contains("origin=AI_EDITOR"), + "URL should contain origin parameter, got: {}", url); + prop_assert!(url.contains("resourceType=AGENTIC_REQUEST"), + "URL should contain resourceType parameter, got: {}", url); + } + + /// Property 4 边界情况: Social auth 缺少 profileArn 应返回错误 + #[test] + fn prop_social_auth_requires_profile_arn(_dummy in 0..10i32) { + // 测试 None + let result = build_usage_api_url("social", None); + prop_assert!(result.is_err(), "Social auth without profileArn should fail"); + + // 测试空字符串 + let result = build_usage_api_url("social", Some("")); + prop_assert!(result.is_err(), "Social auth with empty profileArn should fail"); + } + } + + // ======================================================================== + // Property 5: IdC Auth URL Construction + // **Feature: kiro-usage-api, Property 5: IdC Auth URL Construction** + // **Validates: Requirements 2.2** + // ======================================================================== + + proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// Property 5: IdC Auth URL 构造 + /// + /// *For any* request with auth_method="idc", + /// the request URL SHALL NOT contain `profileArn` as a query parameter. + #[test] + fn prop_idc_auth_url_no_profile_arn(profile_arn in arb_profile_arn()) { + // 即使提供了 profile_arn,IdC 认证也不应该包含它 + let url = build_usage_api_url("idc", Some(&profile_arn)).unwrap(); + + // URL 不应该包含 profileArn 参数 + prop_assert!(!url.contains("profileArn"), + "IdC auth URL should NOT contain profileArn parameter, got: {}", url); + + // URL 应该包含基本参数 + prop_assert!(url.contains("origin=AI_EDITOR"), + "URL should contain origin parameter, got: {}", url); + prop_assert!(url.contains("resourceType=AGENTIC_REQUEST"), + "URL should contain resourceType parameter, got: {}", url); + } + + /// Property 5: IdC auth 不需要 profileArn + #[test] + fn prop_idc_auth_works_without_profile_arn(_dummy in 0..10i32) { + // IdC 认证不需要 profileArn + let result = build_usage_api_url("idc", None); + prop_assert!(result.is_ok(), "IdC auth without profileArn should succeed"); + + let url = result.unwrap(); + prop_assert!(!url.contains("profileArn"), + "IdC auth URL should NOT contain profileArn parameter, got: {}", url); + } + } + + // ======================================================================== + // Property 3: Error Handling Graceful Degradation + // **Feature: kiro-usage-api, Property 3: Error Handling Graceful Degradation** + // **Validates: Requirements 1.4** + // ======================================================================== + + /// 生成各种错误输入场景 + #[derive(Debug, Clone)] + enum ErrorScenario { + EmptyToken, + EmptyMachineId, + MissingProfileArn, + EmptyProfileArn, + InvalidToken, + } + + /// 生成错误场景的策略 + fn arb_error_scenario() -> impl Strategy { + prop_oneof![ + Just(ErrorScenario::EmptyToken), + Just(ErrorScenario::EmptyMachineId), + Just(ErrorScenario::MissingProfileArn), + Just(ErrorScenario::EmptyProfileArn), + Just(ErrorScenario::InvalidToken), + ] + } + + /// 生成随机的有效 token(用于非空 token 场景) + fn arb_valid_token() -> impl Strategy { + "[a-zA-Z0-9]{20,50}".prop_map(|s| s) + } + + /// 生成随机的有效 machine_id(用于非空 machine_id 场景) + fn arb_valid_machine_id() -> impl Strategy { + "[a-f0-9]{32,64}".prop_map(|s| s) + } + + /// 生成随机的有效 kiro_version + fn arb_valid_kiro_version() -> impl Strategy { + "[0-9]{1,2}\\.[0-9]{1,2}\\.[0-9]{1,3}".prop_map(|s| s) + } + + proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// Property 3: 错误处理优雅降级 + /// + /// *For any* error condition (network error, invalid response, missing token), + /// the safe wrapper function SHALL return an empty UsageInfo with zero values + /// instead of panicking. + #[test] + fn prop_error_handling_graceful_degradation( + scenario in arb_error_scenario(), + valid_token in arb_valid_token(), + valid_machine_id in arb_valid_machine_id(), + valid_version in arb_valid_kiro_version(), + ) { + // 创建 tokio runtime 来运行异步代码 + let rt = tokio::runtime::Runtime::new().unwrap(); + + let result = rt.block_on(async { + match scenario { + ErrorScenario::EmptyToken => { + // 空 token 应该导致错误 + get_usage_limits_safe( + "", + "social", + Some("arn:aws:test"), + &valid_machine_id, + &valid_version + ).await + } + ErrorScenario::EmptyMachineId => { + // 空 machine_id 应该导致错误 + get_usage_limits_safe( + &valid_token, + "social", + Some("arn:aws:test"), + "", + &valid_version + ).await + } + ErrorScenario::MissingProfileArn => { + // social 认证缺少 profileArn 应该导致错误 + get_usage_limits_safe( + &valid_token, + "social", + None, + &valid_machine_id, + &valid_version + ).await + } + ErrorScenario::EmptyProfileArn => { + // social 认证空 profileArn 应该导致错误 + get_usage_limits_safe( + &valid_token, + "social", + Some(""), + &valid_machine_id, + &valid_version + ).await + } + ErrorScenario::InvalidToken => { + // 无效 token 会导致网络错误(401/403) + get_usage_limits_safe( + "invalid_token_that_will_fail", + "idc", + None, + &valid_machine_id, + &valid_version + ).await + } + } + }); + + // 无论什么错误场景,safe 函数都应该返回空的 UsageInfo + prop_assert_eq!(result.usage_limit, 0.0, + "Error scenario {:?} should return zero usage_limit", scenario); + prop_assert_eq!(result.current_usage, 0.0, + "Error scenario {:?} should return zero current_usage", scenario); + prop_assert_eq!(result.balance, 0.0, + "Error scenario {:?} should return zero balance", scenario); + prop_assert!(!result.is_low_balance, + "Error scenario {:?} should return false is_low_balance", scenario); + prop_assert!(result.subscription_title.is_empty(), + "Error scenario {:?} should return empty subscription_title", scenario); + } + } + + // 保留原有的单元测试作为补充(快速验证) + #[tokio::test] + async fn test_error_handling_empty_token() { + let result = + get_usage_limits_safe("", "social", Some("arn:aws:test"), "machine123", "1.0.0").await; + assert_eq!(result.usage_limit, 0.0); + assert_eq!(result.balance, 0.0); + } + + #[tokio::test] + async fn test_error_handling_empty_machine_id() { + let result = + get_usage_limits_safe("token123", "social", Some("arn:aws:test"), "", "1.0.0").await; + assert_eq!(result.usage_limit, 0.0); + assert_eq!(result.balance, 0.0); + } + + // ======================================================================== + // Property 6: User-Agent Header Format + // **Feature: kiro-usage-api, Property 6: User-Agent Header Format** + // **Validates: Requirements 4.1, 4.2** + // ======================================================================== + + /// 生成有效的 Kiro 版本号 + fn arb_kiro_version() -> impl Strategy { + "[0-9]{1,2}\\.[0-9]{1,2}\\.[0-9]{1,3}".prop_map(|s| s) + } + + /// 生成有效的 Machine ID (SHA256 哈希) + fn arb_machine_id() -> impl Strategy { + "[a-f0-9]{64}".prop_map(|s| s) + } + + proptest! { + #![proptest_config(ProptestConfig::with_cases(100))] + + /// Property 6: User-Agent 头格式 + /// + /// *For any* kiro_version and machine_id strings, + /// the User-Agent header SHALL match the format: + /// `aws-sdk-js/1.0.0 ua/2.1 os/{os} lang/rust api/codewhispererruntime#1.0.0 m/N,E KiroIDE-{version}-{machineId}` + #[test] + fn prop_user_agent_header_format( + kiro_version in arb_kiro_version(), + machine_id in arb_machine_id() + ) { + let user_agent = build_user_agent(&kiro_version, &machine_id); + + // 验证格式各部分 + prop_assert!(user_agent.starts_with("aws-sdk-js/1.0.0 ua/2.1 os/"), + "User-Agent should start with 'aws-sdk-js/1.0.0 ua/2.1 os/', got: {}", user_agent); + + prop_assert!(user_agent.contains("lang/rust"), + "User-Agent should contain 'lang/rust', got: {}", user_agent); + + prop_assert!(user_agent.contains("api/codewhispererruntime#1.0.0"), + "User-Agent should contain 'api/codewhispererruntime#1.0.0', got: {}", user_agent); + + prop_assert!(user_agent.contains("m/N,E"), + "User-Agent should contain 'm/N,E', got: {}", user_agent); + + // 验证包含 KiroIDE-{version}-{machineId} + let kiro_suffix = format!("KiroIDE-{}-{}", kiro_version, machine_id); + prop_assert!(user_agent.ends_with(&kiro_suffix), + "User-Agent should end with '{}', got: {}", kiro_suffix, user_agent); + } + + /// Property 6: x-amz-user-agent 头格式 + /// + /// *For any* kiro_version and machine_id strings, + /// the x-amz-user-agent header SHALL match the format: + /// `aws-sdk-js/1.0.0 KiroIDE-{version}-{machineId}` + #[test] + fn prop_x_amz_user_agent_header_format( + kiro_version in arb_kiro_version(), + machine_id in arb_machine_id() + ) { + let x_amz_user_agent = build_x_amz_user_agent(&kiro_version, &machine_id); + + // 验证格式 + let expected = format!("aws-sdk-js/1.0.0 KiroIDE-{}-{}", kiro_version, machine_id); + prop_assert_eq!(x_amz_user_agent, expected, + "x-amz-user-agent format mismatch"); + } + } +} diff --git a/src-tauri/tauri.conf.json b/src-tauri/tauri.conf.json index b344b2508..69d3ce296 100644 --- a/src-tauri/tauri.conf.json +++ b/src-tauri/tauri.conf.json @@ -1,7 +1,7 @@ { "$schema": "https://schema.tauri.app/config/2", "productName": "ProxyCast", - "version": "0.12.4", + "version": "0.12.5", "identifier": "com.proxycast.app", "build": { "beforeDevCommand": "npm run dev", diff --git a/src/components/provider-pool/CredentialCard.tsx b/src/components/provider-pool/CredentialCard.tsx index 56fe2be97..edcd5b052 100644 --- a/src/components/provider-pool/CredentialCard.tsx +++ b/src/components/provider-pool/CredentialCard.tsx @@ -1,3 +1,4 @@ +import { useState } from "react"; import { Heart, HeartOff, @@ -14,11 +15,15 @@ import { Lock, User, Globe, + BarChart3, + ChevronUp, } from "lucide-react"; import type { CredentialDisplay, CredentialSource, } from "@/lib/api/providerPool"; +import { usageApi, type UsageInfo } from "@/lib/api/usage"; +import { UsageDisplay } from "./UsageDisplay"; interface CredentialCardProps { credential: CredentialDisplay; @@ -31,6 +36,8 @@ interface CredentialCardProps { deleting: boolean; checkingHealth: boolean; refreshingToken?: boolean; + /** 是否为 Kiro 凭证(支持用量查询) */ + isKiroCredential?: boolean; } export function CredentialCard({ @@ -44,7 +51,36 @@ export function CredentialCard({ deleting, checkingHealth, refreshingToken, + isKiroCredential, }: CredentialCardProps) { + // 用量查询状态 + const [usageExpanded, setUsageExpanded] = useState(false); + const [usageLoading, setUsageLoading] = useState(false); + const [usageInfo, setUsageInfo] = useState(null); + const [usageError, setUsageError] = useState(null); + + // 查询用量 + const handleCheckUsage = async () => { + if (usageExpanded && usageInfo) { + // 已展开且有数据,直接折叠 + setUsageExpanded(false); + return; + } + + setUsageExpanded(true); + setUsageLoading(true); + setUsageError(null); + + try { + const info = await usageApi.getKiroUsage(credential.uuid); + setUsageInfo(info); + } catch (e) { + setUsageError(e instanceof Error ? e.message : String(e)); + } finally { + setUsageLoading(false); + } + }; + const formatDate = (dateStr?: string) => { if (!dateStr) return "从未"; const date = new Date(dateStr); @@ -138,28 +174,30 @@ export function CredentialCard({ {/* Main Info */}
-
+

{credential.name || `凭证 #${credential.uuid.slice(0, 8)}`}

- - {getCredentialTypeLabel(credential.credential_type)} - - - - {sourceInfo.text} - - {credential.proxy_url && ( - - - 代理 +
+ + {getCredentialTypeLabel(credential.credential_type)} - )} + + + {sourceInfo.text} + + {credential.proxy_url && ( + + + 代理 + + )} +

{credential.uuid} @@ -257,6 +295,24 @@ export function CredentialCard({ )} + {/* 用量查询按钮 - 仅 Kiro 凭证显示 */} + {isKiroCredential && ( + + )} +

)} + + {/* 用量信息展示区域 - 仅 Kiro 凭证 */} + {isKiroCredential && usageExpanded && ( +
+
+ + + Kiro 用量 + + +
+ + {usageError ? ( +
+ {usageError} +
+ ) : usageInfo ? ( + + ) : ( + + )} +
+ )}
); } diff --git a/src/components/provider-pool/ProviderPoolPage.tsx b/src/components/provider-pool/ProviderPoolPage.tsx index 2955fc8d1..65b9d35a3 100644 --- a/src/components/provider-pool/ProviderPoolPage.tsx +++ b/src/components/provider-pool/ProviderPoolPage.tsx @@ -568,6 +568,8 @@ export const ProviderPoolPage = forwardRef( // 判断是否为 OAuth 类型(需要刷新 Token 功能) const isOAuthType = credential.credential_type.includes("oauth"); + // 判断是否为 Kiro 凭证(支持用量查询) + const isKiroCredential = activeTab === "kiro"; return ( ( deleting={deletingCredentials.has(credential.uuid)} checkingHealth={checkingHealth === credential.uuid} refreshingToken={refreshingToken === credential.uuid} + isKiroCredential={isKiroCredential} /> ); })} diff --git a/src/components/provider-pool/UsageDisplay.tsx b/src/components/provider-pool/UsageDisplay.tsx new file mode 100644 index 000000000..a671aa2a9 --- /dev/null +++ b/src/components/provider-pool/UsageDisplay.tsx @@ -0,0 +1,142 @@ +import { AlertTriangle, TrendingUp, Zap, Wallet } from "lucide-react"; +import type { UsageInfo } from "@/lib/api/usage"; + +interface UsageDisplayProps { + usage: UsageInfo; + loading?: boolean; +} + +/** + * 用量显示组件 + * + * 显示订阅类型、总额度、已使用、余额 + * 低余额时显示警告样式 + * + * _Requirements: 3.3, 3.4_ + */ +export function UsageDisplay({ usage, loading }: UsageDisplayProps) { + if (loading) { + return ( +
+
+
+
+
+
+
+
+ ); + } + + // 计算使用百分比 + const usagePercent = + usage.usageLimit > 0 + ? Math.round((usage.currentUsage / usage.usageLimit) * 100) + : 0; + + // 格式化数字 + const formatNumber = (num: number) => { + if (num >= 1000000) { + return `${(num / 1000000).toFixed(1)}M`; + } + if (num >= 1000) { + return `${(num / 1000).toFixed(1)}K`; + } + return num.toFixed(1); + }; + + return ( +
+ {/* 标题和警告 */} +
+
+ + + {usage.subscriptionTitle || "用量信息"} + +
+ {usage.isLowBalance && ( +
+ + 余额不足 +
+ )} +
+ + {/* 进度条 */} +
+
+
50 + ? "bg-blue-500" + : "bg-green-500" + }`} + style={{ width: `${Math.min(usagePercent, 100)}%` }} + /> +
+
+ 已使用 {usagePercent}% + 剩余 {100 - usagePercent}% +
+
+ + {/* 数据统计 */} +
+
+
+ + 总额度 +
+
+ {formatNumber(usage.usageLimit)} +
+
+ +
+
+ + 已使用 +
+
+ {formatNumber(usage.currentUsage)} +
+
+ +
+
+ + 余额 +
+
+ {formatNumber(usage.balance)} +
+
+
+
+ ); +} diff --git a/src/components/provider-pool/index.ts b/src/components/provider-pool/index.ts index b02b7583c..18031babd 100644 --- a/src/components/provider-pool/index.ts +++ b/src/components/provider-pool/index.ts @@ -7,3 +7,4 @@ export { VertexAISection } from "./VertexAISection"; export { CodexSection } from "./CodexSection"; export { IFlowSection } from "./IFlowSection"; export { AmpConfigSection } from "./AmpConfigSection"; +export { UsageDisplay } from "./UsageDisplay"; diff --git a/src/lib/api/usage.ts b/src/lib/api/usage.ts new file mode 100644 index 000000000..9bdb30dbe --- /dev/null +++ b/src/lib/api/usage.ts @@ -0,0 +1,34 @@ +import { invoke } from "@tauri-apps/api/core"; + +/** + * 用量信息接口 + * + * 与后端 UsageInfo 结构对应 + * _Requirements: 3.3_ + */ +export interface UsageInfo { + /** 订阅类型名称 */ + subscriptionTitle: string; + /** 总额度 */ + usageLimit: number; + /** 已使用 */ + currentUsage: number; + /** 余额 = usageLimit - currentUsage */ + balance: number; + /** 余额低于 20% */ + isLowBalance: boolean; +} + +/** + * Usage API + */ +export const usageApi = { + /** + * 获取 Kiro 凭证的用量信息 + * + * @param credentialUuid - 凭证的 UUID + * @returns 用量信息 + */ + getKiroUsage: (credentialUuid: string): Promise => + invoke("get_kiro_usage", { credentialUuid }), +};