feat: 新增 Kiro 适配器、OAuth 改进与多项功能增强

- 新增 Kiro provider 适配器(EventStream 协议解析、令牌管理、用量追踪)
- 重构 OAuth 账户管理与统一配额机制
- 重构 Handler 基类(CLI adapter/handler、请求构建器、流处理器)
- 增强缓存监控后端 API 与前端可视化
- 改进 Gemini 格式标准化器与请求头处理
- Antigravity/Codex 适配器更新,移除旧 metadata_collector
- 新增数据库迁移:proxy provider API keys
- 前端 UI 多项优化

Co-Authored-By: AAEE86 <ppk0227@hotmail.com>
This commit is contained in:
fawney19
2026-02-09 01:05:48 +08:00
parent e324ffdcd6
commit b34dd12863
82 changed files with 6498 additions and 840 deletions

View File

@@ -5,6 +5,7 @@ Provider API Keys 管理
from __future__ import annotations
import json
import time
import uuid
from dataclasses import dataclass
from datetime import datetime, timezone
@@ -128,6 +129,26 @@ async def reveal_endpoint_key(
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.get("/keys/{key_id}/export")
async def export_key(
key_id: str,
request: Request,
db: Session = Depends(get_db),
_: User = Depends(require_admin),
) -> dict:
"""
导出 OAuth Key 凭据(用于跨实例迁移)
解密 auth_config返回精简的扁平 JSON去掉 null 和临时字段。
所有 OAuth Provider 格式统一。
**路径参数**:
- `key_id`: Key ID
"""
adapter = AdminExportKeyAdapter(key_id=key_id)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.delete("/keys/{key_id}")
async def delete_endpoint_key(
key_id: str,
@@ -245,6 +266,102 @@ async def add_provider_key(
# -------- Adapters --------
def _normalize_auth_type(raw: str) -> str:
"""将数据库中的 auth_type 归一化为逻辑类型。
Kiro 在数据库中存储为 ``"kiro"`` 或 ``"oauth"``,统一映射为 ``"oauth"``。
"""
t = str(raw or "api_key").strip() or "api_key"
return "oauth" if t == "kiro" else t
def check_duplicate_key(
db: Session,
provider_id: str,
auth_type: str,
new_api_key: str | None = None,
new_auth_config: dict | None = None,
exclude_key_id: str | None = None,
) -> None:
"""
检查密钥是否与其他现有密钥重复
对于不同的认证类型,使用不同的比较方式:
- api_key: 比较 API Key 的哈希值
- vertex_ai: 比较 Service Account 的 client_email
Args:
db: 数据库会话
provider_id: Provider ID
auth_type: 认证类型 (api_key, vertex_ai, oauth)
new_api_key: 新的 API Key用于 api_key 类型)
new_auth_config: 新的认证配置(用于 vertex_ai 类型)
exclude_key_id: 要排除的 Key ID用于更新场景
"""
if auth_type == "api_key" and new_api_key:
# 跳过占位符
if new_api_key == "__placeholder__":
return
# 仅查询同 auth_type 的 Keys减少不必要的解密操作
query = db.query(ProviderAPIKey).filter(
ProviderAPIKey.provider_id == provider_id,
ProviderAPIKey.auth_type == "api_key",
)
if exclude_key_id:
query = query.filter(ProviderAPIKey.id != exclude_key_id)
new_key_hash = crypto_service.hash_api_key(new_api_key)
for existing_key in query:
try:
decrypted_key = crypto_service.decrypt(existing_key.api_key, silent=True)
if decrypted_key == "__placeholder__":
continue
existing_hash = crypto_service.hash_api_key(decrypted_key)
if new_key_hash == existing_hash:
raise InvalidRequestException(
f"该 API Key 已存在于当前 Provider 中(名称: {existing_key.name}"
)
except InvalidRequestException:
raise
except Exception:
# 解密失败时跳过该 Key
continue
elif auth_type == "vertex_ai" and new_auth_config:
new_client_email = (
new_auth_config.get("client_email") if isinstance(new_auth_config, dict) else None
)
if not new_client_email:
return
# 仅查询同 auth_type 且有 auth_config 的 Keys
query = db.query(ProviderAPIKey).filter(
ProviderAPIKey.provider_id == provider_id,
ProviderAPIKey.auth_type == "vertex_ai",
ProviderAPIKey.auth_config.isnot(None),
)
if exclude_key_id:
query = query.filter(ProviderAPIKey.id != exclude_key_id)
for existing_key in query:
try:
decrypted_config = json.loads(
crypto_service.decrypt(existing_key.auth_config, silent=True)
)
existing_email = decrypted_config.get("client_email")
if existing_email and existing_email == new_client_email:
raise InvalidRequestException(
f"该 Service Account ({new_client_email}) 已存在于当前 Provider 中"
f"(名称: {existing_key.name}"
)
except InvalidRequestException:
raise
except Exception:
# 解密失败时跳过该 Key
continue
@dataclass
class AdminUpdateEndpointKeyAdapter(AdminApiAdapter):
key_id: str
@@ -274,7 +391,7 @@ class AdminUpdateEndpointKeyAdapter(AdminApiAdapter):
update_data = self.key_data.model_dump(exclude_unset=True)
# 验证 auth_type 切换
current_auth_type = getattr(key, "auth_type", "api_key") or "api_key"
current_auth_type = _normalize_auth_type(getattr(key, "auth_type", "api_key"))
target_auth_type = update_data.get("auth_type", current_auth_type) or current_auth_type
# auth_type 切换校验 + 字段归一化
@@ -299,7 +416,16 @@ class AdminUpdateEndpointKeyAdapter(AdminApiAdapter):
if "api_key" not in update_data:
update_data["api_key"] = "__placeholder__"
# 加密 api_key非 None 时
# 检查密钥是否与其他现有密钥重复(排除当前正在更新的密钥
check_duplicate_key(
db=db,
provider_id=key.provider_id,
auth_type=target_auth_type,
new_api_key=update_data.get("api_key"),
new_auth_config=update_data.get("auth_config"),
exclude_key_id=self.key_id,
)
if "api_key" in update_data and update_data["api_key"] is not None:
update_data["api_key"] = crypto_service.encrypt(update_data["api_key"])
# 加密 auth_config包含敏感的 Service Account 凭证)
@@ -338,6 +464,13 @@ class AdminUpdateEndpointKeyAdapter(AdminApiAdapter):
if isinstance(patterns, list) and len(patterns) == 0:
update_data["model_exclude_patterns"] = None
# 处理 proxy将 ProxyConfig 转换为 dict 存储null 清除代理
if "proxy" in self.key_data.model_fields_set:
if self.key_data.proxy is None:
update_data["proxy"] = None
else:
update_data["proxy"] = self.key_data.proxy.model_dump(exclude_none=True)
for field, value in update_data.items():
setattr(key, field, value)
key.updated_at = datetime.now(timezone.utc)
@@ -433,7 +566,7 @@ class AdminRevealEndpointKeyAdapter(AdminApiAdapter):
if not key:
raise NotFoundException(f"Key {self.key_id} 不存在")
auth_type = getattr(key, "auth_type", "api_key") or "api_key"
auth_type = _normalize_auth_type(getattr(key, "auth_type", "api_key"))
# Vertex AI 类型返回 auth_config需要解密
if auth_type == "vertex_ai":
@@ -468,7 +601,7 @@ class AdminRevealEndpointKeyAdapter(AdminApiAdapter):
"无法解密认证配置,可能是加密密钥已更改。请重新添加该密钥。"
)
# OAuth 类型:返回 access_token + refresh_token
# OAuth 类型:返回 access_token(导出走 /export 端点)
if auth_type == "oauth":
try:
decrypted_key = crypto_service.decrypt(key.api_key)
@@ -477,19 +610,8 @@ class AdminRevealEndpointKeyAdapter(AdminApiAdapter):
raise InvalidRequestException(
"无法解密 API Key可能是加密密钥已更改。请重新添加该密钥。"
)
result: dict[str, Any] = {"auth_type": "oauth", "api_key": decrypted_key}
encrypted_auth_config = getattr(key, "auth_config", None)
if encrypted_auth_config:
try:
decrypted_config = crypto_service.decrypt(encrypted_auth_config)
auth_config = json.loads(decrypted_config)
refresh_token = auth_config.get("refresh_token")
if refresh_token:
result["refresh_token"] = refresh_token
except Exception as e:
logger.error(f"解密 auth_config 失败: ID={self.key_id}, Error={e}")
logger.info(f"[REVEAL] 查看 OAuth Key: ID={self.key_id}, Name={key.name}")
return result
return {"auth_type": "oauth", "api_key": decrypted_key}
# API Key 类型返回 api_key
try:
@@ -504,6 +626,48 @@ class AdminRevealEndpointKeyAdapter(AdminApiAdapter):
return {"auth_type": "api_key", "api_key": decrypted_key}
@dataclass
class AdminExportKeyAdapter(AdminApiAdapter):
"""导出 OAuth Key 凭据:解密 auth_config委托 provider-specific builder 构建导出数据。"""
key_id: str
async def handle(self, context: ApiRequestContext) -> Any: # type: ignore[override]
from src.services.provider.export import build_export_data
db = context.db
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == self.key_id).first()
if not key:
raise NotFoundException(f"Key {self.key_id} 不存在")
auth_type = _normalize_auth_type(getattr(key, "auth_type", "api_key"))
if auth_type != "oauth":
raise InvalidRequestException("仅 OAuth 类型的 Key 支持导出")
encrypted_auth_config = getattr(key, "auth_config", None)
if not encrypted_auth_config:
raise InvalidRequestException("缺少认证配置,无法导出")
try:
auth_config: dict[str, Any] = json.loads(crypto_service.decrypt(encrypted_auth_config))
except Exception:
raise InvalidRequestException("无法解密认证配置")
if not auth_config.get("refresh_token"):
raise InvalidRequestException("缺少 refresh_token无法导出")
provider_type = str(auth_config.get("provider_type") or "").strip()
upstream = getattr(key, "upstream_metadata", None)
export_data = build_export_data(provider_type, auth_config, upstream)
export_data["name"] = key.name or ""
export_data["exported_at"] = datetime.now(timezone.utc).isoformat()
logger.info("[EXPORT] Key {}... 导出成功", self.key_id[:8])
return export_data
@dataclass
class AdminDeleteEndpointKeyAdapter(AdminApiAdapter):
key_id: str
@@ -586,9 +750,11 @@ class AdminGetKeysGroupedByFormatAdapter(AdminApiAdapter):
if not api_formats:
continue # 跳过没有 API 格式的 Key
auth_type = getattr(key, "auth_type", "api_key") or "api_key"
auth_type = _normalize_auth_type(getattr(key, "auth_type", "api_key"))
if auth_type == "vertex_ai":
masked_key = "[Service Account]"
elif auth_type == "oauth":
masked_key = "[OAuth Token]"
else:
try:
decrypted_key = crypto_service.decrypt(key.api_key)
@@ -665,7 +831,7 @@ def _build_key_response(
key: ProviderAPIKey, api_key_plain: str | None = None
) -> EndpointAPIKeyResponse:
"""构建 Key 响应对象的辅助函数"""
auth_type = getattr(key, "auth_type", "api_key") or "api_key"
auth_type = _normalize_auth_type(getattr(key, "auth_type", "api_key"))
if auth_type == "vertex_ai":
# Vertex AI 使用 Service Account不显示占位符
@@ -688,6 +854,7 @@ def _build_key_response(
key_dict = key.__dict__.copy()
key_dict.pop("_sa_instance_state", None)
key_dict.pop("api_key", None) # 移除敏感字段,避免泄露
key_dict["auth_type"] = auth_type
# 提取 OAuth 元数据(如果是 OAuth 类型)
oauth_expires_at = None
@@ -829,8 +996,14 @@ class AdminCreateProviderKeyAdapter(AdminApiAdapter):
if self.key_data.api_key:
raise InvalidRequestException("OAuth 认证模式下不允许直接填写 api_key")
# 允许同一个 API Key 在同一 Provider 下添加多次
# 用户可以为不同的 API 格式创建独立的配置记录,便于分开管理
# 检查密钥是否已存在(防止重复添加)
check_duplicate_key(
db=db,
provider_id=self.provider_id,
auth_type=auth_type,
new_api_key=self.key_data.api_key,
new_auth_config=self.key_data.auth_config,
)
# 加密 API Key如果有
encrypted_key = (
@@ -929,8 +1102,121 @@ class AdminCreateProviderKeyAdapter(AdminApiAdapter):
# ========== Codex Quota Refresh API ==========
# Codex 限额刷新测试请求使用的模型(选择最小/最便宜的模型
CODEX_QUOTA_REFRESH_MODEL = "gpt-5.1-codex-mini"
# Codex wham/usage API 地址(用于查询限额信息
CODEX_WHAM_USAGE_URL = "https://chatgpt.com/backend-api/wham/usage"
def _parse_codex_wham_usage_response(data: dict) -> dict | None:
"""
解析 Codex wham/usage API 响应,提取限额信息
Free 账号:
- rate_limit.primary_window: 周限额
- code_review_rate_limit.primary_window: 代码审查周限额
Team/Plus/Enterprise 账号:
- rate_limit.primary_window: 5H 限额
- rate_limit.secondary_window: 周限额
- code_review_rate_limit.primary_window: 代码审查周限额
"""
if not data:
return None
result: dict = {}
plan_type = data.get("plan_type")
if plan_type:
result["plan_type"] = plan_type
# 解析 rate_limit
rate_limit = data.get("rate_limit") or {}
primary_window = rate_limit.get("primary_window") or {}
secondary_window = rate_limit.get("secondary_window")
# 根据账号类型解析限额
# Free 账号: primary_window 是周限额,无 secondary_window
# Team/Plus/Enterprise: primary_window 是 5H 限额secondary_window 是周限额
if plan_type == "free":
# Free 账号: primary_window 是周限额
if primary_window:
used_percent = primary_window.get("used_percent")
if used_percent is not None:
result["primary_used_percent"] = float(used_percent)
reset_seconds = primary_window.get("reset_after_seconds")
if reset_seconds is not None:
result["primary_reset_seconds"] = int(reset_seconds)
reset_at = primary_window.get("reset_at")
if reset_at is not None:
result["primary_reset_at"] = int(reset_at)
limit_window_seconds = primary_window.get("limit_window_seconds")
if limit_window_seconds is not None:
result["primary_window_minutes"] = int(limit_window_seconds) // 60
else:
# Team/Plus/Enterprise: primary_window 是 5H 限额, secondary_window 是周限额
if secondary_window:
# 周限额 (secondary_window)
used_percent = secondary_window.get("used_percent")
if used_percent is not None:
result["primary_used_percent"] = float(used_percent)
reset_seconds = secondary_window.get("reset_after_seconds")
if reset_seconds is not None:
result["primary_reset_seconds"] = int(reset_seconds)
reset_at = secondary_window.get("reset_at")
if reset_at is not None:
result["primary_reset_at"] = int(reset_at)
limit_window_seconds = secondary_window.get("limit_window_seconds")
if limit_window_seconds is not None:
result["primary_window_minutes"] = int(limit_window_seconds) // 60
if primary_window:
# 5H 限额 (primary_window)
used_percent = primary_window.get("used_percent")
if used_percent is not None:
result["secondary_used_percent"] = float(used_percent)
reset_seconds = primary_window.get("reset_after_seconds")
if reset_seconds is not None:
result["secondary_reset_seconds"] = int(reset_seconds)
reset_at = primary_window.get("reset_at")
if reset_at is not None:
result["secondary_reset_at"] = int(reset_at)
limit_window_seconds = primary_window.get("limit_window_seconds")
if limit_window_seconds is not None:
result["secondary_window_minutes"] = int(limit_window_seconds) // 60
# 解析 code_review_rate_limit (代码审查限额)
code_review_limit = data.get("code_review_rate_limit") or {}
code_review_primary = code_review_limit.get("primary_window") or {}
if code_review_primary:
used_percent = code_review_primary.get("used_percent")
if used_percent is not None:
result["code_review_used_percent"] = float(used_percent)
reset_seconds = code_review_primary.get("reset_after_seconds")
if reset_seconds is not None:
result["code_review_reset_seconds"] = int(reset_seconds)
reset_at = code_review_primary.get("reset_at")
if reset_at is not None:
result["code_review_reset_at"] = int(reset_at)
limit_window_seconds = code_review_primary.get("limit_window_seconds")
if limit_window_seconds is not None:
result["code_review_window_minutes"] = int(limit_window_seconds) // 60
# 解析 credits
credits = data.get("credits") or {}
has_credits = credits.get("has_credits")
if has_credits is not None:
result["has_credits"] = bool(has_credits)
balance = credits.get("balance")
if balance is not None:
result["credits_balance"] = float(balance)
# 添加更新时间戳
if result:
result["updated_at"] = int(time.time())
return result if result else None
# ========== Kiro Quota Refresh API ==========
@router.post("/providers/{provider_id}/refresh-quota")
@@ -940,10 +1226,12 @@ async def refresh_provider_quota(
db: Session = Depends(get_db),
) -> dict:
"""
刷新 Provider 所有 Keys 的限额信息Codex
刷新 Provider 所有 Keys 的限额信息
向每个 Key 发送一个测试请求,从响应头中获取最新的限额信息。
仅适用于 Codex 类型的 Provider。
支持的 Provider 类型:
- Codex: 调用 wham/usage API 获取限额
- Antigravity: 调用 fetchAvailableModels 获取配额
- Kiro: 调用 getUsageLimits API 获取使用额度
**路径参数**:
- `provider_id`: Provider ID
@@ -969,24 +1257,18 @@ class AdminRefreshProviderQuotaAdapter(AdminApiAdapter):
import httpx
from src.api.handlers.base.request_builder import get_provider_auth
from src.services.provider.metadata_collectors import (
MetadataCollectorRegistry,
ensure_collectors_registered,
)
from src.services.provider.transport import build_provider_url
from src.utils.ssl_utils import get_ssl_context
# 确保 Codex 采集器已注册
ensure_collectors_registered()
db = context.db
provider = db.query(Provider).filter(Provider.id == self.provider_id).first()
if not provider:
raise NotFoundException(f"Provider {self.provider_id} 不存在")
provider_type = str(getattr(provider, "provider_type", "") or "").strip().lower()
if provider_type not in {ProviderType.CODEX, ProviderType.ANTIGRAVITY}:
raise InvalidRequestException("仅支持 Codex / Antigravity 类型的 Provider 刷新限额")
if provider_type not in {ProviderType.CODEX, ProviderType.ANTIGRAVITY, ProviderType.KIRO}:
raise InvalidRequestException(
"仅支持 Codex / Antigravity / Kiro 类型的 Provider 刷新限额"
)
# 获取所有活跃的 Keys
keys = (
@@ -1010,6 +1292,7 @@ class AdminRefreshProviderQuotaAdapter(AdminApiAdapter):
# 获取端点:
# - Codex: openai:cli
# - Antigravity: gemini:chat用于触发 oauth 刷新 + 提供 auth_config.project_id
# - Kiro: 不需要特定端点,直接使用 auth_config 中的凭据
endpoint = None
if provider_type == ProviderType.CODEX:
for ep in provider.endpoints:
@@ -1018,7 +1301,7 @@ class AdminRefreshProviderQuotaAdapter(AdminApiAdapter):
break
if not endpoint:
raise InvalidRequestException("找不到有效的 openai:cli 端点")
else:
elif provider_type == ProviderType.ANTIGRAVITY:
# Prefer the new signature, but keep backward-compat with existing DB rows.
for sig in ("gemini:chat", "gemini:cli"):
for ep in provider.endpoints:
@@ -1029,6 +1312,7 @@ class AdminRefreshProviderQuotaAdapter(AdminApiAdapter):
break
if not endpoint:
raise InvalidRequestException("找不到有效的 gemini:chat/gemini:cli 端点")
# Kiro 不需要端点检查,直接使用 auth_config
results: list[dict] = []
success_count = 0
@@ -1041,15 +1325,11 @@ class AdminRefreshProviderQuotaAdapter(AdminApiAdapter):
async def refresh_single_key(key: ProviderAPIKey) -> dict:
try:
if provider_type == ProviderType.CODEX:
# 获取认证信息
# 获取认证信息(用于刷新 OAuth token
auth_info = await get_provider_auth(endpoint, key)
# 构建请求 URL
url = build_provider_url(endpoint, key=key)
# 构建请求头
headers = {
"Content-Type": "application/json",
"Accept": "application/json",
}
if auth_info:
@@ -1059,31 +1339,53 @@ class AdminRefreshProviderQuotaAdapter(AdminApiAdapter):
decrypted_key = crypto_service.decrypt(key.api_key)
headers["Authorization"] = f"Bearer {decrypted_key}"
# 发送最小的测试请求,使用 Codex Responses API 格式
test_body = {
"model": CODEX_QUOTA_REFRESH_MODEL,
"input": [
{
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": "hi"}],
}
],
"instructions": "",
"stream": True,
"store": False,
}
# 从 auth_config 中解密获取 plan_type 和 account_id
oauth_plan_type = None
oauth_account_id = None
auth_type = _normalize_auth_type(getattr(key, "auth_type", "api_key"))
if auth_type == "oauth" and key.auth_config:
try:
decrypted_config = crypto_service.decrypt(key.auth_config)
auth_config_data = json.loads(decrypted_config)
oauth_plan_type = auth_config_data.get("plan_type")
oauth_account_id = auth_config_data.get("account_id")
except Exception:
pass
# 如果有 account_id 且不是 free 账号,添加 chatgpt-account-id 头
if oauth_account_id and oauth_plan_type and oauth_plan_type.lower() != "free":
headers["chatgpt-account-id"] = oauth_account_id
# 使用 wham/usage API 获取限额信息
async with httpx.AsyncClient(timeout=30.0, verify=get_ssl_context()) as client:
response = await client.post(url, json=test_body, headers=headers)
response = await client.get(CODEX_WHAM_USAGE_URL, headers=headers)
# 解析响应头中的限额信息
response_headers = dict(response.headers)
metadata = MetadataCollectorRegistry.collect("codex", response_headers)
if response.status_code != 200:
return {
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": f"wham/usage API 返回状态码 {response.status_code}",
"status_code": response.status_code,
}
# 解析 JSON 响应
try:
data = response.json()
except Exception:
return {
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "无法解析 wham/usage API 响应",
}
# 解析限额信息
metadata = _parse_codex_wham_usage_response(data)
if metadata:
# 收集元数据,稍后统一更新数据库
metadata_updates[key.id] = metadata
# 收集元数据,稍后统一更新数据库(存储到 codex 子对象)
metadata_updates[key.id] = {"codex": metadata}
return {
"key_id": key.id,
"key_name": key.name,
@@ -1091,7 +1393,7 @@ class AdminRefreshProviderQuotaAdapter(AdminApiAdapter):
"metadata": metadata,
}
# 响应成功但没有限额
# 响应成功但没有限额信息
return {
"key_id": key.id,
"key_name": key.name,
@@ -1176,6 +1478,92 @@ class AdminRefreshProviderQuotaAdapter(AdminApiAdapter):
"message": error_msg,
}
elif provider_type == ProviderType.KIRO:
from src.services.provider.adapters.kiro.usage import (
fetch_kiro_usage_limits as _fetch_kiro_usage_limits,
)
from src.services.provider.adapters.kiro.usage import (
parse_kiro_usage_response as _parse_kiro_usage_response,
)
# Kiro: 直接使用 auth_config 调用 getUsageLimits API
if not key.auth_config:
return {
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "缺少 Kiro 认证配置 (auth_config)",
}
# 解密 auth_config
try:
decrypted_config = crypto_service.decrypt(key.auth_config)
auth_config_data = json.loads(decrypted_config)
except Exception:
return {
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "无法解密 auth_config可能是加密密钥已更改",
}
# 获取代理配置
proxy_config = getattr(provider, "proxy", None)
# 调用 Kiro getUsageLimits API
try:
result = await _fetch_kiro_usage_limits(
auth_config=auth_config_data,
proxy_config=proxy_config,
)
except RuntimeError as e:
error_msg = str(e)
# 检查是否需要标记账号异常
if "401" in error_msg or "认证失败" in error_msg:
key.oauth_invalid_at = datetime.now(timezone.utc)
key.oauth_invalid_reason = "Kiro Token 无效或已过期"
db.commit()
logger.warning("[KIRO_QUOTA] Key {} Token 无效,已标记为异常", key.id)
return {
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": error_msg,
}
usage_data = result.get("usage_data")
updated_auth_config = result.get("updated_auth_config")
# 解析限额信息
metadata = _parse_kiro_usage_response(usage_data)
if metadata:
# 收集元数据,稍后统一更新数据库(存储到 kiro 子对象)
metadata_updates[key.id] = {"kiro": metadata}
# 如果 auth_config 有更新(例如 token 刷新),也需要更新
if updated_auth_config:
try:
new_auth_config_json = json.dumps(updated_auth_config)
key.auth_config = crypto_service.encrypt(new_auth_config_json)
except Exception as exc:
logger.warning("更新 auth_config 失败 (key={}): {}", key.id, exc)
return {
"key_id": key.id,
"key_name": key.name,
"status": "success",
"metadata": metadata,
}
# 响应成功但没有限额信息
return {
"key_id": key.id,
"key_name": key.name,
"status": "no_metadata",
"message": "响应中未包含限额信息",
}
except Exception as e:
logger.error("刷新 Key {} 限额失败: {}", key.id, e)
return {

View File

@@ -1540,6 +1540,224 @@ class AdminModelMappingCacheStatsAdapter(AdminApiAdapter):
raise HTTPException(status_code=500, detail=f"获取统计失败: {exc}")
# ==================== Redis 缓存分类管理 ====================
# 所有已知的 Redis 缓存分类
# 格式: (category_key, display_name, redis_pattern, description)
# 注意: redis_pattern 必须与各模块实际使用的 key 前缀保持一致。
# 新增或修改缓存 key 前缀时,请同步更新此列表。
_CACHE_CATEGORIES: list[tuple[str, str, str, str]] = [
("upstream_models", "上游模型", "upstream_models:*", "Provider 上游获取的模型列表缓存"),
("model_id", "模型 ID", "model:id:*", "Model 按 ID 缓存"),
(
"model_provider_global",
"模型映射",
"model:provider_global:*",
"Provider-GlobalModel 模型映射缓存",
),
("global_model", "全局模型", "global_model:*", "GlobalModel 缓存ID/名称/解析)"),
("models_list", "模型列表", "models:list:*", "/v1/models 端点模型列表缓存"),
("user", "用户", "user:*", "用户信息缓存ID/Email"),
("apikey", "API Key", "apikey:*", "API Key 认证缓存Hash/Auth"),
("api_key_id", "API Key ID", "api_key:id:*", "API Key 按 ID 缓存"),
("cache_affinity", "缓存亲和性", "cache_affinity:*", "请求路由亲和性缓存"),
("provider_billing", "Provider 计费", "provider:billing_type:*", "Provider 计费类型缓存"),
(
"provider_rate",
"Provider 费率",
"provider_api_key:rate_multiplier:*",
"ProviderAPIKey 费率倍数缓存",
),
("provider_balance", "Provider 余额", "provider_ops:balance:*", "Provider 余额查询缓存"),
("health", "健康检查", "health:*", "端点健康状态缓存"),
("endpoint_status", "端点状态", "endpoint_status:*", "用户端点状态缓存"),
("dashboard", "仪表盘", "dashboard:*", "仪表盘统计缓存"),
("activity_heatmap", "活动热力图", "activity_heatmap:*", "用户活动热力图缓存"),
("gemini_files", "Gemini 文件映射", "gemini_files:*", "Gemini Files API 文件-Key 映射缓存"),
("provider_oauth", "OAuth 状态", "provider_oauth_state:*", "Provider OAuth 授权流程临时状态"),
(
"oauth_refresh_lock",
"OAuth 刷新锁",
"provider_oauth_refresh_lock:*",
"OAuth Token 刷新分布式锁",
),
("concurrency_lock", "并发锁", "concurrency:*", "请求并发控制锁"),
]
@router.get("/redis-keys")
async def get_redis_cache_categories(
request: Request,
db: Session = Depends(get_db),
) -> Any:
"""
获取 Redis 缓存分类概览
扫描 Redis 中所有已知的缓存键模式,返回各分类的键数量。
用于管理员全局了解缓存使用情况。
**返回字段**:
- `status`: 状态ok
- `data`: 分类列表
- `categories`: 各分类信息数组
- `key`: 分类标识
- `name`: 显示名称
- `pattern`: Redis 键模式
- `description`: 描述
- `count`: 键数量
- `total_keys`: 总键数
"""
adapter = AdminRedisCacheCategoriesAdapter()
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
@router.delete("/redis-keys/{category}")
async def clear_redis_cache_category(
category: str,
request: Request,
db: Session = Depends(get_db),
) -> Any:
"""
清除指定分类的 Redis 缓存
根据分类标识清除该分类下的所有缓存键。
**路径参数**:
- `category`: 分类标识(如 upstream_models、user、dashboard 等)
**返回字段**:
- `status`: 状态ok
- `message`: 操作结果消息
- `category`: 分类标识
- `deleted_count`: 删除的键数量
"""
adapter = AdminClearRedisCacheCategoryAdapter(category=category)
return await pipeline.run(adapter=adapter, http_request=request, db=db, mode=adapter.mode)
class AdminRedisCacheCategoriesAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
import asyncio
from src.clients.redis_client import get_redis_client
try:
redis = await get_redis_client(require_redis=False)
if not redis:
return {
"status": "ok",
"data": {"available": False, "message": "Redis 未启用"},
}
async def _count_keys(pattern: str) -> int:
count = 0
async for _ in redis.scan_iter(match=pattern, count=500):
count += 1
return count
# 并行扫描所有分类的 key 数量,避免串行 20 次 SCAN
counts = await asyncio.gather(
*[_count_keys(pattern) for _, _, pattern, _ in _CACHE_CATEGORIES]
)
categories = []
total_keys = 0
for (cat_key, name, pattern, description), count in zip(_CACHE_CATEGORIES, counts):
categories.append(
{
"key": cat_key,
"name": name,
"pattern": pattern,
"description": description,
"count": count,
}
)
total_keys += count
context.add_audit_metadata(
action="redis_cache_categories",
total_keys=total_keys,
category_count=len(categories),
)
return {
"status": "ok",
"data": {
"available": True,
"categories": categories,
"total_keys": total_keys,
},
}
except Exception as exc:
logger.exception("获取 Redis 缓存分类失败: {}", exc)
raise HTTPException(status_code=500, detail="获取缓存分类失败,请检查 Redis 连接")
@dataclass
class AdminClearRedisCacheCategoryAdapter(AdminApiAdapter):
category: str
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
from src.clients.redis_client import get_redis_client
try:
# 查找分类
target = None
for cat_key, name, pattern, _desc in _CACHE_CATEGORIES:
if cat_key == self.category:
target = (cat_key, name, pattern)
break
if not target:
raise HTTPException(
status_code=404,
detail=f"未知的缓存分类: {self.category}",
)
cat_key, name, pattern = target
redis = await get_redis_client(require_redis=False)
if not redis:
raise HTTPException(status_code=503, detail="Redis 未启用")
keys_to_delete: list[str] = []
async for key in redis.scan_iter(match=pattern, count=200):
keys_to_delete.append(key)
deleted_count = 0
# 分批删除,避免单次 DELETE 命令阻塞 Redis 事件循环
batch_size = 1000
for i in range(0, len(keys_to_delete), batch_size):
batch = keys_to_delete[i : i + batch_size]
deleted_count += await redis.delete(*batch)
logger.warning(
"已清除 Redis 缓存分类(管理员操作): {} ({}), pattern={}, deleted={}",
name,
cat_key,
pattern,
deleted_count,
)
context.add_audit_metadata(
action="redis_cache_clear_category",
category=cat_key,
category_name=name,
pattern=pattern,
deleted_count=deleted_count,
)
return {
"status": "ok",
"message": f"已清除 {name} 缓存",
"category": cat_key,
"deleted_count": deleted_count,
}
except HTTPException:
raise
except Exception as exc:
logger.exception("清除 Redis 缓存分类失败: {}", exc)
raise HTTPException(status_code=500, detail="清除缓存失败,请检查 Redis 连接")
class AdminClearAllModelMappingCacheAdapter(AdminApiAdapter):
async def handle(self, context: ApiRequestContext) -> dict[str, Any]: # type: ignore[override]
from src.clients.redis_client import get_redis_client

View File

@@ -184,6 +184,149 @@ def _parse_callback_params(callback_url: str) -> dict[str, str]:
return {str(k): str(v) for k, v in merged.items()}
# ==============================================================================
# Shared helpers
# ==============================================================================
def _get_provider_api_formats(provider: Provider) -> list[str]:
"""从 Provider 的活跃 endpoints 中提取所有 api_format。"""
return [
ep.api_format
for ep in provider.endpoints
if getattr(ep, "api_format", None) and getattr(ep, "is_active", False)
]
def _create_oauth_key(
db: Session,
*,
provider_id: str,
name: str,
access_token: str,
auth_config: dict[str, Any],
api_formats: list[str],
flush_only: bool = False,
) -> "ProviderAPIKey":
"""创建 OAuth Key 记录并持久化。
Args:
flush_only: True 时仅 flush批量导入场景False 时 commit + refresh。
"""
from src.models.database import ProviderAPIKey as ProviderAPIKeyModel
new_key = ProviderAPIKeyModel(
provider_id=provider_id,
name=name,
api_key=crypto_service.encrypt(access_token),
auth_type="oauth",
auth_config=crypto_service.encrypt(json.dumps(auth_config)),
api_formats=api_formats,
is_active=True,
)
db.add(new_key)
if flush_only:
db.flush()
else:
db.commit()
db.refresh(new_key)
return new_key
def _generate_kiro_key_name(cfg: Any) -> str:
"""根据 KiroAuthConfig 生成 Key 名称。"""
from src.services.provider.adapters.kiro.constants import DEFAULT_REGION
region = (getattr(cfg, "region", None) or "").strip() or DEFAULT_REGION
auth_method = getattr(cfg, "auth_method", None) or "social"
email = getattr(cfg, "email", None)
profile_arn = getattr(cfg, "profile_arn", None)
if email:
suffix = email
elif isinstance(profile_arn, str) and profile_arn.strip():
suffix = profile_arn.rsplit("/", 1)[-1]
else:
suffix = str(int(time.time()))
name = f"Kiro_{auth_method}_{region}_{suffix}"
return name[:100]
def _check_duplicate_oauth_account(
db: Session,
provider_id: str,
auth_config: dict[str, Any],
exclude_key_id: str | None = None,
) -> None:
"""
检查是否存在重复的 OAuth 账号
通过以下字段判断重复(按优先级):
- account_id: Codex 等使用的账号 ID
- email: OAuth 账号邮箱
- profile_arn: Kiro 使用的 profile ARN
Args:
db: 数据库 session
provider_id: Provider ID
auth_config: 新账号的 auth_config
exclude_key_id: 排除的 Key ID用于更新场景
Raises:
InvalidRequestException: 如果发现重复账号
"""
new_email = auth_config.get("email")
new_account_id = auth_config.get("account_id")
new_profile_arn = auth_config.get("profile_arn")
# 如果没有可用于识别的字段,跳过检查
if not new_email and not new_account_id and not new_profile_arn:
return
# 查询该 Provider 下所有 OAuth 类型的 Keys
query = db.query(ProviderAPIKey).filter(
ProviderAPIKey.provider_id == provider_id,
ProviderAPIKey.auth_type.in_(["oauth", "kiro"]), # kiro 也是 OAuth 类型
)
if exclude_key_id:
query = query.filter(ProviderAPIKey.id != exclude_key_id)
existing_keys = query.all()
for existing_key in existing_keys:
if not existing_key.auth_config:
continue
try:
decrypted_config = json.loads(
crypto_service.decrypt(existing_key.auth_config, silent=True)
)
existing_email = decrypted_config.get("email")
existing_account_id = decrypted_config.get("account_id")
existing_profile_arn = decrypted_config.get("profile_arn")
# 独立检查每个标识字段,避免 elif 链遗漏跨字段匹配
if new_account_id and existing_account_id and new_account_id == existing_account_id:
raise InvalidRequestException(
f"该 OAuth 账号已存在于当前 Provider 中(名称: {existing_key.name}"
)
if new_profile_arn and existing_profile_arn and new_profile_arn == existing_profile_arn:
raise InvalidRequestException(
f"该 Kiro 账号已存在于当前 Provider 中(名称: {existing_key.name}"
)
if new_email and existing_email and new_email == existing_email:
raise InvalidRequestException(
f"该 OAuth 账号 ({new_email}) 已存在于当前 Provider 中"
f"(名称: {existing_key.name}"
)
except InvalidRequestException:
raise
except Exception:
# 解密失败时跳过该 Key
continue
# ==============================================================================
# Routes
# ==============================================================================
@@ -444,6 +587,52 @@ async def refresh_oauth(
raise NotFoundException("Provider 不存在", "provider")
provider_type = _require_fixed_provider(provider)
# Kiro 使用自定义 token refresh 机制
if provider_type == ProviderType.KIRO.value:
from datetime import datetime, timezone
from src.services.provider.adapters.kiro.models.credentials import KiroAuthConfig
from src.services.provider.adapters.kiro.token_manager import refresh_access_token
encrypted_auth_config = getattr(key, "auth_config", None)
if not encrypted_auth_config:
raise InvalidRequestException("缺少 auth_config无法 refresh")
decrypted = crypto_service.decrypt(encrypted_auth_config)
parsed = json.loads(decrypted)
cfg = KiroAuthConfig.from_dict(parsed)
cfg.provider_type = ProviderType.KIRO.value
from src.services.proxy_node.resolver import resolve_effective_proxy
proxy_config = resolve_effective_proxy(
getattr(provider, "proxy", None), getattr(key, "proxy", None)
)
try:
access_token, new_cfg = await refresh_access_token(cfg, proxy_config=proxy_config)
except Exception as e:
# 标记为失效
key.oauth_invalid_at = datetime.now(timezone.utc)
key.oauth_invalid_reason = str(e)
db.commit()
logger.warning("Kiro Key {} token 刷新失败,已标记为失效: {}", key_id, e)
raise InvalidRequestException("Kiro token refresh 失败,请检查凭据是否有效")
# 更新 key
key.api_key = crypto_service.encrypt(access_token)
key.auth_config = crypto_service.encrypt(json.dumps(new_cfg.to_dict()))
key.oauth_invalid_at = None
key.oauth_invalid_reason = None
db.commit()
return CompleteOAuthResponse(
provider_type=provider_type,
expires_at=new_cfg.expires_at or None,
has_refresh_token=bool(new_cfg.refresh_token),
email=None,
)
try:
template = FIXED_PROVIDERS.get(ProviderType(provider_type))
except Exception:
@@ -463,6 +652,7 @@ async def refresh_oauth(
token_url = template.oauth.token_url
is_json = "anthropic.com" in token_url
scope_str = " ".join(template.oauth.scopes) if template.oauth.scopes else ""
if is_json:
body: dict[str, Any] = {
@@ -470,6 +660,8 @@ async def refresh_oauth(
"client_id": template.oauth.client_id,
"refresh_token": refresh_token,
}
if scope_str:
body["scope"] = scope_str
headers = {"Content-Type": "application/json", "Accept": "application/json"}
data = None
json_body = body
@@ -488,7 +680,11 @@ async def refresh_oauth(
data = form
json_body = None
proxy_config = getattr(provider, "proxy", None)
from src.services.proxy_node.resolver import resolve_effective_proxy
proxy_config = resolve_effective_proxy(
getattr(provider, "proxy", None), getattr(key, "proxy", None)
)
resp = await post_oauth_token(
provider_type=provider_type,
@@ -598,6 +794,9 @@ async def start_provider_oauth(
raise NotFoundException("Provider 不存在", "provider")
provider_type = _require_fixed_provider(provider)
if provider_type == ProviderType.KIRO.value:
raise InvalidRequestException("Kiro 不支持 OAuth 授权,请使用导入授权。")
try:
template = FIXED_PROVIDERS.get(ProviderType(provider_type))
except Exception:
@@ -685,6 +884,9 @@ async def complete_provider_oauth(
raise NotFoundException("Provider 不存在", "provider")
provider_type = _require_fixed_provider(provider)
if provider_type == ProviderType.KIRO.value:
raise InvalidRequestException("Kiro 不支持 OAuth 授权,请使用导入授权。")
try:
template = FIXED_PROVIDERS.get(ProviderType(provider_type))
except Exception:
@@ -774,29 +976,22 @@ async def complete_provider_oauth(
proxy_config=proxy_config,
)
# 检查是否存在重复的 OAuth 账号
_check_duplicate_oauth_account(db, provider_id, auth_config)
# 确定账号名称
name = (payload.name or "").strip()
if not name:
name = auth_config.get("email") or f"账号_{int(time.time())}"
# 从 Provider 的 endpoints 中提取所有 api_format 作为 Key 的支持格式
api_formats = [ep.api_format for ep in provider.endpoints if ep.api_format and ep.is_active]
# 创建 key
from src.models.database import ProviderAPIKey as ProviderAPIKeyModel
new_key = ProviderAPIKeyModel(
new_key = _create_oauth_key(
db,
provider_id=provider_id,
name=name,
api_key=crypto_service.encrypt(access_token),
auth_type="oauth",
auth_config=crypto_service.encrypt(json.dumps(auth_config)),
api_formats=api_formats,
is_active=True,
access_token=access_token,
auth_config=auth_config,
api_formats=_get_provider_api_formats(provider),
)
db.add(new_key)
db.commit()
db.refresh(new_key)
return ProviderCompleteOAuthResponse(
key_id=str(new_key.id),
@@ -812,11 +1007,128 @@ async def complete_provider_oauth(
# ==============================================================================
def _parse_tokens_input(raw_input: str) -> list[str]:
"""
解析通用 Token 导入输入,支持多种格式。
支持的格式:
1. 单个 Token 字符串
2. JSON 数组: ["token1", "token2", ...]
3. 纯 Token 导入(一行一个): "token1\\ntoken2\\ntoken3"
返回: Token 字符串列表
"""
raw = raw_input.strip()
if not raw:
return []
result: list[str] = []
# 尝试解析为 JSON 数组
if raw.startswith("["):
try:
parsed = json.loads(raw)
if isinstance(parsed, list):
for item in parsed:
if isinstance(item, str) and item.strip():
result.append(item.strip())
return result
except json.JSONDecodeError:
pass # 不是有效 JSON继续尝试其他格式
# 纯 Token 导入(一行一个)
lines = raw.splitlines()
for line in lines:
token = line.strip()
if token and not token.startswith("#"): # 忽略空行和注释行
result.append(token)
return result
def _parse_kiro_import_input(raw_input: str) -> list[dict[str, Any]]:
"""
解析 Kiro 凭据导入输入。
支持的格式:
1. 扁平 JSON 对象: {"refresh_token": "...", "auth_method": "social", ...}
2. JSON 数组(批量): [{...}, {...}]
3. 纯 Token一行一个: "token1\\ntoken2"
返回: 凭据字典列表
"""
raw = raw_input.strip()
if not raw:
return []
# 尝试解析为 JSON
if raw.startswith("{") or raw.startswith("["):
try:
parsed = json.loads(raw)
if isinstance(parsed, list):
result: list[dict[str, Any]] = []
for item in parsed:
if isinstance(item, dict):
result.append(item)
elif isinstance(item, str) and item.strip():
result.append({"refreshToken": item.strip()})
return result
if isinstance(parsed, dict):
# 兼容嵌套格式: {"auth_config": {...}} / {"authConfig": {...}}
nested = parsed.get("auth_config") or parsed.get("authConfig")
if isinstance(nested, dict):
return [nested]
return [parsed]
except json.JSONDecodeError:
pass
# 纯 Token一行一个
return [
{"refreshToken": line.strip()}
for line in raw.splitlines()
if line.strip() and not line.strip().startswith("#")
]
class ImportRefreshTokenRequest(BaseModel):
refresh_token: str = Field(..., min_length=1, description="Refresh Token")
name: str | None = Field(None, max_length=100, description="账号名称(可选)")
class BatchImportRequest(BaseModel):
"""批量导入 Kiro 凭据请求"""
credentials: str = Field(
...,
min_length=1,
max_length=500_000,
description="凭据数据支持多种格式JSON 对象、JSON 数组、纯 Token一行一个",
)
class BatchImportResultItem(BaseModel):
"""单个凭据导入结果"""
index: int = Field(..., description="凭据在输入中的索引(从 0 开始)")
status: str = Field(..., description="状态success / error")
key_id: str | None = Field(None, description="创建的 Key ID成功时")
key_name: str | None = Field(None, description="创建的 Key 名称(成功时)")
auth_method: str | None = Field(None, description="认证类型(成功时)")
error: str | None = Field(None, description="错误信息(失败时)")
class BatchImportResponse(BaseModel):
"""批量导入响应"""
total: int = Field(..., description="总凭据数")
success: int = Field(..., description="成功导入数")
failed: int = Field(..., description="失败数")
results: list[BatchImportResultItem] = Field(..., description="每个凭据的导入结果")
@router.post(
"/providers/{provider_id}/import-refresh-token",
response_model=ProviderCompleteOAuthResponse,
@@ -836,6 +1148,60 @@ async def import_refresh_token(
raise NotFoundException("Provider 不存在", "provider")
provider_type = _require_fixed_provider(provider)
if provider_type == ProviderType.KIRO.value:
raw_import = payload.refresh_token.strip()
if not raw_import:
raise InvalidRequestException("Refresh Token 不能为空")
# 使用统一的解析函数
credentials = _parse_kiro_import_input(raw_import)
if not credentials:
raise InvalidRequestException("无法解析凭据数据")
# 单条导入只取第一个
raw_cfg = credentials[0]
from src.services.provider.adapters.kiro.models.credentials import KiroAuthConfig
from src.services.provider.adapters.kiro.token_manager import refresh_access_token
# 验证必需字段
is_valid, error_msg = KiroAuthConfig.validate_required_fields(raw_cfg)
if not is_valid:
raise InvalidRequestException(error_msg)
# 解析配置(自动推断 auth_method
cfg = KiroAuthConfig.from_dict(raw_cfg)
cfg.provider_type = ProviderType.KIRO.value
proxy_config = getattr(provider, "proxy", None)
try:
access_token, new_cfg = await refresh_access_token(cfg, proxy_config=proxy_config)
except Exception as e:
logger.warning("Kiro Refresh Token 验证失败: {}", e)
raise InvalidRequestException("Kiro Refresh Token 验证失败,请检查凭据是否有效")
# 检查是否存在重复的 Kiro 账号
_check_duplicate_oauth_account(db, provider_id, new_cfg.to_dict())
name = (payload.name or "").strip() or _generate_kiro_key_name(new_cfg)
new_key = _create_oauth_key(
db,
provider_id=provider_id,
name=name,
access_token=access_token,
auth_config=new_cfg.to_dict(),
api_formats=_get_provider_api_formats(provider),
)
return ProviderCompleteOAuthResponse(
key_id=str(new_key.id),
provider_type=provider_type,
expires_at=new_cfg.expires_at or None,
has_refresh_token=bool(new_cfg.refresh_token),
email=None,
)
try:
template = FIXED_PROVIDERS.get(ProviderType(provider_type))
except Exception:
@@ -847,6 +1213,7 @@ async def import_refresh_token(
refresh_token = payload.refresh_token.strip()
token_url = template.oauth.token_url
is_json = "anthropic.com" in token_url
scope_str = " ".join(template.oauth.scopes) if template.oauth.scopes else ""
if is_json:
body: dict[str, Any] = {
@@ -854,6 +1221,8 @@ async def import_refresh_token(
"client_id": template.oauth.client_id,
"refresh_token": refresh_token,
}
if scope_str:
body["scope"] = scope_str
headers = {"Content-Type": "application/json", "Accept": "application/json"}
data = None
json_body = body
@@ -863,6 +1232,8 @@ async def import_refresh_token(
"client_id": template.oauth.client_id,
"refresh_token": refresh_token,
}
if scope_str:
form["scope"] = scope_str
if template.oauth.client_secret:
form["client_secret"] = template.oauth.client_secret
headers = {
@@ -926,29 +1297,22 @@ async def import_refresh_token(
proxy_config=proxy_config,
)
# 检查是否存在重复的 OAuth 账号
_check_duplicate_oauth_account(db, provider_id, auth_config)
# 确定账号名称
name = (payload.name or "").strip()
if not name:
name = auth_config.get("email") or f"账号_{int(time.time())}"
# 从 Provider 的 endpoints 中提取所有 api_format 作为 Key 的支持格式
api_formats = [ep.api_format for ep in provider.endpoints if ep.api_format and ep.is_active]
# 创建 key
from src.models.database import ProviderAPIKey as ProviderAPIKeyModel
new_key = ProviderAPIKeyModel(
new_key = _create_oauth_key(
db,
provider_id=provider_id,
name=name,
api_key=crypto_service.encrypt(access_token),
auth_type="oauth",
auth_config=crypto_service.encrypt(json.dumps(auth_config)),
api_formats=api_formats,
is_active=True,
access_token=access_token,
auth_config=auth_config,
api_formats=_get_provider_api_formats(provider),
)
db.add(new_key)
db.commit()
db.refresh(new_key)
return ProviderCompleteOAuthResponse(
key_id=str(new_key.id),
@@ -957,3 +1321,399 @@ async def import_refresh_token(
has_refresh_token=bool(new_refresh_token),
email=auth_config.get("email"),
)
# ==============================================================================
# 通用批量导入(支持所有 OAuth Provider
# ==============================================================================
@router.post(
"/providers/{provider_id}/batch-import",
response_model=BatchImportResponse,
)
async def batch_import_oauth(
provider_id: str,
payload: BatchImportRequest,
db: Session = Depends(get_db),
_: User = Depends(require_admin),
) -> BatchImportResponse:
"""批量导入 OAuth 凭据(通用)。
支持的 Provider 类型Codex、Antigravity、GeminiCli、ClaudeCode、Kiro
支持多种格式:
1. JSON 数组: ["token1", "token2", ...]
2. 纯 Token 导入(一行一个)
3. Kiro 专用JSON 对象或对象数组(含 refreshToken/clientId 等字段)
批量导入时自动跳过错误,不中断导入。
"""
provider = db.query(Provider).filter(Provider.id == provider_id).first()
if not provider:
raise NotFoundException("Provider 不存在", "provider")
provider_type = _require_fixed_provider(provider)
# Kiro 使用专用逻辑
if provider_type == ProviderType.KIRO.value:
return await _batch_import_kiro_internal(
provider_id=provider_id,
provider=provider,
raw_credentials=payload.credentials,
db=db,
)
# 标准 OAuth ProviderCodex、Antigravity、GeminiCli、ClaudeCode
try:
template = FIXED_PROVIDERS.get(ProviderType(provider_type))
except Exception:
template = None
if not template:
raise InvalidRequestException(f"不支持的 provider_type: {provider_type}")
# 解析 Token 列表
tokens = _parse_tokens_input(payload.credentials)
if not tokens:
raise InvalidRequestException("未找到有效的 Token 数据")
api_formats = _get_provider_api_formats(provider)
proxy_config = getattr(provider, "proxy", None)
token_url = template.oauth.token_url
is_json = "anthropic.com" in token_url
scope_str = " ".join(template.oauth.scopes) if template.oauth.scopes else ""
results: list[BatchImportResultItem] = []
success_count = 0
failed_count = 0
for idx, refresh_token in enumerate(tokens):
try:
# 验证 Token 非空
if not refresh_token or len(refresh_token) < 10:
results.append(
BatchImportResultItem(
index=idx,
status="error",
error="Token 无效或过短",
)
)
failed_count += 1
continue
# 使用 refresh_token 换取 access_token
if is_json:
body: dict[str, Any] = {
"grant_type": "refresh_token",
"client_id": template.oauth.client_id,
"refresh_token": refresh_token,
}
if scope_str:
body["scope"] = scope_str
headers = {"Content-Type": "application/json", "Accept": "application/json"}
data = None
json_body = body
else:
form: dict[str, str] = {
"grant_type": "refresh_token",
"client_id": template.oauth.client_id,
"refresh_token": refresh_token,
}
if scope_str:
form["scope"] = scope_str
if template.oauth.client_secret:
form["client_secret"] = template.oauth.client_secret
headers = {
"Content-Type": "application/x-www-form-urlencoded",
"Accept": "application/json",
}
data = form
json_body = None
try:
resp = await post_oauth_token(
provider_type=provider_type,
token_url=token_url,
headers=headers,
data=data,
json_body=json_body,
proxy_config=proxy_config,
timeout_seconds=30.0,
)
except Exception as e:
results.append(
BatchImportResultItem(
index=idx,
status="error",
error=f"Token 刷新请求失败: {e}",
)
)
failed_count += 1
continue
if resp.status_code < 200 or resp.status_code >= 300:
error_reason = f"HTTP {resp.status_code}"
try:
error_body = resp.json()
if "error" in error_body:
error_reason = str(
error_body.get("error_description") or error_body.get("error")
)
except Exception:
error_reason = resp.text[:100] if resp.text else f"HTTP {resp.status_code}"
results.append(
BatchImportResultItem(
index=idx,
status="error",
error=f"Token 验证失败: {error_reason}",
)
)
failed_count += 1
continue
token_data = resp.json()
access_token = str(token_data.get("access_token") or "")
new_refresh_token = str(token_data.get("refresh_token") or "") or refresh_token
if not access_token:
results.append(
BatchImportResultItem(
index=idx,
status="error",
error="Token 刷新返回缺少 access_token",
)
)
failed_count += 1
continue
expires_in = token_data.get("expires_in")
expires_at: int | None = None
try:
if expires_in is not None:
expires_at = int(time.time()) + int(expires_in)
except Exception:
expires_at = None
# 构建 auth_config
auth_config: dict[str, Any] = {
"provider_type": provider_type,
"token_type": token_data.get("token_type"),
"refresh_token": new_refresh_token or None,
"expires_at": expires_at,
"scope": token_data.get("scope"),
"updated_at": int(time.time()),
}
# 获取额外信息email 等)
try:
auth_config = await enrich_auth_config(
provider_type=provider_type,
auth_config=auth_config,
token_response=token_data,
access_token=access_token,
proxy_config=proxy_config,
)
except Exception as e:
logger.warning("批量导入: enrich_auth_config 失败 (index={}): {}", idx, e)
# 不中断,继续使用基本 auth_config
# 检查是否存在重复
try:
_check_duplicate_oauth_account(db, provider_id, auth_config)
except InvalidRequestException as e:
results.append(
BatchImportResultItem(
index=idx,
status="error",
error=str(e),
)
)
failed_count += 1
continue
# 生成名称
email = auth_config.get("email")
if email:
name = f"{provider_type}_{email}"
else:
name = f"{provider_type}_{int(time.time())}_{idx}"
if len(name) > 100:
name = name[:100]
new_key = _create_oauth_key(
db,
provider_id=provider_id,
name=name,
access_token=access_token,
auth_config=auth_config,
api_formats=api_formats,
flush_only=True,
)
results.append(
BatchImportResultItem(
index=idx,
status="success",
key_id=str(new_key.id),
key_name=name,
)
)
success_count += 1
except Exception as e:
logger.error("批量导入 OAuth 凭据失败 (index={}): {}", idx, e)
results.append(
BatchImportResultItem(
index=idx,
status="error",
error=f"导入失败: {e}",
)
)
failed_count += 1
# 提交所有成功的记录
if success_count > 0:
db.commit()
logger.info(
"[BATCH_IMPORT] Provider {} ({}): 成功 {}/{}, 失败 {}",
provider_id,
provider_type,
success_count,
len(tokens),
failed_count,
)
return BatchImportResponse(
total=len(tokens),
success=success_count,
failed=failed_count,
results=results,
)
async def _batch_import_kiro_internal(
provider_id: str,
provider: Provider,
raw_credentials: str,
db: Session,
) -> BatchImportResponse:
"""Kiro 批量导入内部实现(供通用端点调用)。"""
from src.services.provider.adapters.kiro.models.credentials import KiroAuthConfig
from src.services.provider.adapters.kiro.token_manager import refresh_access_token
# 解析输入
credentials = _parse_kiro_import_input(raw_credentials)
if not credentials:
raise InvalidRequestException("未找到有效的凭据数据")
api_formats = _get_provider_api_formats(provider)
proxy_config = getattr(provider, "proxy", None)
results: list[BatchImportResultItem] = []
success_count = 0
failed_count = 0
for idx, cred in enumerate(credentials):
try:
# 验证必需字段
is_valid, error_msg = KiroAuthConfig.validate_required_fields(cred)
if not is_valid:
results.append(
BatchImportResultItem(
index=idx,
status="error",
error=error_msg,
)
)
failed_count += 1
continue
# 解析凭据配置
cfg = KiroAuthConfig.from_dict(cred)
cfg.provider_type = ProviderType.KIRO.value
# 刷新 Token 以验证有效性
try:
access_token, new_cfg = await refresh_access_token(cfg, proxy_config=proxy_config)
except Exception as e:
results.append(
BatchImportResultItem(
index=idx,
status="error",
error=f"Token 验证失败: {e}",
)
)
failed_count += 1
continue
# 检查是否存在重复
try:
_check_duplicate_oauth_account(db, provider_id, new_cfg.to_dict())
except InvalidRequestException as e:
results.append(
BatchImportResultItem(
index=idx,
status="error",
error=str(e),
)
)
failed_count += 1
continue
# 生成名称
name = _generate_kiro_key_name(new_cfg)
new_key = _create_oauth_key(
db,
provider_id=provider_id,
name=name,
access_token=access_token,
auth_config=new_cfg.to_dict(),
api_formats=api_formats,
flush_only=True,
)
results.append(
BatchImportResultItem(
index=idx,
status="success",
key_id=str(new_key.id),
key_name=name,
auth_method=new_cfg.auth_method or "social",
)
)
success_count += 1
except Exception as e:
logger.error("批量导入 Kiro 凭据失败 (index={}): {}", idx, e)
results.append(
BatchImportResultItem(
index=idx,
status="error",
error=f"导入失败: {e}",
)
)
failed_count += 1
# 提交所有成功的记录
if success_count > 0:
db.commit()
logger.info(
"[KIRO_BATCH_IMPORT] Provider {}: 成功 {}/{}, 失败 {}",
provider_id,
success_count,
len(credentials),
failed_count,
)
return BatchImportResponse(
total=len(credentials),
success=success_count,
failed=failed_count,
results=results,
)

View File

@@ -68,7 +68,6 @@ async def _resolve_key_auth(
api_key_value: str | None = None
auth_config: dict[str, Any] | None = None
if auth_type == "oauth":
endpoint_api_format = "gemini:chat" if provider_type == ProviderType.ANTIGRAVITY else None
try:

View File

@@ -293,7 +293,11 @@ class AdminCreateProviderAdapter(AdminApiAdapter):
# 有 envelope 包装的 Provider 类型(如 Antigravity、Codex需要格式转换来正确
# 解包上游响应,创建时默认开启 enable_format_conversion。
pt = (validated_data.provider_type or "custom").strip()
envelope_provider_types = {ProviderType.ANTIGRAVITY, ProviderType.CODEX}
envelope_provider_types = {
ProviderType.ANTIGRAVITY,
ProviderType.CODEX,
ProviderType.KIRO,
}
default_enable_format_conversion = pt in envelope_provider_types
# 创建 Provider 对象

View File

@@ -872,6 +872,28 @@ class AdminUsageRecordsAdapter(AdminApiAdapter):
elif self.status == "active":
# 活跃请求pending 或 streaming 状态
query = query.filter(Usage.status.in_(["pending", "streaming"]))
elif self.status == "has_retry":
# 发生重试:存在 retry_index > 0 的已执行候选
retry_subq = (
db.query(RequestCandidate.request_id)
.filter(
RequestCandidate.status.in_(["success", "failed"]),
RequestCandidate.retry_index > 0,
)
.distinct()
.subquery()
)
query = query.filter(Usage.request_id.in_(retry_subq))
elif self.status == "has_fallback":
# 发生转移:同一请求有多个不同 candidate_index 的已执行候选
fallback_subq = (
db.query(RequestCandidate.request_id)
.filter(RequestCandidate.status.in_(["success", "failed"]))
.group_by(RequestCandidate.request_id)
.having(func.count(func.distinct(RequestCandidate.candidate_index)) > 1)
.subquery()
)
query = query.filter(Usage.request_id.in_(fallback_subq))
if self.time_range:
start_utc, end_utc = self.time_range.to_utc_datetime_range()
query = query.filter(Usage.created_at >= start_utc, Usage.created_at < end_utc)