mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 10:57:03 +08:00
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 <[email protected]>
This commit is contained in:
+450
-62
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
+791
-31
@@ -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 Provider(Codex、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,
|
||||
)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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 对象
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -422,6 +422,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
用于根据目标模型的特性对请求体做最终调整,例如:
|
||||
- 图像生成模型需要移除不兼容的 tools/system_instruction 并注入 imageConfig
|
||||
- 特定模型需要注入/移除某些字段
|
||||
- Gemini 格式:清理无效 parts 和合并连续同角色 contents
|
||||
|
||||
此方法在流式和非流式路径中均会被调用,且 mapped_model 已确定。
|
||||
|
||||
@@ -433,6 +434,18 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
Returns:
|
||||
调整后的请求体
|
||||
"""
|
||||
# Gemini 格式请求:清理无效 parts 和合并连续同角色 contents
|
||||
# 跨格式转换(如 Claude → Gemini)可能产生 thinking 等无法表示的块,
|
||||
# 导致 parts 为空或缺少有效 data-oneof 字段,被 Google API 拒绝。
|
||||
if provider_api_format and "gemini" in str(provider_api_format).lower():
|
||||
contents = request_body.get("contents")
|
||||
if isinstance(contents, list):
|
||||
from src.core.api_format.conversion.normalizers.gemini import (
|
||||
compact_gemini_contents,
|
||||
)
|
||||
|
||||
request_body["contents"] = compact_gemini_contents(contents)
|
||||
|
||||
return request_body
|
||||
|
||||
def _set_model_after_conversion(
|
||||
@@ -896,8 +909,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
|
||||
)
|
||||
if upstream_is_stream:
|
||||
# Ensure upstream returns SSE payload when in streaming mode.
|
||||
provider_headers["Accept"] = "text/event-stream"
|
||||
from src.core.api_format.headers import set_accept_if_absent
|
||||
|
||||
set_accept_if_absent(provider_headers)
|
||||
|
||||
ctx.provider_request_headers = provider_headers
|
||||
ctx.provider_request_body = provider_payload
|
||||
@@ -913,10 +927,15 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
# Capture the selected base_url from transport (used by some envelopes for failover).
|
||||
ctx.selected_base_url = envelope.capture_selected_base_url() if envelope else None
|
||||
|
||||
# 记录代理信息
|
||||
from src.services.proxy_node.resolver import get_proxy_label, resolve_proxy_info
|
||||
# 解析有效代理(Key 级别优先于 Provider 级别)
|
||||
from src.services.proxy_node.resolver import (
|
||||
get_proxy_label,
|
||||
resolve_effective_proxy,
|
||||
resolve_proxy_info,
|
||||
)
|
||||
|
||||
ctx.proxy_info = resolve_proxy_info(provider.proxy)
|
||||
effective_proxy = resolve_effective_proxy(provider.proxy, getattr(key, "proxy", None))
|
||||
ctx.proxy_info = resolve_proxy_info(effective_proxy)
|
||||
proxy_label = get_proxy_label(ctx.proxy_info)
|
||||
|
||||
logger.debug(
|
||||
@@ -931,9 +950,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
from src.services.proxy_node.resolver import build_post_kwargs, resolve_delegate_config
|
||||
|
||||
request_timeout_sync = provider.request_timeout or config.http_request_timeout
|
||||
delegate_cfg = resolve_delegate_config(provider.proxy)
|
||||
delegate_cfg = resolve_delegate_config(effective_proxy)
|
||||
http_client = await HTTPClientPool.get_upstream_client(
|
||||
delegate_cfg, proxy_config=provider.proxy
|
||||
delegate_cfg, proxy_config=effective_proxy
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -1125,13 +1144,13 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
# 优先使用 Provider 配置,否则使用全局配置
|
||||
request_timeout = provider.stream_first_byte_timeout or config.stream_first_byte_timeout
|
||||
|
||||
# 创建 HTTP 客户端(支持代理配置,从 Provider 读取)
|
||||
# 创建 HTTP 客户端(支持代理配置,Key 级别优先于 Provider 级别)
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.services.proxy_node.resolver import build_stream_kwargs, resolve_delegate_config
|
||||
|
||||
delegate_cfg = resolve_delegate_config(provider.proxy)
|
||||
delegate_cfg = resolve_delegate_config(effective_proxy)
|
||||
http_client = HTTPClientPool.create_upstream_stream_client(
|
||||
delegate_cfg, proxy_config=provider.proxy, timeout=timeout_config
|
||||
delegate_cfg, proxy_config=effective_proxy, timeout=timeout_config
|
||||
)
|
||||
|
||||
# 用于存储内部函数的结果(必须在函数定义前声明,供 nonlocal 使用)
|
||||
@@ -1546,8 +1565,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
|
||||
)
|
||||
if upstream_is_stream:
|
||||
# Ensure upstream returns SSE payload when forced to streaming mode.
|
||||
provider_hdrs["Accept"] = "text/event-stream"
|
||||
from src.core.api_format.headers import set_accept_if_absent
|
||||
|
||||
set_accept_if_absent(provider_hdrs)
|
||||
|
||||
provider_request_headers = provider_hdrs
|
||||
provider_request_body = provider_payload
|
||||
@@ -1563,10 +1583,15 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
# 非流式:必须在 build_provider_url 调用后立即缓存(避免 contextvar 被后续调用覆盖)
|
||||
selected_base_url_cached = envelope.capture_selected_base_url() if envelope else None
|
||||
|
||||
# 记录代理信息
|
||||
from src.services.proxy_node.resolver import get_proxy_label, resolve_proxy_info
|
||||
# 解析有效代理(Key 级别优先于 Provider 级别)
|
||||
from src.services.proxy_node.resolver import (
|
||||
get_proxy_label,
|
||||
resolve_effective_proxy,
|
||||
resolve_proxy_info,
|
||||
)
|
||||
|
||||
sync_proxy_info = resolve_proxy_info(provider.proxy)
|
||||
_effective_proxy = resolve_effective_proxy(provider.proxy, getattr(key, "proxy", None))
|
||||
sync_proxy_info = resolve_proxy_info(_effective_proxy)
|
||||
_proxy_label = get_proxy_label(sync_proxy_info)
|
||||
|
||||
logger.info(
|
||||
@@ -1576,7 +1601,7 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
)
|
||||
logger.debug(f" [{self.request_id}] 请求URL: {redact_url_for_log(url)}")
|
||||
|
||||
# 获取复用的 HTTP 客户端(支持代理配置,从 Provider 读取)
|
||||
# 获取复用的 HTTP 客户端(支持代理配置,Key 级别优先于 Provider 级别)
|
||||
# 注意:使用 get_proxy_client 复用连接池,不再每次创建新客户端
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.services.proxy_node.resolver import (
|
||||
@@ -1589,9 +1614,9 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
# 优先使用 Provider 配置,否则使用全局配置
|
||||
request_timeout = provider.request_timeout or config.http_request_timeout
|
||||
|
||||
delegate_cfg = resolve_delegate_config(provider.proxy)
|
||||
delegate_cfg = resolve_delegate_config(_effective_proxy)
|
||||
http_client = await HTTPClientPool.get_upstream_client(
|
||||
delegate_cfg, proxy_config=provider.proxy
|
||||
delegate_cfg, proxy_config=_effective_proxy
|
||||
)
|
||||
|
||||
# 注意:不使用 async with,因为复用的客户端不应该被关闭
|
||||
@@ -1643,8 +1668,16 @@ class ChatHandlerBase(BaseMessageHandler, ABC):
|
||||
|
||||
stream_resp.raise_for_status()
|
||||
|
||||
byte_iter = stream_resp.aiter_bytes()
|
||||
if provider_type == "kiro" and envelope and envelope.force_stream_rewrite():
|
||||
from src.services.provider.adapters.kiro.eventstream_rewriter import (
|
||||
apply_kiro_stream_rewrite,
|
||||
)
|
||||
|
||||
byte_iter = apply_kiro_stream_rewrite(byte_iter, model=str(model or ""))
|
||||
|
||||
internal_resp = await aggregate_upstream_stream_to_internal_response(
|
||||
stream_resp.aiter_bytes(),
|
||||
byte_iter,
|
||||
provider_api_format=provider_api_format,
|
||||
provider_name=str(provider.name),
|
||||
model=str(model or ""),
|
||||
|
||||
@@ -643,10 +643,23 @@ class CliAdapterBase(ApiAdapter):
|
||||
from src.core.provider_types import ProviderType
|
||||
|
||||
is_antigravity = provider_type == ProviderType.ANTIGRAVITY
|
||||
is_kiro = provider_type == ProviderType.KIRO
|
||||
is_oauth = auth_type == "oauth"
|
||||
|
||||
# ---- URL ----
|
||||
if is_antigravity:
|
||||
if is_kiro:
|
||||
# Kiro 需要替换 base_url 中的 {region} 占位符,并使用专用路径
|
||||
from src.services.provider.adapters.kiro.constants import (
|
||||
DEFAULT_REGION,
|
||||
KIRO_GENERATE_ASSISTANT_PATH,
|
||||
)
|
||||
|
||||
region = (decrypted_auth_config or {}).get("region") or DEFAULT_REGION
|
||||
effective_base_url = (
|
||||
base_url.replace("{region}", region) if "{region}" in base_url else base_url
|
||||
)
|
||||
url = f"{str(effective_base_url).rstrip('/')}{KIRO_GENERATE_ASSISTANT_PATH}"
|
||||
elif is_antigravity:
|
||||
# Antigravity 走 v1internal 端点,模型名在请求体 envelope 中,不在 URL 路径里
|
||||
from src.services.provider.adapters.antigravity.constants import (
|
||||
V1INTERNAL_PATH_TEMPLATE,
|
||||
@@ -673,6 +686,25 @@ class CliAdapterBase(ApiAdapter):
|
||||
if is_antigravity:
|
||||
merged_extra["User-Agent"] = _get_antigravity_ua()
|
||||
|
||||
# Kiro 需要特定的请求头
|
||||
if is_kiro:
|
||||
from src.services.provider.adapters.kiro.headers import build_generate_assistant_headers
|
||||
from src.services.provider.adapters.kiro.models.credentials import KiroAuthConfig
|
||||
from src.services.provider.adapters.kiro.token_manager import generate_machine_id
|
||||
|
||||
kiro_cfg = KiroAuthConfig.from_dict(decrypted_auth_config or {})
|
||||
region = kiro_cfg.region or DEFAULT_REGION
|
||||
machine_id = generate_machine_id(kiro_cfg)
|
||||
kiro_headers = build_generate_assistant_headers(
|
||||
host=f"q.{region}.amazonaws.com",
|
||||
access_token=api_key,
|
||||
machine_id=machine_id,
|
||||
kiro_version=kiro_cfg.kiro_version,
|
||||
system_version=kiro_cfg.system_version,
|
||||
node_version=kiro_cfg.node_version,
|
||||
)
|
||||
merged_extra.update(kiro_headers)
|
||||
|
||||
headers = cls.build_headers_with_extra(api_key, merged_extra if merged_extra else None)
|
||||
|
||||
# OAuth 统一处理:替换端点默认认证头为 Authorization: Bearer
|
||||
@@ -701,6 +733,21 @@ class CliAdapterBase(ApiAdapter):
|
||||
model=effective_model,
|
||||
)
|
||||
|
||||
# Kiro:用 conversationState envelope 包装请求体
|
||||
if is_kiro:
|
||||
from src.services.provider.adapters.kiro.converter import (
|
||||
convert_claude_messages_to_conversation_state,
|
||||
)
|
||||
|
||||
effective_model = model_name or request_data.get("model", "")
|
||||
conversation_state = convert_claude_messages_to_conversation_state(
|
||||
body,
|
||||
model=effective_model,
|
||||
)
|
||||
body = {"conversationState": conversation_state}
|
||||
if isinstance(kiro_cfg.profile_arn, str) and kiro_cfg.profile_arn.strip():
|
||||
body["profileArn"] = kiro_cfg.profile_arn.strip()
|
||||
|
||||
# ---- Header Rules ----
|
||||
if header_rules:
|
||||
from src.core.api_format import get_auth_config_for_endpoint as _get_auth_cfg
|
||||
|
||||
@@ -381,6 +381,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
用于根据目标模型的特性对请求体做最终调整,例如:
|
||||
- 图像生成模型需要移除不兼容的 tools/system_instruction 并注入 imageConfig
|
||||
- 特定模型需要注入/移除某些字段
|
||||
- Gemini 格式:清理无效 parts 和合并连续同角色 contents
|
||||
|
||||
此方法在流式和非流式路径中均会被调用,且 mapped_model 已确定。
|
||||
|
||||
@@ -392,6 +393,18 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
Returns:
|
||||
调整后的请求体
|
||||
"""
|
||||
# Gemini 格式请求:清理无效 parts 和合并连续同角色 contents
|
||||
# 跨格式转换(如 Claude → Gemini)可能产生 thinking 等无法表示的块,
|
||||
# 导致 parts 为空或缺少有效 data-oneof 字段,被 Google API 拒绝。
|
||||
if provider_api_format and "gemini" in str(provider_api_format).lower():
|
||||
contents = request_body.get("contents")
|
||||
if isinstance(contents, list):
|
||||
from src.core.api_format.conversion.normalizers.gemini import (
|
||||
compact_gemini_contents,
|
||||
)
|
||||
|
||||
request_body["contents"] = compact_gemini_contents(contents)
|
||||
|
||||
return request_body
|
||||
|
||||
@staticmethod
|
||||
@@ -872,8 +885,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
|
||||
)
|
||||
if upstream_is_stream:
|
||||
# Ensure upstream returns SSE payload when in streaming mode.
|
||||
provider_headers["Accept"] = "text/event-stream"
|
||||
from src.core.api_format.headers import set_accept_if_absent
|
||||
|
||||
set_accept_if_absent(provider_headers)
|
||||
|
||||
# 保存发送给 Provider 的请求信息(用于调试和统计)
|
||||
ctx.provider_request_headers = provider_headers
|
||||
@@ -890,11 +904,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
# Capture the selected base_url from transport (used by some envelopes for failover).
|
||||
ctx.selected_base_url = envelope.capture_selected_base_url() if envelope else None
|
||||
|
||||
# 记录代理信息(sync-bridge 路径,早于流式路径执行)
|
||||
# 解析有效代理(Key 级别优先于 Provider 级别)
|
||||
from src.services.proxy_node.resolver import get_proxy_label as _gpl
|
||||
from src.services.proxy_node.resolver import resolve_effective_proxy as _rep
|
||||
from src.services.proxy_node.resolver import resolve_proxy_info as _rpi
|
||||
|
||||
ctx.proxy_info = _rpi(provider.proxy)
|
||||
effective_proxy = _rep(provider.proxy, getattr(key, "proxy", None))
|
||||
ctx.proxy_info = _rpi(effective_proxy)
|
||||
|
||||
# If upstream is forced to non-stream mode, we execute a sync request and then
|
||||
# simulate streaming to the client (sync -> stream bridge).
|
||||
@@ -903,9 +919,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
from src.services.proxy_node.resolver import build_post_kwargs, resolve_delegate_config
|
||||
|
||||
request_timeout_sync = provider.request_timeout or config.http_request_timeout
|
||||
delegate_cfg = resolve_delegate_config(provider.proxy)
|
||||
delegate_cfg = resolve_delegate_config(effective_proxy)
|
||||
http_client = await HTTPClientPool.get_upstream_client(
|
||||
delegate_cfg, proxy_config=provider.proxy
|
||||
delegate_cfg, proxy_config=effective_proxy
|
||||
)
|
||||
|
||||
try:
|
||||
@@ -1096,13 +1112,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
f"timeout={request_timeout}s, 代理={_proxy_label}"
|
||||
)
|
||||
|
||||
# 创建 HTTP 客户端(支持代理配置,从 Provider 读取)
|
||||
# 创建 HTTP 客户端(支持代理配置,Key 级别优先于 Provider 级别)
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.services.proxy_node.resolver import build_stream_kwargs, resolve_delegate_config
|
||||
|
||||
delegate_cfg = resolve_delegate_config(provider.proxy)
|
||||
delegate_cfg = resolve_delegate_config(effective_proxy)
|
||||
http_client = HTTPClientPool.create_upstream_stream_client(
|
||||
delegate_cfg, proxy_config=provider.proxy, timeout=timeout_config
|
||||
delegate_cfg, proxy_config=effective_proxy, timeout=timeout_config
|
||||
)
|
||||
|
||||
# 用于存储内部函数的结果(必须在函数定义前声明,供 nonlocal 使用)
|
||||
@@ -1297,7 +1313,25 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
needs_conversion = True
|
||||
ctx.needs_conversion = True
|
||||
|
||||
async for chunk in stream_response.aiter_bytes():
|
||||
# Kiro 特殊处理:AWS Event Stream 二进制流需要重写为 SSE
|
||||
ctx_provider_type = str(ctx.provider_type or "").strip().lower()
|
||||
if ctx_provider_type == "kiro" and envelope and envelope.force_stream_rewrite():
|
||||
from src.services.provider.adapters.kiro.eventstream_rewriter import (
|
||||
apply_kiro_stream_rewrite,
|
||||
)
|
||||
|
||||
chunk_source: AsyncGenerator[bytes, None] = apply_kiro_stream_rewrite(
|
||||
stream_response.aiter_bytes(),
|
||||
model=str(ctx.model or ""),
|
||||
input_tokens=int(ctx.input_tokens or 0),
|
||||
)
|
||||
# Kiro 重写后输出的是 Claude SSE 格式,不需要再进行格式转换
|
||||
needs_conversion = False
|
||||
ctx.needs_conversion = False
|
||||
else:
|
||||
chunk_source = stream_response.aiter_bytes()
|
||||
|
||||
async for chunk in chunk_source:
|
||||
buffer += chunk
|
||||
# 处理缓冲区中的完整行
|
||||
while b"\n" in buffer:
|
||||
@@ -1307,8 +1341,10 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
line = decoder.decode(line_bytes + b"\n", False).rstrip("\n")
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"[{self.request_id}] UTF-8 解码失败: {e}, "
|
||||
f"bytes={line_bytes[:50]!r}"
|
||||
"[{}] UTF-8 解码失败: {}, bytes={!r}",
|
||||
self.request_id,
|
||||
e,
|
||||
line_bytes[:50],
|
||||
)
|
||||
continue
|
||||
|
||||
@@ -1333,11 +1369,14 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
if ctx.chunk_count > self.EMPTY_CHUNK_THRESHOLD and ctx.data_count == 0:
|
||||
elapsed = time.time() - last_data_time
|
||||
if elapsed > self.DATA_TIMEOUT:
|
||||
logger.warning(f"Provider '{ctx.provider_name}' 流超时且无数据")
|
||||
# 设置错误状态用于后续记录
|
||||
logger.warning("Provider '{}' 流超时且无数据", ctx.provider_name)
|
||||
ctx.status_code = 504
|
||||
ctx.error_message = "流式响应超时,未收到有效数据"
|
||||
ctx.upstream_response = f"流超时: Provider={ctx.provider_name}, elapsed={elapsed:.1f}s, chunk_count={ctx.chunk_count}, data_count=0"
|
||||
ctx.upstream_response = (
|
||||
f"流超时: Provider={ctx.provider_name}, "
|
||||
f"elapsed={elapsed:.1f}s, "
|
||||
f"chunk_count={ctx.chunk_count}, data_count=0"
|
||||
)
|
||||
error_event = {
|
||||
"type": "error",
|
||||
"error": {
|
||||
@@ -1364,16 +1403,16 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
self._mark_first_output(ctx, output_state)
|
||||
yield (line + "\n").encode("utf-8")
|
||||
|
||||
for event in events:
|
||||
self._handle_sse_event(
|
||||
ctx,
|
||||
event.get("event"),
|
||||
event.get("data") or "",
|
||||
record_chunk=not needs_conversion,
|
||||
)
|
||||
for event in events:
|
||||
self._handle_sse_event(
|
||||
ctx,
|
||||
event.get("event"),
|
||||
event.get("data") or "",
|
||||
record_chunk=not needs_conversion,
|
||||
)
|
||||
|
||||
if ctx.data_count > 0:
|
||||
last_data_time = time.time()
|
||||
if ctx.data_count > 0:
|
||||
last_data_time = time.time()
|
||||
|
||||
# 处理剩余事件
|
||||
for event in sse_parser.flush():
|
||||
@@ -1386,13 +1425,13 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
# 检查是否收到数据
|
||||
if ctx.data_count == 0:
|
||||
# 流已开始,无法抛出异常进行故障转移
|
||||
# 发送错误事件并记录日志
|
||||
logger.warning(f"Provider '{ctx.provider_name}' 返回空流式响应")
|
||||
# 设置错误状态用于后续记录
|
||||
logger.warning("Provider '{}' 返回空流式响应", ctx.provider_name)
|
||||
ctx.status_code = 503
|
||||
ctx.error_message = "上游服务返回了空的流式响应"
|
||||
ctx.upstream_response = f"空流式响应: Provider={ctx.provider_name}, chunk_count={ctx.chunk_count}, data_count=0"
|
||||
ctx.upstream_response = (
|
||||
f"空流式响应: Provider={ctx.provider_name}, "
|
||||
f"chunk_count={ctx.chunk_count}, data_count=0"
|
||||
)
|
||||
error_event = {
|
||||
"type": "error",
|
||||
"error": {
|
||||
@@ -1792,6 +1831,26 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
needs_conversion = True
|
||||
ctx.needs_conversion = True
|
||||
|
||||
# Kiro 特殊处理:AWS Event Stream 二进制流需要重写为 SSE
|
||||
ctx_provider_type = str(ctx.provider_type or "").strip().lower()
|
||||
if ctx_provider_type == "kiro" and envelope and envelope.force_stream_rewrite():
|
||||
from src.services.provider.adapters.kiro.eventstream_rewriter import (
|
||||
apply_kiro_stream_rewrite,
|
||||
)
|
||||
|
||||
byte_iterator = apply_kiro_stream_rewrite(
|
||||
byte_iterator,
|
||||
model=str(ctx.model or ""),
|
||||
input_tokens=int(ctx.input_tokens or 0),
|
||||
prefetched_chunks=list(prefetched_chunks) if prefetched_chunks else None,
|
||||
)
|
||||
prefetched_chunks = []
|
||||
|
||||
# Kiro 重写后输出的是 Claude SSE 格式
|
||||
# 客户端也是 Claude CLI,不需要再进行格式转换
|
||||
needs_conversion = False
|
||||
ctx.needs_conversion = False
|
||||
|
||||
# 先处理预读的字节块
|
||||
for chunk in prefetched_chunks:
|
||||
buffer += chunk
|
||||
@@ -2461,10 +2520,6 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
try:
|
||||
from src.models.database import ApiKey as ApiKeyModel
|
||||
|
||||
# 采集上游元数据(仅成功请求)
|
||||
if ctx.is_success():
|
||||
self._collect_upstream_metadata(bg_db, ctx)
|
||||
|
||||
user = bg_db.query(User).filter(User.id == ctx.user_id).first()
|
||||
api_key = bg_db.query(ApiKeyModel).filter(ApiKeyModel.id == ctx.api_key_id).first()
|
||||
|
||||
@@ -2719,19 +2774,6 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
except Exception as e:
|
||||
logger.exception("记录流式统计信息时出错")
|
||||
|
||||
@staticmethod
|
||||
def _collect_upstream_metadata(db: Session, ctx: StreamContext) -> None:
|
||||
"""采集上游元数据并更新 ProviderAPIKey.upstream_metadata(带节流)"""
|
||||
from src.services.provider.metadata_collectors import collect_and_save_upstream_metadata
|
||||
|
||||
collect_and_save_upstream_metadata(
|
||||
db,
|
||||
provider_type=ctx.provider_type or "",
|
||||
key_id=ctx.key_id or "",
|
||||
response_headers=ctx.response_headers or {},
|
||||
request_id=ctx.request_id or "",
|
||||
)
|
||||
|
||||
async def _record_stream_failure(
|
||||
self,
|
||||
ctx: StreamContext,
|
||||
@@ -2960,8 +3002,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
pre_computed_auth=auth_info.as_tuple() if auth_info else None,
|
||||
)
|
||||
if upstream_is_stream:
|
||||
# Ensure upstream returns SSE payload when forced to streaming mode.
|
||||
provider_headers["Accept"] = "text/event-stream"
|
||||
from src.core.api_format.headers import set_accept_if_absent
|
||||
|
||||
set_accept_if_absent(provider_headers)
|
||||
|
||||
# 保存发送给 Provider 的请求信息(用于调试和统计)
|
||||
provider_request_headers = provider_headers
|
||||
@@ -2978,10 +3021,15 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
# 非流式:必须在 build_provider_url 调用后立即缓存(避免 contextvar 被后续调用覆盖)
|
||||
selected_base_url_cached = envelope.capture_selected_base_url() if envelope else None
|
||||
|
||||
# 记录代理信息
|
||||
from src.services.proxy_node.resolver import get_proxy_label, resolve_proxy_info
|
||||
# 解析有效代理(Key 级别优先于 Provider 级别)
|
||||
from src.services.proxy_node.resolver import (
|
||||
get_proxy_label,
|
||||
resolve_effective_proxy,
|
||||
resolve_proxy_info,
|
||||
)
|
||||
|
||||
sync_proxy_info = resolve_proxy_info(provider.proxy)
|
||||
_effective_proxy = resolve_effective_proxy(provider.proxy, getattr(key, "proxy", None))
|
||||
sync_proxy_info = resolve_proxy_info(_effective_proxy)
|
||||
_proxy_label = get_proxy_label(sync_proxy_info)
|
||||
|
||||
logger.info(
|
||||
@@ -2992,7 +3040,7 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
f"代理={_proxy_label}"
|
||||
)
|
||||
|
||||
# 获取复用的 HTTP 客户端(支持代理配置,从 Provider 读取)
|
||||
# 获取复用的 HTTP 客户端(支持代理配置,Key 级别优先于 Provider 级别)
|
||||
# 注意:使用 get_proxy_client 复用连接池,不再每次创建新客户端
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.services.proxy_node.resolver import (
|
||||
@@ -3005,9 +3053,9 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
# 优先使用 Provider 配置,否则使用全局配置
|
||||
request_timeout = provider.request_timeout or config.http_request_timeout
|
||||
|
||||
delegate_cfg = resolve_delegate_config(provider.proxy)
|
||||
delegate_cfg = resolve_delegate_config(_effective_proxy)
|
||||
http_client = await HTTPClientPool.get_upstream_client(
|
||||
delegate_cfg, proxy_config=provider.proxy
|
||||
delegate_cfg, proxy_config=_effective_proxy
|
||||
)
|
||||
|
||||
# 注意:不使用 async with,因为复用的客户端不应该被关闭
|
||||
@@ -3060,8 +3108,16 @@ class CliMessageHandlerBase(BaseMessageHandler):
|
||||
|
||||
stream_resp.raise_for_status()
|
||||
|
||||
byte_iter = stream_resp.aiter_bytes()
|
||||
if provider_type == "kiro" and envelope and envelope.force_stream_rewrite():
|
||||
from src.services.provider.adapters.kiro.eventstream_rewriter import (
|
||||
apply_kiro_stream_rewrite,
|
||||
)
|
||||
|
||||
byte_iter = apply_kiro_stream_rewrite(byte_iter, model=str(model or ""))
|
||||
|
||||
internal_resp = await aggregate_upstream_stream_to_internal_response(
|
||||
stream_resp.aiter_bytes(),
|
||||
byte_iter,
|
||||
provider_api_format=provider_api_format,
|
||||
provider_name=str(provider.name),
|
||||
model=str(model or ""),
|
||||
|
||||
@@ -57,6 +57,73 @@ class ProviderAuthInfo:
|
||||
return (self.auth_header, self.auth_value)
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# OAuth Token Refresh helpers
|
||||
# ==============================================================================
|
||||
|
||||
|
||||
async def _acquire_refresh_lock(key_id: str) -> tuple[Any, bool]:
|
||||
"""尝试获取 OAuth refresh 分布式锁。
|
||||
|
||||
返回 ``(redis_client | None, got_lock)``。调用方在刷新完成后
|
||||
必须调用 :func:`_release_refresh_lock` 释放锁。
|
||||
"""
|
||||
redis = await get_redis_client(require_redis=False)
|
||||
lock_key = f"provider_oauth_refresh_lock:{key_id}"
|
||||
got_lock = False
|
||||
if redis is not None:
|
||||
try:
|
||||
got_lock = bool(await redis.set(lock_key, "1", ex=30, nx=True))
|
||||
except Exception:
|
||||
got_lock = False
|
||||
return redis, got_lock
|
||||
|
||||
|
||||
async def _release_refresh_lock(redis: Any, key_id: str) -> None:
|
||||
"""释放 OAuth refresh 分布式锁(best-effort)。"""
|
||||
if redis is not None:
|
||||
try:
|
||||
await redis.delete(f"provider_oauth_refresh_lock:{key_id}")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _persist_refreshed_token(
|
||||
key: Any,
|
||||
access_token: str,
|
||||
token_meta: dict[str, Any],
|
||||
) -> None:
|
||||
"""将刷新后的 access_token 和 auth_config 持久化到数据库。"""
|
||||
key.api_key = crypto_service.encrypt(access_token)
|
||||
key.auth_config = crypto_service.encrypt(json.dumps(token_meta))
|
||||
|
||||
sess = object_session(key)
|
||||
if sess is not None:
|
||||
sess.add(key)
|
||||
sess.commit()
|
||||
else:
|
||||
logger.warning(
|
||||
"[OAUTH_REFRESH] key {} refreshed but cannot persist (no session); "
|
||||
"next request will refresh again",
|
||||
key.id,
|
||||
)
|
||||
|
||||
|
||||
def _get_proxy_config(key: Any, endpoint: Any = None) -> Any:
|
||||
"""获取有效代理配置(Key 级别优先于 Provider 级别)。"""
|
||||
try:
|
||||
from src.services.proxy_node.resolver import resolve_effective_proxy
|
||||
|
||||
provider = getattr(key, "provider", None) or (
|
||||
getattr(endpoint, "provider", None) if endpoint else None
|
||||
)
|
||||
provider_proxy = getattr(provider, "proxy", None)
|
||||
key_proxy = getattr(key, "proxy", None)
|
||||
return resolve_effective_proxy(provider_proxy, key_proxy)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# 统一的头部配置常量
|
||||
# ==============================================================================
|
||||
@@ -591,6 +658,142 @@ def build_passthrough_request(
|
||||
)
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# OAuth Token Refresh logic (Kiro / Generic)
|
||||
# ==============================================================================
|
||||
|
||||
|
||||
async def _refresh_kiro_token(
|
||||
key: Any,
|
||||
endpoint: Any,
|
||||
token_meta: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""Kiro OAuth refresh: validate + call Kiro-specific refresh endpoint."""
|
||||
from src.core.exceptions import InvalidRequestException
|
||||
from src.services.provider.adapters.kiro.models.credentials import KiroAuthConfig
|
||||
from src.services.provider.adapters.kiro.token_manager import (
|
||||
refresh_access_token,
|
||||
validate_refresh_token,
|
||||
)
|
||||
|
||||
cfg = KiroAuthConfig.from_dict(token_meta or {})
|
||||
if not (cfg.refresh_token or "").strip():
|
||||
raise InvalidRequestException(
|
||||
"Kiro auth_config missing refresh_token; please re-import credentials."
|
||||
)
|
||||
|
||||
proxy_config = _get_proxy_config(key, endpoint)
|
||||
|
||||
validate_refresh_token(cfg.refresh_token)
|
||||
access_token, new_cfg = await refresh_access_token(
|
||||
cfg,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
new_meta = new_cfg.to_dict()
|
||||
new_meta["updated_at"] = int(time.time())
|
||||
|
||||
_persist_refreshed_token(key, access_token, new_meta)
|
||||
return new_meta
|
||||
|
||||
|
||||
async def _refresh_generic_oauth_token(
|
||||
key: Any,
|
||||
endpoint: Any,
|
||||
template: Any,
|
||||
provider_type: str,
|
||||
refresh_token: str,
|
||||
token_meta: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
"""Generic OAuth refresh via template (Codex, Antigravity, ClaudeCode, etc.)."""
|
||||
token_url = template.oauth.token_url
|
||||
is_json = "anthropic.com" in token_url
|
||||
|
||||
scopes = getattr(template.oauth, "scopes", None) or []
|
||||
scope_str = " ".join(scopes) if scopes else ""
|
||||
|
||||
if is_json:
|
||||
body: dict[str, Any] = {
|
||||
"grant_type": "refresh_token",
|
||||
"client_id": template.oauth.client_id,
|
||||
"refresh_token": str(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": str(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
|
||||
|
||||
proxy_config = _get_proxy_config(key, endpoint)
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
if 200 <= resp.status_code < 300:
|
||||
token = resp.json()
|
||||
access_token = str(token.get("access_token") or "")
|
||||
new_refresh_token = str(token.get("refresh_token") or "")
|
||||
expires_in = token.get("expires_in")
|
||||
new_expires_at: int | None = None
|
||||
try:
|
||||
if expires_in is not None:
|
||||
new_expires_at = int(time.time()) + int(expires_in)
|
||||
except Exception:
|
||||
new_expires_at = None
|
||||
|
||||
if access_token:
|
||||
token_meta["token_type"] = token.get("token_type")
|
||||
if new_refresh_token:
|
||||
token_meta["refresh_token"] = new_refresh_token
|
||||
token_meta["expires_at"] = new_expires_at
|
||||
token_meta["scope"] = token.get("scope")
|
||||
token_meta["updated_at"] = int(time.time())
|
||||
|
||||
token_meta = await enrich_auth_config(
|
||||
provider_type=provider_type,
|
||||
auth_config=token_meta,
|
||||
token_response=token,
|
||||
access_token=access_token,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
|
||||
_persist_refreshed_token(key, access_token, token_meta)
|
||||
else:
|
||||
logger.warning(
|
||||
"OAuth token refresh failed: provider={}, key_id={}, status={}",
|
||||
provider_type,
|
||||
getattr(key, "id", "?"),
|
||||
resp.status_code,
|
||||
)
|
||||
|
||||
return token_meta
|
||||
|
||||
|
||||
# ==============================================================================
|
||||
# Service Account 认证支持
|
||||
# ==============================================================================
|
||||
@@ -660,113 +863,19 @@ async def get_provider_auth(
|
||||
template = FIXED_PROVIDERS.get(ProviderType(provider_type))
|
||||
except Exception:
|
||||
template = None
|
||||
if template:
|
||||
redis = await get_redis_client(require_redis=False)
|
||||
lock_key = f"provider_oauth_refresh_lock:{key.id}"
|
||||
got_lock = False
|
||||
if redis is not None:
|
||||
try:
|
||||
got_lock = bool(await redis.set(lock_key, "1", ex=30, nx=True))
|
||||
except Exception:
|
||||
got_lock = False
|
||||
|
||||
if got_lock or redis is None:
|
||||
try:
|
||||
token_url = template.oauth.token_url
|
||||
is_json = "anthropic.com" in token_url
|
||||
|
||||
if is_json:
|
||||
body: dict[str, Any] = {
|
||||
"grant_type": "refresh_token",
|
||||
"client_id": template.oauth.client_id,
|
||||
"refresh_token": str(refresh_token),
|
||||
}
|
||||
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": str(refresh_token),
|
||||
}
|
||||
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
|
||||
|
||||
proxy_config = None
|
||||
try:
|
||||
provider = getattr(key, "provider", None)
|
||||
proxy_config = getattr(provider, "proxy", None)
|
||||
except Exception:
|
||||
proxy_config = None
|
||||
|
||||
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,
|
||||
redis, got_lock = await _acquire_refresh_lock(key.id)
|
||||
if got_lock or redis is None:
|
||||
try:
|
||||
if provider_type == ProviderType.KIRO.value:
|
||||
token_meta = await _refresh_kiro_token(key, endpoint, token_meta)
|
||||
elif template:
|
||||
token_meta = await _refresh_generic_oauth_token(
|
||||
key, endpoint, template, provider_type, refresh_token, token_meta
|
||||
)
|
||||
|
||||
if 200 <= resp.status_code < 300:
|
||||
token = resp.json()
|
||||
access_token = str(token.get("access_token") or "")
|
||||
new_refresh_token = str(token.get("refresh_token") or "")
|
||||
expires_in = token.get("expires_in")
|
||||
new_expires_at: int | None = None
|
||||
try:
|
||||
if expires_in is not None:
|
||||
new_expires_at = int(time.time()) + int(expires_in)
|
||||
except Exception:
|
||||
new_expires_at = None
|
||||
|
||||
if access_token:
|
||||
token_meta["token_type"] = token.get("token_type")
|
||||
if new_refresh_token:
|
||||
token_meta["refresh_token"] = new_refresh_token
|
||||
token_meta["expires_at"] = new_expires_at
|
||||
token_meta["scope"] = token.get("scope")
|
||||
token_meta["updated_at"] = int(time.time())
|
||||
|
||||
token_meta = await enrich_auth_config(
|
||||
provider_type=provider_type,
|
||||
auth_config=token_meta,
|
||||
token_response=token,
|
||||
access_token=access_token,
|
||||
proxy_config=proxy_config,
|
||||
)
|
||||
|
||||
key.api_key = crypto_service.encrypt(access_token)
|
||||
key.auth_config = crypto_service.encrypt(json.dumps(token_meta))
|
||||
|
||||
# 持久化:key 实体来自 DB session 时,尝试直接提交更新。
|
||||
sess = object_session(key)
|
||||
if sess is not None:
|
||||
sess.add(key)
|
||||
sess.commit()
|
||||
else:
|
||||
logger.warning(
|
||||
"[OAUTH_REFRESH] key {} 刷新成功但无法持久化(无绑定 session),"
|
||||
"下次请求将重新刷新",
|
||||
key.id,
|
||||
)
|
||||
finally:
|
||||
if got_lock and redis is not None:
|
||||
try:
|
||||
await redis.delete(lock_key)
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
if got_lock:
|
||||
await _release_refresh_lock(redis, key.id)
|
||||
except Exception:
|
||||
# 刷新失败不阻断请求;后续由上游返回 401 再触发管理端处理
|
||||
pass
|
||||
@@ -782,7 +891,6 @@ async def get_provider_auth(
|
||||
auth_value=f"Bearer {decrypted_key}",
|
||||
decrypted_auth_config=decrypted_auth_config,
|
||||
)
|
||||
|
||||
if auth_type == "vertex_ai":
|
||||
from src.core.vertex_auth import VertexAuthError, VertexAuthService
|
||||
|
||||
|
||||
@@ -383,6 +383,10 @@ class StreamProcessor:
|
||||
endpoint_sig=str(getattr(ctx, "provider_api_format", "") or ""),
|
||||
)
|
||||
envelope = behavior.envelope
|
||||
ctx_provider_type = str(getattr(ctx, "provider_type", "") or "").strip().lower()
|
||||
kiro_binary_stream = (
|
||||
ctx_provider_type == "kiro" and envelope and envelope.force_stream_rewrite()
|
||||
)
|
||||
buffer = b""
|
||||
line_count = 0
|
||||
should_stop = False
|
||||
@@ -402,6 +406,11 @@ class StreamProcessor:
|
||||
)
|
||||
prefetched_chunks.append(first_chunk)
|
||||
total_prefetched_bytes += len(first_chunk)
|
||||
|
||||
# Kiro upstream uses AWS Event Stream (binary). Do not attempt to split/decode lines here;
|
||||
# we only enforce TTFB and let StreamProcessor rewrite bytes later.
|
||||
if kiro_binary_stream:
|
||||
return prefetched_chunks
|
||||
buffer += first_chunk
|
||||
|
||||
# 继续读取剩余的预读数据
|
||||
@@ -619,6 +628,26 @@ class StreamProcessor:
|
||||
needs_conversion = True
|
||||
ctx.needs_conversion = True
|
||||
|
||||
ctx_provider_type = str(getattr(ctx, "provider_type", "") or "").strip().lower()
|
||||
if ctx_provider_type == "kiro" and envelope and envelope.force_stream_rewrite():
|
||||
from src.services.provider.adapters.kiro.eventstream_rewriter import (
|
||||
apply_kiro_stream_rewrite,
|
||||
)
|
||||
|
||||
byte_iterator = apply_kiro_stream_rewrite(
|
||||
byte_iterator,
|
||||
model=str(ctx.model or ""),
|
||||
input_tokens=int(ctx.input_tokens or 0),
|
||||
prefetched_chunks=list(prefetched_chunks) if prefetched_chunks else None,
|
||||
)
|
||||
prefetched_chunks = None
|
||||
|
||||
# Kiro 重写后输出的是 Claude SSE 格式(data: {...}\n\n)
|
||||
# 如果客户端也是 Claude 格式,则不需要再进行格式转换
|
||||
if client_family == "claude":
|
||||
needs_conversion = False
|
||||
ctx.needs_conversion = False
|
||||
|
||||
# 安全检查:needs_conversion 为 True 时,provider_format 必须有值
|
||||
if needs_conversion and not provider_format:
|
||||
logger.warning(
|
||||
|
||||
@@ -97,10 +97,6 @@ class StreamTelemetryRecorder:
|
||||
bg_db = next(db_gen)
|
||||
|
||||
try:
|
||||
# 采集上游元数据(仅成功请求,放在 writer 获取之前以确保执行)
|
||||
if ctx.is_success():
|
||||
self._collect_upstream_metadata(bg_db, ctx)
|
||||
|
||||
writer = await self._get_telemetry_writer(bg_db, ctx, response_time_ms)
|
||||
if writer is None:
|
||||
return
|
||||
@@ -537,19 +533,6 @@ class StreamTelemetryRecorder:
|
||||
error_message=ctx.error_message or f"HTTP {ctx.status_code}",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _collect_upstream_metadata(db: Session, ctx: StreamContext) -> None:
|
||||
"""采集上游元数据并更新 ProviderAPIKey.upstream_metadata(带节流)"""
|
||||
from src.services.provider.metadata_collectors import collect_and_save_upstream_metadata
|
||||
|
||||
collect_and_save_upstream_metadata(
|
||||
db,
|
||||
provider_type=ctx.provider_type or "",
|
||||
key_id=ctx.key_id or "",
|
||||
response_headers=ctx.response_headers or {},
|
||||
request_id=ctx.request_id or "",
|
||||
)
|
||||
|
||||
def _get_status_from_ctx(self, ctx: StreamContext) -> str:
|
||||
"""根据上下文获取状态字符串"""
|
||||
if ctx.is_success():
|
||||
|
||||
@@ -169,6 +169,17 @@ class GeminiChatHandler(ChatHandlerBase):
|
||||
is_image_gen_model,
|
||||
)
|
||||
|
||||
# Sanitize Gemini contents: strip parts without a valid data-oneof
|
||||
# field and merge consecutive same-role entries. This catches cases
|
||||
# missed by the normalizer (passthrough) or the antigravity envelope.
|
||||
from src.core.api_format.conversion.normalizers.gemini import (
|
||||
compact_gemini_contents,
|
||||
)
|
||||
|
||||
contents = request_body.get("contents")
|
||||
if isinstance(contents, list):
|
||||
request_body["contents"] = compact_gemini_contents(contents)
|
||||
|
||||
if not is_image_gen_model(mapped_model):
|
||||
return request_body
|
||||
return adapt_request_for_image_gen(request_body)
|
||||
|
||||
@@ -90,6 +90,17 @@ class GeminiCliMessageHandler(CliMessageHandlerBase):
|
||||
is_image_gen_model,
|
||||
)
|
||||
|
||||
# Sanitize Gemini contents: strip parts without a valid data-oneof
|
||||
# field and merge consecutive same-role entries. This catches cases
|
||||
# missed by the normalizer (passthrough) or the antigravity envelope.
|
||||
from src.core.api_format.conversion.normalizers.gemini import (
|
||||
compact_gemini_contents,
|
||||
)
|
||||
|
||||
contents = request_body.get("contents")
|
||||
if isinstance(contents, list):
|
||||
request_body["contents"] = compact_gemini_contents(contents)
|
||||
|
||||
if not is_image_gen_model(mapped_model):
|
||||
return request_body
|
||||
return adapt_request_for_image_gen(request_body)
|
||||
|
||||
@@ -60,6 +60,74 @@ from src.core.api_format.conversion.stream_events import (
|
||||
ToolCallDeltaEvent,
|
||||
)
|
||||
from src.core.api_format.conversion.stream_state import StreamState
|
||||
from src.core.api_format.schema_utils import clean_gemini_schema as _clean_gemini_schema
|
||||
|
||||
# Valid Gemini Part data-oneof field names (camelCase + snake_case).
|
||||
_VALID_PART_DATA_FIELDS = frozenset(
|
||||
{
|
||||
"text",
|
||||
"inlineData",
|
||||
"inline_data",
|
||||
"functionCall",
|
||||
"function_call",
|
||||
"functionResponse",
|
||||
"function_response",
|
||||
"fileData",
|
||||
"file_data",
|
||||
"executableCode",
|
||||
"executable_code",
|
||||
"codeExecutionResult",
|
||||
"code_execution_result",
|
||||
"videoMetadata",
|
||||
"video_metadata",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _is_valid_gemini_part(part: Any) -> bool:
|
||||
"""Return True if *part* has at least one recognised Gemini data field."""
|
||||
if not isinstance(part, dict) or not part:
|
||||
return False
|
||||
return bool(part.keys() & _VALID_PART_DATA_FIELDS)
|
||||
|
||||
|
||||
def compact_gemini_contents(contents: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
"""Strip invalid parts, drop empty contents, merge consecutive same-role.
|
||||
|
||||
Gemini requires every content to have at least one valid part (with a
|
||||
recognised data-oneof field), and strictly alternating user/model roles.
|
||||
Cross-format conversions (e.g. Responses API reasoning items) may produce
|
||||
contents with no Gemini-compatible parts, or consecutive same-role entries
|
||||
after filtering.
|
||||
"""
|
||||
# 1. Strip invalid parts, then drop contents with no valid parts remaining.
|
||||
# Use shallow copy to avoid mutating the caller's original dicts.
|
||||
non_empty: list[dict[str, Any]] = []
|
||||
for c in contents:
|
||||
parts = c.get("parts")
|
||||
if not isinstance(parts, list):
|
||||
continue
|
||||
valid_parts = [p for p in parts if _is_valid_gemini_part(p)]
|
||||
if valid_parts:
|
||||
c = {**c, "parts": valid_parts}
|
||||
non_empty.append(c)
|
||||
|
||||
# 2. Merge consecutive same-role entries.
|
||||
if not non_empty:
|
||||
return non_empty
|
||||
|
||||
merged: list[dict[str, Any]] = [non_empty[0]]
|
||||
for c in non_empty[1:]:
|
||||
if c.get("role") == merged[-1].get("role"):
|
||||
prev_parts = merged[-1].get("parts")
|
||||
if isinstance(prev_parts, list):
|
||||
prev_parts.extend(c.get("parts") or [])
|
||||
else:
|
||||
merged[-1]["parts"] = list(c.get("parts") or [])
|
||||
else:
|
||||
merged.append(c)
|
||||
|
||||
return merged
|
||||
|
||||
|
||||
class GeminiNormalizer(FormatNormalizer):
|
||||
@@ -206,22 +274,22 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
self._is_antigravity_thinking_enabled(internal) if is_antigravity else False
|
||||
)
|
||||
|
||||
# tools/tool_choice
|
||||
# tools/tool_choice — clean unsupported JSON Schema fields from parameters
|
||||
tools = None
|
||||
if internal.tools:
|
||||
tools = [
|
||||
{
|
||||
"function_declarations": [
|
||||
{
|
||||
"name": t.name,
|
||||
"description": t.description,
|
||||
"parameters": t.parameters or {},
|
||||
**(t.extra.get("gemini_function_declaration") or {}),
|
||||
}
|
||||
for t in internal.tools
|
||||
]
|
||||
func_decls: list[dict[str, Any]] = []
|
||||
for t in internal.tools:
|
||||
params = dict(t.parameters) if t.parameters else {}
|
||||
if params:
|
||||
_clean_gemini_schema(params)
|
||||
decl: dict[str, Any] = {
|
||||
"name": t.name,
|
||||
"description": t.description,
|
||||
"parameters": params,
|
||||
**(t.extra.get("gemini_function_declaration") or {}),
|
||||
}
|
||||
]
|
||||
func_decls.append(decl)
|
||||
tools = [{"function_declarations": func_decls}]
|
||||
|
||||
tool_config = None
|
||||
if internal.tool_choice:
|
||||
@@ -286,7 +354,7 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
if "thinking_config" in orig_gc and "thinkingConfig" not in generation_config:
|
||||
generation_config["thinkingConfig"] = orig_gc["thinking_config"]
|
||||
|
||||
contents: list[dict[str, Any]] = []
|
||||
raw_contents: list[dict[str, Any]] = []
|
||||
last_idx = len(internal.messages) - 1
|
||||
for idx, msg in enumerate(internal.messages):
|
||||
content = self._internal_message_to_content(
|
||||
@@ -327,7 +395,12 @@ class GeminiNormalizer(FormatNormalizer):
|
||||
content = dict(content)
|
||||
content["parts"] = [dummy_part, *parts]
|
||||
|
||||
contents.append(content)
|
||||
raw_contents.append(content)
|
||||
|
||||
# Drop contents with empty parts (e.g. reasoning-only messages that have
|
||||
# no Gemini-compatible representation) and merge consecutive same-role
|
||||
# contents — Gemini requires strictly alternating user/model turns.
|
||||
contents = compact_gemini_contents(raw_contents)
|
||||
|
||||
result: dict[str, Any] = {
|
||||
"contents": contents,
|
||||
|
||||
@@ -21,12 +21,13 @@ from src.core.api_format.metadata import (
|
||||
get_protected_keys_for_endpoint,
|
||||
)
|
||||
from src.core.api_format.signature import EndpointSignature, parse_signature_key
|
||||
from src.core.logger import logger
|
||||
|
||||
# =============================================================================
|
||||
# 头部常量定义
|
||||
# =============================================================================
|
||||
|
||||
# 转发给上游时需要剔除的头部(系统管理 + 认证替换)
|
||||
# 转发给上游时需要剔除的头部(系统管理 + 认证替换 + 客户端/代理元数据)
|
||||
UPSTREAM_DROP_HEADERS: frozenset[str] = frozenset(
|
||||
{
|
||||
# 认证头 - 会被替换为 Provider 的认证
|
||||
@@ -40,6 +41,13 @@ UPSTREAM_DROP_HEADERS: frozenset[str] = frozenset(
|
||||
"connection",
|
||||
# 编码头 - 避免客户端请求 brotli/zstd 但 httpx 不支持
|
||||
"accept-encoding",
|
||||
# 反向代理 / 网关注入的头部 - 属于本站基础设施,不应泄露给上游
|
||||
"x-real-ip",
|
||||
"x-real-proto",
|
||||
"x-forwarded-for",
|
||||
"x-forwarded-proto",
|
||||
"x-forwarded-host",
|
||||
"x-forwarded-port",
|
||||
}
|
||||
)
|
||||
|
||||
@@ -324,8 +332,23 @@ class HeaderBuilder:
|
||||
return self
|
||||
|
||||
def build(self) -> dict[str, str]:
|
||||
"""构建最终的头部字典"""
|
||||
return {original_key: value for original_key, value in self._headers.values()}
|
||||
"""构建最终的头部字典
|
||||
|
||||
Safety net: 跳过值中包含非 ASCII 字符的头部并记录警告,
|
||||
防止 httpx 发送时抛出 ``UnicodeEncodeError``。
|
||||
"""
|
||||
result: dict[str, str] = {}
|
||||
for original_key, value in self._headers.values():
|
||||
try:
|
||||
value.encode("ascii")
|
||||
except (UnicodeEncodeError, UnicodeDecodeError):
|
||||
logger.warning(
|
||||
"Dropping non-ASCII header before upstream request: {}",
|
||||
original_key,
|
||||
)
|
||||
continue
|
||||
result[original_key] = value
|
||||
return result
|
||||
|
||||
|
||||
def build_upstream_headers_for_endpoint(
|
||||
@@ -565,3 +588,18 @@ def get_extra_headers_from_endpoint(endpoint: Any) -> dict[str, str] | None:
|
||||
"""
|
||||
header_rules = getattr(endpoint, "header_rules", None)
|
||||
return extract_set_headers_from_rules(header_rules)
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# 请求头辅助工具
|
||||
# =============================================================================
|
||||
|
||||
|
||||
def set_accept_if_absent(headers: dict[str, str], value: str = "text/event-stream") -> None:
|
||||
"""Set the ``Accept`` header only if not already present (case-insensitive check).
|
||||
|
||||
Used by stream handlers to request SSE format from upstream without overriding
|
||||
provider-specific Accept headers (e.g. Kiro's ``application/vnd.amazon.eventstream``).
|
||||
"""
|
||||
if not any(k.lower() == "accept" for k in headers):
|
||||
headers["Accept"] = value
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
"""JSON Schema cleaning utilities shared across Gemini-compatible providers.
|
||||
|
||||
Google Gemini's function declaration API does not support certain JSON Schema
|
||||
fields. These must be stripped recursively from tool parameter schemas before
|
||||
forwarding to any Gemini-based upstream (native Gemini, Antigravity, etc.).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
# JSON Schema fields unsupported by Google Gemini's function declaration API.
|
||||
GEMINI_FORBIDDEN_SCHEMA_FIELDS: frozenset[str] = frozenset(
|
||||
{
|
||||
"$schema",
|
||||
"additionalProperties",
|
||||
"const",
|
||||
"contentEncoding",
|
||||
"contentMediaType",
|
||||
"default",
|
||||
"exclusiveMaximum",
|
||||
"exclusiveMinimum",
|
||||
"multipleOf",
|
||||
"patternProperties",
|
||||
"propertyNames",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def clean_gemini_schema(schema: dict[str, Any]) -> None:
|
||||
"""Recursively strip JSON Schema fields unsupported by Gemini.
|
||||
|
||||
Modifies *schema* in-place. Recurses into ``properties``, ``items``,
|
||||
and ``anyOf`` / ``oneOf`` / ``allOf`` sub-schemas.
|
||||
"""
|
||||
for field in GEMINI_FORBIDDEN_SCHEMA_FIELDS:
|
||||
schema.pop(field, None)
|
||||
|
||||
props = schema.get("properties")
|
||||
if isinstance(props, dict):
|
||||
for prop_schema in props.values():
|
||||
if isinstance(prop_schema, dict):
|
||||
clean_gemini_schema(prop_schema)
|
||||
|
||||
items = schema.get("items")
|
||||
if isinstance(items, dict):
|
||||
clean_gemini_schema(items)
|
||||
|
||||
for combo_key in ("anyOf", "oneOf", "allOf"):
|
||||
combo = schema.get(combo_key)
|
||||
if isinstance(combo, list):
|
||||
for sub in combo:
|
||||
if isinstance(sub, dict):
|
||||
clean_gemini_schema(sub)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"GEMINI_FORBIDDEN_SCHEMA_FIELDS",
|
||||
"clean_gemini_schema",
|
||||
]
|
||||
@@ -83,6 +83,25 @@ FIXED_PROVIDERS: dict[ProviderType, FixedProviderTemplate] = {
|
||||
use_pkce=True,
|
||||
),
|
||||
),
|
||||
ProviderType.KIRO: FixedProviderTemplate(
|
||||
provider_type=ProviderType.KIRO,
|
||||
display_name="Kiro",
|
||||
# Region is resolved from per-key auth_config (imported credentials).
|
||||
# Keep a templated base_url so endpoints remain fixed/locked.
|
||||
api_base_url="https://q.{region}.amazonaws.com",
|
||||
endpoint_signatures=["claude:cli"],
|
||||
# Kiro does not support Aether's OAuth flow; credentials are imported.
|
||||
# Keep a placeholder OAuth config so the fixed-provider template shape stays stable.
|
||||
oauth=FixedProviderOAuth(
|
||||
authorize_url="",
|
||||
token_url="",
|
||||
client_id="",
|
||||
client_secret="",
|
||||
scopes=[],
|
||||
redirect_uri="",
|
||||
use_pkce=False,
|
||||
),
|
||||
),
|
||||
ProviderType.GEMINI_CLI: FixedProviderTemplate(
|
||||
provider_type=ProviderType.GEMINI_CLI,
|
||||
display_name="GeminiCli",
|
||||
|
||||
@@ -17,6 +17,7 @@ class ProviderType(str, Enum):
|
||||
|
||||
CUSTOM = "custom"
|
||||
CLAUDE_CODE = "claude_code"
|
||||
KIRO = "kiro"
|
||||
CODEX = "codex"
|
||||
GEMINI_CLI = "gemini_cli"
|
||||
ANTIGRAVITY = "antigravity"
|
||||
|
||||
@@ -1487,6 +1487,11 @@ class ProviderAPIKey(Base):
|
||||
oauth_invalid_at = Column(DateTime(timezone=True), nullable=True) # 失效时间
|
||||
oauth_invalid_reason = Column(String(255), nullable=True) # 失效原因
|
||||
|
||||
# Key 级别的代理配置(覆盖 Provider 级别的代理设置)
|
||||
# 结构: {"node_id": "xxx", "enabled": true} 或 {"url": "socks5://...", "enabled": true}
|
||||
# null 表示使用 Provider 级别代理(默认行为)
|
||||
proxy = Column(JSON, nullable=True, default=None)
|
||||
|
||||
# 时间戳
|
||||
created_at = Column(
|
||||
DateTime(timezone=True), default=lambda: datetime.now(timezone.utc), nullable=False
|
||||
|
||||
@@ -418,6 +418,14 @@ class EndpointAPIKeyUpdate(BaseModel):
|
||||
model_exclude_patterns: list[str] | None = Field(
|
||||
default=None, description="模型排除规则(支持 * 和 ? 通配符),空表示不排除"
|
||||
)
|
||||
# Key 级别代理配置(覆盖 Provider 级别代理)
|
||||
# - 不提供:不更新
|
||||
# - 提供 null:清除 Key 级别代理,回退到 Provider 级别代理
|
||||
# - 提供 ProxyConfig:设置 Key 级别代理
|
||||
proxy: ProxyConfig | None = Field(
|
||||
default=None,
|
||||
description="Key 级别代理配置(覆盖 Provider 级别代理),null=使用 Provider 级别代理",
|
||||
)
|
||||
|
||||
@field_validator("api_formats")
|
||||
@classmethod
|
||||
@@ -590,6 +598,11 @@ class EndpointAPIKeyResponse(BaseModel):
|
||||
None, description="上游元数据(如 Codex 额度信息)"
|
||||
)
|
||||
|
||||
# Key 级别代理配置
|
||||
proxy: dict[str, Any] | None = Field(
|
||||
None, description="Key 级别代理配置(覆盖 Provider 级别代理)"
|
||||
)
|
||||
|
||||
# 时间戳
|
||||
last_used_at: datetime | None = None
|
||||
created_at: datetime
|
||||
|
||||
@@ -78,6 +78,11 @@ async def fetch_models_for_key(
|
||||
timeout_seconds: float = 30.0,
|
||||
) -> tuple[list[dict], list[str], bool, dict[str, Any] | None]:
|
||||
"""统一入口:按 provider_type 选择策略获取模型列表(可附带 upstream_metadata)。"""
|
||||
# Ensure provider plugins (including custom model fetchers) are registered.
|
||||
from src.services.provider.envelope import ensure_providers_bootstrapped
|
||||
|
||||
ensure_providers_bootstrapped()
|
||||
|
||||
fetcher = UpstreamModelsFetcherRegistry.get(ctx.provider_type) or _fetch_models_default
|
||||
return await fetcher(ctx, timeout_seconds)
|
||||
|
||||
|
||||
@@ -105,22 +105,10 @@ ANTIGRAVITY_SYSTEM_INSTRUCTION = (
|
||||
"**Proactiveness**"
|
||||
)
|
||||
|
||||
# ============== JSON Schema 禁止字段 ==============
|
||||
FORBIDDEN_SCHEMA_FIELDS = frozenset(
|
||||
{
|
||||
"multipleOf",
|
||||
"exclusiveMinimum",
|
||||
"exclusiveMaximum",
|
||||
"contentEncoding",
|
||||
"contentMediaType",
|
||||
}
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"ANTIGRAVITY_SYSTEM_INSTRUCTION",
|
||||
"DAILY_BASE_URL",
|
||||
"DUMMY_THOUGHT_SIGNATURE",
|
||||
"FORBIDDEN_SCHEMA_FIELDS",
|
||||
"HTTP_USER_AGENT",
|
||||
"MIN_SIGNATURE_LENGTH",
|
||||
"PROD_BASE_URL",
|
||||
|
||||
@@ -22,7 +22,6 @@ from typing import Any
|
||||
|
||||
from src.services.provider.adapters.antigravity.constants import (
|
||||
ANTIGRAVITY_SYSTEM_INSTRUCTION,
|
||||
FORBIDDEN_SCHEMA_FIELDS,
|
||||
)
|
||||
from src.services.provider.adapters.antigravity.constants import (
|
||||
REQUEST_USER_AGENT as ANTIGRAVITY_REQUEST_USER_AGENT,
|
||||
@@ -199,33 +198,27 @@ def _clean_tool_declarations(inner_request: dict[str, Any]) -> None:
|
||||
|
||||
def _clean_json_schema(schema: dict[str, Any]) -> None:
|
||||
"""递归移除 Gemini 不支持的 JSON Schema 字段。"""
|
||||
for field in FORBIDDEN_SCHEMA_FIELDS:
|
||||
schema.pop(field, None)
|
||||
from src.core.api_format.schema_utils import clean_gemini_schema
|
||||
|
||||
# Recurse into properties
|
||||
props = schema.get("properties")
|
||||
if isinstance(props, dict):
|
||||
for prop_schema in props.values():
|
||||
if isinstance(prop_schema, dict):
|
||||
_clean_json_schema(prop_schema)
|
||||
clean_gemini_schema(schema)
|
||||
|
||||
# Recurse into items
|
||||
items = schema.get("items")
|
||||
if isinstance(items, dict):
|
||||
_clean_json_schema(items)
|
||||
|
||||
# Recurse into additionalProperties
|
||||
addl = schema.get("additionalProperties")
|
||||
if isinstance(addl, dict):
|
||||
_clean_json_schema(addl)
|
||||
def _compact_contents(inner_request: dict[str, Any]) -> None:
|
||||
"""Strip invalid parts, drop empty contents, merge consecutive same-role.
|
||||
|
||||
# Recurse into anyOf / oneOf / allOf
|
||||
for combo_key in ("anyOf", "oneOf", "allOf"):
|
||||
combo = schema.get(combo_key)
|
||||
if isinstance(combo, list):
|
||||
for sub_schema in combo:
|
||||
if isinstance(sub_schema, dict):
|
||||
_clean_json_schema(sub_schema)
|
||||
Delegates to :func:`compact_gemini_contents` from the Gemini normalizer to
|
||||
avoid duplicating the validation logic.
|
||||
|
||||
This function modifies *inner_request* in-place.
|
||||
"""
|
||||
from src.core.api_format.conversion.normalizers.gemini import compact_gemini_contents
|
||||
|
||||
contents = inner_request.get("contents")
|
||||
if not isinstance(contents, list):
|
||||
return
|
||||
|
||||
result = compact_gemini_contents(contents)
|
||||
inner_request["contents"] = result
|
||||
|
||||
|
||||
def _inject_system_instruction(inner_request: dict[str, Any]) -> None:
|
||||
@@ -349,8 +342,9 @@ def wrap_v1internal_request(
|
||||
5. Thinking budget 处理
|
||||
6. 工具声明清洗(图像生成模型跳过)
|
||||
7. System Instruction 注入(图像生成模型跳过)
|
||||
8. 注入 sessionId(对齐 CLIProxyAPI)
|
||||
9. 构建 v1internal 信封
|
||||
8. 清理空 parts 的 contents 并合并连续同角色条目
|
||||
9. 注入 sessionId(对齐 CLIProxyAPI)
|
||||
10. 构建 v1internal 信封
|
||||
"""
|
||||
from src.api.handlers.gemini.image_gen import is_image_gen_model
|
||||
|
||||
@@ -385,7 +379,12 @@ def wrap_v1internal_request(
|
||||
inner_request.pop("system_instruction", None)
|
||||
request_type = "image_gen"
|
||||
|
||||
# 6. 注入 sessionId(对齐 CLIProxyAPI/sub2api)
|
||||
# 6. 清理空 parts 的 contents 并合并连续同角色条目
|
||||
# 跨格式转换(如 Responses API reasoning 块)可能产生空 parts 的 content,
|
||||
# Gemini API 要求每个 content 至少有一个有效 part,并且严格交替 user/model 角色。
|
||||
_compact_contents(inner_request)
|
||||
|
||||
# 7. 注入 sessionId(对齐 CLIProxyAPI/sub2api)
|
||||
if "sessionId" not in inner_request:
|
||||
inner_request["sessionId"] = _generate_stable_session_id(inner_request)
|
||||
|
||||
|
||||
@@ -309,6 +309,31 @@ async def fetch_models_antigravity(
|
||||
return models, [], True, upstream_metadata
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Export builder
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_AG_SKIP_KEYS = frozenset(
|
||||
{
|
||||
"access_token",
|
||||
"expires_at",
|
||||
"updated_at",
|
||||
"token_type",
|
||||
"scope",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def antigravity_export_builder(
|
||||
auth_config: dict[str, Any],
|
||||
upstream_metadata: dict[str, Any] | None,
|
||||
) -> dict[str, Any]:
|
||||
"""Antigravity 导出:保留 refresh_token / email / project_id / tier。"""
|
||||
return {
|
||||
k: v for k, v in auth_config.items() if k not in _AG_SKIP_KEYS and v is not None and v != ""
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Unified Registration
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -321,6 +346,7 @@ def register_all() -> None:
|
||||
from src.services.provider.adapters.antigravity.envelope import antigravity_v1internal_envelope
|
||||
from src.services.provider.behavior import register_behavior_variant
|
||||
from src.services.provider.envelope import register_envelope
|
||||
from src.services.provider.export import register_export_builder
|
||||
from src.services.provider.transport import register_transport_hook
|
||||
|
||||
# Envelope
|
||||
@@ -337,6 +363,9 @@ def register_all() -> None:
|
||||
# Auth
|
||||
register_auth_enricher("antigravity", enrich_antigravity)
|
||||
|
||||
# Export
|
||||
register_export_builder("antigravity", antigravity_export_builder)
|
||||
|
||||
# Model Fetcher
|
||||
UpstreamModelsFetcherRegistry.register(
|
||||
provider_types=["antigravity"],
|
||||
|
||||
@@ -1,97 +0,0 @@
|
||||
"""
|
||||
Codex Provider 元数据采集器
|
||||
|
||||
从响应头解析 Codex 额度/限流信息:
|
||||
- x-codex-plan-type: 套餐类型
|
||||
- x-codex-primary-*: 主限额窗口(通常 7 天)
|
||||
- x-codex-secondary-*: 次级限额窗口(通常 5 小时)
|
||||
- x-codex-credits-*: 积分信息
|
||||
"""
|
||||
|
||||
from typing import Any, ClassVar
|
||||
|
||||
from src.services.provider.metadata_collectors import ( # noqa: E501 — core registry
|
||||
MetadataCollector,
|
||||
)
|
||||
|
||||
|
||||
def _safe_float(value: str | None) -> float | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return float(value)
|
||||
except (ValueError, TypeError):
|
||||
return None
|
||||
|
||||
|
||||
def _safe_int(value: str | None) -> int | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return int(float(value))
|
||||
except (ValueError, TypeError):
|
||||
return None
|
||||
|
||||
|
||||
def _safe_bool(value: str | None) -> bool | None:
|
||||
if value is None:
|
||||
return None
|
||||
return value.lower() in ("true", "1", "yes")
|
||||
|
||||
|
||||
class CodexMetadataCollector(MetadataCollector):
|
||||
"""Codex 额度/限流元数据采集器"""
|
||||
|
||||
# 支持 codex 类型,通过响应头判断是否是 Codex
|
||||
PROVIDER_TYPES: ClassVar[list[str]] = ["codex"]
|
||||
|
||||
def parse_headers(self, headers: dict[str, str]) -> dict[str, Any] | None:
|
||||
# 大小写不敏感查找
|
||||
lower_headers = {k.lower(): v for k, v in headers.items()}
|
||||
|
||||
plan_type = lower_headers.get("x-codex-plan-type")
|
||||
if plan_type is None:
|
||||
# 没有 Codex 特征头,跳过
|
||||
return None
|
||||
|
||||
result: dict[str, Any] = {"plan_type": plan_type}
|
||||
|
||||
# 主限额窗口(7 天)
|
||||
primary_used = _safe_float(lower_headers.get("x-codex-primary-used-percent"))
|
||||
if primary_used is not None:
|
||||
result["primary_used_percent"] = primary_used
|
||||
primary_reset_seconds = _safe_int(lower_headers.get("x-codex-primary-reset-after-seconds"))
|
||||
if primary_reset_seconds is not None:
|
||||
result["primary_reset_seconds"] = primary_reset_seconds
|
||||
primary_reset_at = _safe_int(lower_headers.get("x-codex-primary-reset-at"))
|
||||
if primary_reset_at is not None:
|
||||
result["primary_reset_at"] = primary_reset_at
|
||||
primary_window = _safe_int(lower_headers.get("x-codex-primary-window-minutes"))
|
||||
if primary_window is not None:
|
||||
result["primary_window_minutes"] = primary_window
|
||||
|
||||
# 次级限额窗口(5 小时)
|
||||
secondary_used = _safe_float(lower_headers.get("x-codex-secondary-used-percent"))
|
||||
if secondary_used is not None:
|
||||
result["secondary_used_percent"] = secondary_used
|
||||
secondary_reset_seconds = _safe_int(
|
||||
lower_headers.get("x-codex-secondary-reset-after-seconds")
|
||||
)
|
||||
if secondary_reset_seconds is not None:
|
||||
result["secondary_reset_seconds"] = secondary_reset_seconds
|
||||
secondary_reset_at = _safe_int(lower_headers.get("x-codex-secondary-reset-at"))
|
||||
if secondary_reset_at is not None:
|
||||
result["secondary_reset_at"] = secondary_reset_at
|
||||
secondary_window = _safe_int(lower_headers.get("x-codex-secondary-window-minutes"))
|
||||
if secondary_window is not None:
|
||||
result["secondary_window_minutes"] = secondary_window
|
||||
|
||||
# 积分信息
|
||||
has_credits = _safe_bool(lower_headers.get("x-codex-credits-has-credits"))
|
||||
if has_credits is not None:
|
||||
result["has_credits"] = has_credits
|
||||
credits_balance = _safe_float(lower_headers.get("x-codex-credits-balance"))
|
||||
if credits_balance is not None:
|
||||
result["credits_balance"] = credits_balance
|
||||
|
||||
return result
|
||||
@@ -5,6 +5,7 @@
|
||||
- Transport Hook (URL 构建)
|
||||
- Auth Enricher (OAuth enrichment)
|
||||
- Behavior Variants (格式变体)
|
||||
- Model Fetcher (fixed catalog — Codex has no /v1/models endpoint)
|
||||
|
||||
新增 provider 时参照此文件创建对应的 plugin.py 即可。
|
||||
"""
|
||||
@@ -16,6 +17,40 @@ from urllib.parse import urlencode
|
||||
|
||||
from src.core.logger import logger
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixed model catalog
|
||||
# ---------------------------------------------------------------------------
|
||||
# Codex upstream (chatgpt.com/backend-api/codex) has no /v1/models endpoint.
|
||||
# Return a static list of known models.
|
||||
_CODEX_MODELS: list[dict[str, Any]] = [
|
||||
{
|
||||
"id": "gpt-5.2",
|
||||
"object": "model",
|
||||
"owned_by": "openai",
|
||||
"display_name": "gpt-5.2",
|
||||
},
|
||||
{
|
||||
"id": "gpt-5.2-codex",
|
||||
"object": "model",
|
||||
"owned_by": "openai",
|
||||
"display_name": "gpt-5.2-codex",
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
async def fetch_models_codex(
|
||||
ctx: Any,
|
||||
timeout_seconds: float, # noqa: ARG001
|
||||
) -> tuple[list[dict], list[str], bool, dict[str, Any] | None]:
|
||||
"""Return a fixed model catalog for Codex.
|
||||
|
||||
Codex upstream does not expose a ``/v1/models`` endpoint, so we skip the
|
||||
HTTP call entirely and return a hardcoded list.
|
||||
"""
|
||||
_ = ctx
|
||||
return list(_CODEX_MODELS), [], True, None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Transport Hook
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -87,6 +122,7 @@ async def enrich_codex(
|
||||
def register_all() -> None:
|
||||
"""一次性注册 Codex 的所有 hooks 到各通用 registry。"""
|
||||
from src.core.provider_oauth_utils import register_auth_enricher
|
||||
from src.services.model.upstream_fetcher import UpstreamModelsFetcherRegistry
|
||||
from src.services.provider.adapters.codex.envelope import codex_oauth_envelope
|
||||
from src.services.provider.behavior import register_behavior_variant
|
||||
from src.services.provider.envelope import register_envelope
|
||||
@@ -104,3 +140,12 @@ def register_all() -> None:
|
||||
|
||||
# Behavior
|
||||
register_behavior_variant("codex", same_format=True, cross_format=True)
|
||||
|
||||
# Export: Codex uses the default export builder (strip null + temp fields)
|
||||
# No need to register a custom one — the default in export.py suffices.
|
||||
|
||||
# Model Fetcher
|
||||
UpstreamModelsFetcherRegistry.register(
|
||||
provider_types=["codex"],
|
||||
fetcher=fetch_models_codex,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
"""Kiro provider adapter."""
|
||||
|
||||
__all__ = []
|
||||
@@ -0,0 +1,95 @@
|
||||
"""Kiro adapter constants.
|
||||
|
||||
Kiro upstream uses AWS Event Stream (binary frames) for streaming responses.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import platform
|
||||
|
||||
AWS_EVENTSTREAM_CONTENT_TYPE = "application/vnd.amazon.eventstream"
|
||||
|
||||
# Kiro API endpoints
|
||||
KIRO_GENERATE_ASSISTANT_PATH = "/generateAssistantResponse"
|
||||
KIRO_USAGE_LIMITS_PATH = "/getUsageLimits"
|
||||
|
||||
# Default AWS region when not specified in credentials
|
||||
DEFAULT_REGION = "us-east-1"
|
||||
|
||||
# Default client fingerprints used in headers (best-effort)
|
||||
DEFAULT_KIRO_VERSION = "0.8.0"
|
||||
DEFAULT_NODE_VERSION = "22.21.1"
|
||||
|
||||
|
||||
def _detect_system_version() -> str:
|
||||
system = platform.system().lower() or "other"
|
||||
release = platform.release() or "unknown"
|
||||
# Match KiroIDE style: darwin#24.6.0, windows#10.0.22631, linux#6.8.0-...
|
||||
return f"{system}#{release}"
|
||||
|
||||
|
||||
DEFAULT_SYSTEM_VERSION = _detect_system_version()
|
||||
|
||||
# Header constants
|
||||
KIRO_AGENT_MODE = "vibe"
|
||||
CODEWHISPERER_OPTOUT = "true"
|
||||
|
||||
# aws-sdk-js versions observed in kiro.rs
|
||||
AWS_SDK_JS_MAIN_VERSION = "1.0.27"
|
||||
AWS_SDK_JS_USAGE_VERSION = "1.0.0"
|
||||
|
||||
# Claude model context window used by kiro.rs to convert contextUsage percentage -> tokens
|
||||
CONTEXT_WINDOW_TOKENS = 200_000
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Chunked-write policy injected into tool descriptions and system prompt
|
||||
# ---------------------------------------------------------------------------
|
||||
# Kiro upstream has lower per-message size limits than standard Claude.
|
||||
# We inject instructions for Write/Edit tools and a system-level policy
|
||||
# so the model splits large writes into smaller chunks automatically.
|
||||
|
||||
WRITE_TOOL_DESCRIPTION_SUFFIX = (
|
||||
"- IMPORTANT: If the content to write exceeds 150 lines, you MUST only write "
|
||||
"the first 50 lines using this tool, then use `Edit` tool to append the "
|
||||
"remaining content in chunks of no more than 50 lines each. If needed, leave "
|
||||
"a unique placeholder to help append content. Do NOT attempt to write all "
|
||||
"content at once."
|
||||
)
|
||||
|
||||
EDIT_TOOL_DESCRIPTION_SUFFIX = (
|
||||
"- IMPORTANT: If the `new_string` content exceeds 50 lines, you MUST split "
|
||||
"it into multiple Edit calls, each replacing no more than 50 lines at a time. "
|
||||
"If used to append content, leave a unique placeholder to help append content. "
|
||||
"On the final chunk, do NOT include the placeholder."
|
||||
)
|
||||
|
||||
TOOL_DESCRIPTION_SUFFIXES: dict[str, str] = {
|
||||
"Write": WRITE_TOOL_DESCRIPTION_SUFFIX,
|
||||
"Edit": EDIT_TOOL_DESCRIPTION_SUFFIX,
|
||||
}
|
||||
|
||||
SYSTEM_CHUNKED_POLICY = (
|
||||
"When the Write or Edit tool has content size limits, always comply silently. "
|
||||
"Never suggest bypassing these limits via alternative tools. "
|
||||
"Never ask the user whether to switch approaches. "
|
||||
"Complete all chunked operations without commentary."
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"AWS_EVENTSTREAM_CONTENT_TYPE",
|
||||
"AWS_SDK_JS_MAIN_VERSION",
|
||||
"AWS_SDK_JS_USAGE_VERSION",
|
||||
"CODEWHISPERER_OPTOUT",
|
||||
"CONTEXT_WINDOW_TOKENS",
|
||||
"DEFAULT_KIRO_VERSION",
|
||||
"DEFAULT_NODE_VERSION",
|
||||
"DEFAULT_REGION",
|
||||
"DEFAULT_SYSTEM_VERSION",
|
||||
"EDIT_TOOL_DESCRIPTION_SUFFIX",
|
||||
"KIRO_AGENT_MODE",
|
||||
"KIRO_GENERATE_ASSISTANT_PATH",
|
||||
"KIRO_USAGE_LIMITS_PATH",
|
||||
"SYSTEM_CHUNKED_POLICY",
|
||||
"TOOL_DESCRIPTION_SUFFIXES",
|
||||
"WRITE_TOOL_DESCRIPTION_SUFFIX",
|
||||
]
|
||||
@@ -0,0 +1,42 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import contextvars
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class KiroRequestContext:
|
||||
"""Per-request context for the Kiro adapter.
|
||||
|
||||
This bridges data from `KiroEnvelope.wrap_request()` (which receives the
|
||||
decrypted auth_config + original request body) to other layers that only
|
||||
expose parameterless hooks (extra_headers) or transport hooks.
|
||||
"""
|
||||
|
||||
region: str
|
||||
machine_id: str
|
||||
kiro_version: str | None = None
|
||||
system_version: str | None = None
|
||||
node_version: str | None = None
|
||||
thinking_enabled: bool = False
|
||||
|
||||
|
||||
_kiro_request_context: contextvars.ContextVar[KiroRequestContext | None] = contextvars.ContextVar(
|
||||
"kiro_request_context",
|
||||
default=None,
|
||||
)
|
||||
|
||||
|
||||
def set_kiro_request_context(ctx: KiroRequestContext | None) -> None:
|
||||
_kiro_request_context.set(ctx)
|
||||
|
||||
|
||||
def get_kiro_request_context() -> KiroRequestContext | None:
|
||||
return _kiro_request_context.get()
|
||||
|
||||
|
||||
__all__ = [
|
||||
"KiroRequestContext",
|
||||
"get_kiro_request_context",
|
||||
"set_kiro_request_context",
|
||||
]
|
||||
@@ -0,0 +1,620 @@
|
||||
"""Claude Messages -> Kiro ConversationState converter (best-effort).
|
||||
|
||||
This mirrors `kiro.rs/src/anthropic/converter.rs` but focuses on the fields
|
||||
needed by generateAssistantResponse.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
from src.core.logger import logger
|
||||
from src.services.provider.adapters.kiro.constants import (
|
||||
SYSTEM_CHUNKED_POLICY as _SYSTEM_CHUNKED_POLICY,
|
||||
)
|
||||
from src.services.provider.adapters.kiro.constants import (
|
||||
TOOL_DESCRIPTION_SUFFIXES as _TOOL_DESCRIPTION_SUFFIXES,
|
||||
)
|
||||
|
||||
|
||||
def map_model(model: str) -> str | None:
|
||||
"""Map an Anthropic model name to a Kiro-compatible model ID.
|
||||
|
||||
Kiro upstream expects specific model IDs (e.g. ``claude-sonnet-4.5``).
|
||||
If *model* already matches a Kiro ID it is returned as-is; otherwise we
|
||||
attempt a best-effort fuzzy mapping. Returns ``None`` when the model is
|
||||
unrecognised.
|
||||
"""
|
||||
raw = str(model or "").strip()
|
||||
if not raw:
|
||||
return None
|
||||
|
||||
model_lower = raw.lower()
|
||||
|
||||
# Already a Kiro-native ID?
|
||||
_KIRO_NATIVE_IDS = {
|
||||
"claude-sonnet-4.5",
|
||||
"claude-opus-4.5",
|
||||
"claude-opus-4.6",
|
||||
"claude-haiku-4.5",
|
||||
}
|
||||
if model_lower in _KIRO_NATIVE_IDS:
|
||||
return model_lower
|
||||
|
||||
# Fuzzy mapping from Anthropic-style model names
|
||||
if "sonnet" in model_lower:
|
||||
return "claude-sonnet-4.5"
|
||||
if "opus" in model_lower:
|
||||
if "4-5" in model_lower or "4.5" in model_lower:
|
||||
return "claude-opus-4.5"
|
||||
return "claude-opus-4.6"
|
||||
if "haiku" in model_lower:
|
||||
return "claude-haiku-4.5"
|
||||
|
||||
# Unrecognised — pass through as-is (let the upstream decide)
|
||||
return raw
|
||||
|
||||
|
||||
def _extract_session_id(user_id: str) -> str | None:
|
||||
text = str(user_id or "")
|
||||
pos = text.find("session_")
|
||||
if pos < 0:
|
||||
return None
|
||||
session_part = text[pos + 8 :]
|
||||
if len(session_part) < 36:
|
||||
return None
|
||||
candidate = session_part[:36]
|
||||
if candidate.count("-") != 4:
|
||||
return None
|
||||
return candidate
|
||||
|
||||
|
||||
def _generate_thinking_prefix(request_body: dict[str, Any]) -> str | None:
|
||||
thinking = request_body.get("thinking")
|
||||
if not isinstance(thinking, dict):
|
||||
return None
|
||||
|
||||
thinking_type = str(thinking.get("type") or "").strip()
|
||||
if thinking_type == "enabled":
|
||||
budget = thinking.get("budget_tokens")
|
||||
try:
|
||||
budget_i = int(budget) if budget is not None else 0
|
||||
except Exception:
|
||||
budget_i = 0
|
||||
return (
|
||||
f"<thinking_mode>enabled</thinking_mode>"
|
||||
f"<max_thinking_length>{budget_i}</max_thinking_length>"
|
||||
)
|
||||
|
||||
if thinking_type == "adaptive":
|
||||
output_cfg = request_body.get("output_config")
|
||||
effort = "high"
|
||||
if isinstance(output_cfg, dict):
|
||||
eff = output_cfg.get("effort")
|
||||
if isinstance(eff, str) and eff.strip():
|
||||
effort = eff.strip()
|
||||
return (
|
||||
f"<thinking_mode>adaptive</thinking_mode>"
|
||||
f"<thinking_effort>{effort}</thinking_effort>"
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def _has_thinking_tags(content: str) -> bool:
|
||||
return "<thinking_mode>" in content or "<max_thinking_length>" in content
|
||||
|
||||
|
||||
def _system_to_text(system: Any) -> str:
|
||||
if system is None:
|
||||
return ""
|
||||
if isinstance(system, str):
|
||||
return system
|
||||
if isinstance(system, list):
|
||||
parts: list[str] = []
|
||||
for item in system:
|
||||
if isinstance(item, dict):
|
||||
t = item.get("text")
|
||||
if isinstance(t, str) and t:
|
||||
parts.append(t)
|
||||
else:
|
||||
# best-effort
|
||||
try:
|
||||
parts.append(str(item))
|
||||
except Exception:
|
||||
pass
|
||||
return "\n".join([p for p in parts if p])
|
||||
return ""
|
||||
|
||||
|
||||
def _get_image_format(media_type: str | None) -> str | None:
|
||||
if not isinstance(media_type, str) or "/" not in media_type:
|
||||
return None
|
||||
prefix, suffix = media_type.split("/", 1)
|
||||
if prefix != "image":
|
||||
return None
|
||||
suffix = suffix.strip().lower()
|
||||
if suffix in {"jpeg", "png", "gif", "webp"}:
|
||||
return suffix
|
||||
if suffix == "jpg":
|
||||
return "jpeg"
|
||||
return None
|
||||
|
||||
|
||||
def _process_message_content(
|
||||
content: Any,
|
||||
) -> tuple[str, list[dict[str, Any]], list[dict[str, Any]]]:
|
||||
"""Extract text/images/tool_results from a Claude content field."""
|
||||
text_parts: list[str] = []
|
||||
images: list[dict[str, Any]] = []
|
||||
tool_results: list[dict[str, Any]] = []
|
||||
|
||||
if isinstance(content, str):
|
||||
if content:
|
||||
text_parts.append(content)
|
||||
return "".join(text_parts), images, tool_results
|
||||
|
||||
if not isinstance(content, list):
|
||||
return "".join(text_parts), images, tool_results
|
||||
|
||||
for block in content:
|
||||
if not isinstance(block, dict):
|
||||
continue
|
||||
|
||||
btype = str(block.get("type") or "").strip()
|
||||
|
||||
if btype == "text":
|
||||
text = block.get("text")
|
||||
if isinstance(text, str) and text:
|
||||
text_parts.append(text)
|
||||
continue
|
||||
|
||||
if btype == "image":
|
||||
source = block.get("source")
|
||||
if not isinstance(source, dict):
|
||||
continue
|
||||
media_type = source.get("media_type") or source.get("mediaType")
|
||||
fmt = _get_image_format(media_type if isinstance(media_type, str) else None)
|
||||
data = source.get("data")
|
||||
if fmt and isinstance(data, str) and data:
|
||||
images.append({"format": fmt, "source": {"bytes": data}})
|
||||
continue
|
||||
|
||||
if btype == "tool_result":
|
||||
tool_use_id = block.get("tool_use_id") or block.get("toolUseId")
|
||||
if not isinstance(tool_use_id, str) or not tool_use_id.strip():
|
||||
continue
|
||||
|
||||
raw_content = block.get("content")
|
||||
if isinstance(raw_content, str):
|
||||
text = raw_content
|
||||
elif isinstance(raw_content, list):
|
||||
# Claude tool_result content blocks; keep only text parts.
|
||||
parts: list[str] = []
|
||||
for item in raw_content:
|
||||
if isinstance(item, dict) and item.get("type") == "text":
|
||||
t = item.get("text")
|
||||
if isinstance(t, str) and t:
|
||||
parts.append(t)
|
||||
text = "\n".join(parts)
|
||||
else:
|
||||
try:
|
||||
text = json.dumps(raw_content, ensure_ascii=False)
|
||||
except Exception:
|
||||
text = str(raw_content)
|
||||
|
||||
is_error = bool(block.get("is_error") or block.get("isError") or False)
|
||||
status = "error" if is_error else "success"
|
||||
|
||||
tool_results.append(
|
||||
{
|
||||
"toolUseId": tool_use_id.strip(),
|
||||
"content": [{"text": text or ""}],
|
||||
"status": status,
|
||||
"isError": bool(is_error),
|
||||
}
|
||||
)
|
||||
continue
|
||||
|
||||
return "".join(text_parts), images, tool_results
|
||||
|
||||
|
||||
def _convert_tools(tools: Any) -> list[dict[str, Any]]:
|
||||
if not isinstance(tools, list):
|
||||
return []
|
||||
|
||||
out: list[dict[str, Any]] = []
|
||||
for t in tools:
|
||||
if not isinstance(t, dict):
|
||||
continue
|
||||
name = t.get("name")
|
||||
if not isinstance(name, str) or not name.strip():
|
||||
continue
|
||||
|
||||
description = t.get("description")
|
||||
description_str = description if isinstance(description, str) else ""
|
||||
|
||||
# Inject chunked-write instructions for Write/Edit tools.
|
||||
suffix = _TOOL_DESCRIPTION_SUFFIXES.get(name.strip())
|
||||
if suffix:
|
||||
description_str = f"{description_str}\n{suffix}" if description_str else suffix
|
||||
|
||||
if len(description_str) > 10000:
|
||||
description_str = description_str[:10000]
|
||||
|
||||
input_schema = t.get("input_schema") or t.get("inputSchema") or {}
|
||||
if not isinstance(input_schema, dict):
|
||||
input_schema = {}
|
||||
|
||||
out.append(
|
||||
{
|
||||
"toolSpecification": {
|
||||
"name": name.strip(),
|
||||
"description": description_str,
|
||||
"inputSchema": {"json": input_schema},
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
return out
|
||||
|
||||
|
||||
def _create_placeholder_tool(name: str) -> dict[str, Any]:
|
||||
return {
|
||||
"toolSpecification": {
|
||||
"name": name,
|
||||
"description": "Tool used in conversation history",
|
||||
"inputSchema": {
|
||||
"json": {
|
||||
"$schema": "http://json-schema.org/draft-07/schema#",
|
||||
"type": "object",
|
||||
"properties": {},
|
||||
"required": [],
|
||||
"additionalProperties": True,
|
||||
}
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
def _convert_assistant_message(message: dict[str, Any]) -> dict[str, Any] | None:
|
||||
content = message.get("content")
|
||||
|
||||
tool_uses: list[dict[str, Any]] = []
|
||||
thinking_parts: list[str] = []
|
||||
text_parts: list[str] = []
|
||||
|
||||
if isinstance(content, str):
|
||||
if content:
|
||||
text_parts.append(content)
|
||||
elif isinstance(content, list):
|
||||
for block in content:
|
||||
if not isinstance(block, dict):
|
||||
continue
|
||||
btype = str(block.get("type") or "")
|
||||
if btype == "thinking":
|
||||
# Preserve thinking content so multi-turn context is not lost.
|
||||
t = block.get("thinking")
|
||||
if isinstance(t, str) and t:
|
||||
thinking_parts.append(t)
|
||||
elif btype == "text":
|
||||
t = block.get("text")
|
||||
if isinstance(t, str) and t:
|
||||
text_parts.append(t)
|
||||
elif btype == "tool_use":
|
||||
tool_use_id = block.get("id")
|
||||
name = block.get("name")
|
||||
if not isinstance(tool_use_id, str) or not tool_use_id.strip():
|
||||
continue
|
||||
if not isinstance(name, str) or not name.strip():
|
||||
continue
|
||||
inp = block.get("input")
|
||||
if not isinstance(inp, dict):
|
||||
inp = {}
|
||||
tool_uses.append(
|
||||
{
|
||||
"toolUseId": tool_use_id.strip(),
|
||||
"name": name.strip(),
|
||||
"input": inp,
|
||||
}
|
||||
)
|
||||
|
||||
# Combine thinking + text into final content.
|
||||
# Format: <thinking>...</thinking>\n\ntext
|
||||
thinking_str = "".join(thinking_parts)
|
||||
text_str = "".join(text_parts)
|
||||
|
||||
if thinking_str:
|
||||
if text_str:
|
||||
content_str = f"<thinking>{thinking_str}</thinking>\n\n{text_str}"
|
||||
else:
|
||||
content_str = f"<thinking>{thinking_str}</thinking>"
|
||||
else:
|
||||
content_str = text_str
|
||||
|
||||
if not content_str and tool_uses:
|
||||
content_str = " " # Kiro API requires non-empty content.
|
||||
|
||||
if not content_str and not tool_uses:
|
||||
return None
|
||||
|
||||
out: dict[str, Any] = {"content": content_str}
|
||||
if tool_uses:
|
||||
out["toolUses"] = tool_uses
|
||||
return out
|
||||
|
||||
|
||||
def convert_claude_messages_to_conversation_state(
|
||||
request_body: dict[str, Any],
|
||||
*,
|
||||
model: str,
|
||||
) -> dict[str, Any]:
|
||||
model_id = map_model(model)
|
||||
if not model_id:
|
||||
raise ValueError(f"kiro: model is required (got {model!r})")
|
||||
|
||||
messages = request_body.get("messages")
|
||||
if not isinstance(messages, list) or not messages:
|
||||
raise ValueError("kiro: empty messages")
|
||||
|
||||
conversation_id = None
|
||||
metadata = request_body.get("metadata")
|
||||
if isinstance(metadata, dict):
|
||||
user_id = metadata.get("user_id") or metadata.get("userId")
|
||||
if isinstance(user_id, str) and user_id:
|
||||
conversation_id = _extract_session_id(user_id)
|
||||
|
||||
if not conversation_id:
|
||||
conversation_id = str(uuid.uuid4())
|
||||
|
||||
agent_continuation_id = str(uuid.uuid4())
|
||||
|
||||
thinking_prefix = _generate_thinking_prefix(request_body)
|
||||
|
||||
history: list[dict[str, Any]] = []
|
||||
|
||||
# System injection: add as (user, assistant) pair.
|
||||
system_text = _system_to_text(request_body.get("system"))
|
||||
if system_text:
|
||||
# Append chunked-write policy so the model silently obeys tool limits.
|
||||
final_system = f"{system_text}\n{_SYSTEM_CHUNKED_POLICY}"
|
||||
if thinking_prefix and not _has_thinking_tags(system_text):
|
||||
final_system = f"{thinking_prefix}\n{final_system}"
|
||||
history.append(
|
||||
{
|
||||
"userInputMessage": {
|
||||
"content": final_system,
|
||||
"modelId": model_id,
|
||||
"origin": "AI_EDITOR",
|
||||
}
|
||||
}
|
||||
)
|
||||
history.append(
|
||||
{"assistantResponseMessage": {"content": "I will follow these instructions."}}
|
||||
)
|
||||
elif thinking_prefix:
|
||||
history.append(
|
||||
{
|
||||
"userInputMessage": {
|
||||
"content": thinking_prefix,
|
||||
"modelId": model_id,
|
||||
"origin": "AI_EDITOR",
|
||||
}
|
||||
}
|
||||
)
|
||||
history.append(
|
||||
{"assistantResponseMessage": {"content": "I will follow these instructions."}}
|
||||
)
|
||||
|
||||
# Build history from messages.
|
||||
# If the last message is assistant, include it in history (Kiro currentMessage
|
||||
# must be user; we synthesise one). Otherwise the last user message becomes
|
||||
# currentMessage and everything before it goes into history.
|
||||
last_msg = messages[-1]
|
||||
last_is_assistant = (
|
||||
isinstance(last_msg, dict) and str(last_msg.get("role") or "") == "assistant"
|
||||
)
|
||||
|
||||
if last_is_assistant:
|
||||
# All messages go into history; we'll synthesise a currentMessage later.
|
||||
history_end_index = len(messages)
|
||||
else:
|
||||
history_end_index = max(len(messages) - 1, 0)
|
||||
|
||||
user_buffer: list[dict[str, Any]] = []
|
||||
|
||||
def _flush_user_buffer() -> dict[str, Any] | None:
|
||||
nonlocal user_buffer
|
||||
if not user_buffer:
|
||||
return None
|
||||
|
||||
parts: list[str] = []
|
||||
images: list[dict[str, Any]] = []
|
||||
tool_results: list[dict[str, Any]] = []
|
||||
|
||||
for msg in user_buffer:
|
||||
text, imgs, results = _process_message_content(msg.get("content"))
|
||||
if text:
|
||||
parts.append(text)
|
||||
images.extend(imgs)
|
||||
tool_results.extend(results)
|
||||
|
||||
user_buffer = []
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
"content": "\n".join(parts),
|
||||
"modelId": model_id,
|
||||
"origin": "AI_EDITOR",
|
||||
}
|
||||
|
||||
if images:
|
||||
payload["images"] = images
|
||||
|
||||
if tool_results:
|
||||
payload["userInputMessageContext"] = {"toolResults": tool_results}
|
||||
|
||||
return {"userInputMessage": payload}
|
||||
|
||||
for i in range(history_end_index):
|
||||
msg = messages[i]
|
||||
if not isinstance(msg, dict):
|
||||
continue
|
||||
role = str(msg.get("role") or "")
|
||||
if role == "user":
|
||||
user_buffer.append(msg)
|
||||
continue
|
||||
if role == "assistant":
|
||||
user_item = _flush_user_buffer()
|
||||
if user_item is not None:
|
||||
history.append(user_item)
|
||||
assistant_item = _convert_assistant_message(msg)
|
||||
if assistant_item is not None:
|
||||
history.append({"assistantResponseMessage": assistant_item})
|
||||
continue
|
||||
|
||||
# trailing unpaired user messages in history
|
||||
tail_user = _flush_user_buffer()
|
||||
if tail_user is not None:
|
||||
history.append(tail_user)
|
||||
history.append({"assistantResponseMessage": {"content": "OK"}})
|
||||
|
||||
# Current message: last message as user input.
|
||||
if last_is_assistant:
|
||||
# Synthesise a minimal user continuation message.
|
||||
text_content = "Continue."
|
||||
images: list[dict[str, Any]] = []
|
||||
tool_results: list[dict[str, Any]] = []
|
||||
else:
|
||||
last = messages[-1]
|
||||
if not isinstance(last, dict) or str(last.get("role") or "") != "user":
|
||||
raise ValueError("kiro: last message must be user")
|
||||
text_content, images, tool_results = _process_message_content(last.get("content"))
|
||||
|
||||
tools = _convert_tools(request_body.get("tools"))
|
||||
|
||||
# Ensure tools referenced in history assistant toolUses are defined.
|
||||
# Also collect ids for tool_use / tool_result pairing validation.
|
||||
history_tool_names: set[str] = set()
|
||||
history_tool_results_ids: set[str] = set()
|
||||
history_tool_use_ids: set[str] = set()
|
||||
|
||||
for item in history:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
u = item.get("userInputMessage")
|
||||
if isinstance(u, dict):
|
||||
ctx = u.get("userInputMessageContext")
|
||||
if isinstance(ctx, dict):
|
||||
results = ctx.get("toolResults")
|
||||
if isinstance(results, list):
|
||||
for r in results:
|
||||
if isinstance(r, dict):
|
||||
tid = r.get("toolUseId")
|
||||
if isinstance(tid, str) and tid:
|
||||
history_tool_results_ids.add(tid)
|
||||
a = item.get("assistantResponseMessage")
|
||||
if isinstance(a, dict):
|
||||
uses = a.get("toolUses")
|
||||
if isinstance(uses, list):
|
||||
for tu in uses:
|
||||
if not isinstance(tu, dict):
|
||||
continue
|
||||
nm = tu.get("name")
|
||||
if isinstance(nm, str) and nm:
|
||||
history_tool_names.add(nm)
|
||||
tid = tu.get("toolUseId")
|
||||
if isinstance(tid, str) and tid:
|
||||
history_tool_use_ids.add(tid)
|
||||
|
||||
existing_tool_names = {
|
||||
str(t.get("toolSpecification", {}).get("name", "")).lower() for t in tools
|
||||
}
|
||||
|
||||
for tool_name in sorted(history_tool_names):
|
||||
if tool_name.lower() not in existing_tool_names:
|
||||
tools.append(_create_placeholder_tool(tool_name))
|
||||
|
||||
# Filter tool_results: only keep those with matching tool_use in history, and not duplicated.
|
||||
validated_tool_results: list[dict[str, Any]] = []
|
||||
current_tool_result_ids: set[str] = set()
|
||||
for r in tool_results:
|
||||
if not isinstance(r, dict):
|
||||
continue
|
||||
tid = r.get("toolUseId")
|
||||
if not isinstance(tid, str) or not tid:
|
||||
continue
|
||||
if tid not in history_tool_use_ids:
|
||||
continue
|
||||
if tid in history_tool_results_ids:
|
||||
continue
|
||||
validated_tool_results.append(r)
|
||||
current_tool_result_ids.add(tid)
|
||||
|
||||
# Remove orphaned tool_uses from history.
|
||||
# Kiro API requires every tool_use to have a matching tool_result; otherwise
|
||||
# it returns 400 Bad Request.
|
||||
orphaned_tool_use_ids = (
|
||||
history_tool_use_ids - history_tool_results_ids - current_tool_result_ids
|
||||
)
|
||||
if orphaned_tool_use_ids:
|
||||
logger.warning(
|
||||
"kiro: removing {} orphaned tool_use(s) from history: {}",
|
||||
len(orphaned_tool_use_ids),
|
||||
orphaned_tool_use_ids,
|
||||
)
|
||||
for item in history:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
a = item.get("assistantResponseMessage")
|
||||
if not isinstance(a, dict):
|
||||
continue
|
||||
uses = a.get("toolUses")
|
||||
if not isinstance(uses, list):
|
||||
continue
|
||||
filtered = [
|
||||
u
|
||||
for u in uses
|
||||
if not (
|
||||
isinstance(u, dict)
|
||||
and isinstance(u.get("toolUseId"), str)
|
||||
and u["toolUseId"] in orphaned_tool_use_ids
|
||||
)
|
||||
]
|
||||
if not filtered:
|
||||
a.pop("toolUses", None)
|
||||
elif len(filtered) != len(uses):
|
||||
a["toolUses"] = filtered
|
||||
|
||||
user_ctx: dict[str, Any] = {}
|
||||
if tools:
|
||||
user_ctx["tools"] = tools
|
||||
if validated_tool_results:
|
||||
user_ctx["toolResults"] = validated_tool_results
|
||||
|
||||
user_input: dict[str, Any] = {
|
||||
"userInputMessageContext": user_ctx,
|
||||
"content": text_content,
|
||||
"modelId": model_id,
|
||||
"origin": "AI_EDITOR",
|
||||
}
|
||||
if images:
|
||||
user_input["images"] = images
|
||||
|
||||
conversation_state = {
|
||||
"agentContinuationId": agent_continuation_id,
|
||||
"agentTaskType": "vibe",
|
||||
"chatTriggerType": "MANUAL",
|
||||
"currentMessage": {"userInputMessage": user_input},
|
||||
"conversationId": conversation_id,
|
||||
"history": history,
|
||||
}
|
||||
|
||||
return conversation_state
|
||||
|
||||
|
||||
__all__ = [
|
||||
"convert_claude_messages_to_conversation_state",
|
||||
"map_model",
|
||||
]
|
||||
@@ -0,0 +1,121 @@
|
||||
"""Kiro provider envelope.
|
||||
|
||||
Kiro upstream is not Claude wire-compatible:
|
||||
- Request: wrap Claude Messages body into Kiro `conversationState` request.
|
||||
- Stream response: handled by StreamProcessor via binary EventStream rewrite.
|
||||
|
||||
We use contextvars to pass request-scoped values (region, machine_id, thinking)
|
||||
from wrap_request() to extra_headers() and transport hook.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from src.services.provider.adapters.kiro.context import KiroRequestContext, set_kiro_request_context
|
||||
from src.services.provider.adapters.kiro.converter import (
|
||||
convert_claude_messages_to_conversation_state,
|
||||
)
|
||||
from src.services.provider.adapters.kiro.headers import build_generate_assistant_headers
|
||||
from src.services.provider.adapters.kiro.models.credentials import KiroAuthConfig
|
||||
from src.services.provider.adapters.kiro.token_manager import generate_machine_id
|
||||
|
||||
|
||||
def _resolve_region(cfg: KiroAuthConfig) -> str:
|
||||
from src.services.provider.adapters.kiro.constants import DEFAULT_REGION
|
||||
|
||||
region = str(cfg.region or "").strip()
|
||||
return region or DEFAULT_REGION
|
||||
|
||||
|
||||
def _is_thinking_enabled(request_body: dict[str, Any]) -> bool:
|
||||
thinking = request_body.get("thinking")
|
||||
if not isinstance(thinking, dict):
|
||||
return False
|
||||
ttype = str(thinking.get("type") or "").strip().lower()
|
||||
return ttype in {"enabled", "adaptive"}
|
||||
|
||||
|
||||
class KiroEnvelope:
|
||||
name = "kiro:generateAssistantResponse"
|
||||
|
||||
def extra_headers(self) -> dict[str, str] | None:
|
||||
# Called after wrap_request(); relies on KiroRequestContext.
|
||||
from src.services.provider.adapters.kiro.context import get_kiro_request_context
|
||||
|
||||
ctx = get_kiro_request_context()
|
||||
if ctx is None:
|
||||
return None
|
||||
|
||||
host = f"q.{ctx.region}.amazonaws.com"
|
||||
return build_generate_assistant_headers(
|
||||
host=host,
|
||||
machine_id=ctx.machine_id,
|
||||
kiro_version=ctx.kiro_version,
|
||||
system_version=ctx.system_version,
|
||||
node_version=ctx.node_version,
|
||||
)
|
||||
|
||||
def wrap_request(
|
||||
self,
|
||||
request_body: dict[str, Any],
|
||||
*,
|
||||
model: str,
|
||||
url_model: str | None,
|
||||
decrypted_auth_config: dict[str, Any] | None,
|
||||
) -> tuple[dict[str, Any], str | None]:
|
||||
cfg = KiroAuthConfig.from_dict(decrypted_auth_config or {})
|
||||
|
||||
region = _resolve_region(cfg)
|
||||
machine_id = generate_machine_id(cfg)
|
||||
|
||||
thinking_enabled = _is_thinking_enabled(request_body)
|
||||
|
||||
set_kiro_request_context(
|
||||
KiroRequestContext(
|
||||
region=region,
|
||||
machine_id=machine_id,
|
||||
kiro_version=cfg.kiro_version,
|
||||
system_version=cfg.system_version,
|
||||
node_version=cfg.node_version,
|
||||
thinking_enabled=thinking_enabled,
|
||||
)
|
||||
)
|
||||
|
||||
conversation_state = convert_claude_messages_to_conversation_state(
|
||||
request_body,
|
||||
model=model,
|
||||
)
|
||||
|
||||
wrapped: dict[str, Any] = {
|
||||
"conversationState": conversation_state,
|
||||
}
|
||||
if isinstance(cfg.profile_arn, str) and cfg.profile_arn.strip():
|
||||
wrapped["profileArn"] = cfg.profile_arn.strip()
|
||||
|
||||
return wrapped, url_model
|
||||
|
||||
def unwrap_response(self, data: Any) -> Any:
|
||||
return data
|
||||
|
||||
def postprocess_unwrapped_response(self, *, model: str, data: Any) -> None: # noqa: ARG002
|
||||
return
|
||||
|
||||
def capture_selected_base_url(self) -> str | None:
|
||||
return None
|
||||
|
||||
def on_http_status(self, *, base_url: str | None, status_code: int) -> None: # noqa: ARG002
|
||||
return
|
||||
|
||||
def on_connection_error(self, *, base_url: str | None, exc: Exception) -> None: # noqa: ARG002
|
||||
return
|
||||
|
||||
def force_stream_rewrite(self) -> bool:
|
||||
# Kiro streaming is binary AWS Event Stream and must be rewritten.
|
||||
return True
|
||||
|
||||
|
||||
kiro_envelope = KiroEnvelope()
|
||||
|
||||
|
||||
__all__ = ["KiroEnvelope", "kiro_envelope"]
|
||||
@@ -0,0 +1,741 @@
|
||||
"""AWS Event Stream -> Claude SSE rewriter for Kiro.
|
||||
|
||||
Kiro streaming responses are returned as `application/vnd.amazon.eventstream`
|
||||
(binary framed). This module decodes frames and emits Claude-style streaming
|
||||
SSE events (as UTF-8 bytes).
|
||||
|
||||
The output format uses ``event: {type}\\ndata: {...}\\n\\n`` for typed events and
|
||||
plain ``data: {...}\\n\\n`` for untyped events, matching how Aether parses Claude
|
||||
streams.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import uuid
|
||||
from collections.abc import AsyncGenerator
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from src.core.logger import logger
|
||||
from src.services.provider.adapters.kiro.constants import CONTEXT_WINDOW_TOKENS
|
||||
from src.services.provider.adapters.kiro.parser.decoder import EventStreamDecoder
|
||||
|
||||
# Safety limit for thinking_buffer to prevent memory exhaustion from
|
||||
# pathological upstream responses that never close the thinking tag.
|
||||
_MAX_THINKING_BUFFER = 1024 * 1024 # 1 MiB
|
||||
|
||||
_QUOTE_CHARS: frozenset[str] = frozenset("`\"'\\#!@$%^&*()-_=+[]{};:<>,.?/")
|
||||
|
||||
|
||||
def _is_quote_char(buffer: str, pos: int) -> bool:
|
||||
if pos < 0 or pos >= len(buffer):
|
||||
return False
|
||||
return buffer[pos] in _QUOTE_CHARS
|
||||
|
||||
|
||||
def _find_real_thinking_start_tag(buffer: str) -> int | None:
|
||||
tag = "<thinking>"
|
||||
search = 0
|
||||
while True:
|
||||
pos = buffer.find(tag, search)
|
||||
if pos < 0:
|
||||
return None
|
||||
has_before = pos > 0 and _is_quote_char(buffer, pos - 1)
|
||||
after_pos = pos + len(tag)
|
||||
has_after = _is_quote_char(buffer, after_pos)
|
||||
if not has_before and not has_after:
|
||||
return pos
|
||||
search = pos + 1
|
||||
|
||||
|
||||
def _find_real_thinking_end_tag(buffer: str) -> int | None:
|
||||
tag = "</thinking>"
|
||||
search = 0
|
||||
while True:
|
||||
pos = buffer.find(tag, search)
|
||||
if pos < 0:
|
||||
return None
|
||||
|
||||
has_before = pos > 0 and _is_quote_char(buffer, pos - 1)
|
||||
after_pos = pos + len(tag)
|
||||
has_after = _is_quote_char(buffer, after_pos)
|
||||
if has_before or has_after:
|
||||
search = pos + 1
|
||||
continue
|
||||
|
||||
after = buffer[after_pos:]
|
||||
if len(after) < 2:
|
||||
return None
|
||||
if after.startswith("\n\n"):
|
||||
return pos
|
||||
|
||||
search = pos + 1
|
||||
|
||||
|
||||
def _find_real_thinking_end_tag_at_buffer_end(buffer: str) -> int | None:
|
||||
tag = "</thinking>"
|
||||
search = 0
|
||||
while True:
|
||||
pos = buffer.find(tag, search)
|
||||
if pos < 0:
|
||||
return None
|
||||
|
||||
has_before = pos > 0 and _is_quote_char(buffer, pos - 1)
|
||||
after_pos = pos + len(tag)
|
||||
has_after = _is_quote_char(buffer, after_pos)
|
||||
if has_before or has_after:
|
||||
search = pos + 1
|
||||
continue
|
||||
|
||||
if buffer[after_pos:].strip() == "":
|
||||
return pos
|
||||
|
||||
search = pos + 1
|
||||
|
||||
|
||||
def _estimate_tokens(text: str) -> int:
|
||||
if not text:
|
||||
return 0
|
||||
chinese = 0
|
||||
other = 0
|
||||
for c in text:
|
||||
if "\u4e00" <= c <= "\u9fff":
|
||||
chinese += 1
|
||||
else:
|
||||
other += 1
|
||||
chinese_tokens = (chinese * 2 + 2) // 3
|
||||
other_tokens = (other + 3) // 4
|
||||
return max(chinese_tokens + other_tokens, 1)
|
||||
|
||||
|
||||
def _sse_data_bytes(obj: dict[str, Any]) -> bytes:
|
||||
data = json.dumps(obj, ensure_ascii=False)
|
||||
event_type = obj.get("type", "")
|
||||
if event_type:
|
||||
return f"event: {event_type}\ndata: {data}\n\n".encode("utf-8")
|
||||
return f"data: {data}\n\n".encode("utf-8")
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class _KiroStreamState:
|
||||
model: str
|
||||
thinking_enabled: bool
|
||||
estimated_input_tokens: int = 0
|
||||
|
||||
message_id: str = field(default_factory=lambda: f"msg_{uuid.uuid4().hex}")
|
||||
output_tokens: int = 0
|
||||
context_input_tokens: int | None = None
|
||||
|
||||
next_block_index: int = 0
|
||||
open_blocks: dict[int, str] = field(default_factory=dict)
|
||||
|
||||
text_block_index: int | None = None
|
||||
thinking_block_index: int | None = None
|
||||
tool_block_indices: dict[str, int] = field(default_factory=dict)
|
||||
|
||||
thinking_buffer: str = ""
|
||||
in_thinking_block: bool = False
|
||||
thinking_extracted: bool = False
|
||||
strip_thinking_leading_newline: bool = False
|
||||
|
||||
has_tool_use: bool = False
|
||||
stop_reason_override: str | None = None
|
||||
had_error: bool = False
|
||||
|
||||
def generate_initial_events(self) -> list[dict[str, Any]]:
|
||||
events: list[dict[str, Any]] = []
|
||||
|
||||
# message_start
|
||||
events.append(
|
||||
{
|
||||
"type": "message_start",
|
||||
"message": {
|
||||
"id": self.message_id,
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"content": [],
|
||||
"model": self.model,
|
||||
"stop_reason": None,
|
||||
"stop_sequence": None,
|
||||
# Claude CLI clients expect usage to exist.
|
||||
"usage": {
|
||||
"input_tokens": int(self.estimated_input_tokens or 0),
|
||||
"output_tokens": 1,
|
||||
},
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
if not self.thinking_enabled:
|
||||
events.extend(self._ensure_text_block_open())
|
||||
|
||||
return events
|
||||
|
||||
def _ensure_text_block_open(self) -> list[dict[str, Any]]:
|
||||
if self.text_block_index is not None:
|
||||
if (
|
||||
self.text_block_index in self.open_blocks
|
||||
and self.open_blocks[self.text_block_index] == "text"
|
||||
):
|
||||
return []
|
||||
self.text_block_index = None
|
||||
|
||||
idx = self.next_block_index
|
||||
self.next_block_index += 1
|
||||
self.text_block_index = idx
|
||||
self.open_blocks[idx] = "text"
|
||||
|
||||
return [
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": idx,
|
||||
"content_block": {"type": "text", "text": ""},
|
||||
}
|
||||
]
|
||||
|
||||
def _close_block(self, idx: int) -> list[dict[str, Any]]:
|
||||
if idx not in self.open_blocks:
|
||||
return []
|
||||
self.open_blocks.pop(idx, None)
|
||||
return [{"type": "content_block_stop", "index": idx}]
|
||||
|
||||
def _ensure_thinking_block_open(self) -> list[dict[str, Any]]:
|
||||
if self.thinking_block_index is not None:
|
||||
if (
|
||||
self.thinking_block_index in self.open_blocks
|
||||
and self.open_blocks[self.thinking_block_index] == "thinking"
|
||||
):
|
||||
return []
|
||||
|
||||
idx = self.next_block_index
|
||||
self.next_block_index += 1
|
||||
self.thinking_block_index = idx
|
||||
self.open_blocks[idx] = "thinking"
|
||||
|
||||
return [
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": idx,
|
||||
"content_block": {"type": "thinking", "thinking": ""},
|
||||
}
|
||||
]
|
||||
|
||||
def _emit_text_delta(self, text: str) -> list[dict[str, Any]]:
|
||||
if not text:
|
||||
return []
|
||||
events: list[dict[str, Any]] = []
|
||||
events.extend(self._ensure_text_block_open())
|
||||
idx = int(self.text_block_index or 0)
|
||||
events.append(
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": idx,
|
||||
"delta": {"type": "text_delta", "text": text},
|
||||
}
|
||||
)
|
||||
return events
|
||||
|
||||
def _emit_thinking_delta(self, thinking: str) -> list[dict[str, Any]]:
|
||||
if not thinking:
|
||||
return []
|
||||
events: list[dict[str, Any]] = []
|
||||
events.extend(self._ensure_thinking_block_open())
|
||||
idx = int(self.thinking_block_index or 0)
|
||||
events.append(
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": idx,
|
||||
"delta": {"type": "thinking_delta", "thinking": thinking},
|
||||
}
|
||||
)
|
||||
return events
|
||||
|
||||
def _close_thinking_block(self) -> list[dict[str, Any]]:
|
||||
"""Send an empty thinking_delta sentinel and close the thinking block."""
|
||||
if self.thinking_block_index is None:
|
||||
return []
|
||||
idx = int(self.thinking_block_index)
|
||||
events: list[dict[str, Any]] = [
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": idx,
|
||||
"delta": {"type": "thinking_delta", "thinking": ""},
|
||||
}
|
||||
]
|
||||
events.extend(self._close_block(idx))
|
||||
return events
|
||||
|
||||
def process_context_usage(self, percentage: float) -> None:
|
||||
try:
|
||||
pct = float(percentage)
|
||||
except Exception:
|
||||
return
|
||||
# percentage * CONTEXT_WINDOW_TOKENS / 100
|
||||
self.context_input_tokens = int(pct * float(CONTEXT_WINDOW_TOKENS) / 100.0)
|
||||
|
||||
def process_exception(self, exception_type: str) -> None:
|
||||
if exception_type == "ContentLengthExceededException":
|
||||
# ContentLengthExceededException is a normal completion signal (output
|
||||
# exceeded size limit), not a fatal error. We record the stop_reason
|
||||
# but do NOT set had_error so that finalize() still emits message_delta
|
||||
# with stop_reason="max_tokens" and message_stop.
|
||||
self.stop_reason_override = "max_tokens"
|
||||
return
|
||||
|
||||
def process_assistant_response(self, content: str) -> list[dict[str, Any]]:
|
||||
if not content:
|
||||
return []
|
||||
|
||||
self.output_tokens += _estimate_tokens(content)
|
||||
|
||||
if not self.thinking_enabled:
|
||||
return self._emit_text_delta(content)
|
||||
|
||||
self.thinking_buffer += content
|
||||
|
||||
# Safety: flush as text if thinking_buffer grows too large without closing tag
|
||||
if len(self.thinking_buffer) > _MAX_THINKING_BUFFER:
|
||||
logger.warning(
|
||||
"kiro thinking_buffer exceeded {} bytes, force-flushing as text",
|
||||
_MAX_THINKING_BUFFER,
|
||||
)
|
||||
overflow = self.thinking_buffer
|
||||
self.thinking_buffer = ""
|
||||
if self.in_thinking_block:
|
||||
result = self._emit_thinking_delta(overflow)
|
||||
result.extend(self._close_thinking_block())
|
||||
self.in_thinking_block = False
|
||||
self.thinking_extracted = True
|
||||
return result
|
||||
return self._emit_text_delta(overflow)
|
||||
|
||||
events: list[dict[str, Any]] = []
|
||||
|
||||
while True:
|
||||
if not self.in_thinking_block and not self.thinking_extracted:
|
||||
start_pos = _find_real_thinking_start_tag(self.thinking_buffer)
|
||||
if start_pos is not None:
|
||||
before = self.thinking_buffer[:start_pos]
|
||||
if before and before.strip():
|
||||
events.extend(self._emit_text_delta(before))
|
||||
|
||||
self.in_thinking_block = True
|
||||
self.strip_thinking_leading_newline = True
|
||||
self.thinking_buffer = self.thinking_buffer[start_pos + len("<thinking>") :]
|
||||
events.extend(self._ensure_thinking_block_open())
|
||||
continue
|
||||
|
||||
# Keep a short suffix in buffer for partial tag detection.
|
||||
keep = len("<thinking>")
|
||||
if len(self.thinking_buffer) > keep:
|
||||
safe = self.thinking_buffer[:-keep]
|
||||
if safe and safe.strip():
|
||||
events.extend(self._emit_text_delta(safe))
|
||||
self.thinking_buffer = self.thinking_buffer[-keep:]
|
||||
break
|
||||
|
||||
if self.in_thinking_block:
|
||||
# Strip a single leading \n after <thinking> tag.
|
||||
# The model outputs `<thinking>\n` and the \n may arrive in the
|
||||
# same chunk or the next one; we drop it for cleaner output.
|
||||
if self.strip_thinking_leading_newline:
|
||||
if self.thinking_buffer.startswith("\n"):
|
||||
self.thinking_buffer = self.thinking_buffer[1:]
|
||||
self.strip_thinking_leading_newline = False
|
||||
elif self.thinking_buffer:
|
||||
# Buffer is non-empty but doesn't start with \n; stop waiting.
|
||||
self.strip_thinking_leading_newline = False
|
||||
# else: buffer is empty, keep the flag for the next chunk.
|
||||
|
||||
end_pos = _find_real_thinking_end_tag(self.thinking_buffer)
|
||||
if end_pos is not None:
|
||||
thinking_text = self.thinking_buffer[:end_pos]
|
||||
if thinking_text:
|
||||
events.extend(self._emit_thinking_delta(thinking_text))
|
||||
|
||||
events.extend(self._close_thinking_block())
|
||||
|
||||
self.in_thinking_block = False
|
||||
self.thinking_extracted = True
|
||||
self.thinking_buffer = self.thinking_buffer[end_pos + len("</thinking>") :]
|
||||
continue
|
||||
|
||||
keep = len("</thinking>")
|
||||
if len(self.thinking_buffer) > keep:
|
||||
safe = self.thinking_buffer[:-keep]
|
||||
if safe:
|
||||
events.extend(self._emit_thinking_delta(safe))
|
||||
self.thinking_buffer = self.thinking_buffer[-keep:]
|
||||
break
|
||||
|
||||
# thinking extracted: remaining buffer is text
|
||||
if self.thinking_buffer:
|
||||
remaining = self.thinking_buffer
|
||||
self.thinking_buffer = ""
|
||||
events.extend(self._emit_text_delta(remaining))
|
||||
break
|
||||
|
||||
return events
|
||||
|
||||
def process_tool_use(
|
||||
self,
|
||||
*,
|
||||
name: str,
|
||||
tool_use_id: str,
|
||||
input_json: str,
|
||||
stop: bool,
|
||||
) -> list[dict[str, Any]]:
|
||||
if not tool_use_id:
|
||||
return []
|
||||
|
||||
self.has_tool_use = True
|
||||
|
||||
events: list[dict[str, Any]] = []
|
||||
|
||||
# Boundary: close thinking block if needed, filtering a dangling </thinking>.
|
||||
if self.thinking_enabled and self.in_thinking_block and self.thinking_buffer:
|
||||
end_pos = _find_real_thinking_end_tag_at_buffer_end(self.thinking_buffer)
|
||||
if end_pos is not None:
|
||||
thinking_text = self.thinking_buffer[:end_pos]
|
||||
if thinking_text:
|
||||
events.extend(self._emit_thinking_delta(thinking_text))
|
||||
|
||||
events.extend(self._close_thinking_block())
|
||||
|
||||
after_pos = end_pos + len("</thinking>")
|
||||
remaining = self.thinking_buffer[after_pos:]
|
||||
self.thinking_buffer = ""
|
||||
self.in_thinking_block = False
|
||||
self.thinking_extracted = True
|
||||
if remaining:
|
||||
events.extend(self._emit_text_delta(remaining))
|
||||
else:
|
||||
# Best-effort flush all as thinking
|
||||
events.extend(self._emit_thinking_delta(self.thinking_buffer))
|
||||
events.extend(self._close_thinking_block())
|
||||
self.thinking_buffer = ""
|
||||
self.in_thinking_block = False
|
||||
self.thinking_extracted = True
|
||||
|
||||
# Flush any buffered pre-thinking tail so tool_use doesn't swallow it.
|
||||
if (
|
||||
self.thinking_enabled
|
||||
and not self.in_thinking_block
|
||||
and not self.thinking_extracted
|
||||
and self.thinking_buffer
|
||||
):
|
||||
buffered = self.thinking_buffer
|
||||
self.thinking_buffer = ""
|
||||
events.extend(self._emit_text_delta(buffered))
|
||||
|
||||
# Close current text block before tool_use.
|
||||
if self.text_block_index is not None:
|
||||
idx = int(self.text_block_index)
|
||||
events.extend(self._close_block(idx))
|
||||
|
||||
block_index = self.tool_block_indices.get(tool_use_id)
|
||||
if block_index is None:
|
||||
block_index = self.next_block_index
|
||||
self.next_block_index += 1
|
||||
self.tool_block_indices[tool_use_id] = block_index
|
||||
|
||||
# Start tool block if not open.
|
||||
if block_index not in self.open_blocks:
|
||||
self.open_blocks[block_index] = "tool_use"
|
||||
events.append(
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": block_index,
|
||||
"content_block": {
|
||||
"type": "tool_use",
|
||||
"id": tool_use_id,
|
||||
"name": name,
|
||||
"input": {},
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
if input_json:
|
||||
self.output_tokens += _estimate_tokens(input_json)
|
||||
events.append(
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": block_index,
|
||||
"delta": {"type": "input_json_delta", "partial_json": input_json},
|
||||
}
|
||||
)
|
||||
|
||||
if stop:
|
||||
events.extend(self._close_block(block_index))
|
||||
|
||||
return events
|
||||
|
||||
def finalize(self) -> list[dict[str, Any]]:
|
||||
events: list[dict[str, Any]] = []
|
||||
|
||||
# Flush remaining thinking/text buffer.
|
||||
if self.thinking_enabled and self.thinking_buffer:
|
||||
if self.in_thinking_block:
|
||||
end_pos = _find_real_thinking_end_tag_at_buffer_end(self.thinking_buffer)
|
||||
if end_pos is not None:
|
||||
thinking_text = self.thinking_buffer[:end_pos]
|
||||
if thinking_text:
|
||||
events.extend(self._emit_thinking_delta(thinking_text))
|
||||
|
||||
events.extend(self._close_thinking_block())
|
||||
|
||||
after_pos = end_pos + len("</thinking>")
|
||||
remaining = self.thinking_buffer[after_pos:]
|
||||
if remaining:
|
||||
events.extend(self._emit_text_delta(remaining))
|
||||
else:
|
||||
events.extend(self._emit_thinking_delta(self.thinking_buffer))
|
||||
events.extend(self._close_thinking_block())
|
||||
|
||||
else:
|
||||
events.extend(self._emit_text_delta(self.thinking_buffer))
|
||||
|
||||
self.thinking_buffer = ""
|
||||
self.in_thinking_block = False
|
||||
self.thinking_extracted = True
|
||||
|
||||
# Close any open blocks (best-effort).
|
||||
for idx in sorted(list(self.open_blocks.keys()), reverse=True):
|
||||
events.extend(self._close_block(idx))
|
||||
|
||||
stop_reason = self.stop_reason_override
|
||||
if not stop_reason:
|
||||
stop_reason = "tool_use" if self.has_tool_use else "end_turn"
|
||||
|
||||
input_tokens = (
|
||||
int(self.context_input_tokens)
|
||||
if self.context_input_tokens is not None
|
||||
else int(self.estimated_input_tokens or 0)
|
||||
)
|
||||
|
||||
events.append(
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": stop_reason, "stop_sequence": None},
|
||||
"usage": {"input_tokens": input_tokens, "output_tokens": int(self.output_tokens)},
|
||||
}
|
||||
)
|
||||
events.append({"type": "message_stop"})
|
||||
|
||||
return events
|
||||
|
||||
|
||||
async def rewrite_eventstream_to_sse(
|
||||
byte_iterator: Any,
|
||||
*,
|
||||
model: str,
|
||||
thinking_enabled: bool,
|
||||
estimated_input_tokens: int = 0,
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""Rewrite Kiro AWS Event Stream bytes to Claude SSE bytes."""
|
||||
decoder = EventStreamDecoder()
|
||||
state = _KiroStreamState(
|
||||
model=str(model or ""),
|
||||
thinking_enabled=bool(thinking_enabled),
|
||||
estimated_input_tokens=int(estimated_input_tokens or 0),
|
||||
)
|
||||
|
||||
# 收集原始字节用于错误诊断
|
||||
raw_bytes_buffer = b""
|
||||
|
||||
# Initial events
|
||||
for evt in state.generate_initial_events():
|
||||
yield _sse_data_bytes(evt)
|
||||
|
||||
async for chunk in byte_iterator:
|
||||
if not chunk:
|
||||
continue
|
||||
|
||||
# 保留原始字节用于错误诊断(限制大小)
|
||||
if len(raw_bytes_buffer) < 4096:
|
||||
raw_bytes_buffer += chunk
|
||||
|
||||
try:
|
||||
decoder.feed(chunk)
|
||||
frames = decoder.decode_available()
|
||||
except Exception as e:
|
||||
logger.warning("kiro eventstream decode error: {}", e)
|
||||
# 尝试解析原始响应为 JSON 错误
|
||||
error_message = f"kiro eventstream decode failed: {type(e).__name__}"
|
||||
try:
|
||||
raw_text = raw_bytes_buffer.decode("utf-8", errors="replace")
|
||||
# 尝试解析为 JSON
|
||||
error_json = json.loads(raw_text)
|
||||
if isinstance(error_json, dict):
|
||||
# 提取上游错误信息
|
||||
upstream_msg = error_json.get("message") or error_json.get("error", {}).get(
|
||||
"message"
|
||||
)
|
||||
if upstream_msg:
|
||||
error_message = f"Kiro API error: {upstream_msg}"
|
||||
except Exception:
|
||||
pass
|
||||
yield _sse_data_bytes(
|
||||
{
|
||||
"type": "error",
|
||||
"error": {
|
||||
"type": "upstream_stream_error",
|
||||
"message": error_message,
|
||||
},
|
||||
}
|
||||
)
|
||||
break
|
||||
|
||||
for frame in frames:
|
||||
mtype = (frame.message_type() or "event").strip().lower()
|
||||
etype = (frame.event_type() or "").strip()
|
||||
payload_text = frame.payload_as_text()
|
||||
|
||||
if mtype == "event":
|
||||
try:
|
||||
payload = json.loads(payload_text) if payload_text else {}
|
||||
except Exception:
|
||||
payload = {}
|
||||
|
||||
if etype == "assistantResponseEvent":
|
||||
content = payload.get("content") if isinstance(payload, dict) else None
|
||||
if isinstance(content, str) and content:
|
||||
for evt in state.process_assistant_response(content):
|
||||
yield _sse_data_bytes(evt)
|
||||
continue
|
||||
|
||||
if etype == "toolUseEvent":
|
||||
if isinstance(payload, dict):
|
||||
name = str(payload.get("name") or "")
|
||||
tool_use_id = payload.get("toolUseId") or payload.get("tool_use_id")
|
||||
tool_use_id = str(tool_use_id or "")
|
||||
raw_input = payload.get("input")
|
||||
if raw_input is None:
|
||||
input_json = ""
|
||||
elif isinstance(raw_input, str):
|
||||
input_json = raw_input
|
||||
else:
|
||||
try:
|
||||
input_json = json.dumps(raw_input, ensure_ascii=False)
|
||||
except Exception:
|
||||
input_json = str(raw_input)
|
||||
stop = bool(payload.get("stop", False))
|
||||
for evt in state.process_tool_use(
|
||||
name=name,
|
||||
tool_use_id=tool_use_id,
|
||||
input_json=input_json,
|
||||
stop=stop,
|
||||
):
|
||||
yield _sse_data_bytes(evt)
|
||||
continue
|
||||
|
||||
if etype == "contextUsageEvent":
|
||||
if isinstance(payload, dict):
|
||||
pct = payload.get("contextUsagePercentage")
|
||||
if pct is not None:
|
||||
try:
|
||||
state.process_context_usage(float(pct))
|
||||
except (ValueError, TypeError):
|
||||
logger.debug(
|
||||
"kiro: failed to parse contextUsagePercentage: {!r}", pct
|
||||
)
|
||||
continue
|
||||
|
||||
# meteringEvent / unknown: ignore
|
||||
continue
|
||||
|
||||
if mtype == "exception":
|
||||
ex_type = frame.headers.exception_type() or "UnknownException"
|
||||
state.process_exception(ex_type)
|
||||
# ContentLengthExceededException is handled by process_exception
|
||||
# (sets stop_reason_override) and should NOT prevent finalize().
|
||||
if not state.stop_reason_override:
|
||||
state.had_error = True
|
||||
logger.debug("kiro upstream exception: {} | {}", ex_type, payload_text[:200])
|
||||
if state.had_error:
|
||||
yield _sse_data_bytes(
|
||||
{
|
||||
"type": "error",
|
||||
"error": {
|
||||
"type": "upstream_exception",
|
||||
"message": ex_type,
|
||||
},
|
||||
}
|
||||
)
|
||||
continue
|
||||
|
||||
if mtype == "error":
|
||||
err_code = frame.headers.error_code() or "UnknownError"
|
||||
state.had_error = True
|
||||
logger.debug("kiro upstream error: {} | {}", err_code, payload_text[:200])
|
||||
yield _sse_data_bytes(
|
||||
{
|
||||
"type": "error",
|
||||
"error": {
|
||||
"type": "upstream_error",
|
||||
"message": err_code,
|
||||
},
|
||||
}
|
||||
)
|
||||
continue
|
||||
|
||||
if not state.had_error:
|
||||
for evt in state.finalize():
|
||||
yield _sse_data_bytes(evt)
|
||||
|
||||
|
||||
def apply_kiro_stream_rewrite(
|
||||
byte_iter: Any,
|
||||
*,
|
||||
model: str = "",
|
||||
input_tokens: int = 0,
|
||||
prefetched_chunks: list[bytes] | None = None,
|
||||
) -> AsyncGenerator[bytes]:
|
||||
"""Apply Kiro EventStream->SSE rewrite if context is available.
|
||||
|
||||
Consolidates the repeated import-context-rewrite pattern used across
|
||||
``chat_handler_base``, ``cli_handler_base``, and ``stream_processor``.
|
||||
|
||||
Args:
|
||||
byte_iter: Upstream byte iterator (raw AWS Event Stream).
|
||||
model: Model name for SSE events.
|
||||
input_tokens: Estimated input token count.
|
||||
prefetched_chunks: Optional pre-fetched bytes to prepend.
|
||||
|
||||
Returns:
|
||||
An async generator of Claude-compatible SSE bytes.
|
||||
"""
|
||||
from src.services.provider.adapters.kiro.context import get_kiro_request_context
|
||||
|
||||
kiro_ctx = get_kiro_request_context()
|
||||
thinking_enabled = bool(getattr(kiro_ctx, "thinking_enabled", False)) if kiro_ctx else False
|
||||
|
||||
if prefetched_chunks:
|
||||
upstream = byte_iter
|
||||
prefix = list(prefetched_chunks)
|
||||
|
||||
async def _combined() -> AsyncGenerator[bytes, None]:
|
||||
for c in prefix:
|
||||
if c:
|
||||
yield c
|
||||
async for c in upstream:
|
||||
if c:
|
||||
yield c
|
||||
|
||||
source: Any = _combined()
|
||||
else:
|
||||
source = byte_iter
|
||||
|
||||
return rewrite_eventstream_to_sse(
|
||||
source,
|
||||
model=str(model or ""),
|
||||
thinking_enabled=thinking_enabled,
|
||||
estimated_input_tokens=int(input_tokens or 0),
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"apply_kiro_stream_rewrite",
|
||||
"rewrite_eventstream_to_sse",
|
||||
]
|
||||
@@ -0,0 +1,102 @@
|
||||
"""Kiro header builders."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
|
||||
from src.services.provider.adapters.kiro.constants import (
|
||||
AWS_EVENTSTREAM_CONTENT_TYPE,
|
||||
AWS_SDK_JS_MAIN_VERSION,
|
||||
AWS_SDK_JS_USAGE_VERSION,
|
||||
CODEWHISPERER_OPTOUT,
|
||||
DEFAULT_KIRO_VERSION,
|
||||
DEFAULT_NODE_VERSION,
|
||||
DEFAULT_SYSTEM_VERSION,
|
||||
KIRO_AGENT_MODE,
|
||||
)
|
||||
|
||||
|
||||
def build_kiro_ide_tag(*, kiro_version: str, machine_id: str) -> str:
|
||||
version = (kiro_version or DEFAULT_KIRO_VERSION).strip() or DEFAULT_KIRO_VERSION
|
||||
mid = (machine_id or "").strip()
|
||||
return f"KiroIDE-{version}-{mid}" if mid else f"KiroIDE-{version}"
|
||||
|
||||
|
||||
def build_x_amz_user_agent_main(*, kiro_version: str, machine_id: str) -> str:
|
||||
return f"aws-sdk-js/{AWS_SDK_JS_MAIN_VERSION} {build_kiro_ide_tag(kiro_version=kiro_version, machine_id=machine_id)}"
|
||||
|
||||
|
||||
def build_user_agent_main(
|
||||
*, system_version: str, node_version: str, kiro_version: str, machine_id: str
|
||||
) -> str:
|
||||
os_tag = (system_version or DEFAULT_SYSTEM_VERSION).strip() or DEFAULT_SYSTEM_VERSION
|
||||
node_tag = (node_version or DEFAULT_NODE_VERSION).strip() or DEFAULT_NODE_VERSION
|
||||
ide = build_kiro_ide_tag(kiro_version=kiro_version, machine_id=machine_id)
|
||||
return (
|
||||
f"aws-sdk-js/{AWS_SDK_JS_MAIN_VERSION} ua/2.1 os/{os_tag} lang/js "
|
||||
f"md/nodejs#{node_tag} api/codewhispererstreaming#{AWS_SDK_JS_MAIN_VERSION} m/E {ide}"
|
||||
)
|
||||
|
||||
|
||||
def build_x_amz_user_agent_usage(*, kiro_version: str, machine_id: str) -> str:
|
||||
ide = build_kiro_ide_tag(kiro_version=kiro_version, machine_id=machine_id)
|
||||
return f"aws-sdk-js/{AWS_SDK_JS_USAGE_VERSION} {ide}"
|
||||
|
||||
|
||||
def build_user_agent_usage(*, kiro_version: str, machine_id: str) -> str:
|
||||
ide = build_kiro_ide_tag(kiro_version=kiro_version, machine_id=machine_id)
|
||||
os_tag = DEFAULT_SYSTEM_VERSION
|
||||
node_tag = DEFAULT_NODE_VERSION
|
||||
return (
|
||||
f"aws-sdk-js/{AWS_SDK_JS_USAGE_VERSION} ua/2.1 os/{os_tag} lang/js "
|
||||
f"md/nodejs#{node_tag} api/codewhispererruntime#1.0.0 m/N,E {ide}"
|
||||
)
|
||||
|
||||
|
||||
def build_generate_assistant_headers(
|
||||
*,
|
||||
host: str,
|
||||
access_token: str | None = None,
|
||||
machine_id: str,
|
||||
kiro_version: str | None = None,
|
||||
system_version: str | None = None,
|
||||
node_version: str | None = None,
|
||||
) -> dict[str, str]:
|
||||
version = (kiro_version or DEFAULT_KIRO_VERSION).strip() or DEFAULT_KIRO_VERSION
|
||||
sys_ver = (system_version or DEFAULT_SYSTEM_VERSION).strip() or DEFAULT_SYSTEM_VERSION
|
||||
node_ver = (node_version or DEFAULT_NODE_VERSION).strip() or DEFAULT_NODE_VERSION
|
||||
|
||||
headers: dict[str, str] = {
|
||||
"Content-Type": "application/json",
|
||||
"Accept": AWS_EVENTSTREAM_CONTENT_TYPE,
|
||||
"host": host,
|
||||
"Connection": "close",
|
||||
"x-amzn-codewhisperer-optout": CODEWHISPERER_OPTOUT,
|
||||
"x-amzn-kiro-agent-mode": KIRO_AGENT_MODE,
|
||||
"x-amz-user-agent": build_x_amz_user_agent_main(
|
||||
kiro_version=version, machine_id=machine_id
|
||||
),
|
||||
"User-Agent": build_user_agent_main(
|
||||
system_version=sys_ver,
|
||||
node_version=node_ver,
|
||||
kiro_version=version,
|
||||
machine_id=machine_id,
|
||||
),
|
||||
"amz-sdk-invocation-id": str(uuid.uuid4()),
|
||||
"amz-sdk-request": "attempt=1; max=3",
|
||||
}
|
||||
|
||||
if access_token:
|
||||
headers["Authorization"] = f"Bearer {access_token}"
|
||||
|
||||
return headers
|
||||
|
||||
|
||||
__all__ = [
|
||||
"build_generate_assistant_headers",
|
||||
"build_kiro_ide_tag",
|
||||
"build_user_agent_main",
|
||||
"build_user_agent_usage",
|
||||
"build_x_amz_user_agent_main",
|
||||
"build_x_amz_user_agent_usage",
|
||||
]
|
||||
@@ -0,0 +1,21 @@
|
||||
from .credentials import KiroAuthConfig
|
||||
from .usage_limits import (
|
||||
Bonus,
|
||||
FreeTrialInfo,
|
||||
SubscriptionInfo,
|
||||
UsageBreakdown,
|
||||
UsageLimitsResponse,
|
||||
calculate_current_usage,
|
||||
calculate_total_usage_limit,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"Bonus",
|
||||
"FreeTrialInfo",
|
||||
"KiroAuthConfig",
|
||||
"SubscriptionInfo",
|
||||
"UsageBreakdown",
|
||||
"UsageLimitsResponse",
|
||||
"calculate_current_usage",
|
||||
"calculate_total_usage_limit",
|
||||
]
|
||||
@@ -0,0 +1,178 @@
|
||||
"""Internal Kiro credential schema (stored in ProviderAPIKey.auth_config)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
|
||||
def _parse_epoch_seconds(value: object) -> int | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
if isinstance(value, (int, float)):
|
||||
return int(value)
|
||||
if isinstance(value, str) and value.strip().isdigit():
|
||||
return int(value.strip())
|
||||
except Exception:
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def _get_str(raw: dict[str, Any], *keys: str) -> str | None:
|
||||
"""Return the first non-empty stripped string for *keys*, or ``None``."""
|
||||
for k in keys:
|
||||
v = raw.get(k)
|
||||
if isinstance(v, str) and v.strip():
|
||||
return v.strip()
|
||||
return None
|
||||
|
||||
|
||||
def _parse_iso_to_epoch_seconds(value: object) -> int | None:
|
||||
if not isinstance(value, str) or not value.strip():
|
||||
return None
|
||||
text = value.strip()
|
||||
# Support RFC3339 with Z suffix.
|
||||
if text.endswith("Z"):
|
||||
text = text[:-1] + "+00:00"
|
||||
try:
|
||||
dt = datetime.fromisoformat(text)
|
||||
except Exception:
|
||||
return None
|
||||
if dt.tzinfo is None:
|
||||
dt = dt.replace(tzinfo=timezone.utc)
|
||||
return int(dt.timestamp())
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class KiroAuthConfig:
|
||||
provider_type: str = "kiro"
|
||||
|
||||
auth_method: str = "social" # social | idc
|
||||
refresh_token: str = ""
|
||||
expires_at: int = 0
|
||||
|
||||
profile_arn: str | None = None
|
||||
region: str | None = None
|
||||
|
||||
client_id: str | None = None
|
||||
client_secret: str | None = None
|
||||
|
||||
machine_id: str | None = None
|
||||
kiro_version: str | None = None
|
||||
system_version: str | None = None
|
||||
node_version: str | None = None
|
||||
|
||||
email: str | None = None # 账号邮箱
|
||||
|
||||
# 缓存的 access_token(可选,用于避免频繁刷新)
|
||||
access_token: str | None = None
|
||||
|
||||
@staticmethod
|
||||
def infer_auth_method(raw: dict[str, Any]) -> str:
|
||||
"""
|
||||
根据凭据字段自动推断认证类型。
|
||||
|
||||
规则:
|
||||
- 包含 clientId + clientSecret -> IdC
|
||||
- 仅含 refreshToken -> Social
|
||||
"""
|
||||
client_id = raw.get("client_id") or raw.get("clientId")
|
||||
client_secret = raw.get("client_secret") or raw.get("clientSecret")
|
||||
|
||||
if client_id and client_secret:
|
||||
return "idc"
|
||||
return "social"
|
||||
|
||||
@staticmethod
|
||||
def validate_required_fields(raw: dict[str, Any]) -> tuple[bool, str]:
|
||||
"""
|
||||
验证凭据是否包含必需字段。
|
||||
|
||||
返回: (is_valid, error_message)
|
||||
"""
|
||||
refresh_token = raw.get("refresh_token") or raw.get("refreshToken") or ""
|
||||
refresh_token = str(refresh_token).strip()
|
||||
|
||||
if not refresh_token:
|
||||
return False, "refreshToken 为必填字段"
|
||||
|
||||
# refreshToken 不能含有 ...(表示被截断)
|
||||
if "..." in refresh_token:
|
||||
return False, "refreshToken 不完整(含有 ...),请导出完整的 Token"
|
||||
|
||||
# IdC 类型需要 clientId 和 clientSecret
|
||||
auth_method = KiroAuthConfig.infer_auth_method(raw)
|
||||
if auth_method == "idc":
|
||||
client_id = raw.get("client_id") or raw.get("clientId")
|
||||
client_secret = raw.get("client_secret") or raw.get("clientSecret")
|
||||
if not client_id:
|
||||
return False, "IdC 类型需要 clientId"
|
||||
if not client_secret:
|
||||
return False, "IdC 类型需要 clientSecret"
|
||||
|
||||
return True, ""
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, raw: dict[str, Any]) -> "KiroAuthConfig":
|
||||
if not isinstance(raw, dict):
|
||||
raw = {}
|
||||
|
||||
provider_type = _get_str(raw, "provider_type", "providerType") or "kiro"
|
||||
|
||||
# 自动推断 auth_method(如果未显式指定)
|
||||
explicit_method = _get_str(raw, "auth_method", "authMethod")
|
||||
auth_method = explicit_method.lower() if explicit_method else cls.infer_auth_method(raw)
|
||||
|
||||
refresh_token = (_get_str(raw, "refresh_token", "refreshToken") or "").strip()
|
||||
|
||||
expires_at = _parse_epoch_seconds(raw.get("expires_at"))
|
||||
if expires_at is None:
|
||||
expires_at = _parse_iso_to_epoch_seconds(raw.get("expiresAt"))
|
||||
if expires_at is None:
|
||||
expires_at = 0
|
||||
|
||||
cfg = cls(
|
||||
provider_type=provider_type,
|
||||
auth_method=(auth_method or "social").lower(),
|
||||
refresh_token=refresh_token,
|
||||
expires_at=int(expires_at),
|
||||
profile_arn=_get_str(raw, "profile_arn", "profileArn"),
|
||||
region=_get_str(raw, "region"),
|
||||
client_id=_get_str(raw, "client_id", "clientId"),
|
||||
client_secret=_get_str(raw, "client_secret", "clientSecret"),
|
||||
machine_id=_get_str(raw, "machine_id", "machineId"),
|
||||
kiro_version=_get_str(raw, "kiro_version", "kiroVersion"),
|
||||
system_version=_get_str(raw, "system_version", "systemVersion"),
|
||||
node_version=_get_str(raw, "node_version", "nodeVersion"),
|
||||
email=_get_str(raw, "email"),
|
||||
access_token=_get_str(raw, "access_token", "accessToken"),
|
||||
)
|
||||
|
||||
# Normalize auth_method aliases.
|
||||
if cfg.auth_method in {"builder-id", "builder_id", "iam"}:
|
||||
cfg.auth_method = "idc"
|
||||
|
||||
return cfg
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"provider_type": self.provider_type,
|
||||
"auth_method": self.auth_method,
|
||||
"refresh_token": self.refresh_token,
|
||||
"expires_at": self.expires_at,
|
||||
"profile_arn": self.profile_arn,
|
||||
"region": self.region,
|
||||
"client_id": self.client_id,
|
||||
"client_secret": self.client_secret,
|
||||
"machine_id": self.machine_id,
|
||||
"kiro_version": self.kiro_version,
|
||||
"system_version": self.system_version,
|
||||
"node_version": self.node_version,
|
||||
"email": self.email,
|
||||
"access_token": self.access_token,
|
||||
}
|
||||
|
||||
|
||||
__all__ = ["KiroAuthConfig"]
|
||||
@@ -0,0 +1,222 @@
|
||||
"""Kiro getUsageLimits response models (best-effort).
|
||||
|
||||
The AWS API uses camelCase fields; we parse defensively.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class SubscriptionInfo:
|
||||
subscription_title: str | None = None
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, raw: Any) -> "SubscriptionInfo | None":
|
||||
if not isinstance(raw, dict):
|
||||
return None
|
||||
title = raw.get("subscriptionTitle")
|
||||
if isinstance(title, str) and title.strip():
|
||||
return cls(subscription_title=title.strip())
|
||||
return cls(subscription_title=None)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class Bonus:
|
||||
current_usage: float = 0.0
|
||||
usage_limit: float = 0.0
|
||||
status: str | None = None
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, raw: Any) -> "Bonus | None":
|
||||
if not isinstance(raw, dict):
|
||||
return None
|
||||
status = raw.get("status")
|
||||
status_val = status.strip() if isinstance(status, str) and status.strip() else None
|
||||
cu = raw.get("currentUsage")
|
||||
ul = raw.get("usageLimit")
|
||||
try:
|
||||
cu_f = float(cu) if cu is not None else 0.0
|
||||
except Exception:
|
||||
cu_f = 0.0
|
||||
try:
|
||||
ul_f = float(ul) if ul is not None else 0.0
|
||||
except Exception:
|
||||
ul_f = 0.0
|
||||
return cls(current_usage=cu_f, usage_limit=ul_f, status=status_val)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class FreeTrialInfo:
|
||||
current_usage: int = 0
|
||||
current_usage_with_precision: float = 0.0
|
||||
usage_limit: int = 0
|
||||
usage_limit_with_precision: float = 0.0
|
||||
free_trial_expiry: float | None = None
|
||||
free_trial_status: str | None = None
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, raw: Any) -> "FreeTrialInfo | None":
|
||||
if not isinstance(raw, dict):
|
||||
return None
|
||||
|
||||
def _int(v: Any) -> int:
|
||||
try:
|
||||
return int(v)
|
||||
except Exception:
|
||||
return 0
|
||||
|
||||
def _float(v: Any) -> float:
|
||||
try:
|
||||
return float(v)
|
||||
except Exception:
|
||||
return 0.0
|
||||
|
||||
expiry = raw.get("freeTrialExpiry")
|
||||
try:
|
||||
expiry_f = float(expiry) if expiry is not None else None
|
||||
except Exception:
|
||||
expiry_f = None
|
||||
|
||||
status = raw.get("freeTrialStatus")
|
||||
status_val = status.strip() if isinstance(status, str) and status.strip() else None
|
||||
|
||||
return cls(
|
||||
current_usage=_int(raw.get("currentUsage")),
|
||||
current_usage_with_precision=_float(raw.get("currentUsageWithPrecision")),
|
||||
usage_limit=_int(raw.get("usageLimit")),
|
||||
usage_limit_with_precision=_float(raw.get("usageLimitWithPrecision")),
|
||||
free_trial_expiry=expiry_f,
|
||||
free_trial_status=status_val,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class UsageBreakdown:
|
||||
current_usage: int = 0
|
||||
current_usage_with_precision: float = 0.0
|
||||
usage_limit: int = 0
|
||||
usage_limit_with_precision: float = 0.0
|
||||
next_date_reset: float | None = None
|
||||
bonuses: list[Bonus] = field(default_factory=list)
|
||||
free_trial_info: FreeTrialInfo | None = None
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, raw: Any) -> "UsageBreakdown | None":
|
||||
if not isinstance(raw, dict):
|
||||
return None
|
||||
|
||||
def _int(v: Any) -> int:
|
||||
try:
|
||||
return int(v)
|
||||
except Exception:
|
||||
return 0
|
||||
|
||||
def _float(v: Any) -> float:
|
||||
try:
|
||||
return float(v)
|
||||
except Exception:
|
||||
return 0.0
|
||||
|
||||
next_reset = raw.get("nextDateReset")
|
||||
try:
|
||||
next_reset_f = float(next_reset) if next_reset is not None else None
|
||||
except Exception:
|
||||
next_reset_f = None
|
||||
|
||||
bonuses_raw = raw.get("bonuses")
|
||||
bonuses: list[Bonus] = []
|
||||
if isinstance(bonuses_raw, list):
|
||||
for b in bonuses_raw:
|
||||
parsed = Bonus.from_dict(b)
|
||||
if parsed is not None:
|
||||
bonuses.append(parsed)
|
||||
|
||||
return cls(
|
||||
current_usage=_int(raw.get("currentUsage")),
|
||||
current_usage_with_precision=_float(raw.get("currentUsageWithPrecision")),
|
||||
usage_limit=_int(raw.get("usageLimit")),
|
||||
usage_limit_with_precision=_float(raw.get("usageLimitWithPrecision")),
|
||||
next_date_reset=next_reset_f,
|
||||
bonuses=bonuses,
|
||||
free_trial_info=FreeTrialInfo.from_dict(raw.get("freeTrialInfo")),
|
||||
)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class UsageLimitsResponse:
|
||||
next_date_reset: float | None = None
|
||||
subscription_info: SubscriptionInfo | None = None
|
||||
usage_breakdown_list: list[UsageBreakdown] = field(default_factory=list)
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, raw: Any) -> "UsageLimitsResponse":
|
||||
if not isinstance(raw, dict):
|
||||
raw = {}
|
||||
|
||||
next_reset = raw.get("nextDateReset")
|
||||
try:
|
||||
next_reset_f = float(next_reset) if next_reset is not None else None
|
||||
except Exception:
|
||||
next_reset_f = None
|
||||
|
||||
breakdown_raw = raw.get("usageBreakdownList")
|
||||
breakdowns: list[UsageBreakdown] = []
|
||||
if isinstance(breakdown_raw, list):
|
||||
for b in breakdown_raw:
|
||||
parsed = UsageBreakdown.from_dict(b)
|
||||
if parsed is not None:
|
||||
breakdowns.append(parsed)
|
||||
|
||||
return cls(
|
||||
next_date_reset=next_reset_f,
|
||||
subscription_info=SubscriptionInfo.from_dict(raw.get("subscriptionInfo")),
|
||||
usage_breakdown_list=breakdowns,
|
||||
)
|
||||
|
||||
|
||||
def calculate_total_usage_limit(response: UsageLimitsResponse) -> float:
|
||||
if not response.usage_breakdown_list:
|
||||
return 0.0
|
||||
|
||||
breakdown = response.usage_breakdown_list[0]
|
||||
total = breakdown.usage_limit_with_precision
|
||||
|
||||
if breakdown.free_trial_info and breakdown.free_trial_info.free_trial_status == "ACTIVE":
|
||||
total += breakdown.free_trial_info.usage_limit_with_precision
|
||||
|
||||
for bonus in breakdown.bonuses:
|
||||
if bonus.status == "ACTIVE":
|
||||
total += bonus.usage_limit
|
||||
|
||||
return total
|
||||
|
||||
|
||||
def calculate_current_usage(response: UsageLimitsResponse) -> float:
|
||||
if not response.usage_breakdown_list:
|
||||
return 0.0
|
||||
|
||||
breakdown = response.usage_breakdown_list[0]
|
||||
total = breakdown.current_usage_with_precision
|
||||
|
||||
if breakdown.free_trial_info and breakdown.free_trial_info.free_trial_status == "ACTIVE":
|
||||
total += breakdown.free_trial_info.current_usage_with_precision
|
||||
|
||||
for bonus in breakdown.bonuses:
|
||||
if bonus.status == "ACTIVE":
|
||||
total += bonus.current_usage
|
||||
|
||||
return total
|
||||
|
||||
|
||||
__all__ = [
|
||||
"Bonus",
|
||||
"FreeTrialInfo",
|
||||
"SubscriptionInfo",
|
||||
"UsageBreakdown",
|
||||
"UsageLimitsResponse",
|
||||
"calculate_current_usage",
|
||||
"calculate_total_usage_limit",
|
||||
]
|
||||
@@ -0,0 +1,6 @@
|
||||
"""AWS Event Stream parser for Kiro."""
|
||||
|
||||
from .decoder import EventStreamDecoder
|
||||
from .frame import Frame
|
||||
|
||||
__all__ = ["EventStreamDecoder", "Frame"]
|
||||
@@ -0,0 +1,13 @@
|
||||
"""CRC helpers for AWS Event Stream frames."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import binascii
|
||||
|
||||
|
||||
def crc32(data: bytes) -> int:
|
||||
"""Compute unsigned CRC32 (IEEE)."""
|
||||
return binascii.crc32(data) & 0xFFFFFFFF
|
||||
|
||||
|
||||
__all__ = ["crc32"]
|
||||
@@ -0,0 +1,91 @@
|
||||
"""Incremental AWS Event Stream decoder."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from .error import BufferOverflowError, EventStreamParseError
|
||||
from .frame import MAX_MESSAGE_SIZE, Frame, parse_frame
|
||||
|
||||
DEFAULT_MAX_BUFFER_SIZE = MAX_MESSAGE_SIZE
|
||||
DEFAULT_MAX_ERRORS = 5
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class DecoderStats:
|
||||
frames_decoded: int = 0
|
||||
bytes_skipped: int = 0
|
||||
error_count: int = 0
|
||||
|
||||
|
||||
class EventStreamDecoder:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
max_buffer_size: int = DEFAULT_MAX_BUFFER_SIZE,
|
||||
max_errors: int = DEFAULT_MAX_ERRORS,
|
||||
) -> None:
|
||||
self._buffer = bytearray()
|
||||
self._max_buffer_size = int(max_buffer_size)
|
||||
self._max_errors = int(max_errors)
|
||||
self._stopped = False
|
||||
self.stats = DecoderStats()
|
||||
|
||||
@property
|
||||
def stopped(self) -> bool:
|
||||
return self._stopped
|
||||
|
||||
def feed(self, data: bytes) -> None:
|
||||
if self._stopped:
|
||||
return
|
||||
if not data:
|
||||
return
|
||||
new_size = len(self._buffer) + len(data)
|
||||
if new_size > self._max_buffer_size:
|
||||
self._stopped = True
|
||||
raise BufferOverflowError(size=new_size, max_size=self._max_buffer_size)
|
||||
self._buffer.extend(data)
|
||||
|
||||
def decode_available(self) -> list[Frame]:
|
||||
"""Decode all complete frames currently in buffer."""
|
||||
out: list[Frame] = []
|
||||
if self._stopped:
|
||||
return out
|
||||
|
||||
while True:
|
||||
try:
|
||||
# Use memoryview to avoid full buffer copy on each iteration
|
||||
parsed = parse_frame(memoryview(self._buffer))
|
||||
except EventStreamParseError:
|
||||
self.stats.error_count += 1
|
||||
if self.stats.error_count >= self._max_errors:
|
||||
self._stopped = True
|
||||
raise
|
||||
|
||||
# Recovery: skip a byte and keep scanning.
|
||||
if self._buffer:
|
||||
del self._buffer[0]
|
||||
self.stats.bytes_skipped += 1
|
||||
else:
|
||||
break
|
||||
continue
|
||||
|
||||
if parsed is None:
|
||||
break
|
||||
|
||||
frame, consumed = parsed
|
||||
if consumed <= 0:
|
||||
break
|
||||
|
||||
out.append(frame)
|
||||
del self._buffer[:consumed]
|
||||
self.stats.frames_decoded += 1
|
||||
self.stats.error_count = 0
|
||||
|
||||
return out
|
||||
|
||||
|
||||
__all__ = [
|
||||
"DecoderStats",
|
||||
"EventStreamDecoder",
|
||||
]
|
||||
@@ -0,0 +1,72 @@
|
||||
"""AWS Event Stream parsing errors."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
class EventStreamParseError(Exception):
|
||||
"""Base error for AWS Event Stream decoding."""
|
||||
|
||||
|
||||
class IncompleteFrameError(EventStreamParseError):
|
||||
def __init__(self, *, needed: int, available: int) -> None:
|
||||
super().__init__(f"incomplete frame: needed={needed} available={available}")
|
||||
self.needed = needed
|
||||
self.available = available
|
||||
|
||||
|
||||
class MessageTooSmallError(EventStreamParseError):
|
||||
def __init__(self, *, length: int, min_length: int) -> None:
|
||||
super().__init__(f"message too small: length={length} min={min_length}")
|
||||
self.length = length
|
||||
self.min_length = min_length
|
||||
|
||||
|
||||
class MessageTooLargeError(EventStreamParseError):
|
||||
def __init__(self, *, length: int, max_length: int) -> None:
|
||||
super().__init__(f"message too large: length={length} max={max_length}")
|
||||
self.length = length
|
||||
self.max_length = max_length
|
||||
|
||||
|
||||
class PreludeCrcMismatchError(EventStreamParseError):
|
||||
def __init__(self, *, expected: int, actual: int) -> None:
|
||||
super().__init__(f"prelude crc mismatch: expected={expected} actual={actual}")
|
||||
self.expected = expected
|
||||
self.actual = actual
|
||||
|
||||
|
||||
class MessageCrcMismatchError(EventStreamParseError):
|
||||
def __init__(self, *, expected: int, actual: int) -> None:
|
||||
super().__init__(f"message crc mismatch: expected={expected} actual={actual}")
|
||||
self.expected = expected
|
||||
self.actual = actual
|
||||
|
||||
|
||||
class InvalidHeaderTypeError(EventStreamParseError):
|
||||
def __init__(self, type_id: int) -> None:
|
||||
super().__init__(f"invalid header type: {type_id}")
|
||||
self.type_id = type_id
|
||||
|
||||
|
||||
class HeaderParseError(EventStreamParseError):
|
||||
pass
|
||||
|
||||
|
||||
class BufferOverflowError(EventStreamParseError):
|
||||
def __init__(self, *, size: int, max_size: int) -> None:
|
||||
super().__init__(f"buffer overflow: size={size} max={max_size}")
|
||||
self.size = size
|
||||
self.max_size = max_size
|
||||
|
||||
|
||||
__all__ = [
|
||||
"BufferOverflowError",
|
||||
"EventStreamParseError",
|
||||
"HeaderParseError",
|
||||
"IncompleteFrameError",
|
||||
"InvalidHeaderTypeError",
|
||||
"MessageCrcMismatchError",
|
||||
"MessageTooLargeError",
|
||||
"MessageTooSmallError",
|
||||
"PreludeCrcMismatchError",
|
||||
]
|
||||
@@ -0,0 +1,95 @@
|
||||
"""AWS Event Stream message frame parsing."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from .crc import crc32
|
||||
from .error import (
|
||||
HeaderParseError,
|
||||
IncompleteFrameError,
|
||||
MessageCrcMismatchError,
|
||||
MessageTooLargeError,
|
||||
MessageTooSmallError,
|
||||
PreludeCrcMismatchError,
|
||||
)
|
||||
from .header import Headers, parse_headers
|
||||
|
||||
PRELUDE_SIZE = 12
|
||||
MIN_MESSAGE_SIZE = PRELUDE_SIZE + 4
|
||||
MAX_MESSAGE_SIZE = 16 * 1024 * 1024
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class Frame:
|
||||
headers: Headers
|
||||
payload: bytes
|
||||
|
||||
def message_type(self) -> str | None:
|
||||
return self.headers.message_type()
|
||||
|
||||
def event_type(self) -> str | None:
|
||||
return self.headers.event_type()
|
||||
|
||||
def payload_as_text(self) -> str:
|
||||
return self.payload.decode("utf-8", errors="replace")
|
||||
|
||||
|
||||
def parse_frame(buffer: bytes | memoryview) -> tuple[Frame, int] | None:
|
||||
"""Parse a single frame from the front of buffer.
|
||||
|
||||
Returns:
|
||||
(frame, consumed_bytes) if a full frame is available, otherwise None.
|
||||
|
||||
Raises:
|
||||
EventStreamParseError subclasses on validation errors.
|
||||
"""
|
||||
if len(buffer) < PRELUDE_SIZE:
|
||||
return None
|
||||
|
||||
total_length = int.from_bytes(buffer[0:4], "big", signed=False)
|
||||
header_length = int.from_bytes(buffer[4:8], "big", signed=False)
|
||||
prelude_crc = int.from_bytes(buffer[8:12], "big", signed=False)
|
||||
|
||||
if total_length < MIN_MESSAGE_SIZE:
|
||||
raise MessageTooSmallError(length=total_length, min_length=MIN_MESSAGE_SIZE)
|
||||
if total_length > MAX_MESSAGE_SIZE:
|
||||
raise MessageTooLargeError(length=total_length, max_length=MAX_MESSAGE_SIZE)
|
||||
|
||||
if len(buffer) < total_length:
|
||||
return None
|
||||
|
||||
actual_prelude_crc = crc32(buffer[0:8])
|
||||
if actual_prelude_crc != prelude_crc:
|
||||
raise PreludeCrcMismatchError(expected=prelude_crc, actual=actual_prelude_crc)
|
||||
|
||||
message_crc = int.from_bytes(buffer[total_length - 4 : total_length], "big", signed=False)
|
||||
actual_message_crc = crc32(buffer[0 : total_length - 4])
|
||||
if actual_message_crc != message_crc:
|
||||
raise MessageCrcMismatchError(expected=message_crc, actual=actual_message_crc)
|
||||
|
||||
headers_start = PRELUDE_SIZE
|
||||
headers_end = headers_start + header_length
|
||||
|
||||
if headers_end > total_length - 4:
|
||||
raise HeaderParseError("header length exceeds frame boundary")
|
||||
|
||||
headers = parse_headers(bytes(buffer[headers_start:headers_end]), header_length)
|
||||
|
||||
payload_start = headers_end
|
||||
payload_end = total_length - 4
|
||||
if payload_end < payload_start:
|
||||
raise IncompleteFrameError(needed=payload_start, available=payload_end)
|
||||
|
||||
payload = bytes(buffer[payload_start:payload_end])
|
||||
|
||||
return Frame(headers=headers, payload=payload), total_length
|
||||
|
||||
|
||||
__all__ = [
|
||||
"Frame",
|
||||
"MAX_MESSAGE_SIZE",
|
||||
"MIN_MESSAGE_SIZE",
|
||||
"PRELUDE_SIZE",
|
||||
"parse_frame",
|
||||
]
|
||||
@@ -0,0 +1,144 @@
|
||||
"""AWS Event Stream header parsing."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import IntEnum
|
||||
|
||||
from .error import HeaderParseError, IncompleteFrameError, InvalidHeaderTypeError
|
||||
|
||||
|
||||
class HeaderValueType(IntEnum):
|
||||
BOOL_TRUE = 0
|
||||
BOOL_FALSE = 1
|
||||
BYTE = 2
|
||||
SHORT = 3
|
||||
INTEGER = 4
|
||||
LONG = 5
|
||||
BYTE_ARRAY = 6
|
||||
STRING = 7
|
||||
TIMESTAMP = 8
|
||||
UUID = 9
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class Headers:
|
||||
values: dict[str, object]
|
||||
|
||||
def get(self, name: str) -> object | None:
|
||||
return self.values.get(name)
|
||||
|
||||
def get_string(self, name: str) -> str | None:
|
||||
v = self.values.get(name)
|
||||
return v if isinstance(v, str) else None
|
||||
|
||||
def message_type(self) -> str | None:
|
||||
return self.get_string(":message-type")
|
||||
|
||||
def event_type(self) -> str | None:
|
||||
return self.get_string(":event-type")
|
||||
|
||||
def exception_type(self) -> str | None:
|
||||
return self.get_string(":exception-type")
|
||||
|
||||
def error_code(self) -> str | None:
|
||||
return self.get_string(":error-code")
|
||||
|
||||
|
||||
def _ensure_bytes(data: bytes, offset: int, needed: int) -> None:
|
||||
available = len(data) - offset
|
||||
if available < needed:
|
||||
raise IncompleteFrameError(needed=needed, available=available)
|
||||
|
||||
|
||||
def parse_headers(data: bytes, header_length: int) -> Headers:
|
||||
if len(data) < header_length:
|
||||
raise IncompleteFrameError(needed=header_length, available=len(data))
|
||||
|
||||
values: dict[str, object] = {}
|
||||
offset = 0
|
||||
|
||||
while offset < header_length:
|
||||
_ensure_bytes(data, offset, 1)
|
||||
name_len = data[offset]
|
||||
offset += 1
|
||||
if name_len == 0:
|
||||
raise HeaderParseError("header name length cannot be 0")
|
||||
|
||||
_ensure_bytes(data, offset, name_len)
|
||||
name = data[offset : offset + name_len].decode("utf-8", errors="replace")
|
||||
offset += name_len
|
||||
|
||||
_ensure_bytes(data, offset, 1)
|
||||
type_id = data[offset]
|
||||
offset += 1
|
||||
try:
|
||||
value_type = HeaderValueType(type_id)
|
||||
except ValueError as e:
|
||||
raise InvalidHeaderTypeError(type_id) from e
|
||||
|
||||
if value_type == HeaderValueType.BOOL_TRUE:
|
||||
values[name] = True
|
||||
continue
|
||||
if value_type == HeaderValueType.BOOL_FALSE:
|
||||
values[name] = False
|
||||
continue
|
||||
|
||||
if value_type == HeaderValueType.BYTE:
|
||||
_ensure_bytes(data, offset, 1)
|
||||
values[name] = int.from_bytes(data[offset : offset + 1], "big", signed=True)
|
||||
offset += 1
|
||||
continue
|
||||
|
||||
if value_type == HeaderValueType.SHORT:
|
||||
_ensure_bytes(data, offset, 2)
|
||||
values[name] = int.from_bytes(data[offset : offset + 2], "big", signed=True)
|
||||
offset += 2
|
||||
continue
|
||||
|
||||
if value_type == HeaderValueType.INTEGER:
|
||||
_ensure_bytes(data, offset, 4)
|
||||
values[name] = int.from_bytes(data[offset : offset + 4], "big", signed=True)
|
||||
offset += 4
|
||||
continue
|
||||
|
||||
if value_type in (HeaderValueType.LONG, HeaderValueType.TIMESTAMP):
|
||||
_ensure_bytes(data, offset, 8)
|
||||
values[name] = int.from_bytes(data[offset : offset + 8], "big", signed=True)
|
||||
offset += 8
|
||||
continue
|
||||
|
||||
if value_type == HeaderValueType.BYTE_ARRAY:
|
||||
_ensure_bytes(data, offset, 2)
|
||||
length = int.from_bytes(data[offset : offset + 2], "big", signed=False)
|
||||
offset += 2
|
||||
_ensure_bytes(data, offset, length)
|
||||
values[name] = data[offset : offset + length]
|
||||
offset += length
|
||||
continue
|
||||
|
||||
if value_type == HeaderValueType.STRING:
|
||||
_ensure_bytes(data, offset, 2)
|
||||
length = int.from_bytes(data[offset : offset + 2], "big", signed=False)
|
||||
offset += 2
|
||||
_ensure_bytes(data, offset, length)
|
||||
values[name] = data[offset : offset + length].decode("utf-8", errors="replace")
|
||||
offset += length
|
||||
continue
|
||||
|
||||
if value_type == HeaderValueType.UUID:
|
||||
_ensure_bytes(data, offset, 16)
|
||||
values[name] = bytes(data[offset : offset + 16])
|
||||
offset += 16
|
||||
continue
|
||||
|
||||
raise HeaderParseError(f"unhandled header type: {value_type}")
|
||||
|
||||
return Headers(values=values)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"HeaderValueType",
|
||||
"Headers",
|
||||
"parse_headers",
|
||||
]
|
||||
@@ -0,0 +1,166 @@
|
||||
"""Kiro provider plugin — unified registration entry.
|
||||
|
||||
Kiro upstream looks like Claude CLI (Bearer token) from the outside, but uses a
|
||||
custom wire protocol:
|
||||
- Request: Claude Messages API -> Kiro generateAssistantResponse envelope
|
||||
- Response (stream): AWS Event Stream (binary) -> Claude SSE events
|
||||
|
||||
This plugin registers:
|
||||
- Envelope
|
||||
- Transport hook (dynamic region base_url)
|
||||
- Model fetcher (fixed model catalog — Kiro has no /v1/models endpoint)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from urllib.parse import urlencode
|
||||
|
||||
from src.services.provider.adapters.kiro.constants import (
|
||||
DEFAULT_REGION,
|
||||
KIRO_GENERATE_ASSISTANT_PATH,
|
||||
)
|
||||
from src.services.provider.adapters.kiro.context import get_kiro_request_context
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixed model catalog
|
||||
# ---------------------------------------------------------------------------
|
||||
# Kiro upstream has no /v1/models endpoint. We return a static list matching
|
||||
# the models accepted by map_model() in converter.py.
|
||||
_KIRO_MODELS: list[dict[str, Any]] = [
|
||||
{
|
||||
"id": "claude-sonnet-4.5",
|
||||
"object": "model",
|
||||
"owned_by": "anthropic",
|
||||
"display_name": "Claude Sonnet 4.5",
|
||||
},
|
||||
{
|
||||
"id": "claude-opus-4.5",
|
||||
"object": "model",
|
||||
"owned_by": "anthropic",
|
||||
"display_name": "Claude Opus 4.5",
|
||||
},
|
||||
{
|
||||
"id": "claude-opus-4.6",
|
||||
"object": "model",
|
||||
"owned_by": "anthropic",
|
||||
"display_name": "Claude Opus 4.6",
|
||||
},
|
||||
{
|
||||
"id": "claude-haiku-4.5",
|
||||
"object": "model",
|
||||
"owned_by": "anthropic",
|
||||
"display_name": "Claude Haiku 4.5",
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
async def fetch_models_kiro(
|
||||
ctx: Any,
|
||||
timeout_seconds: float, # noqa: ARG001
|
||||
) -> tuple[list[dict], list[str], bool, dict[str, Any] | None]:
|
||||
"""Return a fixed model catalog for Kiro.
|
||||
|
||||
Kiro upstream does not expose a ``/v1/models`` endpoint, so we skip the
|
||||
HTTP call entirely and return a hardcoded list.
|
||||
"""
|
||||
_ = ctx # not needed — no upstream call
|
||||
return list(_KIRO_MODELS), [], True, None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Transport hook
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def build_kiro_url(
|
||||
endpoint: Any,
|
||||
*,
|
||||
is_stream: bool,
|
||||
effective_query_params: dict[str, Any],
|
||||
) -> str:
|
||||
"""Build Kiro generateAssistantResponse URL.
|
||||
|
||||
Endpoint base_url may contain a `{region}` placeholder. The actual region is
|
||||
resolved from per-request context (set by the envelope).
|
||||
"""
|
||||
_ = is_stream
|
||||
|
||||
base = str(getattr(endpoint, "base_url", "") or "").rstrip("/")
|
||||
|
||||
ctx = get_kiro_request_context()
|
||||
region = (ctx.region if ctx else "") or DEFAULT_REGION
|
||||
if "{region}" in base:
|
||||
base = base.replace("{region}", region)
|
||||
|
||||
path = KIRO_GENERATE_ASSISTANT_PATH
|
||||
url = base if base.endswith(path) else f"{base}{path}"
|
||||
|
||||
if effective_query_params:
|
||||
query_string = urlencode(effective_query_params, doseq=True)
|
||||
if query_string:
|
||||
url = f"{url}?{query_string}"
|
||||
|
||||
return url
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Export builder
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_KIRO_SKIP_KEYS = frozenset(
|
||||
{
|
||||
"access_token",
|
||||
"expires_at",
|
||||
"updated_at",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def kiro_export_builder(
|
||||
auth_config: dict[str, Any],
|
||||
upstream_metadata: dict[str, Any] | None,
|
||||
) -> dict[str, Any]:
|
||||
"""Kiro 导出:保留 auth_method / refresh_token / machine_id / profile_arn 等,
|
||||
IdC 模式额外保留 client_id / client_secret / region。"""
|
||||
data = {
|
||||
k: v
|
||||
for k, v in auth_config.items()
|
||||
if k not in _KIRO_SKIP_KEYS and v is not None and v != ""
|
||||
}
|
||||
# email 可能仅在 upstream_metadata.kiro 中
|
||||
if not data.get("email"):
|
||||
kiro_meta = (upstream_metadata or {}).get("kiro") or {}
|
||||
if kiro_meta.get("email"):
|
||||
data["email"] = kiro_meta["email"]
|
||||
return data
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Registration
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def register_all() -> None:
|
||||
"""Register all Kiro hooks into shared registries."""
|
||||
|
||||
from src.services.model.upstream_fetcher import UpstreamModelsFetcherRegistry
|
||||
from src.services.provider.adapters.kiro.envelope import kiro_envelope
|
||||
from src.services.provider.envelope import register_envelope
|
||||
from src.services.provider.export import register_export_builder
|
||||
from src.services.provider.transport import register_transport_hook
|
||||
|
||||
register_envelope("kiro", "claude:cli", kiro_envelope)
|
||||
register_envelope("kiro", "", kiro_envelope)
|
||||
|
||||
register_transport_hook("kiro", "claude:cli", build_kiro_url)
|
||||
|
||||
register_export_builder("kiro", kiro_export_builder)
|
||||
|
||||
UpstreamModelsFetcherRegistry.register(
|
||||
provider_types=["kiro"],
|
||||
fetcher=fetch_models_kiro,
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["build_kiro_url", "fetch_models_kiro", "kiro_export_builder", "register_all"]
|
||||
@@ -0,0 +1,312 @@
|
||||
"""Kiro token refresh helpers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
import json
|
||||
import re
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.core.logger import logger
|
||||
from src.services.provider.adapters.kiro.headers import build_kiro_ide_tag
|
||||
from src.services.provider.adapters.kiro.models.credentials import KiroAuthConfig
|
||||
|
||||
IDC_AMZ_USER_AGENT = (
|
||||
"aws-sdk-js/3.738.0 ua/2.1 os/other lang/js md/browser#unknown_unknown "
|
||||
"api/sso-oidc#3.738.0 m/E KiroIDE"
|
||||
)
|
||||
|
||||
_REGION_RE = re.compile(r"^[a-z]{2}-[a-z0-9-]+-\d+$")
|
||||
_HEX64_RE = re.compile(r"^[0-9a-fA-F]{64}$")
|
||||
_UUID_RE = re.compile(
|
||||
r"^[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}$"
|
||||
)
|
||||
|
||||
|
||||
def validate_refresh_token(refresh_token: str) -> None:
|
||||
token = str(refresh_token or "").strip()
|
||||
if not token:
|
||||
raise ValueError("missing refresh_token")
|
||||
|
||||
# kiro.rs: length < 100 or contains "..." is considered truncated.
|
||||
if len(token) < 100 or token.endswith("...") or "..." in token:
|
||||
raise ValueError(
|
||||
"refresh_token appears truncated; please export the full token from Kiro IDE"
|
||||
)
|
||||
|
||||
|
||||
def normalize_machine_id(machine_id: str) -> str | None:
|
||||
raw = str(machine_id or "").strip()
|
||||
if not raw:
|
||||
return None
|
||||
|
||||
if _HEX64_RE.fullmatch(raw):
|
||||
return raw.lower()
|
||||
|
||||
if _UUID_RE.fullmatch(raw):
|
||||
without = raw.replace("-", "").lower()
|
||||
return without + without
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def generate_machine_id(cfg: KiroAuthConfig) -> str:
|
||||
normalized = normalize_machine_id(cfg.machine_id or "")
|
||||
if normalized:
|
||||
return normalized
|
||||
|
||||
validate_refresh_token(cfg.refresh_token)
|
||||
seed = f"KotlinNativeAPI/{cfg.refresh_token}".encode("utf-8")
|
||||
return hashlib.sha256(seed).hexdigest()
|
||||
|
||||
|
||||
def is_token_expired(expires_at: int | None, *, skew_seconds: int = 120) -> bool:
|
||||
try:
|
||||
ts = int(expires_at or 0)
|
||||
except Exception:
|
||||
ts = 0
|
||||
if ts <= 0:
|
||||
return True
|
||||
return int(time.time()) >= ts - int(skew_seconds)
|
||||
|
||||
|
||||
def _resolve_region(cfg: KiroAuthConfig) -> str:
|
||||
region = str(cfg.region or "").strip()
|
||||
if region and _REGION_RE.fullmatch(region):
|
||||
return region
|
||||
# Keep best-effort fallback; actual host parsing happens in transport hook.
|
||||
from src.services.provider.adapters.kiro.constants import DEFAULT_REGION
|
||||
|
||||
return region or DEFAULT_REGION
|
||||
|
||||
|
||||
def _try_extract_email_from_jwt(token: str) -> str | None:
|
||||
"""尝试从 JWT access_token 中提取 email。
|
||||
|
||||
Kiro Social / IdC 返回的 accessToken 可能是 JWT 格式,
|
||||
payload 中可能包含 email 字段。仅做 base64 解码,不验证签名。
|
||||
失败时静默返回 None。
|
||||
"""
|
||||
try:
|
||||
parts = token.split(".")
|
||||
if len(parts) != 3:
|
||||
return None
|
||||
# base64url decode the payload (second segment)
|
||||
payload_b64 = parts[1]
|
||||
# Add padding
|
||||
padding = 4 - len(payload_b64) % 4
|
||||
if padding != 4:
|
||||
payload_b64 += "=" * padding
|
||||
payload_bytes = base64.urlsafe_b64decode(payload_b64)
|
||||
claims = json.loads(payload_bytes)
|
||||
# Try common email claim keys
|
||||
for key in ("email", "Email", "mail", "upn"):
|
||||
val = claims.get(key)
|
||||
if isinstance(val, str) and "@" in val:
|
||||
return val.strip()
|
||||
except Exception:
|
||||
pass
|
||||
return None
|
||||
|
||||
|
||||
async def refresh_social_token(
|
||||
cfg: KiroAuthConfig,
|
||||
*,
|
||||
proxy_config: dict[str, Any] | None,
|
||||
timeout_seconds: float = 30.0,
|
||||
) -> tuple[str, KiroAuthConfig]:
|
||||
"""Refresh access token via Kiro Social refresh endpoint."""
|
||||
validate_refresh_token(cfg.refresh_token)
|
||||
|
||||
region = _resolve_region(cfg)
|
||||
url = f"https://prod.{region}.auth.desktop.kiro.dev/refreshToken"
|
||||
host = f"prod.{region}.auth.desktop.kiro.dev"
|
||||
|
||||
machine_id = generate_machine_id(cfg)
|
||||
kiro_version = (cfg.kiro_version or "").strip() or "0.8.0"
|
||||
ua = build_kiro_ide_tag(kiro_version=kiro_version, machine_id=machine_id)
|
||||
|
||||
body = {"refreshToken": cfg.refresh_token}
|
||||
headers = {
|
||||
"User-Agent": ua,
|
||||
"Host": host,
|
||||
"Accept": "application/json, text/plain, */*",
|
||||
"Content-Type": "application/json",
|
||||
"Connection": "close",
|
||||
"Accept-Encoding": "gzip, compress, deflate, br",
|
||||
}
|
||||
|
||||
client = await HTTPClientPool.get_proxy_client(proxy_config=proxy_config)
|
||||
resp = await client.post(
|
||||
url,
|
||||
headers=headers,
|
||||
json=body,
|
||||
timeout=httpx.Timeout(timeout_seconds),
|
||||
)
|
||||
|
||||
if resp.status_code < 200 or resp.status_code >= 300:
|
||||
logger.debug(
|
||||
"kiro social refresh error: HTTP {} | {}",
|
||||
resp.status_code,
|
||||
(resp.text or "").strip()[:200],
|
||||
)
|
||||
raise RuntimeError(f"kiro social refresh failed: HTTP {resp.status_code}")
|
||||
|
||||
data: dict[str, Any]
|
||||
try:
|
||||
data = resp.json()
|
||||
except Exception as e:
|
||||
raise RuntimeError("kiro social refresh: invalid json response") from e
|
||||
|
||||
access_token = str(data.get("accessToken") or "").strip()
|
||||
if not access_token:
|
||||
raise RuntimeError("kiro social refresh returned empty accessToken")
|
||||
|
||||
new_cfg = KiroAuthConfig.from_dict(cfg.to_dict())
|
||||
|
||||
# refreshToken/profileArn may rotate
|
||||
rt = data.get("refreshToken")
|
||||
if isinstance(rt, str) and rt.strip():
|
||||
new_cfg.refresh_token = rt.strip()
|
||||
|
||||
profile_arn = data.get("profileArn")
|
||||
if isinstance(profile_arn, str) and profile_arn.strip():
|
||||
new_cfg.profile_arn = profile_arn.strip()
|
||||
|
||||
expires_in = data.get("expiresIn")
|
||||
try:
|
||||
if expires_in is not None:
|
||||
new_cfg.expires_at = int(time.time()) + int(expires_in)
|
||||
except Exception:
|
||||
new_cfg.expires_at = int(time.time()) + 3600
|
||||
|
||||
# Persist computed machine_id if user didn't provide one.
|
||||
if not (cfg.machine_id or "").strip():
|
||||
new_cfg.machine_id = machine_id
|
||||
|
||||
# 尝试从 accessToken 中提取 email(如果尚未设置)
|
||||
if not (new_cfg.email or "").strip():
|
||||
extracted_email = _try_extract_email_from_jwt(access_token)
|
||||
if extracted_email:
|
||||
new_cfg.email = extracted_email
|
||||
logger.debug("kiro social: extracted email from accessToken: {}", extracted_email)
|
||||
|
||||
# 缓存 access_token
|
||||
new_cfg.access_token = access_token
|
||||
|
||||
return access_token, new_cfg
|
||||
|
||||
|
||||
async def refresh_idc_token(
|
||||
cfg: KiroAuthConfig,
|
||||
*,
|
||||
proxy_config: dict[str, Any] | None,
|
||||
timeout_seconds: float = 30.0,
|
||||
) -> tuple[str, KiroAuthConfig]:
|
||||
"""Refresh access token via AWS SSO OIDC endpoint (IdC)."""
|
||||
validate_refresh_token(cfg.refresh_token)
|
||||
|
||||
if not (cfg.client_id or "").strip() or not (cfg.client_secret or "").strip():
|
||||
raise ValueError("idc refresh requires client_id and client_secret")
|
||||
|
||||
region = _resolve_region(cfg)
|
||||
url = f"https://oidc.{region}.amazonaws.com/token"
|
||||
host = f"oidc.{region}.amazonaws.com"
|
||||
|
||||
body = {
|
||||
"clientId": cfg.client_id,
|
||||
"clientSecret": cfg.client_secret,
|
||||
"refreshToken": cfg.refresh_token,
|
||||
"grantType": "refresh_token",
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Host": host,
|
||||
"x-amz-user-agent": IDC_AMZ_USER_AGENT,
|
||||
"User-Agent": "node",
|
||||
"Accept": "*/*",
|
||||
}
|
||||
|
||||
client = await HTTPClientPool.get_proxy_client(proxy_config=proxy_config)
|
||||
resp = await client.post(
|
||||
url,
|
||||
headers=headers,
|
||||
json=body,
|
||||
timeout=httpx.Timeout(timeout_seconds),
|
||||
)
|
||||
|
||||
if resp.status_code < 200 or resp.status_code >= 300:
|
||||
logger.debug(
|
||||
"kiro idc refresh error: HTTP {} | {}",
|
||||
resp.status_code,
|
||||
(resp.text or "").strip()[:200],
|
||||
)
|
||||
raise RuntimeError(f"kiro idc refresh failed: HTTP {resp.status_code}")
|
||||
|
||||
data: dict[str, Any]
|
||||
try:
|
||||
data = resp.json()
|
||||
except Exception as e:
|
||||
raise RuntimeError("kiro idc refresh: invalid json response") from e
|
||||
|
||||
access_token = str(data.get("accessToken") or "").strip()
|
||||
if not access_token:
|
||||
raise RuntimeError("kiro idc refresh returned empty accessToken")
|
||||
|
||||
new_cfg = KiroAuthConfig.from_dict(cfg.to_dict())
|
||||
|
||||
rt = data.get("refreshToken")
|
||||
if isinstance(rt, str) and rt.strip():
|
||||
new_cfg.refresh_token = rt.strip()
|
||||
|
||||
expires_in = data.get("expiresIn")
|
||||
try:
|
||||
if expires_in is not None:
|
||||
new_cfg.expires_at = int(time.time()) + int(expires_in)
|
||||
except Exception:
|
||||
new_cfg.expires_at = int(time.time()) + 3600
|
||||
|
||||
# Persist computed machine_id if user didn't provide one.
|
||||
if not (cfg.machine_id or "").strip():
|
||||
new_cfg.machine_id = generate_machine_id(cfg)
|
||||
|
||||
# 尝试从 accessToken 中提取 email(如果尚未设置)
|
||||
if not (new_cfg.email or "").strip():
|
||||
extracted_email = _try_extract_email_from_jwt(access_token)
|
||||
if extracted_email:
|
||||
new_cfg.email = extracted_email
|
||||
logger.debug("kiro idc: extracted email from accessToken: {}", extracted_email)
|
||||
|
||||
# 缓存 access_token
|
||||
new_cfg.access_token = access_token
|
||||
|
||||
return access_token, new_cfg
|
||||
|
||||
|
||||
async def refresh_access_token(
|
||||
cfg: KiroAuthConfig,
|
||||
*,
|
||||
proxy_config: dict[str, Any] | None,
|
||||
) -> tuple[str, KiroAuthConfig]:
|
||||
method = (cfg.auth_method or "social").strip().lower()
|
||||
if method == "idc":
|
||||
return await refresh_idc_token(cfg, proxy_config=proxy_config)
|
||||
return await refresh_social_token(cfg, proxy_config=proxy_config)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"IDC_AMZ_USER_AGENT",
|
||||
"generate_machine_id",
|
||||
"is_token_expired",
|
||||
"normalize_machine_id",
|
||||
"refresh_access_token",
|
||||
"refresh_idc_token",
|
||||
"refresh_social_token",
|
||||
"validate_refresh_token",
|
||||
]
|
||||
@@ -0,0 +1,183 @@
|
||||
"""Kiro usage/quota fetching utilities."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
from src.clients.http_client import HTTPClientPool
|
||||
from src.core.logger import logger
|
||||
from src.services.provider.adapters.kiro.headers import (
|
||||
build_user_agent_usage,
|
||||
build_x_amz_user_agent_usage,
|
||||
)
|
||||
from src.services.provider.adapters.kiro.models.credentials import KiroAuthConfig
|
||||
from src.services.provider.adapters.kiro.models.usage_limits import (
|
||||
UsageLimitsResponse,
|
||||
calculate_current_usage,
|
||||
calculate_total_usage_limit,
|
||||
)
|
||||
from src.services.provider.adapters.kiro.token_manager import (
|
||||
generate_machine_id,
|
||||
is_token_expired,
|
||||
refresh_access_token,
|
||||
)
|
||||
|
||||
|
||||
async def fetch_kiro_usage_limits(
|
||||
auth_config: dict[str, Any],
|
||||
proxy_config: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
调用 Kiro getUsageLimits API 获取使用额度信息
|
||||
|
||||
Args:
|
||||
auth_config: 解密后的 KiroAuthConfig 数据
|
||||
proxy_config: 代理配置(可选)
|
||||
|
||||
Returns:
|
||||
包含 usage_data 和 updated_auth_config 的字典
|
||||
|
||||
Raises:
|
||||
RuntimeError: 请求失败时抛出
|
||||
"""
|
||||
cfg = KiroAuthConfig.from_dict(auth_config)
|
||||
|
||||
# 检查是否有缓存的 access_token 且未过期
|
||||
access_token: str | None = None
|
||||
updated_cfg: KiroAuthConfig | None = None
|
||||
|
||||
if cfg.access_token and not is_token_expired(cfg.expires_at):
|
||||
# 使用缓存的 token
|
||||
access_token = cfg.access_token
|
||||
updated_cfg = cfg
|
||||
logger.debug("[KIRO_QUOTA] 使用缓存的 access_token")
|
||||
else:
|
||||
# token 过期或不存在,需要刷新
|
||||
logger.debug("[KIRO_QUOTA] Token 已过期或不存在,正在刷新...")
|
||||
access_token, updated_cfg = await refresh_access_token(cfg, proxy_config=proxy_config)
|
||||
|
||||
if not access_token:
|
||||
raise RuntimeError("无法获取 Kiro access_token")
|
||||
|
||||
# 构建请求
|
||||
from src.services.provider.adapters.kiro.constants import DEFAULT_REGION
|
||||
|
||||
region = (updated_cfg.region if updated_cfg else cfg.region) or DEFAULT_REGION
|
||||
host = f"q.{region}.amazonaws.com"
|
||||
machine_id = generate_machine_id(updated_cfg or cfg)
|
||||
kiro_version = (updated_cfg.kiro_version if updated_cfg else cfg.kiro_version) or "0.8.0"
|
||||
|
||||
# 构建 URL(添加 isEmailRequired=true 获取邮箱)
|
||||
url = f"https://{host}/getUsageLimits?origin=AI_EDITOR&resourceType=AGENTIC_REQUEST&isEmailRequired=true"
|
||||
|
||||
profile_arn = updated_cfg.profile_arn if updated_cfg else cfg.profile_arn
|
||||
if profile_arn:
|
||||
from urllib.parse import quote
|
||||
|
||||
url += f"&profileArn={quote(profile_arn, safe='')}"
|
||||
|
||||
# 构建 headers
|
||||
headers = {
|
||||
"x-amz-user-agent": build_x_amz_user_agent_usage(
|
||||
kiro_version=kiro_version, machine_id=machine_id
|
||||
),
|
||||
"User-Agent": build_user_agent_usage(kiro_version=kiro_version, machine_id=machine_id),
|
||||
"host": host,
|
||||
"amz-sdk-invocation-id": str(uuid.uuid4()),
|
||||
"amz-sdk-request": "attempt=1; max=1",
|
||||
"Authorization": f"Bearer {access_token}",
|
||||
"Connection": "close",
|
||||
}
|
||||
|
||||
client = await HTTPClientPool.get_proxy_client(proxy_config=proxy_config)
|
||||
response = await client.get(url, headers=headers, timeout=httpx.Timeout(30.0))
|
||||
|
||||
if response.status_code != 200:
|
||||
error_msg = {
|
||||
401: "认证失败,Token 无效或已过期",
|
||||
403: "权限不足,无法获取使用额度",
|
||||
429: "请求过于频繁,已被限流",
|
||||
}.get(response.status_code, "获取使用额度失败")
|
||||
if 500 <= response.status_code < 600:
|
||||
error_msg = "服务器错误,AWS 服务暂时不可用"
|
||||
logger.debug(
|
||||
"kiro usage API error: HTTP {} | {}",
|
||||
response.status_code,
|
||||
(response.text or "").strip()[:200],
|
||||
)
|
||||
raise RuntimeError(f"{error_msg}: HTTP {response.status_code}")
|
||||
|
||||
try:
|
||||
data = response.json()
|
||||
except Exception as exc:
|
||||
raise RuntimeError(f"获取使用额度成功但响应解析失败: HTTP {response.status_code}") from exc
|
||||
|
||||
# 返回刷新后的配置(用于更新 auth_config)
|
||||
return {
|
||||
"usage_data": data,
|
||||
"updated_auth_config": updated_cfg.to_dict() if updated_cfg else None,
|
||||
}
|
||||
|
||||
|
||||
def parse_kiro_usage_response(data: dict) -> dict | None:
|
||||
"""
|
||||
解析 Kiro getUsageLimits API 响应,提取限额信息和用户邮箱
|
||||
|
||||
返回格式与 kiro.rs BalanceResponse 类似:
|
||||
- subscription_title: 订阅类型(如 "KIRO PRO+")
|
||||
- current_usage: 当前使用量
|
||||
- usage_limit: 使用限额
|
||||
- remaining: 剩余额度
|
||||
- usage_percentage: 使用百分比
|
||||
- next_reset_at: 下次重置时间(Unix 时间戳)
|
||||
- email: 用户邮箱(通过 isEmailRequired=true 获取)
|
||||
"""
|
||||
if not data:
|
||||
return None
|
||||
|
||||
usage_resp = UsageLimitsResponse.from_dict(data)
|
||||
|
||||
current_usage = calculate_current_usage(usage_resp)
|
||||
usage_limit = calculate_total_usage_limit(usage_resp)
|
||||
remaining = max(usage_limit - current_usage, 0.0)
|
||||
usage_percentage = (current_usage / usage_limit * 100.0) if usage_limit > 0 else 0.0
|
||||
usage_percentage = min(usage_percentage, 100.0)
|
||||
|
||||
result: dict[str, Any] = {
|
||||
"current_usage": current_usage,
|
||||
"usage_limit": usage_limit,
|
||||
"remaining": remaining,
|
||||
"usage_percentage": usage_percentage,
|
||||
}
|
||||
|
||||
# 订阅类型
|
||||
if usage_resp.subscription_info and usage_resp.subscription_info.subscription_title:
|
||||
result["subscription_title"] = usage_resp.subscription_info.subscription_title
|
||||
|
||||
# 下次重置时间
|
||||
if usage_resp.next_date_reset is not None:
|
||||
result["next_reset_at"] = usage_resp.next_date_reset
|
||||
elif usage_resp.usage_breakdown_list and usage_resp.usage_breakdown_list[0].next_date_reset:
|
||||
result["next_reset_at"] = usage_resp.usage_breakdown_list[0].next_date_reset
|
||||
|
||||
# 解析用户邮箱(从 desktopUserInfo 或 userInfo 中获取)
|
||||
user_info = data.get("desktopUserInfo") or data.get("userInfo") or {}
|
||||
if isinstance(user_info, dict):
|
||||
email = user_info.get("email")
|
||||
if isinstance(email, str) and email.strip():
|
||||
result["email"] = email.strip()
|
||||
|
||||
# 添加更新时间戳
|
||||
result["updated_at"] = int(time.time())
|
||||
|
||||
return result
|
||||
|
||||
|
||||
__all__ = [
|
||||
"fetch_kiro_usage_limits",
|
||||
"parse_kiro_usage_response",
|
||||
]
|
||||
@@ -120,9 +120,11 @@ def ensure_providers_bootstrapped() -> None:
|
||||
register_all as _reg_antigravity,
|
||||
)
|
||||
from src.services.provider.adapters.codex.plugin import register_all as _reg_codex
|
||||
from src.services.provider.adapters.kiro.plugin import register_all as _reg_kiro
|
||||
|
||||
_reg_antigravity()
|
||||
_reg_codex()
|
||||
_reg_kiro()
|
||||
|
||||
|
||||
__all__ = [
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
"""OAuth Key 导出:provider-specific export builders.
|
||||
|
||||
每个 Provider adapter 在 ``register_all()`` 时注册自己的 ``build_export_data``,
|
||||
导出端点通过 ``build_export_data()`` 分发。
|
||||
|
||||
未注册的 provider_type 使用默认实现(strip null + 临时字段)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Callable
|
||||
|
||||
# 所有 provider 共用的临时/无用字段
|
||||
_DEFAULT_SKIP_KEYS = frozenset(
|
||||
{
|
||||
"access_token",
|
||||
"expires_at",
|
||||
"updated_at",
|
||||
"token_type",
|
||||
"scope",
|
||||
}
|
||||
)
|
||||
|
||||
ExportBuilder = Callable[[dict[str, Any], dict[str, Any] | None], dict[str, Any]]
|
||||
|
||||
_BUILDERS: dict[str, ExportBuilder] = {}
|
||||
|
||||
|
||||
def register_export_builder(provider_type: str, builder: ExportBuilder) -> None:
|
||||
_BUILDERS[provider_type] = builder
|
||||
|
||||
|
||||
def _default_builder(
|
||||
auth_config: dict[str, Any],
|
||||
upstream_metadata: dict[str, Any] | None,
|
||||
) -> dict[str, Any]:
|
||||
"""默认导出:去掉 null、空字符串、临时字段。"""
|
||||
return {
|
||||
k: v
|
||||
for k, v in auth_config.items()
|
||||
if k not in _DEFAULT_SKIP_KEYS and v is not None and v != ""
|
||||
}
|
||||
|
||||
|
||||
def build_export_data(
|
||||
provider_type: str,
|
||||
auth_config: dict[str, Any],
|
||||
upstream_metadata: dict[str, Any] | None,
|
||||
) -> dict[str, Any]:
|
||||
"""调用 provider-specific builder 构建导出数据。"""
|
||||
builder = _BUILDERS.get(provider_type, _default_builder)
|
||||
return builder(auth_config, upstream_metadata)
|
||||
@@ -1,156 +0,0 @@
|
||||
"""
|
||||
上游元数据采集器(MetadataCollector)
|
||||
|
||||
可扩展注册表模式:
|
||||
- 每个 Provider 类型可注册一个 MetadataCollector
|
||||
- 从响应头解析有价值的元数据(额度、限流等)
|
||||
- 解析结果存入 ProviderAPIKey.upstream_metadata
|
||||
|
||||
扩展方式:
|
||||
1. 创建新文件实现 MetadataCollector
|
||||
2. 在本文件底部注册
|
||||
"""
|
||||
|
||||
import time
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, ClassVar
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from src.core.logger import logger
|
||||
|
||||
# 节流:每个 key_id 至少间隔 _THROTTLE_SECONDS 秒才写入一次
|
||||
_THROTTLE_SECONDS = 30
|
||||
_last_write_ts: dict[str, float] = {}
|
||||
|
||||
|
||||
class MetadataCollector(ABC):
|
||||
"""元数据采集器基类"""
|
||||
|
||||
# 支持的 provider_type 列表(小写)
|
||||
PROVIDER_TYPES: ClassVar[list[str]] = []
|
||||
|
||||
@abstractmethod
|
||||
def parse_headers(self, headers: dict[str, str]) -> dict[str, Any] | None:
|
||||
"""解析响应头,返回结构化元数据。返回 None 表示无可用数据。"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class MetadataCollectorRegistry:
|
||||
"""元数据采集器注册表"""
|
||||
|
||||
_collectors: ClassVar[list[MetadataCollector]] = []
|
||||
_type_index: ClassVar[dict[str, MetadataCollector]] = {}
|
||||
|
||||
@classmethod
|
||||
def register(cls, collector: MetadataCollector) -> None:
|
||||
cls._collectors.append(collector)
|
||||
for pt in collector.PROVIDER_TYPES:
|
||||
cls._type_index[pt.lower()] = collector
|
||||
logger.debug(
|
||||
"[MetadataCollectorRegistry] 注册: {} -> {}",
|
||||
collector.__class__.__name__,
|
||||
collector.PROVIDER_TYPES,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def collect(cls, provider_type: str, headers: dict[str, str]) -> dict[str, Any] | None:
|
||||
"""根据 provider_type 查找采集器并解析响应头"""
|
||||
collector = cls._type_index.get(provider_type.lower())
|
||||
if collector is None:
|
||||
return None
|
||||
try:
|
||||
return collector.parse_headers(headers)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"[MetadataCollectorRegistry] {} 解析失败", collector.__class__.__name__
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
_initialized = False
|
||||
|
||||
|
||||
def _ensure_collectors_registered() -> None:
|
||||
"""惰性注册所有采集器(首次调用时执行,避免循环导入)"""
|
||||
global _initialized
|
||||
if _initialized:
|
||||
return
|
||||
_initialized = True
|
||||
|
||||
# 延迟导入,避免模块加载时的循环依赖
|
||||
from src.services.provider.adapters.codex.metadata_collector import CodexMetadataCollector
|
||||
|
||||
MetadataCollectorRegistry.register(CodexMetadataCollector())
|
||||
|
||||
|
||||
def ensure_collectors_registered() -> None:
|
||||
"""Ensure metadata collectors are registered (idempotent)."""
|
||||
_ensure_collectors_registered()
|
||||
|
||||
|
||||
def collect_and_save_upstream_metadata(
|
||||
db: Session,
|
||||
*,
|
||||
provider_type: str,
|
||||
key_id: str,
|
||||
response_headers: dict[str, str],
|
||||
request_id: str,
|
||||
) -> None:
|
||||
"""采集上游元数据并更新 ProviderAPIKey.upstream_metadata(带节流)
|
||||
|
||||
每个 key_id 至少间隔 _THROTTLE_SECONDS 秒才执行一次数据库写入,
|
||||
避免高并发时频繁更新同一行。
|
||||
|
||||
Args:
|
||||
db: 数据库 Session
|
||||
provider_type: Provider 类型(如 "codex")
|
||||
key_id: ProviderAPIKey.id
|
||||
response_headers: 上游响应头
|
||||
request_id: 请求 ID(用于日志)
|
||||
"""
|
||||
if not provider_type or not key_id or not response_headers:
|
||||
return
|
||||
|
||||
# 确保采集器已注册
|
||||
_ensure_collectors_registered()
|
||||
|
||||
# 节流检查
|
||||
now = time.monotonic()
|
||||
last_ts = _last_write_ts.get(key_id, 0.0)
|
||||
if now - last_ts < _THROTTLE_SECONDS:
|
||||
return
|
||||
|
||||
try:
|
||||
metadata = MetadataCollectorRegistry.collect(provider_type, response_headers)
|
||||
if metadata is None:
|
||||
return
|
||||
|
||||
from src.models.database import ProviderAPIKey
|
||||
|
||||
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
|
||||
if key is None:
|
||||
return
|
||||
|
||||
key.upstream_metadata = metadata
|
||||
db.commit()
|
||||
_last_write_ts[key_id] = now
|
||||
logger.debug(
|
||||
"[{}] 已更新 ProviderAPIKey({}) upstream_metadata",
|
||||
request_id,
|
||||
key_id,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("[{}] 采集上游元数据失败", request_id)
|
||||
try:
|
||||
db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
__all__ = [
|
||||
"MetadataCollector",
|
||||
"MetadataCollectorRegistry",
|
||||
"collect_and_save_upstream_metadata",
|
||||
"ensure_collectors_registered",
|
||||
]
|
||||
@@ -78,12 +78,18 @@ def get_upstream_stream_policy(
|
||||
and parsed == UpstreamStreamPolicy.FORCE_NON_STREAM
|
||||
):
|
||||
return UpstreamStreamPolicy.FORCE_STREAM
|
||||
if pt == ProviderType.KIRO and parsed == UpstreamStreamPolicy.FORCE_NON_STREAM:
|
||||
return UpstreamStreamPolicy.FORCE_STREAM
|
||||
return parsed
|
||||
|
||||
# Safe-by-default: Codex Responses OAuth behaves like SSE-only.
|
||||
if pt == ProviderType.CODEX and sig == "openai:cli":
|
||||
return UpstreamStreamPolicy.FORCE_STREAM
|
||||
|
||||
# Kiro upstream streams binary AWS Event Stream; treat as stream-only.
|
||||
if pt == ProviderType.KIRO:
|
||||
return UpstreamStreamPolicy.FORCE_STREAM
|
||||
|
||||
return UpstreamStreamPolicy.AUTO
|
||||
|
||||
|
||||
|
||||
@@ -101,8 +101,18 @@ def get_antigravity_base_url() -> str | None:
|
||||
return get_selected_base_url()
|
||||
|
||||
|
||||
def _get_provider_type(endpoint: Any, key: "ProviderAPIKey" | None = None) -> str | None:
|
||||
"""尽力获取 Provider.provider_type(用于 Antigravity 等 Provider 特判)。"""
|
||||
def _get_provider_type(
|
||||
endpoint: Any,
|
||||
key: "ProviderAPIKey" | None = None,
|
||||
decrypted_auth_config: dict[str, Any] | None = None,
|
||||
) -> str | None:
|
||||
"""尽力获取 Provider.provider_type(用于 Antigravity 等 Provider 特判)。
|
||||
|
||||
优先级:
|
||||
1. endpoint.provider.provider_type
|
||||
2. key.provider.provider_type
|
||||
3. decrypted_auth_config["provider_type"](OAuth 导入的凭证)
|
||||
"""
|
||||
try:
|
||||
provider = getattr(endpoint, "provider", None)
|
||||
if provider is not None:
|
||||
@@ -122,6 +132,12 @@ def _get_provider_type(endpoint: Any, key: "ProviderAPIKey" | None = None) -> st
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Fallback: OAuth 导入的凭证可能包含 provider_type(如 Kiro)
|
||||
if decrypted_auth_config:
|
||||
pt = decrypted_auth_config.get("provider_type")
|
||||
if isinstance(pt, str) and pt.strip():
|
||||
return pt.strip().lower()
|
||||
|
||||
return None
|
||||
|
||||
|
||||
@@ -179,7 +195,7 @@ def build_provider_url(
|
||||
# endpoint_sig 为空时保持为空(更安全:默认路径回退到 "/",避免误判为 claude:chat)
|
||||
endpoint_sig = normalize_endpoint_signature(endpoint_sig) if endpoint_sig else ""
|
||||
|
||||
provider_type = _get_provider_type(endpoint, key)
|
||||
provider_type = _get_provider_type(endpoint, key, decrypted_auth_config)
|
||||
|
||||
# 合并查询参数(部分逻辑需要先拿到 query_params)
|
||||
effective_query_params = dict(query_params) if query_params else {}
|
||||
|
||||
@@ -268,6 +268,30 @@ def resolve_ops_proxy(
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Key 级别代理优先解析
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def resolve_effective_proxy(
|
||||
provider_proxy: dict[str, Any] | None,
|
||||
key_proxy: dict[str, Any] | None = None,
|
||||
) -> dict[str, Any] | None:
|
||||
"""
|
||||
解析有效的代理配置,Key 级别代理优先于 Provider 级别代理。
|
||||
|
||||
Args:
|
||||
provider_proxy: Provider 级别代理配置
|
||||
key_proxy: Key 级别代理配置(可选),非 None 且 enabled 时覆盖 Provider 级别
|
||||
|
||||
Returns:
|
||||
有效的代理配置字典,或 None(无代理)
|
||||
"""
|
||||
if key_proxy and key_proxy.get("enabled", True):
|
||||
return key_proxy
|
||||
return provider_proxy
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 代理 URL 构建
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
Reference in New Issue
Block a user