refactor(provider-keys): 拆分 keys 端点逻辑并补充配额刷新测试

将 admin keys 接口的创建、更新、查询、删除、导出与配额刷新逻辑迁移到 provider_keys 服务层

新增 auth_type 归一化、重复校验、响应构建与 quota_refresh(codex/kiro/antigravity)模块,并补充对应单元测试
This commit is contained in:
AAEE86
2026-02-28 14:09:36 +08:00
parent 4b02078b60
commit 08b89b7ef8
18 changed files with 3465 additions and 1368 deletions

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,114 @@
"""
Provider Keys 领域服务模块。
"""
from __future__ import annotations
from typing import TYPE_CHECKING, Any
if TYPE_CHECKING:
from sqlalchemy.orm import Session
from src.models.endpoint_models import (
EndpointAPIKeyCreate,
EndpointAPIKeyResponse,
EndpointAPIKeyUpdate,
)
__all__ = [
"clear_oauth_invalid_response",
"create_provider_key_response",
"delete_endpoint_key_response",
"update_endpoint_key_response",
"reveal_endpoint_key_payload",
"export_oauth_key_data",
"get_keys_grouped_by_format",
"list_provider_keys_responses",
"refresh_provider_quota_for_provider",
]
def clear_oauth_invalid_response(db: Session, key_id: str) -> dict[str, str]:
"""清除 OAuth 失效标记并返回统一响应(惰性导入实现)。"""
from src.services.provider_keys.key_command_service import clear_oauth_invalid_response as _impl
return _impl(db=db, key_id=key_id)
async def create_provider_key_response(
db: Session,
provider_id: str,
key_data: EndpointAPIKeyCreate,
) -> EndpointAPIKeyResponse:
"""创建 Provider Key 并返回响应对象(惰性导入实现)。"""
from src.services.provider_keys.key_command_service import create_provider_key_response as _impl
return await _impl(db=db, provider_id=provider_id, key_data=key_data)
async def delete_endpoint_key_response(db: Session, key_id: str) -> dict[str, str]:
"""删除 Key 并返回统一响应(惰性导入实现)。"""
from src.services.provider_keys.key_command_service import delete_endpoint_key_response as _impl
return await _impl(db=db, key_id=key_id)
async def update_endpoint_key_response(
db: Session,
key_id: str,
key_data: EndpointAPIKeyUpdate,
) -> EndpointAPIKeyResponse:
"""更新 Key 并返回响应对象(惰性导入实现)。"""
from src.services.provider_keys.key_command_service import update_endpoint_key_response as _impl
return await _impl(db=db, key_id=key_id, key_data=key_data)
def reveal_endpoint_key_payload(db: Session, key_id: str) -> dict[str, Any]:
"""获取完整的 API Key 或 Auth Config惰性导入实现"""
from src.services.provider_keys.key_query_service import reveal_endpoint_key_payload as _impl
return _impl(db=db, key_id=key_id)
def export_oauth_key_data(db: Session, key_id: str) -> dict[str, Any]:
"""导出 OAuth Key 凭据(惰性导入实现)。"""
from src.services.provider_keys.key_query_service import export_oauth_key_data as _impl
return _impl(db=db, key_id=key_id)
def get_keys_grouped_by_format(db: Session) -> dict:
"""按 API 格式分组查询所有 Key惰性导入实现"""
from src.services.provider_keys.key_query_service import get_keys_grouped_by_format as _impl
return _impl(db=db)
def list_provider_keys_responses(
db: Session,
provider_id: str,
skip: int,
limit: int,
) -> list[EndpointAPIKeyResponse]:
"""查询 Provider 下的 Key 列表(惰性导入实现)。"""
from src.services.provider_keys.key_query_service import list_provider_keys_responses as _impl
return _impl(db=db, provider_id=provider_id, skip=skip, limit=limit)
async def refresh_provider_quota_for_provider(
db: Session,
provider_id: str,
codex_wham_usage_url: str,
) -> dict:
"""刷新 Provider 限额信息(惰性导入实现)。"""
from src.services.provider_keys.key_quota_service import (
refresh_provider_quota_for_provider as _impl,
)
return await _impl(
db=db,
provider_id=provider_id,
codex_wham_usage_url=codex_wham_usage_url,
)

View File

@@ -0,0 +1,12 @@
"""
Provider Key 认证类型相关规则。
"""
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

View File

@@ -0,0 +1,210 @@
"""
Codex 配额响应解析器。
"""
from __future__ import annotations
import time
from typing import Any
class CodexUsageParseError(ValueError):
"""Codex 配额响应结构异常。"""
def _raise_type_error(field: str, expected: str, value: Any) -> None:
raise CodexUsageParseError(
f"{field} 类型错误,期望 {expected},实际 {type(value).__name__}: {value!r}"
)
def _as_dict(value: Any, field: str) -> dict[str, Any]:
if value is None:
return {}
if not isinstance(value, dict):
_raise_type_error(field, "object", value)
return value
def _coerce_float(value: Any, field: str) -> float:
if isinstance(value, bool):
_raise_type_error(field, "number", value)
if isinstance(value, (int, float)):
return float(value)
if isinstance(value, str):
raw = value.strip()
if not raw:
_raise_type_error(field, "number", value)
try:
return float(raw)
except ValueError as exc: # pragma: no cover - 仅防御
raise CodexUsageParseError(f"{field} 不是合法数字: {value!r}") from exc
_raise_type_error(field, "number", value)
def _coerce_int(value: Any, field: str) -> int:
if isinstance(value, bool):
_raise_type_error(field, "integer", value)
if isinstance(value, int):
return value
if isinstance(value, float):
if not value.is_integer():
raise CodexUsageParseError(f"{field} 必须是整数,实际为小数: {value!r}")
return int(value)
if isinstance(value, str):
raw = value.strip()
if not raw:
_raise_type_error(field, "integer", value)
try:
if "." in raw or "e" in raw.lower():
parsed = float(raw)
if not parsed.is_integer():
raise CodexUsageParseError(f"{field} 必须是整数,实际为小数: {value!r}")
return int(parsed)
return int(raw)
except ValueError as exc: # pragma: no cover - 仅防御
raise CodexUsageParseError(f"{field} 不是合法整数: {value!r}") from exc
_raise_type_error(field, "integer", value)
def _coerce_bool(value: Any, field: str) -> bool:
if isinstance(value, bool):
return value
if isinstance(value, int):
if value in (0, 1):
return bool(value)
raise CodexUsageParseError(f"{field} 仅支持 0/1 整数,实际为: {value!r}")
if isinstance(value, str):
normalized = value.strip().lower()
if normalized in {"true", "1", "yes"}:
return True
if normalized in {"false", "0", "no"}:
return False
raise CodexUsageParseError(f"{field} 不是合法布尔值: {value!r}")
_raise_type_error(field, "boolean", value)
def _write_window(
result: dict[str, Any],
*,
source: dict[str, Any],
source_field: str,
target_prefix: str,
) -> None:
if not source:
return
used_percent = source.get("used_percent")
if used_percent is not None:
result[f"{target_prefix}_used_percent"] = _coerce_float(
used_percent, f"{source_field}.used_percent"
)
reset_seconds = source.get("reset_after_seconds")
if reset_seconds is not None:
result[f"{target_prefix}_reset_seconds"] = _coerce_int(
reset_seconds, f"{source_field}.reset_after_seconds"
)
reset_at = source.get("reset_at")
if reset_at is not None:
result[f"{target_prefix}_reset_at"] = _coerce_int(reset_at, f"{source_field}.reset_at")
limit_window_seconds = source.get("limit_window_seconds")
if limit_window_seconds is not None:
result[f"{target_prefix}_window_minutes"] = (
_coerce_int(limit_window_seconds, f"{source_field}.limit_window_seconds") // 60
)
def parse_codex_wham_usage_response(data: dict[str, Any]) -> dict[str, Any] | 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 data is None:
return None
if not isinstance(data, dict):
_raise_type_error("root", "object", data)
if not data:
return None
result: dict[str, Any] = {}
plan_type: str | None = None
raw_plan_type = data.get("plan_type")
if raw_plan_type is not None:
if not isinstance(raw_plan_type, str):
_raise_type_error("plan_type", "string", raw_plan_type)
normalized_plan_type = raw_plan_type.strip().lower()
if normalized_plan_type:
plan_type = normalized_plan_type
result["plan_type"] = normalized_plan_type
# 解析 rate_limit
rate_limit = _as_dict(data.get("rate_limit"), "rate_limit")
primary_window = _as_dict(rate_limit.get("primary_window"), "rate_limit.primary_window")
secondary_window = _as_dict(rate_limit.get("secondary_window"), "rate_limit.secondary_window")
# 根据账号类型解析限额
# Free 账号: primary_window 是周限额,无 secondary_window
# Team/Plus/Enterprise: primary_window 是 5H 限额secondary_window 是周限额
use_paid_windows = bool(secondary_window) and plan_type != "free"
if use_paid_windows:
# 周限额 (secondary_window)
_write_window(
result,
source=secondary_window,
source_field="rate_limit.secondary_window",
target_prefix="primary",
)
# 5H 限额 (primary_window)
_write_window(
result,
source=primary_window,
source_field="rate_limit.primary_window",
target_prefix="secondary",
)
else:
# Free / 或 secondary_window 缺失时primary_window 视为主窗口
_write_window(
result,
source=primary_window,
source_field="rate_limit.primary_window",
target_prefix="primary",
)
# 解析 code_review_rate_limit (代码审查限额)
code_review_limit = _as_dict(data.get("code_review_rate_limit"), "code_review_rate_limit")
code_review_primary = _as_dict(
code_review_limit.get("primary_window"), "code_review_rate_limit.primary_window"
)
_write_window(
result,
source=code_review_primary,
source_field="code_review_rate_limit.primary_window",
target_prefix="code_review",
)
# 解析 credits
credits = _as_dict(data.get("credits"), "credits")
has_credits = credits.get("has_credits")
if has_credits is not None:
result["has_credits"] = _coerce_bool(has_credits, "credits.has_credits")
balance = credits.get("balance")
if balance is not None:
result["credits_balance"] = _coerce_float(balance, "credits.balance")
# 添加更新时间戳
if result:
result["updated_at"] = int(time.time())
return result if result else None

View File

@@ -0,0 +1,100 @@
"""
Provider Key 重复校验规则。
"""
from __future__ import annotations
import json
from sqlalchemy.orm import Session
from src.core.crypto import crypto_service
from src.core.exceptions import InvalidRequestException
from src.models.database import ProviderAPIKey
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

View File

@@ -0,0 +1,441 @@
"""
Provider Key 写操作命令服务。
"""
from __future__ import annotations
import asyncio
import json
import uuid
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any
from sqlalchemy.orm import Session
from src.core.crypto import crypto_service
from src.core.exceptions import InvalidRequestException, NotFoundException
from src.core.logger import logger
from src.models.database import Provider, ProviderAPIKey
from src.models.endpoint_models import (
EndpointAPIKeyCreate,
EndpointAPIKeyResponse,
EndpointAPIKeyUpdate,
)
from src.services.provider_keys.auth_type import normalize_auth_type
from src.services.provider_keys.duplicate_check import check_duplicate_key
from src.services.provider_keys.key_side_effects import (
run_create_key_side_effects,
run_delete_key_side_effects,
run_update_key_side_effects,
)
from src.services.provider_keys.response_builder import build_key_response
@dataclass
class _UpdateKeyPreparation:
"""更新 Key 前置准备结果。"""
update_data: dict[str, Any]
auto_fetch_enabled_before: bool
auto_fetch_enabled_after: bool
allowed_models_before: set[str]
include_patterns_before: list[str] | None
exclude_patterns_before: list[str] | None
@dataclass
class _DeleteKeyResult:
"""删除 Key 的执行结果。"""
provider_id: str | None
deleted_key_allowed_models: list[str] | None
def _run_async_with_fallback(coro: Any) -> None:
"""在同步上下文中执行异步任务(有事件循环则调度,无则阻塞执行)。"""
try:
loop = asyncio.get_running_loop()
except RuntimeError:
asyncio.run(coro)
return
task = loop.create_task(coro)
def _log_task_error(done_task: asyncio.Task[Any]) -> None:
try:
done_task.result()
except Exception as exc:
logger.warning("异步缓存失效任务执行失败: {}", exc)
task.add_done_callback(_log_task_error)
async def _invalidate_cache_after_clear_oauth_invalid(key_id: str) -> None:
"""清除 OAuth 失效标记后同步失效相关缓存。"""
from src.api.base.models_service import invalidate_models_list_cache
from src.services.cache.provider_cache import ProviderCacheService
await ProviderCacheService.invalidate_provider_api_key_cache(key_id)
await invalidate_models_list_cache()
def _clear_oauth_invalid_marker(db: Session, key_id: str) -> dict[str, str]:
"""清除 Key 的 OAuth 失效标记。"""
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
if not key:
raise NotFoundException(f"Key {key_id} 不存在")
if not key.oauth_invalid_at:
return {"message": "该 Key 当前无失效标记,无需清除"}
old_reason = key.oauth_invalid_reason
key.oauth_invalid_at = None
key.oauth_invalid_reason = None
key.is_active = True
db.commit()
_run_async_with_fallback(_invalidate_cache_after_clear_oauth_invalid(key_id))
logger.info(
"[OK] 手动清除 Key {}... 的 OAuth 失效标记并自动启用 (原因: {})", key_id[:8], old_reason
)
return {"message": "已清除 OAuth 失效标记Key 已自动启用"}
def clear_oauth_invalid_response(db: Session, key_id: str) -> dict[str, str]:
"""清除 OAuth 失效标记并返回统一响应。"""
return _clear_oauth_invalid_marker(db=db, key_id=key_id)
def _prepare_update_key_payload(
db: Session,
key: ProviderAPIKey,
key_id: str,
key_data: EndpointAPIKeyUpdate,
) -> _UpdateKeyPreparation:
"""准备更新 Key 的数据,并执行规则校验。"""
# 检查是否开启了 auto_fetch_models用于后续立即获取模型
auto_fetch_enabled_before = key.auto_fetch_models
auto_fetch_enabled_after = (
key_data.auto_fetch_models
if "auto_fetch_models" in key_data.model_fields_set
else auto_fetch_enabled_before
)
# 记录 allowed_models 变化前的值
allowed_models_before = set(key.allowed_models or [])
# 记录过滤规则变化前的值(用于检测是否需要重新应用过滤)
include_patterns_before = key.model_include_patterns
exclude_patterns_before = key.model_exclude_patterns
update_data = key_data.model_dump(exclude_unset=True)
# 显式传 null 等价于“不更新 auth_type”避免写入 NULL 触发数据库约束错误。
if update_data.get("auth_type") is None:
update_data.pop("auth_type", None)
if "api_key" in update_data and isinstance(update_data["api_key"], str):
update_data["api_key"] = update_data["api_key"].strip()
# 验证 auth_type
current_auth_type = normalize_auth_type(getattr(key, "auth_type", "api_key"))
target_auth_type = normalize_auth_type(update_data.get("auth_type", current_auth_type))
is_auth_type_switch = "auth_type" in update_data and target_auth_type != current_auth_type
api_key_in_payload = "api_key" in update_data
api_key_value = update_data.get("api_key")
if api_key_in_payload and api_key_value == "":
raise InvalidRequestException("api_key 不能为空")
# auth_type 校验 + 字段归一化
if target_auth_type == "api_key":
if is_auth_type_switch and (not api_key_value or api_key_value == "__placeholder__"):
raise InvalidRequestException("切换到 API Key 认证模式时,必须提供新的 API Key")
if api_key_in_payload and (api_key_value is None or api_key_value == "__placeholder__"):
raise InvalidRequestException("API Key 认证模式下 api_key 不能为空")
# 切换回 API Key清理非本模式配置
update_data["auth_config"] = None
elif target_auth_type == "vertex_ai":
if is_auth_type_switch and not update_data.get("auth_config"):
raise InvalidRequestException(
"从 API Key 切换到 Vertex AI 认证模式时,必须提供 Service Account JSON"
)
# Vertex AI 不允许手工写入 api_key仅保留占位符
if api_key_in_payload and api_key_value not in {None, "__placeholder__"}:
raise InvalidRequestException("Vertex AI 认证模式下不允许直接填写 api_key")
if is_auth_type_switch or api_key_in_payload:
update_data["api_key"] = "__placeholder__"
elif target_auth_type == "oauth":
# OAuth 的 token 不允许在 key 更新接口里手工写入
if api_key_in_payload and api_key_value not in {None, "__placeholder__"}:
raise InvalidRequestException("OAuth 认证模式下不允许直接填写 api_key")
if is_auth_type_switch:
update_data["api_key"] = "__placeholder__"
# 从非 OAuth 切换到 OAuth 时,清理旧认证配置(如 Vertex SA 凭证)。
update_data["auth_config"] = None
elif api_key_in_payload:
# 避免把 null 写入 DB 或意外覆盖现有 OAuth token。
update_data.pop("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=key_id,
)
if "api_key" in update_data:
api_key_raw = update_data["api_key"]
if api_key_raw is None:
# 防御式处理:避免将 NULL 写入 NOT NULL 字段导致 500。
update_data.pop("api_key", None)
else:
update_data["api_key"] = crypto_service.encrypt(api_key_raw)
# 加密 auth_config包含敏感凭证即便是 {} 也必须加密存储。
if "auth_config" in update_data:
auth_config_raw = update_data["auth_config"]
if auth_config_raw is None:
pass
elif isinstance(auth_config_raw, dict):
update_data["auth_config"] = crypto_service.encrypt(json.dumps(auth_config_raw))
else:
raise InvalidRequestException("auth_config 必须是 JSON 对象")
# 特殊处理 rpm_limit需要区分"未提供"和"显式设置为 null"
if "rpm_limit" in key_data.model_fields_set:
update_data["rpm_limit"] = key_data.rpm_limit
if key_data.rpm_limit is None:
update_data["learned_rpm_limit"] = None
logger.info("Key {} 切换为自适应 RPM 模式", key_id)
# 统一处理 allowed_models空列表 -> None表示不限制
if "allowed_models" in update_data:
am = update_data["allowed_models"]
if isinstance(am, list) and len(am) == 0:
update_data["allowed_models"] = None
# 统一处理 locked_models空列表 -> None
if "locked_models" in update_data:
lm = update_data["locked_models"]
if isinstance(lm, list) and len(lm) == 0:
update_data["locked_models"] = None
# 处理模型过滤规则:空字符串 -> None
if "model_include_patterns" in update_data:
patterns = update_data["model_include_patterns"]
if isinstance(patterns, list) and len(patterns) == 0:
update_data["model_include_patterns"] = None
if "model_exclude_patterns" in update_data:
patterns = update_data["model_exclude_patterns"]
if isinstance(patterns, list) and len(patterns) == 0:
update_data["model_exclude_patterns"] = None
# 处理 proxy将 ProxyConfig 转换为 dict 存储null 清除代理
if "proxy" in key_data.model_fields_set:
if key_data.proxy is None:
update_data["proxy"] = None
else:
update_data["proxy"] = key_data.proxy.model_dump(exclude_none=True)
return _UpdateKeyPreparation(
update_data=update_data,
auto_fetch_enabled_before=auto_fetch_enabled_before,
auto_fetch_enabled_after=auto_fetch_enabled_after,
allowed_models_before=allowed_models_before,
include_patterns_before=include_patterns_before,
exclude_patterns_before=exclude_patterns_before,
)
def _prepare_create_key_payload(
db: Session,
provider_id: str,
key_data: EndpointAPIKeyCreate,
) -> tuple[str, ProviderAPIKey]:
"""准备创建 Key 的认证类型校验、重复校验与实体构造。"""
auth_type = key_data.auth_type or "api_key"
if auth_type == "api_key":
if not key_data.api_key:
raise InvalidRequestException("API Key 认证模式下 api_key 为必填字段")
elif auth_type == "vertex_ai":
if not key_data.auth_config:
raise InvalidRequestException("Service Account 认证模式下 auth_config 为必填字段")
elif auth_type == "oauth":
# OAuth key 的 token 通过 provider-oauth 授权流程写入(此处不允许手填)
if key_data.api_key:
raise InvalidRequestException("OAuth 认证模式下不允许直接填写 api_key")
# 检查密钥是否已存在(防止重复添加)
check_duplicate_key(
db=db,
provider_id=provider_id,
auth_type=auth_type,
new_api_key=key_data.api_key,
new_auth_config=key_data.auth_config,
)
# 加密 API Key如果有
encrypted_key = (
crypto_service.encrypt(key_data.api_key)
if key_data.api_key
else crypto_service.encrypt("__placeholder__") # 占位符,保持 NOT NULL 约束
)
# OAuth 类型 key 初始写入占位符token 由 provider-oauth 流程写入)
if auth_type == "oauth":
encrypted_key = crypto_service.encrypt("__placeholder__")
now = datetime.now(timezone.utc)
# 加密 auth_config包含敏感的 Service Account 凭证)
encrypted_auth_config = None
if key_data.auth_config:
encrypted_auth_config = crypto_service.encrypt(json.dumps(key_data.auth_config))
new_key = ProviderAPIKey(
id=str(uuid.uuid4()),
provider_id=provider_id,
api_formats=key_data.api_formats,
auth_type=auth_type,
api_key=encrypted_key,
auth_config=encrypted_auth_config,
name=key_data.name,
note=key_data.note,
rate_multipliers=key_data.rate_multipliers,
internal_priority=key_data.internal_priority,
rpm_limit=key_data.rpm_limit,
allowed_models=key_data.allowed_models if key_data.allowed_models else None,
capabilities=key_data.capabilities if key_data.capabilities else None,
cache_ttl_minutes=key_data.cache_ttl_minutes,
max_probe_interval_minutes=key_data.max_probe_interval_minutes,
auto_fetch_models=key_data.auto_fetch_models,
locked_models=key_data.locked_models if key_data.locked_models else None,
model_include_patterns=(
key_data.model_include_patterns if key_data.model_include_patterns else None
),
model_exclude_patterns=(
key_data.model_exclude_patterns if key_data.model_exclude_patterns else None
),
request_count=0,
success_count=0,
error_count=0,
total_response_time_ms=0,
health_by_format={}, # 按格式存储健康度
circuit_breaker_by_format={}, # 按格式存储熔断器状态
is_active=True,
last_used_at=None,
created_at=now,
updated_at=now,
)
return auth_type, new_key
async def update_endpoint_key_response(
db: Session,
key_id: str,
key_data: EndpointAPIKeyUpdate,
) -> EndpointAPIKeyResponse:
"""更新 Key 并返回响应对象。"""
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
if not key:
raise NotFoundException(f"Key {key_id} 不存在")
prepared = _prepare_update_key_payload(
db=db,
key=key,
key_id=key_id,
key_data=key_data,
)
for field, value in prepared.update_data.items():
setattr(key, field, value)
key.updated_at = datetime.now(timezone.utc)
db.commit()
db.refresh(key)
await run_update_key_side_effects(
db=db,
key=key,
key_id=key_id,
auto_fetch_enabled_before=prepared.auto_fetch_enabled_before,
auto_fetch_enabled_after=prepared.auto_fetch_enabled_after,
include_patterns_before=prepared.include_patterns_before,
exclude_patterns_before=prepared.exclude_patterns_before,
allowed_models_before=prepared.allowed_models_before,
)
logger.info("[OK] 更新 Key: ID={}, Updates={}", key_id, list(prepared.update_data.keys()))
return build_key_response(key)
async def create_provider_key_response(
db: Session,
provider_id: str,
key_data: EndpointAPIKeyCreate,
) -> EndpointAPIKeyResponse:
"""创建 Provider Key 并返回响应对象。"""
provider = db.query(Provider).filter(Provider.id == provider_id).first()
if not provider:
raise NotFoundException(f"Provider {provider_id} 不存在")
if not key_data.api_formats:
raise InvalidRequestException("api_formats 为必填字段")
_, new_key = _prepare_create_key_payload(
db=db,
provider_id=provider_id,
key_data=key_data,
)
db.add(new_key)
db.commit()
db.refresh(new_key)
key_tail = (key_data.api_key or "")[-4:]
logger.info(
"[OK] 添加 Key: Provider={}, Formats={}, Key=***{}, ID={}",
provider_id,
key_data.api_formats,
key_tail,
new_key.id,
)
await run_create_key_side_effects(db=db, provider_id=provider_id, key=new_key)
return build_key_response(new_key, api_key_plain=key_data.api_key)
def _delete_endpoint_key(db: Session, key_id: str) -> _DeleteKeyResult:
"""删除指定 Key 并返回删除上下文。"""
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
if not key:
raise NotFoundException(f"Key {key_id} 不存在")
provider_id = key.provider_id
deleted_key_allowed_models = key.allowed_models # 保存被删除 Key 的 allowed_models
try:
db.delete(key)
db.commit()
except Exception as exc:
db.rollback()
logger.error(f"删除 Key 失败: ID={key_id}, Error={exc}")
raise
return _DeleteKeyResult(
provider_id=provider_id,
deleted_key_allowed_models=deleted_key_allowed_models,
)
async def delete_endpoint_key_response(db: Session, key_id: str) -> dict[str, str]:
"""删除 Key执行副作用并返回统一响应。"""
delete_result = _delete_endpoint_key(db, key_id)
await run_delete_key_side_effects(
db=db,
provider_id=delete_result.provider_id,
deleted_key_allowed_models=delete_result.deleted_key_allowed_models,
)
logger.warning("[DELETE] 删除 Key: ID={}, Provider={}", key_id, delete_result.provider_id)
return {"message": f"Key {key_id} 已删除"}

View File

@@ -0,0 +1,259 @@
"""
Provider Key 查询服务。
"""
from __future__ import annotations
import json
from datetime import datetime, timezone
from typing import Any
from sqlalchemy.orm import Session
from src.core.crypto import crypto_service
from src.core.exceptions import InvalidRequestException, NotFoundException
from src.core.key_capabilities import get_capability
from src.core.logger import logger
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.models.endpoint_models import EndpointAPIKeyResponse
from src.services.provider_keys.auth_type import normalize_auth_type
from src.services.provider_keys.response_builder import build_key_response
def get_keys_grouped_by_format(db: Session) -> dict:
"""查询所有 Key并按 API 格式分组返回。"""
# Key 属于 Provider按 key.api_formats 分组展示
# 包含所有 Key含停用的 Key 和停用的 Provider前端可显示停用标签和快捷开关
keys = (
db.query(ProviderAPIKey, Provider)
.join(Provider, ProviderAPIKey.provider_id == Provider.id)
.order_by(
ProviderAPIKey.internal_priority.asc(),
)
.all()
)
provider_ids = {str(provider.id) for _key, provider in keys}
endpoints = (
db.query(
ProviderEndpoint.provider_id,
ProviderEndpoint.api_format,
ProviderEndpoint.base_url,
)
.filter(
ProviderEndpoint.provider_id.in_(provider_ids),
ProviderEndpoint.is_active.is_(True),
)
.all()
)
endpoint_base_url_map: dict[tuple[str, str], str] = {}
for provider_id, api_format, base_url in endpoints:
fmt = api_format.value if hasattr(api_format, "value") else str(api_format)
endpoint_base_url_map[(str(provider_id), fmt)] = base_url
grouped: dict[str, list[dict]] = {}
for key, provider in keys:
api_formats = key.api_formats or []
if not api_formats:
continue # 跳过没有 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)
masked_key = f"{decrypted_key[:8]}***{decrypted_key[-4:]}"
except Exception as e:
logger.error(f"解密 Key 失败: key_id={key.id}, error={e}")
masked_key = "***ERROR***"
# 计算健康度指标
success_rate = key.success_count / key.request_count if key.request_count > 0 else None
avg_response_time_ms = (
round(key.total_response_time_ms / key.success_count, 2)
if key.success_count > 0
else None
)
# 将 capabilities dict 转换为启用的能力简短名称列表
caps_list = []
if key.capabilities:
for cap_name, enabled in key.capabilities.items():
if enabled:
cap_def = get_capability(cap_name)
caps_list.append(cap_def.short_name if cap_def else cap_name)
# 构建 Key 信息(基础数据)
key_info = {
"id": key.id,
"name": key.name,
"auth_type": auth_type,
"api_key_masked": masked_key,
"internal_priority": key.internal_priority,
"global_priority_by_format": key.global_priority_by_format,
"rate_multipliers": key.rate_multipliers,
"is_active": key.is_active,
"provider_active": provider.is_active,
"provider_name": provider.name,
"api_formats": api_formats,
"capabilities": caps_list,
"success_rate": success_rate,
"avg_response_time_ms": avg_response_time_ms,
"request_count": key.request_count,
}
# 将 Key 添加到每个支持的格式分组中,并附加格式特定的数据
health_by_format = key.health_by_format or {}
circuit_by_format = key.circuit_breaker_by_format or {}
priority_by_format = key.global_priority_by_format or {}
provider_id = str(provider.id)
for api_format in api_formats:
if api_format not in grouped:
grouped[api_format] = []
# 为每个格式创建副本,设置当前格式
format_key_info = key_info.copy()
format_key_info["api_format"] = api_format
format_key_info["endpoint_base_url"] = endpoint_base_url_map.get(
(provider_id, api_format)
)
# 添加格式特定的优先级
format_key_info["format_priority"] = priority_by_format.get(api_format)
# 添加格式特定的健康度数据
format_health = health_by_format.get(api_format, {})
format_circuit = circuit_by_format.get(api_format, {})
format_key_info["health_score"] = float(format_health.get("health_score") or 1.0)
format_key_info["circuit_breaker_open"] = bool(format_circuit.get("open", False))
grouped[api_format].append(format_key_info)
# 直接返回分组对象,供前端使用
return grouped
def list_provider_keys_responses(
db: Session,
provider_id: str,
skip: int,
limit: int,
) -> list[EndpointAPIKeyResponse]:
"""查询 Provider 下的 Key 列表并构建响应。"""
provider = db.query(Provider).filter(Provider.id == provider_id).first()
if not provider:
raise NotFoundException(f"Provider {provider_id} 不存在")
keys = (
db.query(ProviderAPIKey)
.filter(ProviderAPIKey.provider_id == provider_id)
.order_by(ProviderAPIKey.internal_priority.asc(), ProviderAPIKey.created_at.asc())
.offset(skip)
.limit(limit)
.all()
)
return [build_key_response(key) for key in keys]
def reveal_endpoint_key_payload(
db: Session,
key_id: str,
) -> dict[str, Any]:
"""获取完整的 API Key 或 Auth Config用于查看和复制"""
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
if not key:
raise NotFoundException(f"Key {key_id} 不存在")
auth_type = normalize_auth_type(getattr(key, "auth_type", "api_key"))
# Vertex AI 类型返回 auth_config需要解密
if auth_type == "vertex_ai":
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)
logger.info(f"[REVEAL] 查看 Auth Config: ID={key_id}, Name={key.name}")
return {"auth_type": "vertex_ai", "auth_config": auth_config}
except Exception as e:
logger.error(f"解密 Auth Config 失败: ID={key_id}, Error={e}")
raise InvalidRequestException(
"无法解密认证配置,可能是加密密钥已更改。请重新添加该密钥。"
)
# 兼容auth_config 为空时尝试从 api_key 解密(仅对迁移前的旧数据有效)
try:
decrypted_key = crypto_service.decrypt(key.api_key)
# 检查是否是新格式的占位符(表示 auth_config 丢失)
if decrypted_key == "__placeholder__":
logger.error(f"Vertex AI Key 缺少 auth_config: ID={key_id}")
raise InvalidRequestException("认证配置丢失,请重新添加该密钥。")
logger.info(f"[REVEAL] 查看完整 Key (legacy vertex_ai): ID={key_id}, Name={key.name}")
return {"auth_type": "vertex_ai", "auth_config": decrypted_key}
except InvalidRequestException:
raise
except Exception as e:
logger.error(f"解密 Key 失败: ID={key_id}, Error={e}")
raise InvalidRequestException(
"无法解密认证配置,可能是加密密钥已更改。请重新添加该密钥。"
)
# OAuth 类型:返回 access_token导出走 /export 端点)
if auth_type == "oauth":
try:
decrypted_key = crypto_service.decrypt(key.api_key)
except Exception as e:
logger.error(f"解密 Key 失败: ID={key_id}, Error={e}")
raise InvalidRequestException(
"无法解密 API Key可能是加密密钥已更改。请重新添加该密钥。"
)
logger.info(f"[REVEAL] 查看 OAuth Key: ID={key_id}, Name={key.name}")
return {"auth_type": "oauth", "api_key": decrypted_key}
# API Key 类型返回 api_key
try:
decrypted_key = crypto_service.decrypt(key.api_key)
except Exception as e:
logger.error(f"解密 Key 失败: ID={key_id}, Error={e}")
raise InvalidRequestException("无法解密 API Key可能是加密密钥已更改。请重新添加该密钥。")
logger.info(f"[REVEAL] 查看完整 Key: ID={key_id}, Name={key.name}")
return {"auth_type": "api_key", "api_key": decrypted_key}
def export_oauth_key_data(
db: Session,
key_id: str,
) -> dict[str, Any]:
"""导出 OAuth Key 凭据。"""
from src.services.provider.export import build_export_data
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first()
if not key:
raise NotFoundException(f"Key {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 {}... 导出成功", key_id[:8])
return export_data

View File

@@ -0,0 +1,186 @@
"""Provider Key 配额刷新编排服务。"""
from __future__ import annotations
import asyncio
from collections.abc import Awaitable, Callable
from typing import Any
from sqlalchemy.orm import Session
from src.core.exceptions import InvalidRequestException, NotFoundException
from src.core.logger import logger
from src.core.provider_types import ProviderType, normalize_provider_type
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.services.model.upstream_fetcher import merge_upstream_metadata
from src.services.provider_keys.quota_refresh import (
refresh_antigravity_key_quota,
refresh_codex_key_quota,
refresh_kiro_key_quota,
)
QuotaRefreshHandler = Callable[..., Awaitable[dict]]
def _normalize_api_format(api_format: Any) -> str:
"""规范化 api_format兼容大小写和首尾空白。"""
if not isinstance(api_format, str):
return ""
return api_format.strip().lower()
def _select_refresh_endpoint(provider: Provider, provider_type: str) -> ProviderEndpoint | None:
"""为配额刷新选择端点。"""
if provider_type == ProviderType.CODEX:
for ep in provider.endpoints:
if _normalize_api_format(ep.api_format) == "openai:cli" and ep.is_active:
return ep
raise InvalidRequestException("找不到有效的 openai:cli 端点")
if 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:
if _normalize_api_format(ep.api_format) == sig and ep.is_active:
return ep
raise InvalidRequestException("找不到有效的 gemini:chat/gemini:cli 端点")
# Kiro 不需要端点检查,直接使用 auth_config
return None
def _resolve_quota_refresh_handler(provider_type: str) -> QuotaRefreshHandler:
"""按 provider 类型返回刷新策略。"""
if provider_type == ProviderType.CODEX:
return refresh_codex_key_quota
if provider_type == ProviderType.ANTIGRAVITY:
return refresh_antigravity_key_quota
if provider_type == ProviderType.KIRO:
return refresh_kiro_key_quota
raise InvalidRequestException("仅支持 Codex / Antigravity / Kiro 类型的 Provider 刷新限额")
async def refresh_provider_quota_for_provider(
db: Session,
provider_id: str,
codex_wham_usage_url: str,
) -> dict:
"""刷新指定 Provider 下所有活跃 Key 的限额信息。"""
provider = db.query(Provider).filter(Provider.id == provider_id).first()
if not provider:
raise NotFoundException(f"Provider {provider_id} 不存在")
provider_type = normalize_provider_type(getattr(provider, "provider_type", ""))
if provider_type not in {ProviderType.CODEX, ProviderType.ANTIGRAVITY, ProviderType.KIRO}:
raise InvalidRequestException("仅支持 Codex / Antigravity / Kiro 类型的 Provider 刷新限额")
keys = (
db.query(ProviderAPIKey)
.filter(
ProviderAPIKey.provider_id == provider_id,
ProviderAPIKey.is_active.is_(True),
)
.all()
)
if not keys:
return {
"success": 0,
"failed": 0,
"total": 0,
"results": [],
"message": "没有活跃的 Key",
}
endpoint = _select_refresh_endpoint(provider, provider_type)
handler = _resolve_quota_refresh_handler(provider_type)
results: list[dict[str, Any]] = []
success_count = 0
failed_count = 0
metadata_updates: dict[str, dict] = {} # key_id -> metadata
state_updates: dict[str, dict[str, Any]] = {} # key_id -> model field updates
async def refresh_single_key(key: ProviderAPIKey) -> dict:
try:
return await handler(
db=db,
provider=provider,
key=key,
endpoint=endpoint,
codex_wham_usage_url=codex_wham_usage_url,
metadata_updates=metadata_updates,
state_updates=state_updates,
)
except Exception as e:
error_msg = str(e) or type(e).__name__
logger.error("刷新 Key {} 限额失败: {}", key.id, error_msg)
return {
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": error_msg,
}
# 分批执行,每批最多 5 个并发
batch_size = 5
for i in range(0, len(keys), batch_size):
batch = keys[i : i + batch_size]
batch_tasks = [refresh_single_key(key) for key in batch]
batch_results = await asyncio.gather(*batch_tasks)
results.extend(batch_results)
# 统计本批次结果
for result in batch_results:
if result["status"] == "success":
success_count += 1
else:
failed_count += 1
# 统一更新数据库(避免在并发任务中操作 session
if metadata_updates or state_updates:
for key in keys:
key_dirty = False
if key.id in metadata_updates:
updates = metadata_updates[key.id]
if isinstance(updates, dict):
key.upstream_metadata = merge_upstream_metadata(key.upstream_metadata, updates)
key_dirty = True
if key.id in state_updates:
updates = state_updates[key.id]
if isinstance(updates, dict):
for field_name, field_value in updates.items():
setattr(key, field_name, field_value)
key_dirty = True
if key_dirty:
db.add(key)
db.commit()
failed_details = [
f"{r.get('key_name', r.get('key_id', '?'))}: {r.get('message', 'unknown')}"
for r in results
if r["status"] != "success"
]
if failed_details:
logger.info(
"[QUOTA_REFRESH] Provider {}: 成功 {}/{}, 失败 {} [{}]",
provider_id,
success_count,
len(keys),
failed_count,
"; ".join(failed_details),
)
else:
logger.info(
"[QUOTA_REFRESH] Provider {}: 成功 {}/{}",
provider_id,
success_count,
len(keys),
)
return {
"success": success_count,
"failed": failed_count,
"total": len(keys),
"results": results,
}

View File

@@ -0,0 +1,150 @@
"""
Provider Key 写操作后的副作用处理。
"""
from __future__ import annotations
from sqlalchemy.orm import Session
from src.api.base.models_service import invalidate_models_list_cache
from src.core.logger import logger
from src.models.database import ProviderAPIKey
from src.services.cache.provider_cache import ProviderCacheService
async def run_update_key_side_effects(
db: Session,
key: ProviderAPIKey,
key_id: str,
auto_fetch_enabled_before: bool,
auto_fetch_enabled_after: bool,
include_patterns_before: list[str] | None,
exclude_patterns_before: list[str] | None,
allowed_models_before: set[str],
) -> None:
"""执行更新 Key 后的副作用。"""
if not auto_fetch_enabled_before and auto_fetch_enabled_after:
# 刚刚开启了 auto_fetch_models同步执行模型获取
logger.info("[AUTO_FETCH] Key {} 开启自动获取模型,同步执行模型获取", key_id)
try:
from src.services.model.fetch_scheduler import get_model_fetch_scheduler
scheduler = get_model_fetch_scheduler()
# 同步等待模型获取完成,确保前端刷新时能看到最新数据
await scheduler._fetch_models_for_key_by_id(key_id)
# fetch_scheduler 可能在独立 session 更新 allowed_models需刷新当前对象避免后续比较使用旧值。
db.refresh(key)
except Exception as e:
logger.error(f"触发模型获取失败: {e}")
# 不抛出异常,避免影响 Key 更新操作
elif auto_fetch_enabled_before and not auto_fetch_enabled_after:
# 关闭了 auto_fetch_models只保留锁定的模型清除自动获取的模型
locked = key.locked_models or []
if locked:
key.allowed_models = locked
logger.info(
"[AUTO_FETCH] Key {} 关闭自动获取模型,保留 {} 个锁定模型",
key_id,
len(locked),
)
else:
key.allowed_models = None
logger.info(
"[AUTO_FETCH] Key {} 关闭自动获取模型,无锁定模型,清空 allowed_models",
key_id,
)
db.commit()
db.refresh(key)
elif auto_fetch_enabled_after:
# auto_fetch_models 保持开启状态,检查过滤规则是否变更
include_patterns_after = key.model_include_patterns
exclude_patterns_after = key.model_exclude_patterns
patterns_changed = (
include_patterns_before != include_patterns_after
or exclude_patterns_before != exclude_patterns_after
)
if patterns_changed:
# 过滤规则变更,重新应用过滤(使用缓存的上游模型数据)
logger.info("[AUTO_FETCH] Key {} 过滤规则变更,重新应用过滤", key_id)
try:
from src.services.model.fetch_scheduler import get_model_fetch_scheduler
scheduler = get_model_fetch_scheduler()
await scheduler._fetch_models_for_key_by_id(key_id)
# 重新应用过滤后,刷新当前对象以读取最新 allowed_models。
db.refresh(key)
except Exception as e:
logger.error(f"重新应用过滤规则失败: {e}")
# 任何字段更新都清除缓存,确保缓存一致性
# 包括 is_active、allowed_models、capabilities 等影响权限和行为的字段
await ProviderCacheService.invalidate_provider_api_key_cache(key_id)
# 检查 allowed_models 是否有变化,触发缓存失效和自动关联
allowed_models_after = set(key.allowed_models or [])
if allowed_models_before != allowed_models_after and key.provider_id:
from src.services.model.global_model import on_key_allowed_models_changed
await on_key_allowed_models_changed(
db=db,
provider_id=key.provider_id,
allowed_models=list(key.allowed_models or []),
)
else:
# allowed_models 未变化时,仍需清除 /v1/models 缓存is_active、api_formats 变更会影响模型可用性)
await invalidate_models_list_cache()
async def run_create_key_side_effects(
db: Session,
provider_id: str,
key: ProviderAPIKey,
) -> None:
"""执行创建 Key 后的副作用。"""
# 如果开启了 auto_fetch_models同步执行模型获取
if key.auto_fetch_models:
logger.info("[AUTO_FETCH] 新 Key {} 开启自动获取模型,同步执行模型获取", key.id)
try:
from src.services.model.fetch_scheduler import get_model_fetch_scheduler
scheduler = get_model_fetch_scheduler()
# 同步等待模型获取完成,确保前端刷新时能看到最新数据
await scheduler._fetch_models_for_key_by_id(key.id)
except Exception as e:
logger.error(f"触发模型获取失败: {e}")
# 不抛出异常,避免影响 Key 创建操作
# 如果创建时指定了 allowed_models触发自动关联检查内部会清除 /v1/models 缓存)
if key.allowed_models:
from src.services.model.global_model import on_key_allowed_models_changed
await on_key_allowed_models_changed(
db=db,
provider_id=provider_id,
allowed_models=list(key.allowed_models),
)
else:
# 没有 allowed_models 时,仍需清除 /v1/models 缓存
await invalidate_models_list_cache()
async def run_delete_key_side_effects(
db: Session,
provider_id: str | None,
deleted_key_allowed_models: list[str] | None,
) -> None:
"""执行删除 Key 后的副作用。"""
# 触发缓存失效和自动解除关联检查
# 注意:删除后是否需要解除关联,应基于“删除后的活跃 Key 集合”判断。
# 不能仅凭被删除 Key 的 allowed_models 是否为 null 来跳过 disassociate。
_ = deleted_key_allowed_models
if provider_id:
from src.services.model.global_model import on_key_allowed_models_changed
await on_key_allowed_models_changed(
db=db,
provider_id=provider_id,
)
else:
# 无 provider_id 时仅清除缓存
await invalidate_models_list_cache()

View File

@@ -0,0 +1,15 @@
"""
Provider Key 配额刷新策略模块。
"""
from src.services.provider_keys.quota_refresh.antigravity_refresher import (
refresh_antigravity_key_quota,
)
from src.services.provider_keys.quota_refresh.codex_refresher import refresh_codex_key_quota
from src.services.provider_keys.quota_refresh.kiro_refresher import refresh_kiro_key_quota
__all__ = [
"refresh_codex_key_quota",
"refresh_antigravity_key_quota",
"refresh_kiro_key_quota",
]

View File

@@ -0,0 +1,139 @@
"""
Antigravity 配额刷新策略。
"""
from __future__ import annotations
import time
from datetime import datetime, timezone
from sqlalchemy.orm import Session
from src.api.handlers.base.request_builder import get_provider_auth
from src.core.logger import logger
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
async def refresh_antigravity_key_quota(
*,
db: Session,
provider: Provider,
key: ProviderAPIKey,
endpoint: ProviderEndpoint | None,
codex_wham_usage_url: str,
metadata_updates: dict[str, dict],
state_updates: dict[str, dict],
) -> dict:
"""刷新单个 Antigravity Key 的配额信息。"""
_ = db
_ = codex_wham_usage_url
if endpoint is None:
return {
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "找不到有效的 gemini:chat/gemini:cli 端点",
}
# 直接调用 /v1internal:fetchAvailableModels 获取 quotaInfo无需发送真实对话请求
auth_info = await get_provider_auth(endpoint, key)
if not auth_info:
return {
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "缺少 OAuth 认证信息,请先授权/刷新 Token",
}
access_token = str(auth_info.auth_value).removeprefix("Bearer ").strip()
from src.services.model.upstream_fetcher import (
UpstreamModelsFetchContext,
fetch_models_for_key,
)
from src.services.provider.adapters.antigravity.client import (
AntigravityAccountForbiddenException,
)
from src.services.proxy_node.resolver import resolve_effective_proxy
effective_proxy = resolve_effective_proxy(
getattr(provider, "proxy", None),
getattr(key, "proxy", None),
)
fetch_ctx = UpstreamModelsFetchContext(
provider_type="antigravity",
api_key_value=access_token,
# antigravity fetcher 不依赖 endpoint mapping
format_to_endpoint={},
proxy_config=effective_proxy,
auth_config=auth_info.decrypted_auth_config,
)
try:
_models, errors, ok, upstream_meta = await fetch_models_for_key(
fetch_ctx, timeout_seconds=10.0
)
except AntigravityAccountForbiddenException as e:
# 对齐 AM所有 403 一律标记 is_forbidden 并停用
state_updates[key.id] = {
"is_active": False,
"oauth_invalid_at": datetime.now(timezone.utc),
"oauth_invalid_reason": f"账户访问被禁止: {e.reason or e.message}",
}
# 更新 upstream_metadata 标记封禁状态
metadata_updates[key.id] = {
"antigravity": {
"is_forbidden": True,
"forbidden_reason": e.reason or e.message,
"forbidden_at": int(time.time()),
"updated_at": int(time.time()),
}
}
logger.warning(
"[ANTIGRAVITY_QUOTA] Key {} 账户访问被禁止,已自动停用: {}",
key.id,
e.reason or e.message,
)
return {
"key_id": key.id,
"key_name": key.name,
"status": "forbidden",
"message": f"账户访问被禁止: {e.reason or e.message}",
"is_forbidden": True,
"auto_disabled": True,
}
if ok and upstream_meta:
# 刷新成功时清除之前的封禁标记(如果账户已恢复)
if "antigravity" in upstream_meta:
upstream_meta["antigravity"]["is_forbidden"] = False
upstream_meta["antigravity"]["forbidden_reason"] = None
upstream_meta["antigravity"]["forbidden_at"] = None
metadata_updates[key.id] = upstream_meta
state_updates[key.id] = {
"oauth_invalid_at": None,
"oauth_invalid_reason": None,
}
return {
"key_id": key.id,
"key_name": key.name,
"status": "success",
"metadata": upstream_meta,
}
if ok and not upstream_meta:
return {
"key_id": key.id,
"key_name": key.name,
"status": "no_metadata",
"message": "响应中未包含配额信息",
}
error_msg = "; ".join(errors) if errors else "fetchAvailableModels failed"
return {
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": error_msg,
}

View File

@@ -0,0 +1,151 @@
"""
Codex 配额刷新策略。
"""
from __future__ import annotations
import json
from typing import Any
import httpx
from sqlalchemy.orm import Session
from src.api.handlers.base.request_builder import get_provider_auth
from src.core.crypto import crypto_service
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.services.provider_keys.auth_type import normalize_auth_type
from src.services.provider_keys.codex_usage_parser import parse_codex_wham_usage_response
def _normalize_plan_type(value: Any) -> str | None:
if not isinstance(value, str):
return None
normalized = value.strip().lower()
return normalized or None
async def refresh_codex_key_quota(
*,
db: Session,
provider: Provider,
key: ProviderAPIKey,
endpoint: ProviderEndpoint | None,
codex_wham_usage_url: str,
metadata_updates: dict[str, dict],
state_updates: dict[str, dict],
) -> dict:
"""刷新单个 Codex Key 的限额信息。"""
_ = db
if endpoint is None:
return {
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": "找不到有效的 openai:cli 端点",
}
# 获取认证信息(用于刷新 OAuth token
auth_info = await get_provider_auth(endpoint, key)
# 构建请求头
headers: dict[str, Any] = {
"Accept": "application/json",
}
if auth_info:
headers[auth_info.auth_header] = auth_info.auth_value
else:
# 标准 API Key
decrypted_key = crypto_service.decrypt(key.api_key)
headers["Authorization"] = f"Bearer {decrypted_key}"
# 从 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)
if isinstance(auth_config_data, dict):
oauth_plan_type = _normalize_plan_type(auth_config_data.get("plan_type"))
raw_account_id = auth_config_data.get("account_id")
if isinstance(raw_account_id, str):
oauth_account_id = raw_account_id.strip() or None
except Exception:
pass
# 如果有 account_id 且不是 free 账号plan_type 缺失时默认携带,增强兼容性)
if oauth_account_id and oauth_plan_type != "free":
headers["chatgpt-account-id"] = oauth_account_id
# 解析代理配置key 级别 > provider 级别 > 系统默认)
from src.services.proxy_node.resolver import (
build_proxy_client_kwargs,
resolve_effective_proxy,
)
effective_proxy = resolve_effective_proxy(
getattr(provider, "proxy", None),
getattr(key, "proxy", None),
)
# 使用 wham/usage API 获取限额信息
async with httpx.AsyncClient(
**build_proxy_client_kwargs(effective_proxy, timeout=30.0)
) as client:
response = await client.get(codex_wham_usage_url, headers=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 响应",
}
# 解析限额信息
try:
metadata = parse_codex_wham_usage_response(data)
except Exception as exc:
return {
"key_id": key.id,
"key_name": key.name,
"status": "error",
"message": f"wham/usage 响应结构异常: {exc}",
"status_code": response.status_code,
}
if metadata:
# 收集元数据,稍后统一更新数据库(存储到 codex 子对象)
metadata_updates[key.id] = {"codex": metadata}
state_updates[key.id] = {
"oauth_invalid_at": None,
"oauth_invalid_reason": None,
}
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": "响应中未包含限额信息",
"status_code": response.status_code,
}

View File

@@ -0,0 +1,163 @@
"""
Kiro 配额刷新策略。
"""
from __future__ import annotations
import json
import time
from datetime import datetime, timezone
from sqlalchemy.orm import Session
from src.core.crypto import crypto_service
from src.core.logger import logger
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
async def refresh_kiro_key_quota(
*,
db: Session,
provider: Provider,
key: ProviderAPIKey,
endpoint: ProviderEndpoint | None,
codex_wham_usage_url: str,
metadata_updates: dict[str, dict],
state_updates: dict[str, dict],
) -> dict:
"""刷新单个 Kiro Key 的配额信息。"""
_ = db
_ = endpoint
_ = codex_wham_usage_url
from src.services.provider.adapters.kiro.usage import (
KiroAccountBannedException,
)
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可能是加密密钥已更改",
}
# 获取代理配置key 级别 > provider 级别)
from src.services.proxy_node.resolver import resolve_effective_proxy
proxy_config = resolve_effective_proxy(
getattr(provider, "proxy", None),
getattr(key, "proxy", None),
)
# 调用 Kiro getUsageLimits API
try:
result = await _fetch_kiro_usage_limits(
auth_config=auth_config_data,
proxy_config=proxy_config,
)
except KiroAccountBannedException as e:
# 账户被封禁,自动停用并标记
state_updates[key.id] = {
"is_active": False,
"oauth_invalid_at": datetime.now(timezone.utc),
"oauth_invalid_reason": f"账户已封禁: {e.reason or e.message}",
}
# 更新 upstream_metadata 标记封禁状态
metadata_updates[key.id] = {
"kiro": {
"is_banned": True,
"ban_reason": e.reason or e.message,
"banned_at": int(time.time()),
"updated_at": int(time.time()),
}
}
logger.warning(
"[KIRO_QUOTA] Key {} 账户已封禁,已自动停用: {}",
key.id,
e.reason or e.message,
)
return {
"key_id": key.id,
"key_name": key.name,
"status": "banned",
"message": f"账户已封禁: {e.reason or e.message}",
"is_banned": True,
"auto_disabled": True,
}
except RuntimeError as e:
error_msg = str(e)
# 检查是否需要标记账号异常
if "401" in error_msg or "认证失败" in error_msg:
state_updates[key.id] = {
"is_active": False,
"oauth_invalid_at": datetime.now(timezone.utc),
"oauth_invalid_reason": "Kiro Token 无效或已过期",
}
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:
# 刷新成功时清除之前的封禁标记(如果账户已恢复)
metadata["is_banned"] = False
metadata["ban_reason"] = None
metadata["banned_at"] = None
# 收集元数据,稍后统一更新数据库(存储到 kiro 子对象)
metadata_updates[key.id] = {"kiro": metadata}
state_updates[key.id] = {
"oauth_invalid_at": None,
"oauth_invalid_reason": None,
}
# 如果 auth_config 有更新(例如 token 刷新),也需要更新
if updated_auth_config:
try:
new_auth_config_json = json.dumps(updated_auth_config)
state_updates[key.id]["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": "响应中未包含限额信息",
}

View File

@@ -0,0 +1,126 @@
"""
Provider Key 响应对象构建器。
"""
from __future__ import annotations
import json
from src.core.crypto import crypto_service
from src.core.logger import logger
from src.models.database import ProviderAPIKey
from src.models.endpoint_models import EndpointAPIKeyResponse
from src.services.provider_keys.auth_type import normalize_auth_type
def build_key_response(
key: ProviderAPIKey, api_key_plain: str | None = None
) -> EndpointAPIKeyResponse:
"""构建 Key 响应对象。"""
auth_type = normalize_auth_type(getattr(key, "auth_type", "api_key"))
if auth_type == "vertex_ai":
# Vertex AI 使用 Service Account不显示占位符
masked_key = "[Service Account]"
elif auth_type == "oauth":
masked_key = "[OAuth Token]"
else:
try:
decrypted_key = crypto_service.decrypt(key.api_key)
masked_key = f"{decrypted_key[:8]}***{decrypted_key[-4:]}"
except Exception:
masked_key = "***ERROR***"
success_rate = key.success_count / key.request_count if key.request_count > 0 else 0.0
avg_response_time_ms = (
key.total_response_time_ms / key.success_count if key.success_count > 0 else 0.0
)
is_adaptive = key.rpm_limit is None
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
oauth_email = None
oauth_plan_type = None
oauth_account_id = None
encrypted_auth_config = key_dict.pop("auth_config", None) # 移除敏感字段,避免泄露
if auth_type == "oauth" and encrypted_auth_config:
try:
decrypted_config = crypto_service.decrypt(encrypted_auth_config)
auth_config = json.loads(decrypted_config)
oauth_expires_at = auth_config.get("expires_at")
oauth_email = auth_config.get("email")
oauth_plan_type = auth_config.get("plan_type") # Codex: plus/free/team/enterprise
# Antigravity 使用 "tier" 字段(如 "PAID"/"FREE"),做小写化 fallback
if not oauth_plan_type:
ag_tier = auth_config.get("tier")
if ag_tier and isinstance(ag_tier, str):
oauth_plan_type = ag_tier.lower()
oauth_account_id = auth_config.get("account_id") # Codex: chatgpt_account_id
except Exception as e:
logger.error("Failed to decrypt auth_config for key {}: {}", key.id, e)
# 从 health_by_format 计算汇总字段(便于列表展示)
health_by_format = key.health_by_format or {}
circuit_by_format = key.circuit_breaker_by_format or {}
# 计算整体健康度(取所有格式中的最低值)
if health_by_format:
health_scores = [float(h.get("health_score") or 1.0) for h in health_by_format.values()]
min_health_score = min(health_scores) if health_scores else 1.0
# 取最大的连续失败次数
max_consecutive = max(
(int(h.get("consecutive_failures") or 0) for h in health_by_format.values()),
default=0,
)
# 取最近的失败时间
failure_times = [
h.get("last_failure_at") for h in health_by_format.values() if h.get("last_failure_at")
]
last_failure = max(failure_times) if failure_times else None
else:
min_health_score = 1.0
max_consecutive = 0
last_failure = None
# 检查是否有任何格式的熔断器打开
any_circuit_open = any(c.get("open", False) for c in circuit_by_format.values())
key_dict.update(
{
"api_key_masked": masked_key,
"api_key_plain": api_key_plain,
"success_rate": success_rate,
"avg_response_time_ms": round(avg_response_time_ms, 2),
"is_adaptive": is_adaptive,
"effective_limit": (
key.learned_rpm_limit # 自适应模式:使用学习值,未学习时为 None不限制
if is_adaptive
else key.rpm_limit
),
# 汇总字段
"health_score": min_health_score,
"consecutive_failures": max_consecutive,
"last_failure_at": last_failure,
"circuit_breaker_open": any_circuit_open,
# OAuth 相关
"oauth_expires_at": oauth_expires_at,
"oauth_email": oauth_email,
"oauth_plan_type": oauth_plan_type,
"oauth_account_id": oauth_account_id,
"oauth_invalid_at": (
int(key.oauth_invalid_at.timestamp()) if key.oauth_invalid_at else None
),
"oauth_invalid_reason": key.oauth_invalid_reason,
}
)
# 防御性:确保 api_formats 存在(历史数据可能为空/缺失)
if "api_formats" not in key_dict or key_dict["api_formats"] is None:
key_dict["api_formats"] = []
return EndpointAPIKeyResponse(**key_dict)