fix: Kiro refreshToken 截断检测与降级处理

- 添加凭证时检测 refreshToken 是否被截断(长度<100或包含...)
- 刷新 Token 时增加截断检测,给出清晰的错误提示
- 用量查询失败时降级使用源文件中的 accessToken
- 修复凭证卡片标签重叠的 UI 问题
- 版本更新至 0.12.5
This commit is contained in:
coso
2025-12-19 00:23:03 +08:00
parent ef222eb961
commit 1fe662d1e3
17 changed files with 1836 additions and 81 deletions
+1 -1
View File
@@ -1,7 +1,7 @@
{
"name": "proxycast",
"private": true,
"version": "0.12.4",
"version": "0.12.5",
"type": "module",
"scripts": {
"dev": "vite",
+2 -1
View File
@@ -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",
+2 -1
View File
@@ -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"
+1
View File
@@ -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;
+44 -3
View File
@@ -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 认证方式");
}
}
// 写入合并后的凭证到副本文件
+349
View File
@@ -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<UsageInfo, String> {
// 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>), 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<String, String> {
// 尝试获取系统 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<String, String> {
#[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);
}
}
+2
View File
@@ -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");
+163 -55
View File
@@ -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<String> {
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
+1
View File
@@ -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;
+64 -1
View File
@@ -43,6 +43,7 @@ impl TokenCacheService {
/// 1. 检查数据库缓存是否有效
/// 2. 如果缓存有效且未过期,直接返回
/// 3. 如果缓存无效或即将过期,执行刷新
/// 4. 如果刷新失败(如 refreshToken 被截断),尝试使用源文件中的 accessToken
pub async fn get_valid_token(&self, db: &DbConnection, uuid: &str) -> Result<String, String> {
// 首先检查缓存
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 并缓存到数据库
+915
View File
@@ -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<UsageBreakdown>,
}
/// 订阅信息结构
#[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<FreeTrialInfo>,
#[serde(default)]
pub bonuses: Option<Vec<Bonus>>,
}
/// 免费试用信息
#[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<String, Box<dyn Error + Send + Sync>> {
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<HeaderMap, Box<dyn Error + Send + Sync>> {
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<UsageInfo, Box<dyn Error + Send + Sync>> {
// 验证参数
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<Value = Bonus> {
(
"[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<Value = FreeTrialInfo> {
(
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<Value = UsageBreakdown> {
(
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<Value = SubscriptionInfo> {
(
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<Value = UsageLimitsResponse> {
(
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<Value = (f64, f64)> {
// 生成 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<Value = String> {
"[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<Value = ErrorScenario> {
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<Value = String> {
"[a-zA-Z0-9]{20,50}".prop_map(|s| s)
}
/// 生成随机的有效 machine_id(用于非空 machine_id 场景)
fn arb_valid_machine_id() -> impl Strategy<Value = String> {
"[a-f0-9]{32,64}".prop_map(|s| s)
}
/// 生成随机的有效 kiro_version
fn arb_valid_kiro_version() -> impl Strategy<Value = String> {
"[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<Value = String> {
"[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<Value = String> {
"[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");
}
}
}
+1 -1
View File
@@ -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",
+111 -18
View File
@@ -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<UsageInfo | null>(null);
const [usageError, setUsageError] = useState<string | null>(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 */}
<div className="flex-1 min-w-0">
<div className="flex items-center gap-2 mb-1">
<div className="flex flex-col gap-1.5 mb-1">
<h4 className="font-semibold text-base truncate">
{credential.name || `凭证 #${credential.uuid.slice(0, 8)}`}
</h4>
<span className="rounded-full bg-muted px-2 py-0.5 text-xs font-medium">
{getCredentialTypeLabel(credential.credential_type)}
</span>
<span
className={`rounded-full px-2 py-0.5 text-xs font-medium inline-flex items-center gap-1 whitespace-nowrap shrink-0 ${sourceInfo.color}`}
>
<SourceIcon className="h-3 w-3 shrink-0" />
{sourceInfo.text}
</span>
{credential.proxy_url && (
<span
className="rounded-full px-2 py-0.5 text-xs font-medium inline-flex items-center gap-1 whitespace-nowrap shrink-0 bg-orange-100 text-orange-700 dark:bg-orange-900/30 dark:text-orange-400"
title={`代理: ${credential.proxy_url}`}
>
<Globe className="h-3 w-3 shrink-0" />
代理
<div className="flex flex-wrap items-center gap-1.5">
<span className="rounded-full bg-muted px-2 py-0.5 text-xs font-medium">
{getCredentialTypeLabel(credential.credential_type)}
</span>
)}
<span
className={`rounded-full px-2 py-0.5 text-xs font-medium inline-flex items-center gap-1 whitespace-nowrap ${sourceInfo.color}`}
>
<SourceIcon className="h-3 w-3 shrink-0" />
{sourceInfo.text}
</span>
{credential.proxy_url && (
<span
className="rounded-full px-2 py-0.5 text-xs font-medium inline-flex items-center gap-1 whitespace-nowrap bg-orange-100 text-orange-700 dark:bg-orange-900/30 dark:text-orange-400"
title={`代理: ${credential.proxy_url}`}
>
<Globe className="h-3 w-3 shrink-0" />
代理
</span>
)}
</div>
</div>
<p className="text-xs text-muted-foreground font-mono truncate">
{credential.uuid}
@@ -257,6 +295,24 @@ export function CredentialCard({
</button>
)}
{/* 用量查询按钮 - 仅 Kiro 凭证显示 */}
{isKiroCredential && (
<button
onClick={handleCheckUsage}
disabled={usageLoading}
className={`rounded-lg p-2 transition-colors ${
usageExpanded
? "bg-cyan-200 text-cyan-800 dark:bg-cyan-800 dark:text-cyan-200"
: "bg-cyan-100 text-cyan-700 hover:bg-cyan-200 dark:bg-cyan-900/30 dark:text-cyan-400"
} disabled:opacity-50`}
title="查看用量"
>
<BarChart3
className={`h-4 w-4 ${usageLoading ? "animate-pulse" : ""}`}
/>
</button>
)}
<button
onClick={onReset}
className="rounded-lg bg-orange-100 p-2 text-orange-700 hover:bg-orange-200 dark:bg-orange-900/30 dark:text-orange-400 transition-colors"
@@ -305,6 +361,43 @@ export function CredentialCard({
{credential.last_error_message.length > 150 && "..."}
</div>
)}
{/* 用量信息展示区域 - 仅 Kiro 凭证 */}
{isKiroCredential && usageExpanded && (
<div className="mt-3 pt-3 border-t border-border/30">
<div className="flex items-center justify-between mb-2">
<span className="text-xs font-medium text-muted-foreground flex items-center gap-1">
<BarChart3 className="h-3 w-3" />
Kiro 用量
</span>
<button
onClick={() => setUsageExpanded(false)}
className="text-muted-foreground hover:text-foreground"
>
<ChevronUp className="h-4 w-4" />
</button>
</div>
{usageError ? (
<div className="rounded-lg bg-red-100 p-2 text-xs text-red-700 dark:bg-red-900/30 dark:text-red-300">
{usageError}
</div>
) : usageInfo ? (
<UsageDisplay usage={usageInfo} loading={usageLoading} />
) : (
<UsageDisplay
usage={{
subscriptionTitle: "",
usageLimit: 0,
currentUsage: 0,
balance: 0,
isLowBalance: false,
}}
loading={true}
/>
)}
</div>
)}
</div>
);
}
@@ -568,6 +568,8 @@ export const ProviderPoolPage = forwardRef<ProviderPoolPageRef>(
// 判断是否为 OAuth 类型(需要刷新 Token 功能)
const isOAuthType =
credential.credential_type.includes("oauth");
// 判断是否为 Kiro 凭证(支持用量查询)
const isKiroCredential = activeTab === "kiro";
return (
<CredentialCard
key={credential.uuid}
@@ -585,6 +587,7 @@ export const ProviderPoolPage = forwardRef<ProviderPoolPageRef>(
deleting={deletingCredentials.has(credential.uuid)}
checkingHealth={checkingHealth === credential.uuid}
refreshingToken={refreshingToken === credential.uuid}
isKiroCredential={isKiroCredential}
/>
);
})}
@@ -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 (
<div className="rounded-lg border p-4 animate-pulse">
<div className="h-4 bg-muted rounded w-1/3 mb-3" />
<div className="grid grid-cols-3 gap-4">
<div className="h-12 bg-muted rounded" />
<div className="h-12 bg-muted rounded" />
<div className="h-12 bg-muted rounded" />
</div>
</div>
);
}
// 计算使用百分比
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 (
<div
className={`rounded-lg border p-4 ${
usage.isLowBalance
? "border-amber-300 bg-amber-50/50 dark:border-amber-700 dark:bg-amber-950/30"
: "border-border bg-card"
}`}
>
{/* 标题和警告 */}
<div className="flex items-center justify-between mb-3">
<div className="flex items-center gap-2">
<Zap className="h-4 w-4 text-primary" />
<span className="font-medium text-sm">
{usage.subscriptionTitle || "用量信息"}
</span>
</div>
{usage.isLowBalance && (
<div className="flex items-center gap-1 text-amber-600 dark:text-amber-400">
<AlertTriangle className="h-4 w-4" />
<span className="text-xs font-medium">余额不足</span>
</div>
)}
</div>
{/* 进度条 */}
<div className="mb-4">
<div className="h-2 bg-muted rounded-full overflow-hidden">
<div
className={`h-full transition-all ${
usage.isLowBalance
? "bg-amber-500"
: usagePercent > 50
? "bg-blue-500"
: "bg-green-500"
}`}
style={{ width: `${Math.min(usagePercent, 100)}%` }}
/>
</div>
<div className="flex justify-between mt-1 text-xs text-muted-foreground">
<span>已使用 {usagePercent}%</span>
<span>剩余 {100 - usagePercent}%</span>
</div>
</div>
{/* 数据统计 */}
<div className="grid grid-cols-3 gap-3">
<div className="text-center p-2 rounded-lg bg-muted/50">
<div className="flex items-center justify-center gap-1 text-muted-foreground mb-1">
<TrendingUp className="h-3 w-3" />
<span className="text-xs">总额度</span>
</div>
<div className="font-semibold text-sm">
{formatNumber(usage.usageLimit)}
</div>
</div>
<div className="text-center p-2 rounded-lg bg-muted/50">
<div className="flex items-center justify-center gap-1 text-muted-foreground mb-1">
<Zap className="h-3 w-3" />
<span className="text-xs">已使用</span>
</div>
<div className="font-semibold text-sm">
{formatNumber(usage.currentUsage)}
</div>
</div>
<div
className={`text-center p-2 rounded-lg ${
usage.isLowBalance
? "bg-amber-100 dark:bg-amber-900/30"
: "bg-muted/50"
}`}
>
<div
className={`flex items-center justify-center gap-1 mb-1 ${
usage.isLowBalance
? "text-amber-600 dark:text-amber-400"
: "text-muted-foreground"
}`}
>
<Wallet className="h-3 w-3" />
<span className="text-xs">余额</span>
</div>
<div
className={`font-semibold text-sm ${
usage.isLowBalance ? "text-amber-600 dark:text-amber-400" : ""
}`}
>
{formatNumber(usage.balance)}
</div>
</div>
</div>
</div>
);
}
+1
View File
@@ -7,3 +7,4 @@ export { VertexAISection } from "./VertexAISection";
export { CodexSection } from "./CodexSection";
export { IFlowSection } from "./IFlowSection";
export { AmpConfigSection } from "./AmpConfigSection";
export { UsageDisplay } from "./UsageDisplay";
+34
View File
@@ -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<UsageInfo> =>
invoke("get_kiro_usage", { credentialUuid }),
};