mirror of
https://github.com/aiclientproxy/proxycast.git
synced 2026-09-24 23:10:56 +08:00
fix: Kiro refreshToken 截断检测与降级处理
- 添加凭证时检测 refreshToken 是否被截断(长度<100或包含...) - 刷新 Token 时增加截断检测,给出清晰的错误提示 - 用量查询失败时降级使用源文件中的 accessToken - 修复凭证卡片标签重叠的 UI 问题 - 版本更新至 0.12.5
This commit is contained in:
+1
-1
@@ -1,7 +1,7 @@
|
||||
{
|
||||
"name": "proxycast",
|
||||
"private": true,
|
||||
"version": "0.12.4",
|
||||
"version": "0.12.5",
|
||||
"type": "module",
|
||||
"scripts": {
|
||||
"dev": "vite",
|
||||
|
||||
Generated
+2
-1
@@ -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",
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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 认证方式");
|
||||
}
|
||||
}
|
||||
|
||||
// 写入合并后的凭证到副本文件
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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 并缓存到数据库
|
||||
|
||||
@@ -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,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",
|
||||
|
||||
@@ -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>
|
||||
);
|
||||
}
|
||||
@@ -7,3 +7,4 @@ export { VertexAISection } from "./VertexAISection";
|
||||
export { CodexSection } from "./CodexSection";
|
||||
export { IFlowSection } from "./IFlowSection";
|
||||
export { AmpConfigSection } from "./AmpConfigSection";
|
||||
export { UsageDisplay } from "./UsageDisplay";
|
||||
|
||||
@@ -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 }),
|
||||
};
|
||||
Reference in New Issue
Block a user