mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
refactor(provider-keys): 拆分 keys 端点逻辑并补充配额刷新测试
将 admin keys 接口的创建、更新、查询、删除、导出与配额刷新逻辑迁移到 provider_keys 服务层 新增 auth_type 归一化、重复校验、响应构建与 quota_refresh(codex/kiro/antigravity)模块,并补充对应单元测试
This commit is contained in:
File diff suppressed because it is too large
Load Diff
114
src/services/provider_keys/__init__.py
Normal file
114
src/services/provider_keys/__init__.py
Normal 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,
|
||||
)
|
||||
12
src/services/provider_keys/auth_type.py
Normal file
12
src/services/provider_keys/auth_type.py
Normal 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
|
||||
210
src/services/provider_keys/codex_usage_parser.py
Normal file
210
src/services/provider_keys/codex_usage_parser.py
Normal 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
|
||||
100
src/services/provider_keys/duplicate_check.py
Normal file
100
src/services/provider_keys/duplicate_check.py
Normal 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
|
||||
441
src/services/provider_keys/key_command_service.py
Normal file
441
src/services/provider_keys/key_command_service.py
Normal 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} 已删除"}
|
||||
259
src/services/provider_keys/key_query_service.py
Normal file
259
src/services/provider_keys/key_query_service.py
Normal 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
|
||||
186
src/services/provider_keys/key_quota_service.py
Normal file
186
src/services/provider_keys/key_quota_service.py
Normal 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,
|
||||
}
|
||||
150
src/services/provider_keys/key_side_effects.py
Normal file
150
src/services/provider_keys/key_side_effects.py
Normal 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()
|
||||
15
src/services/provider_keys/quota_refresh/__init__.py
Normal file
15
src/services/provider_keys/quota_refresh/__init__.py
Normal 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",
|
||||
]
|
||||
@@ -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,
|
||||
}
|
||||
151
src/services/provider_keys/quota_refresh/codex_refresher.py
Normal file
151
src/services/provider_keys/quota_refresh/codex_refresher.py
Normal 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,
|
||||
}
|
||||
163
src/services/provider_keys/quota_refresh/kiro_refresher.py
Normal file
163
src/services/provider_keys/quota_refresh/kiro_refresher.py
Normal 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": "响应中未包含限额信息",
|
||||
}
|
||||
126
src/services/provider_keys/response_builder.py
Normal file
126
src/services/provider_keys/response_builder.py
Normal 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)
|
||||
207
tests/services/test_provider_keys_key_command_service.py
Normal file
207
tests/services/test_provider_keys_key_command_service.py
Normal file
@@ -0,0 +1,207 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
import types
|
||||
from datetime import datetime, timezone
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
|
||||
from src.core.exceptions import InvalidRequestException
|
||||
from src.models.endpoint_models import EndpointAPIKeyUpdate
|
||||
|
||||
|
||||
async def _noop_invalidate_models_list_cache() -> None:
|
||||
return None
|
||||
|
||||
|
||||
_fake_models_service_module = types.ModuleType("src.api.base.models_service")
|
||||
setattr(
|
||||
_fake_models_service_module, "invalidate_models_list_cache", _noop_invalidate_models_list_cache
|
||||
)
|
||||
sys.modules.setdefault("src.api.base.models_service", _fake_models_service_module)
|
||||
|
||||
from src.services.provider_keys import key_command_service as command_module
|
||||
from src.services.provider_keys import key_side_effects as side_effects_module
|
||||
|
||||
|
||||
class _NoQueryDB:
|
||||
def query(self, *args: Any, **kwargs: Any) -> Any: # pragma: no cover - 防御断言
|
||||
_ = args, kwargs
|
||||
raise AssertionError("unexpected query call")
|
||||
|
||||
|
||||
class _FakeQuery:
|
||||
def __init__(self, key: Any) -> None:
|
||||
self._key = key
|
||||
|
||||
def filter(self, *args: Any, **kwargs: Any) -> _FakeQuery:
|
||||
_ = args, kwargs
|
||||
return self
|
||||
|
||||
def first(self) -> Any:
|
||||
return self._key
|
||||
|
||||
|
||||
class _FakeClearOAuthDB:
|
||||
def __init__(self, key: Any) -> None:
|
||||
self._key = key
|
||||
self.commit_count = 0
|
||||
|
||||
def query(self, model: Any) -> _FakeQuery:
|
||||
_ = model
|
||||
return _FakeQuery(self._key)
|
||||
|
||||
def commit(self) -> None:
|
||||
self.commit_count += 1
|
||||
|
||||
|
||||
def _build_key(**overrides: Any) -> SimpleNamespace:
|
||||
base: dict[str, Any] = {
|
||||
"auto_fetch_models": False,
|
||||
"allowed_models": None,
|
||||
"model_include_patterns": None,
|
||||
"model_exclude_patterns": None,
|
||||
"provider_id": "provider-1",
|
||||
"auth_type": "api_key",
|
||||
}
|
||||
base.update(overrides)
|
||||
return SimpleNamespace(**base)
|
||||
|
||||
|
||||
def test_prepare_update_payload_auth_type_null_ignored() -> None:
|
||||
key = _build_key(auth_type="api_key")
|
||||
key_data = EndpointAPIKeyUpdate.model_validate({"auth_type": None})
|
||||
|
||||
prepared = command_module._prepare_update_key_payload(
|
||||
db=cast(Any, _NoQueryDB()),
|
||||
key=cast(Any, key),
|
||||
key_id="key-1",
|
||||
key_data=key_data,
|
||||
)
|
||||
|
||||
assert "auth_type" not in prepared.update_data
|
||||
|
||||
|
||||
def test_prepare_update_payload_rejects_empty_api_key() -> None:
|
||||
key = _build_key(auth_type="api_key")
|
||||
key_data = EndpointAPIKeyUpdate.model_validate({"api_key": " "})
|
||||
|
||||
with pytest.raises(InvalidRequestException, match="api_key 不能为空"):
|
||||
command_module._prepare_update_key_payload(
|
||||
db=cast(Any, _NoQueryDB()),
|
||||
key=cast(Any, key),
|
||||
key_id="key-1",
|
||||
key_data=key_data,
|
||||
)
|
||||
|
||||
|
||||
def test_prepare_update_payload_encrypts_empty_auth_config_dict(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
key = _build_key(auth_type="vertex_ai")
|
||||
key_data = EndpointAPIKeyUpdate.model_validate({"auth_config": {}})
|
||||
|
||||
monkeypatch.setattr(command_module.crypto_service, "encrypt", lambda raw: f"ENC:{raw}")
|
||||
|
||||
prepared = command_module._prepare_update_key_payload(
|
||||
db=cast(Any, _NoQueryDB()),
|
||||
key=cast(Any, key),
|
||||
key_id="key-1",
|
||||
key_data=key_data,
|
||||
)
|
||||
|
||||
assert prepared.update_data["auth_config"] == "ENC:{}"
|
||||
|
||||
|
||||
def test_prepare_update_payload_vertex_to_oauth_clears_auth_config(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
key = _build_key(auth_type="vertex_ai")
|
||||
key_data = EndpointAPIKeyUpdate.model_validate({"auth_type": "oauth"})
|
||||
|
||||
monkeypatch.setattr(command_module.crypto_service, "encrypt", lambda raw: f"ENC:{raw}")
|
||||
|
||||
prepared = command_module._prepare_update_key_payload(
|
||||
db=cast(Any, _NoQueryDB()),
|
||||
key=cast(Any, key),
|
||||
key_id="key-1",
|
||||
key_data=key_data,
|
||||
)
|
||||
|
||||
assert prepared.update_data["auth_config"] is None
|
||||
assert prepared.update_data["api_key"] == "ENC:__placeholder__"
|
||||
|
||||
|
||||
def test_clear_oauth_invalid_response_invalidates_caches(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
cache_calls: list[tuple[str, str | None]] = []
|
||||
|
||||
async def _fake_invalidate_key_cache(key_id: str) -> None:
|
||||
cache_calls.append(("key", key_id))
|
||||
|
||||
async def _fake_invalidate_models_cache() -> None:
|
||||
cache_calls.append(("models", None))
|
||||
|
||||
fake_provider_cache_module = types.ModuleType("src.services.cache.provider_cache")
|
||||
|
||||
class _FakeProviderCacheService:
|
||||
@staticmethod
|
||||
async def invalidate_provider_api_key_cache(key_id: str) -> None:
|
||||
await _fake_invalidate_key_cache(key_id)
|
||||
|
||||
setattr(fake_provider_cache_module, "ProviderCacheService", _FakeProviderCacheService)
|
||||
monkeypatch.setitem(
|
||||
sys.modules, "src.services.cache.provider_cache", fake_provider_cache_module
|
||||
)
|
||||
|
||||
fake_models_service_module = types.ModuleType("src.api.base.models_service")
|
||||
setattr(
|
||||
fake_models_service_module, "invalidate_models_list_cache", _fake_invalidate_models_cache
|
||||
)
|
||||
monkeypatch.setitem(sys.modules, "src.api.base.models_service", fake_models_service_module)
|
||||
|
||||
key = SimpleNamespace(
|
||||
oauth_invalid_at=datetime.now(timezone.utc),
|
||||
oauth_invalid_reason="forbidden",
|
||||
is_active=False,
|
||||
)
|
||||
db = _FakeClearOAuthDB(key=key)
|
||||
|
||||
result = command_module.clear_oauth_invalid_response(cast(Any, db), key_id="key-1")
|
||||
|
||||
assert result["message"] == "已清除 OAuth 失效标记,Key 已自动启用"
|
||||
assert key.oauth_invalid_at is None
|
||||
assert key.oauth_invalid_reason is None
|
||||
assert key.is_active is True
|
||||
assert db.commit_count == 1
|
||||
assert cache_calls == [("key", "key-1"), ("models", None)]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_delete_key_side_effects_not_skip_disassociate(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
captured: dict[str, Any] = {}
|
||||
|
||||
async def _fake_on_key_allowed_models_changed(**kwargs: Any) -> None:
|
||||
captured.update(kwargs)
|
||||
|
||||
from src.services.model import global_model as global_model_module
|
||||
|
||||
monkeypatch.setattr(
|
||||
global_model_module,
|
||||
"on_key_allowed_models_changed",
|
||||
_fake_on_key_allowed_models_changed,
|
||||
)
|
||||
|
||||
await side_effects_module.run_delete_key_side_effects(
|
||||
db=cast(Any, object()),
|
||||
provider_id="provider-1",
|
||||
deleted_key_allowed_models=None,
|
||||
)
|
||||
|
||||
assert captured["provider_id"] == "provider-1"
|
||||
assert "skip_disassociate" not in captured
|
||||
155
tests/services/test_provider_keys_query_service.py
Normal file
155
tests/services/test_provider_keys_query_service.py
Normal file
@@ -0,0 +1,155 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
|
||||
from src.core.exceptions import NotFoundException
|
||||
from src.services.provider_keys import key_query_service as query_service_module
|
||||
from src.services.provider_keys.key_query_service import (
|
||||
get_keys_grouped_by_format,
|
||||
list_provider_keys_responses,
|
||||
)
|
||||
|
||||
|
||||
class _FakeQuery:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
first_result: Any = None,
|
||||
all_result: list[Any] | None = None,
|
||||
) -> None:
|
||||
self._first_result = first_result
|
||||
self._all_result = all_result or []
|
||||
|
||||
def join(self, *args: Any, **kwargs: Any) -> _FakeQuery:
|
||||
return self
|
||||
|
||||
def filter(self, *args: Any, **kwargs: Any) -> _FakeQuery:
|
||||
return self
|
||||
|
||||
def order_by(self, *args: Any, **kwargs: Any) -> _FakeQuery:
|
||||
return self
|
||||
|
||||
def offset(self, *args: Any, **kwargs: Any) -> _FakeQuery:
|
||||
return self
|
||||
|
||||
def limit(self, *args: Any, **kwargs: Any) -> _FakeQuery:
|
||||
return self
|
||||
|
||||
def first(self) -> Any:
|
||||
return self._first_result
|
||||
|
||||
def all(self) -> list[Any]:
|
||||
return self._all_result
|
||||
|
||||
|
||||
class _FakeGroupedDB:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
key_provider_rows: list[tuple[Any, Any]],
|
||||
endpoint_rows: list[tuple[str, str, str]],
|
||||
) -> None:
|
||||
self._key_provider_rows = key_provider_rows
|
||||
self._endpoint_rows = endpoint_rows
|
||||
|
||||
def query(self, *models: Any) -> _FakeQuery:
|
||||
if len(models) == 2:
|
||||
return _FakeQuery(all_result=self._key_provider_rows)
|
||||
if len(models) == 3:
|
||||
return _FakeQuery(all_result=self._endpoint_rows)
|
||||
raise AssertionError(f"unexpected query models: {models}")
|
||||
|
||||
|
||||
class _FakeListDB:
|
||||
def __init__(self, *, provider: Any, keys: list[SimpleNamespace]) -> None:
|
||||
self._provider = provider
|
||||
self._keys = keys
|
||||
|
||||
def query(self, model: Any) -> _FakeQuery:
|
||||
model_name = getattr(model, "__name__", "")
|
||||
if model_name == "Provider":
|
||||
return _FakeQuery(first_result=self._provider)
|
||||
if model_name == "ProviderAPIKey":
|
||||
return _FakeQuery(all_result=self._keys)
|
||||
raise AssertionError(f"unexpected query model: {model}")
|
||||
|
||||
|
||||
def test_get_keys_grouped_by_format_builds_expected_shape(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
provider = SimpleNamespace(id="p1", is_active=True, name="Provider-1")
|
||||
key = SimpleNamespace(
|
||||
id="k1",
|
||||
name="Key-1",
|
||||
api_formats=["openai:chat", "openai:cli"],
|
||||
auth_type="api_key",
|
||||
api_key="enc-key",
|
||||
internal_priority=3,
|
||||
global_priority_by_format={"openai:chat": 8},
|
||||
rate_multipliers={"openai:chat": 1.1},
|
||||
is_active=True,
|
||||
capabilities={"cache_1h": True, "ctx_1m": False},
|
||||
success_count=8,
|
||||
request_count=10,
|
||||
total_response_time_ms=800,
|
||||
health_by_format={"openai:chat": {"health_score": 0.7}},
|
||||
circuit_breaker_by_format={"openai:chat": {"open": True}},
|
||||
)
|
||||
db = _FakeGroupedDB(
|
||||
key_provider_rows=[(key, provider)],
|
||||
endpoint_rows=[
|
||||
("p1", "openai:chat", "https://chat.example"),
|
||||
("p1", "openai:cli", "https://cli.example"),
|
||||
],
|
||||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
query_service_module.crypto_service, "decrypt", lambda _v: "sk-1234567890abcd"
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
query_service_module,
|
||||
"get_capability",
|
||||
lambda name: SimpleNamespace(short_name="缓存1h") if name == "cache_1h" else None,
|
||||
)
|
||||
|
||||
result = get_keys_grouped_by_format(cast(Any, db))
|
||||
|
||||
assert set(result.keys()) == {"openai:chat", "openai:cli"}
|
||||
chat_item = result["openai:chat"][0]
|
||||
assert chat_item["id"] == "k1"
|
||||
assert chat_item["provider_name"] == "Provider-1"
|
||||
assert chat_item["endpoint_base_url"] == "https://chat.example"
|
||||
assert chat_item["format_priority"] == 8
|
||||
assert chat_item["circuit_breaker_open"] is True
|
||||
assert chat_item["health_score"] == 0.7
|
||||
assert chat_item["capabilities"] == ["缓存1h"]
|
||||
assert chat_item["api_key_masked"].startswith("sk-12345")
|
||||
assert chat_item["api_key_masked"].endswith("abcd")
|
||||
|
||||
cli_item = result["openai:cli"][0]
|
||||
assert cli_item["endpoint_base_url"] == "https://cli.example"
|
||||
assert cli_item["format_priority"] is None
|
||||
assert cli_item["circuit_breaker_open"] is False
|
||||
assert cli_item["health_score"] == 1.0
|
||||
|
||||
|
||||
def test_list_provider_keys_responses_provider_not_found_raises() -> None:
|
||||
db = _FakeListDB(provider=None, keys=[])
|
||||
with pytest.raises(NotFoundException, match="Provider p1 不存在"):
|
||||
list_provider_keys_responses(cast(Any, db), provider_id="p1", skip=0, limit=10)
|
||||
|
||||
|
||||
def test_list_provider_keys_responses_uses_response_builder(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
provider = SimpleNamespace(id="p1")
|
||||
keys = [SimpleNamespace(id="k1"), SimpleNamespace(id="k2")]
|
||||
db = _FakeListDB(provider=provider, keys=keys)
|
||||
|
||||
monkeypatch.setattr(query_service_module, "build_key_response", lambda key: {"id": key.id})
|
||||
|
||||
result = list_provider_keys_responses(cast(Any, db), provider_id="p1", skip=0, limit=10)
|
||||
assert result == [{"id": "k1"}, {"id": "k2"}]
|
||||
699
tests/services/test_provider_keys_quota_refresh_strategies.py
Normal file
699
tests/services/test_provider_keys_quota_refresh_strategies.py
Normal file
@@ -0,0 +1,699 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sys
|
||||
import types
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
|
||||
from src.services.provider_keys.codex_usage_parser import (
|
||||
CodexUsageParseError,
|
||||
parse_codex_wham_usage_response,
|
||||
)
|
||||
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
|
||||
|
||||
|
||||
class _FakeDB:
|
||||
def __init__(self) -> None:
|
||||
self.commit_count = 0
|
||||
|
||||
def commit(self) -> None:
|
||||
self.commit_count += 1
|
||||
|
||||
|
||||
class _FakeResponse:
|
||||
def __init__(
|
||||
self,
|
||||
status_code: int,
|
||||
payload: Any = None,
|
||||
json_exc: Exception | None = None,
|
||||
) -> None:
|
||||
self.status_code = status_code
|
||||
self._payload = payload
|
||||
self._json_exc = json_exc
|
||||
|
||||
def json(self) -> Any:
|
||||
if self._json_exc:
|
||||
raise self._json_exc
|
||||
return self._payload
|
||||
|
||||
|
||||
class _FakeAsyncClient:
|
||||
def __init__(self, response: _FakeResponse, **kwargs: Any) -> None:
|
||||
self._response = response
|
||||
self.kwargs = kwargs
|
||||
self.last_url: str | None = None
|
||||
self.last_headers: dict[str, str] | None = None
|
||||
|
||||
async def __aenter__(self) -> "_FakeAsyncClient":
|
||||
return self
|
||||
|
||||
async def __aexit__(
|
||||
self,
|
||||
exc_type: type[BaseException] | None,
|
||||
exc: BaseException | None,
|
||||
tb: Any,
|
||||
) -> bool:
|
||||
_ = exc_type, exc, tb
|
||||
return False
|
||||
|
||||
async def get(self, url: str, headers: dict[str, str]) -> _FakeResponse:
|
||||
self.last_url = url
|
||||
self.last_headers = headers
|
||||
return self._response
|
||||
|
||||
|
||||
def _install_module(monkeypatch: pytest.MonkeyPatch, name: str, attrs: dict[str, Any]) -> None:
|
||||
module = types.ModuleType(name)
|
||||
for key, value in attrs.items():
|
||||
setattr(module, key, value)
|
||||
monkeypatch.setitem(sys.modules, name, module)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_codex_refresher_endpoint_missing_returns_error() -> None:
|
||||
key = SimpleNamespace(id="k1", name="K1")
|
||||
provider = SimpleNamespace(proxy=None)
|
||||
|
||||
result = await refresh_codex_key_quota(
|
||||
db=cast(Any, _FakeDB()),
|
||||
provider=cast(Any, provider),
|
||||
key=cast(Any, key),
|
||||
endpoint=None,
|
||||
codex_wham_usage_url="https://example.test",
|
||||
metadata_updates={},
|
||||
state_updates={},
|
||||
)
|
||||
|
||||
assert result["status"] == "error"
|
||||
assert "openai:cli" in result["message"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_codex_refresher_http_non_200_returns_error(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from src.services.provider_keys.quota_refresh import codex_refresher as module
|
||||
|
||||
key = SimpleNamespace(
|
||||
id="k1", name="K1", api_key="enc", auth_type="api_key", auth_config=None, proxy=None
|
||||
)
|
||||
provider = SimpleNamespace(proxy=None)
|
||||
endpoint = SimpleNamespace()
|
||||
|
||||
async def _fake_auth_info(_endpoint: Any, _key: Any) -> Any:
|
||||
return None
|
||||
|
||||
_install_module(
|
||||
monkeypatch,
|
||||
"src.services.proxy_node.resolver",
|
||||
{
|
||||
"resolve_effective_proxy": lambda provider_proxy, key_proxy: None,
|
||||
"build_proxy_client_kwargs": lambda proxy, timeout: {"timeout": timeout},
|
||||
},
|
||||
)
|
||||
monkeypatch.setattr(module, "get_provider_auth", _fake_auth_info)
|
||||
monkeypatch.setattr(module.crypto_service, "decrypt", lambda _v: "sk-test")
|
||||
response = _FakeResponse(status_code=503, payload={"x": 1})
|
||||
monkeypatch.setattr(
|
||||
module.httpx, "AsyncClient", lambda **kwargs: _FakeAsyncClient(response, **kwargs)
|
||||
)
|
||||
|
||||
result = await refresh_codex_key_quota(
|
||||
db=cast(Any, _FakeDB()),
|
||||
provider=cast(Any, provider),
|
||||
key=cast(Any, key),
|
||||
endpoint=cast(Any, endpoint),
|
||||
codex_wham_usage_url="https://example.test",
|
||||
metadata_updates={},
|
||||
state_updates={},
|
||||
)
|
||||
|
||||
assert result["status"] == "error"
|
||||
assert result["status_code"] == 503
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_codex_refresher_success_updates_metadata(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from src.services.provider_keys.quota_refresh import codex_refresher as module
|
||||
|
||||
key = SimpleNamespace(
|
||||
id="k1",
|
||||
name="K1",
|
||||
api_key="enc",
|
||||
auth_type="api_key",
|
||||
auth_config=None,
|
||||
proxy=None,
|
||||
oauth_invalid_at="old",
|
||||
oauth_invalid_reason="old",
|
||||
)
|
||||
provider = SimpleNamespace(proxy=None)
|
||||
endpoint = SimpleNamespace()
|
||||
metadata_updates: dict[str, dict[str, Any]] = {}
|
||||
state_updates: dict[str, dict[str, Any]] = {}
|
||||
|
||||
async def _fake_auth_info(_endpoint: Any, _key: Any) -> Any:
|
||||
return None
|
||||
|
||||
_install_module(
|
||||
monkeypatch,
|
||||
"src.services.proxy_node.resolver",
|
||||
{
|
||||
"resolve_effective_proxy": lambda provider_proxy, key_proxy: None,
|
||||
"build_proxy_client_kwargs": lambda proxy, timeout: {"timeout": timeout},
|
||||
},
|
||||
)
|
||||
monkeypatch.setattr(module, "get_provider_auth", _fake_auth_info)
|
||||
monkeypatch.setattr(module.crypto_service, "decrypt", lambda _v: "sk-test")
|
||||
monkeypatch.setattr(
|
||||
module, "parse_codex_wham_usage_response", lambda _data: {"used_percent": 10.0}
|
||||
)
|
||||
response = _FakeResponse(status_code=200, payload={"ok": True})
|
||||
monkeypatch.setattr(
|
||||
module.httpx, "AsyncClient", lambda **kwargs: _FakeAsyncClient(response, **kwargs)
|
||||
)
|
||||
|
||||
result = await refresh_codex_key_quota(
|
||||
db=cast(Any, _FakeDB()),
|
||||
provider=cast(Any, provider),
|
||||
key=cast(Any, key),
|
||||
endpoint=cast(Any, endpoint),
|
||||
codex_wham_usage_url="https://example.test",
|
||||
metadata_updates=metadata_updates,
|
||||
state_updates=state_updates,
|
||||
)
|
||||
|
||||
assert result["status"] == "success"
|
||||
assert metadata_updates == {"k1": {"codex": {"used_percent": 10.0}}}
|
||||
assert state_updates == {"k1": {"oauth_invalid_at": None, "oauth_invalid_reason": None}}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_codex_refresher_parse_error_is_diagnostic(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from src.services.provider_keys.quota_refresh import codex_refresher as module
|
||||
|
||||
key = SimpleNamespace(
|
||||
id="k1", name="K1", api_key="enc", auth_type="api_key", auth_config=None, proxy=None
|
||||
)
|
||||
provider = SimpleNamespace(proxy=None)
|
||||
endpoint = SimpleNamespace()
|
||||
|
||||
async def _fake_auth_info(_endpoint: Any, _key: Any) -> Any:
|
||||
return None
|
||||
|
||||
_install_module(
|
||||
monkeypatch,
|
||||
"src.services.proxy_node.resolver",
|
||||
{
|
||||
"resolve_effective_proxy": lambda provider_proxy, key_proxy: None,
|
||||
"build_proxy_client_kwargs": lambda proxy, timeout: {"timeout": timeout},
|
||||
},
|
||||
)
|
||||
monkeypatch.setattr(module, "get_provider_auth", _fake_auth_info)
|
||||
monkeypatch.setattr(module.crypto_service, "decrypt", lambda _v: "sk-test")
|
||||
monkeypatch.setattr(
|
||||
module,
|
||||
"parse_codex_wham_usage_response",
|
||||
lambda _data: (_ for _ in ()).throw(
|
||||
CodexUsageParseError("rate_limit.primary_window 类型错误")
|
||||
),
|
||||
)
|
||||
response = _FakeResponse(status_code=200, payload={"ok": True})
|
||||
monkeypatch.setattr(
|
||||
module.httpx, "AsyncClient", lambda **kwargs: _FakeAsyncClient(response, **kwargs)
|
||||
)
|
||||
|
||||
result = await refresh_codex_key_quota(
|
||||
db=cast(Any, _FakeDB()),
|
||||
provider=cast(Any, provider),
|
||||
key=cast(Any, key),
|
||||
endpoint=cast(Any, endpoint),
|
||||
codex_wham_usage_url="https://example.test",
|
||||
metadata_updates={},
|
||||
state_updates={},
|
||||
)
|
||||
|
||||
assert result["status"] == "error"
|
||||
assert "响应结构异常" in result["message"]
|
||||
assert "rate_limit.primary_window 类型错误" in result["message"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_codex_refresher_oauth_missing_plan_type_adds_account_header(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from src.services.provider_keys.quota_refresh import codex_refresher as module
|
||||
|
||||
key = SimpleNamespace(
|
||||
id="k1",
|
||||
name="K1",
|
||||
api_key="enc",
|
||||
auth_type="oauth",
|
||||
auth_config="enc-config",
|
||||
proxy=None,
|
||||
)
|
||||
provider = SimpleNamespace(proxy=None)
|
||||
endpoint = SimpleNamespace()
|
||||
response = _FakeResponse(status_code=200, payload={"ok": True})
|
||||
client_ref: dict[str, _FakeAsyncClient] = {}
|
||||
|
||||
async def _fake_auth_info(_endpoint: Any, _key: Any) -> Any:
|
||||
return SimpleNamespace(auth_header="Authorization", auth_value="Bearer oauth-token")
|
||||
|
||||
_install_module(
|
||||
monkeypatch,
|
||||
"src.services.proxy_node.resolver",
|
||||
{
|
||||
"resolve_effective_proxy": lambda provider_proxy, key_proxy: None,
|
||||
"build_proxy_client_kwargs": lambda proxy, timeout: {"timeout": timeout},
|
||||
},
|
||||
)
|
||||
monkeypatch.setattr(module, "get_provider_auth", _fake_auth_info)
|
||||
monkeypatch.setattr(
|
||||
module, "parse_codex_wham_usage_response", lambda _data: {"used_percent": 1.0}
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
module.crypto_service, "decrypt", lambda _v: json.dumps({"account_id": "acc-1"})
|
||||
)
|
||||
|
||||
def _client_factory(**kwargs: Any) -> _FakeAsyncClient:
|
||||
client = _FakeAsyncClient(response, **kwargs)
|
||||
client_ref["client"] = client
|
||||
return client
|
||||
|
||||
monkeypatch.setattr(module.httpx, "AsyncClient", _client_factory)
|
||||
|
||||
result = await refresh_codex_key_quota(
|
||||
db=cast(Any, _FakeDB()),
|
||||
provider=cast(Any, provider),
|
||||
key=cast(Any, key),
|
||||
endpoint=cast(Any, endpoint),
|
||||
codex_wham_usage_url="https://example.test",
|
||||
metadata_updates={},
|
||||
state_updates={},
|
||||
)
|
||||
|
||||
assert result["status"] == "success"
|
||||
assert client_ref["client"].last_headers is not None
|
||||
assert client_ref["client"].last_headers.get("chatgpt-account-id") == "acc-1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_codex_refresher_oauth_uppercase_free_does_not_add_account_header(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from src.services.provider_keys.quota_refresh import codex_refresher as module
|
||||
|
||||
key = SimpleNamespace(
|
||||
id="k1",
|
||||
name="K1",
|
||||
api_key="enc",
|
||||
auth_type="oauth",
|
||||
auth_config="enc-config",
|
||||
proxy=None,
|
||||
)
|
||||
provider = SimpleNamespace(proxy=None)
|
||||
endpoint = SimpleNamespace()
|
||||
response = _FakeResponse(status_code=200, payload={"ok": True})
|
||||
client_ref: dict[str, _FakeAsyncClient] = {}
|
||||
|
||||
async def _fake_auth_info(_endpoint: Any, _key: Any) -> Any:
|
||||
return SimpleNamespace(auth_header="Authorization", auth_value="Bearer oauth-token")
|
||||
|
||||
_install_module(
|
||||
monkeypatch,
|
||||
"src.services.proxy_node.resolver",
|
||||
{
|
||||
"resolve_effective_proxy": lambda provider_proxy, key_proxy: None,
|
||||
"build_proxy_client_kwargs": lambda proxy, timeout: {"timeout": timeout},
|
||||
},
|
||||
)
|
||||
monkeypatch.setattr(module, "get_provider_auth", _fake_auth_info)
|
||||
monkeypatch.setattr(
|
||||
module, "parse_codex_wham_usage_response", lambda _data: {"used_percent": 1.0}
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
module.crypto_service,
|
||||
"decrypt",
|
||||
lambda _v: json.dumps({"account_id": "acc-1", "plan_type": "FREE"}),
|
||||
)
|
||||
|
||||
def _client_factory(**kwargs: Any) -> _FakeAsyncClient:
|
||||
client = _FakeAsyncClient(response, **kwargs)
|
||||
client_ref["client"] = client
|
||||
return client
|
||||
|
||||
monkeypatch.setattr(module.httpx, "AsyncClient", _client_factory)
|
||||
|
||||
result = await refresh_codex_key_quota(
|
||||
db=cast(Any, _FakeDB()),
|
||||
provider=cast(Any, provider),
|
||||
key=cast(Any, key),
|
||||
endpoint=cast(Any, endpoint),
|
||||
codex_wham_usage_url="https://example.test",
|
||||
metadata_updates={},
|
||||
state_updates={},
|
||||
)
|
||||
|
||||
assert result["status"] == "success"
|
||||
assert client_ref["client"].last_headers is not None
|
||||
assert "chatgpt-account-id" not in client_ref["client"].last_headers
|
||||
|
||||
|
||||
def test_parse_codex_usage_plan_type_case_insensitive_free_window_semantics() -> None:
|
||||
parsed = parse_codex_wham_usage_response(
|
||||
{
|
||||
"plan_type": "FREE",
|
||||
"rate_limit": {
|
||||
"primary_window": {
|
||||
"used_percent": "12.5",
|
||||
"reset_after_seconds": "120",
|
||||
"reset_at": "1700000000",
|
||||
"limit_window_seconds": "604800",
|
||||
}
|
||||
},
|
||||
}
|
||||
)
|
||||
assert parsed is not None
|
||||
assert parsed["plan_type"] == "free"
|
||||
assert parsed["primary_used_percent"] == 12.5
|
||||
assert parsed["primary_window_minutes"] == 10080
|
||||
assert "secondary_used_percent" not in parsed
|
||||
|
||||
|
||||
def test_parse_codex_usage_missing_plan_type_infers_paid_windows() -> None:
|
||||
parsed = parse_codex_wham_usage_response(
|
||||
{
|
||||
"rate_limit": {
|
||||
"primary_window": {
|
||||
"used_percent": 25,
|
||||
"reset_after_seconds": 600,
|
||||
"reset_at": 1700000000,
|
||||
"limit_window_seconds": 18000,
|
||||
},
|
||||
"secondary_window": {
|
||||
"used_percent": 80,
|
||||
"reset_after_seconds": 3600,
|
||||
"reset_at": 1700003600,
|
||||
"limit_window_seconds": 604800,
|
||||
},
|
||||
}
|
||||
}
|
||||
)
|
||||
assert parsed is not None
|
||||
assert parsed["primary_used_percent"] == 80.0
|
||||
assert parsed["secondary_used_percent"] == 25.0
|
||||
assert parsed["primary_window_minutes"] == 10080
|
||||
assert parsed["secondary_window_minutes"] == 300
|
||||
|
||||
|
||||
def test_parse_codex_usage_invalid_type_raises_diagnostic_error() -> None:
|
||||
with pytest.raises(CodexUsageParseError, match="rate_limit.primary_window 类型错误"):
|
||||
parse_codex_wham_usage_response({"rate_limit": {"primary_window": []}})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_antigravity_refresher_forbidden_collects_updates_without_commit(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from src.services.provider_keys.quota_refresh import antigravity_refresher as module
|
||||
|
||||
class _Forbidden(Exception):
|
||||
def __init__(self, reason: str) -> None:
|
||||
super().__init__(reason)
|
||||
self.reason = reason
|
||||
self.message = reason
|
||||
|
||||
async def _fake_auth_info(_endpoint: Any, _key: Any) -> Any:
|
||||
return SimpleNamespace(auth_value="Bearer tk", decrypted_auth_config={"pid": "p1"})
|
||||
|
||||
async def _fetch_models_for_key(_ctx: Any, timeout_seconds: float) -> Any:
|
||||
_ = timeout_seconds
|
||||
raise _Forbidden("forbidden-by-test")
|
||||
|
||||
class _UpstreamModelsFetchContext: # noqa: D101
|
||||
def __init__(self, **kwargs: Any) -> None:
|
||||
self.kwargs = kwargs
|
||||
|
||||
_install_module(
|
||||
monkeypatch,
|
||||
"src.services.model.upstream_fetcher",
|
||||
{
|
||||
"UpstreamModelsFetchContext": _UpstreamModelsFetchContext,
|
||||
"fetch_models_for_key": _fetch_models_for_key,
|
||||
},
|
||||
)
|
||||
_install_module(
|
||||
monkeypatch,
|
||||
"src.services.provider.adapters.antigravity.client",
|
||||
{"AntigravityAccountForbiddenException": _Forbidden},
|
||||
)
|
||||
_install_module(
|
||||
monkeypatch,
|
||||
"src.services.proxy_node.resolver",
|
||||
{"resolve_effective_proxy": lambda provider_proxy, key_proxy: None},
|
||||
)
|
||||
monkeypatch.setattr(module, "get_provider_auth", _fake_auth_info)
|
||||
|
||||
db = _FakeDB()
|
||||
provider = SimpleNamespace(proxy=None)
|
||||
key = SimpleNamespace(
|
||||
id="k1",
|
||||
name="K1",
|
||||
proxy=None,
|
||||
is_active=True,
|
||||
oauth_invalid_at=None,
|
||||
oauth_invalid_reason=None,
|
||||
upstream_metadata={},
|
||||
)
|
||||
endpoint = SimpleNamespace()
|
||||
metadata_updates: dict[str, dict[str, Any]] = {}
|
||||
state_updates: dict[str, dict[str, Any]] = {}
|
||||
|
||||
result = await refresh_antigravity_key_quota(
|
||||
db=cast(Any, db),
|
||||
provider=cast(Any, provider),
|
||||
key=cast(Any, key),
|
||||
endpoint=cast(Any, endpoint),
|
||||
codex_wham_usage_url="https://example.test",
|
||||
metadata_updates=metadata_updates,
|
||||
state_updates=state_updates,
|
||||
)
|
||||
|
||||
assert result["status"] == "forbidden"
|
||||
assert result["auto_disabled"] is True
|
||||
assert key.is_active is True
|
||||
assert key.oauth_invalid_reason is None
|
||||
assert state_updates["k1"]["is_active"] is False
|
||||
assert state_updates["k1"]["oauth_invalid_reason"].startswith("账户访问被禁止")
|
||||
assert metadata_updates["k1"]["antigravity"]["is_forbidden"] is True
|
||||
assert db.commit_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_antigravity_refresher_success_resets_forbidden_flag(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from src.services.provider_keys.quota_refresh import antigravity_refresher as module
|
||||
|
||||
async def _fake_auth_info(_endpoint: Any, _key: Any) -> Any:
|
||||
return SimpleNamespace(auth_value="Bearer tk", decrypted_auth_config={})
|
||||
|
||||
async def _fetch_models_for_key(_ctx: Any, timeout_seconds: float) -> Any:
|
||||
_ = timeout_seconds
|
||||
return (
|
||||
[],
|
||||
[],
|
||||
True,
|
||||
{"antigravity": {"is_forbidden": True, "forbidden_reason": "x", "forbidden_at": 1}},
|
||||
)
|
||||
|
||||
class _UpstreamModelsFetchContext: # noqa: D101
|
||||
def __init__(self, **kwargs: Any) -> None:
|
||||
self.kwargs = kwargs
|
||||
|
||||
_install_module(
|
||||
monkeypatch,
|
||||
"src.services.model.upstream_fetcher",
|
||||
{
|
||||
"UpstreamModelsFetchContext": _UpstreamModelsFetchContext,
|
||||
"fetch_models_for_key": _fetch_models_for_key,
|
||||
},
|
||||
)
|
||||
_install_module(
|
||||
monkeypatch,
|
||||
"src.services.provider.adapters.antigravity.client",
|
||||
{"AntigravityAccountForbiddenException": RuntimeError},
|
||||
)
|
||||
_install_module(
|
||||
monkeypatch,
|
||||
"src.services.proxy_node.resolver",
|
||||
{"resolve_effective_proxy": lambda provider_proxy, key_proxy: None},
|
||||
)
|
||||
monkeypatch.setattr(module, "get_provider_auth", _fake_auth_info)
|
||||
|
||||
metadata_updates: dict[str, dict[str, Any]] = {}
|
||||
state_updates: dict[str, dict[str, Any]] = {}
|
||||
result = await refresh_antigravity_key_quota(
|
||||
db=cast(Any, _FakeDB()),
|
||||
provider=cast(Any, SimpleNamespace(proxy=None)),
|
||||
key=cast(
|
||||
Any,
|
||||
SimpleNamespace(
|
||||
id="k1",
|
||||
name="K1",
|
||||
proxy=None,
|
||||
oauth_invalid_at="old",
|
||||
oauth_invalid_reason="old",
|
||||
),
|
||||
),
|
||||
endpoint=cast(Any, SimpleNamespace()),
|
||||
codex_wham_usage_url="https://example.test",
|
||||
metadata_updates=metadata_updates,
|
||||
state_updates=state_updates,
|
||||
)
|
||||
|
||||
assert result["status"] == "success"
|
||||
assert metadata_updates["k1"]["antigravity"]["is_forbidden"] is False
|
||||
assert metadata_updates["k1"]["antigravity"]["forbidden_reason"] is None
|
||||
assert metadata_updates["k1"]["antigravity"]["forbidden_at"] is None
|
||||
assert state_updates["k1"]["oauth_invalid_at"] is None
|
||||
assert state_updates["k1"]["oauth_invalid_reason"] is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_kiro_refresher_runtime_401_marks_key_invalid(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from src.services.provider_keys.quota_refresh import kiro_refresher as module
|
||||
|
||||
class _Banned(Exception):
|
||||
pass
|
||||
|
||||
async def _fetch_limits(auth_config: dict[str, Any], proxy_config: object) -> Any:
|
||||
_ = auth_config, proxy_config
|
||||
raise RuntimeError("401 token expired")
|
||||
|
||||
def _parse_usage(_usage: Any) -> dict[str, Any]:
|
||||
return {"quota": 1}
|
||||
|
||||
_install_module(
|
||||
monkeypatch,
|
||||
"src.services.provider.adapters.kiro.usage",
|
||||
{
|
||||
"KiroAccountBannedException": _Banned,
|
||||
"fetch_kiro_usage_limits": _fetch_limits,
|
||||
"parse_kiro_usage_response": _parse_usage,
|
||||
},
|
||||
)
|
||||
_install_module(
|
||||
monkeypatch,
|
||||
"src.services.proxy_node.resolver",
|
||||
{"resolve_effective_proxy": lambda provider_proxy, key_proxy: None},
|
||||
)
|
||||
monkeypatch.setattr(module.crypto_service, "decrypt", lambda _v: "{}")
|
||||
|
||||
db = _FakeDB()
|
||||
key = SimpleNamespace(
|
||||
id="k1",
|
||||
name="K1",
|
||||
auth_config="enc",
|
||||
proxy=None,
|
||||
is_active=True,
|
||||
oauth_invalid_at=None,
|
||||
oauth_invalid_reason=None,
|
||||
upstream_metadata={},
|
||||
)
|
||||
state_updates: dict[str, dict[str, Any]] = {}
|
||||
|
||||
result = await refresh_kiro_key_quota(
|
||||
db=cast(Any, db),
|
||||
provider=cast(Any, SimpleNamespace(proxy=None)),
|
||||
key=cast(Any, key),
|
||||
endpoint=None,
|
||||
codex_wham_usage_url="https://example.test",
|
||||
metadata_updates={},
|
||||
state_updates=state_updates,
|
||||
)
|
||||
|
||||
assert result["status"] == "error"
|
||||
assert "401" in result["message"]
|
||||
assert key.is_active is True
|
||||
assert key.oauth_invalid_reason is None
|
||||
assert state_updates["k1"]["is_active"] is False
|
||||
assert state_updates["k1"]["oauth_invalid_reason"] == "Kiro Token 无效或已过期"
|
||||
assert db.commit_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_kiro_refresher_success_updates_metadata_and_auth_config(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
from src.services.provider_keys.quota_refresh import kiro_refresher as module
|
||||
|
||||
class _Banned(Exception):
|
||||
pass
|
||||
|
||||
async def _fetch_limits(auth_config: dict[str, Any], proxy_config: object) -> Any:
|
||||
_ = auth_config, proxy_config
|
||||
return {"usage_data": {"x": 1}, "updated_auth_config": {"token": "new"}}
|
||||
|
||||
def _parse_usage(_usage: Any) -> dict[str, Any]:
|
||||
return {"quota": 1}
|
||||
|
||||
_install_module(
|
||||
monkeypatch,
|
||||
"src.services.provider.adapters.kiro.usage",
|
||||
{
|
||||
"KiroAccountBannedException": _Banned,
|
||||
"fetch_kiro_usage_limits": _fetch_limits,
|
||||
"parse_kiro_usage_response": _parse_usage,
|
||||
},
|
||||
)
|
||||
_install_module(
|
||||
monkeypatch,
|
||||
"src.services.proxy_node.resolver",
|
||||
{"resolve_effective_proxy": lambda provider_proxy, key_proxy: None},
|
||||
)
|
||||
monkeypatch.setattr(module.crypto_service, "decrypt", lambda _v: json.dumps({"seed": 1}))
|
||||
monkeypatch.setattr(module.crypto_service, "encrypt", lambda raw: f"ENC:{raw}")
|
||||
|
||||
key = SimpleNamespace(
|
||||
id="k1",
|
||||
name="K1",
|
||||
auth_config="enc",
|
||||
proxy=None,
|
||||
oauth_invalid_at="old",
|
||||
oauth_invalid_reason="old",
|
||||
upstream_metadata={},
|
||||
)
|
||||
metadata_updates: dict[str, dict[str, Any]] = {}
|
||||
state_updates: dict[str, dict[str, Any]] = {}
|
||||
|
||||
result = await refresh_kiro_key_quota(
|
||||
db=cast(Any, _FakeDB()),
|
||||
provider=cast(Any, SimpleNamespace(proxy=None)),
|
||||
key=cast(Any, key),
|
||||
endpoint=None,
|
||||
codex_wham_usage_url="https://example.test",
|
||||
metadata_updates=metadata_updates,
|
||||
state_updates=state_updates,
|
||||
)
|
||||
|
||||
assert result["status"] == "success"
|
||||
assert metadata_updates["k1"]["kiro"]["is_banned"] is False
|
||||
assert metadata_updates["k1"]["kiro"]["quota"] == 1
|
||||
assert key.auth_config == "enc"
|
||||
assert state_updates["k1"]["oauth_invalid_at"] is None
|
||||
assert state_updates["k1"]["oauth_invalid_reason"] is None
|
||||
assert state_updates["k1"]["auth_config"].startswith("ENC:")
|
||||
300
tests/services/test_provider_keys_quota_service.py
Normal file
300
tests/services/test_provider_keys_quota_service.py
Normal file
@@ -0,0 +1,300 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
|
||||
from src.core.exceptions import InvalidRequestException
|
||||
from src.core.provider_types import ProviderType
|
||||
from src.services.provider_keys import key_quota_service as quota_service_module
|
||||
from src.services.provider_keys.key_quota_service import (
|
||||
_resolve_quota_refresh_handler,
|
||||
_select_refresh_endpoint,
|
||||
refresh_provider_quota_for_provider,
|
||||
)
|
||||
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
|
||||
|
||||
|
||||
def _provider_with_endpoints(*endpoints: tuple[str, bool]) -> SimpleNamespace:
|
||||
eps = [SimpleNamespace(api_format=fmt, is_active=active) for fmt, active in endpoints]
|
||||
return SimpleNamespace(endpoints=eps)
|
||||
|
||||
|
||||
class _FakeQuery:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
first_result: Any = None,
|
||||
all_result: list[Any] | None = None,
|
||||
) -> None:
|
||||
self._first_result = first_result
|
||||
self._all_result = all_result or []
|
||||
|
||||
def filter(self, *args: Any, **kwargs: Any) -> _FakeQuery:
|
||||
return self
|
||||
|
||||
def first(self) -> Any:
|
||||
return self._first_result
|
||||
|
||||
def all(self) -> list[Any]:
|
||||
return self._all_result
|
||||
|
||||
|
||||
class _FakeDB:
|
||||
def __init__(self, *, provider: Any, keys: list[SimpleNamespace]) -> None:
|
||||
self._provider = provider
|
||||
self._keys = keys
|
||||
self.added: list[object] = []
|
||||
self.commit_count = 0
|
||||
|
||||
def query(self, model: Any) -> _FakeQuery:
|
||||
model_name = getattr(model, "__name__", "")
|
||||
if model_name == "Provider":
|
||||
return _FakeQuery(first_result=self._provider)
|
||||
if model_name == "ProviderAPIKey":
|
||||
return _FakeQuery(all_result=self._keys)
|
||||
raise AssertionError(f"unexpected query model: {model}")
|
||||
|
||||
def add(self, obj: object) -> None:
|
||||
self.added.append(obj)
|
||||
|
||||
def commit(self) -> None:
|
||||
self.commit_count += 1
|
||||
|
||||
|
||||
def test_select_refresh_endpoint_codex() -> None:
|
||||
provider = _provider_with_endpoints(("openai:chat", True), ("openai:cli", True))
|
||||
endpoint = _select_refresh_endpoint(cast(Any, provider), ProviderType.CODEX)
|
||||
assert endpoint is not None
|
||||
assert endpoint.api_format == "openai:cli"
|
||||
|
||||
|
||||
def test_select_refresh_endpoint_codex_normalized_api_format() -> None:
|
||||
provider = _provider_with_endpoints(("openai:chat", True), (" OpenAI:CLI ", True))
|
||||
endpoint = _select_refresh_endpoint(cast(Any, provider), ProviderType.CODEX)
|
||||
assert endpoint is not None
|
||||
assert endpoint.api_format == " OpenAI:CLI "
|
||||
|
||||
|
||||
def test_select_refresh_endpoint_antigravity_prefers_chat() -> None:
|
||||
provider = _provider_with_endpoints(("gemini:cli", True), ("gemini:chat", True))
|
||||
endpoint = _select_refresh_endpoint(cast(Any, provider), ProviderType.ANTIGRAVITY)
|
||||
assert endpoint is not None
|
||||
assert endpoint.api_format == "gemini:chat"
|
||||
|
||||
|
||||
def test_select_refresh_endpoint_antigravity_fallback_cli() -> None:
|
||||
provider = _provider_with_endpoints(("gemini:chat", False), ("gemini:cli", True))
|
||||
endpoint = _select_refresh_endpoint(cast(Any, provider), ProviderType.ANTIGRAVITY)
|
||||
assert endpoint is not None
|
||||
assert endpoint.api_format == "gemini:cli"
|
||||
|
||||
|
||||
def test_select_refresh_endpoint_kiro_returns_none() -> None:
|
||||
provider = _provider_with_endpoints(("openai:cli", True))
|
||||
endpoint = _select_refresh_endpoint(cast(Any, provider), ProviderType.KIRO)
|
||||
assert endpoint is None
|
||||
|
||||
|
||||
def test_select_refresh_endpoint_codex_missing_raises() -> None:
|
||||
provider = _provider_with_endpoints(("openai:chat", True))
|
||||
with pytest.raises(InvalidRequestException, match="找不到有效的 openai:cli 端点"):
|
||||
_select_refresh_endpoint(cast(Any, provider), ProviderType.CODEX)
|
||||
|
||||
|
||||
def test_select_refresh_endpoint_antigravity_missing_raises() -> None:
|
||||
provider = _provider_with_endpoints(("gemini:chat", False), ("gemini:cli", False))
|
||||
with pytest.raises(InvalidRequestException, match="找不到有效的 gemini:chat/gemini:cli 端点"):
|
||||
_select_refresh_endpoint(cast(Any, provider), ProviderType.ANTIGRAVITY)
|
||||
|
||||
|
||||
def test_resolve_quota_refresh_handler() -> None:
|
||||
assert _resolve_quota_refresh_handler(ProviderType.CODEX) is refresh_codex_key_quota
|
||||
assert _resolve_quota_refresh_handler(ProviderType.ANTIGRAVITY) is refresh_antigravity_key_quota
|
||||
assert _resolve_quota_refresh_handler(ProviderType.KIRO) is refresh_kiro_key_quota
|
||||
|
||||
|
||||
def test_resolve_quota_refresh_handler_unsupported_raises() -> None:
|
||||
with pytest.raises(
|
||||
InvalidRequestException,
|
||||
match="仅支持 Codex / Antigravity / Kiro 类型的 Provider 刷新限额",
|
||||
):
|
||||
_resolve_quota_refresh_handler("unknown")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_provider_quota_no_active_keys_returns_empty() -> None:
|
||||
provider = SimpleNamespace(
|
||||
id="p1",
|
||||
provider_type=ProviderType.CODEX,
|
||||
endpoints=[SimpleNamespace(api_format="openai:cli", is_active=True)],
|
||||
)
|
||||
db = _FakeDB(provider=provider, keys=[])
|
||||
|
||||
result = await refresh_provider_quota_for_provider(
|
||||
db=cast(Any, db),
|
||||
provider_id="p1",
|
||||
codex_wham_usage_url="https://example.test/wham/usage",
|
||||
)
|
||||
|
||||
assert result == {
|
||||
"success": 0,
|
||||
"failed": 0,
|
||||
"total": 0,
|
||||
"results": [],
|
||||
"message": "没有活跃的 Key",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_provider_quota_aggregates_and_merges_metadata(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
provider = SimpleNamespace(
|
||||
id="p1",
|
||||
provider_type=ProviderType.CODEX,
|
||||
endpoints=[SimpleNamespace(api_format="openai:cli", is_active=True)],
|
||||
)
|
||||
key1 = SimpleNamespace(id="k1", name="K1", upstream_metadata={"old": True})
|
||||
key2 = SimpleNamespace(id="k2", name="K2", upstream_metadata={})
|
||||
db = _FakeDB(provider=provider, keys=[key1, key2])
|
||||
|
||||
async def _fake_handler(**kwargs: Any) -> dict[str, Any]:
|
||||
key = kwargs["key"]
|
||||
metadata_updates = kwargs["metadata_updates"]
|
||||
if key.id == "k1":
|
||||
metadata_updates[key.id] = {"codex": {"used": 10}}
|
||||
return {"key_id": key.id, "key_name": key.name, "status": "success"}
|
||||
return {"key_id": key.id, "key_name": key.name, "status": "error", "message": "boom"}
|
||||
|
||||
monkeypatch.setattr(
|
||||
quota_service_module,
|
||||
"_select_refresh_endpoint",
|
||||
lambda provider, provider_type: provider.endpoints[0],
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
quota_service_module, "_resolve_quota_refresh_handler", lambda _: _fake_handler
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
quota_service_module,
|
||||
"merge_upstream_metadata",
|
||||
lambda current, updates: {**(current or {}), **updates},
|
||||
)
|
||||
|
||||
result = await refresh_provider_quota_for_provider(
|
||||
db=cast(Any, db),
|
||||
provider_id="p1",
|
||||
codex_wham_usage_url="https://example.test/wham/usage",
|
||||
)
|
||||
|
||||
assert result["success"] == 1
|
||||
assert result["failed"] == 1
|
||||
assert result["total"] == 2
|
||||
assert len(result["results"]) == 2
|
||||
assert key1.upstream_metadata == {"old": True, "codex": {"used": 10}}
|
||||
assert db.commit_count == 1
|
||||
assert db.added == [key1]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_provider_quota_applies_state_updates_and_single_commit(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
provider = SimpleNamespace(
|
||||
id="p1",
|
||||
provider_type=ProviderType.CODEX,
|
||||
endpoints=[SimpleNamespace(api_format="openai:cli", is_active=True)],
|
||||
)
|
||||
key1 = SimpleNamespace(
|
||||
id="k1",
|
||||
name="K1",
|
||||
upstream_metadata={},
|
||||
is_active=True,
|
||||
oauth_invalid_at="old",
|
||||
oauth_invalid_reason="old-reason",
|
||||
)
|
||||
key2 = SimpleNamespace(
|
||||
id="k2",
|
||||
name="K2",
|
||||
upstream_metadata={},
|
||||
is_active=True,
|
||||
oauth_invalid_at=None,
|
||||
oauth_invalid_reason=None,
|
||||
)
|
||||
db = _FakeDB(provider=provider, keys=[key1, key2])
|
||||
|
||||
async def _fake_handler(**kwargs: Any) -> dict[str, Any]:
|
||||
key = kwargs["key"]
|
||||
state_updates = kwargs["state_updates"]
|
||||
if key.id == "k1":
|
||||
state_updates[key.id] = {"oauth_invalid_at": None, "oauth_invalid_reason": None}
|
||||
return {"key_id": key.id, "key_name": key.name, "status": "success"}
|
||||
state_updates[key.id] = {"is_active": False, "oauth_invalid_reason": "401"}
|
||||
return {"key_id": key.id, "key_name": key.name, "status": "error", "message": "401"}
|
||||
|
||||
monkeypatch.setattr(
|
||||
quota_service_module,
|
||||
"_select_refresh_endpoint",
|
||||
lambda provider, provider_type: provider.endpoints[0],
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
quota_service_module, "_resolve_quota_refresh_handler", lambda _: _fake_handler
|
||||
)
|
||||
|
||||
result = await refresh_provider_quota_for_provider(
|
||||
db=cast(Any, db),
|
||||
provider_id="p1",
|
||||
codex_wham_usage_url="https://example.test/wham/usage",
|
||||
)
|
||||
|
||||
assert result["success"] == 1
|
||||
assert result["failed"] == 1
|
||||
assert key1.oauth_invalid_at is None
|
||||
assert key1.oauth_invalid_reason is None
|
||||
assert key2.is_active is False
|
||||
assert key2.oauth_invalid_reason == "401"
|
||||
assert db.commit_count == 1
|
||||
assert db.added == [key1, key2]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_provider_quota_handler_exception_returns_error(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
provider = SimpleNamespace(
|
||||
id="p1",
|
||||
provider_type=ProviderType.CODEX,
|
||||
endpoints=[SimpleNamespace(api_format="openai:cli", is_active=True)],
|
||||
)
|
||||
key = SimpleNamespace(id="k1", name="K1", upstream_metadata={})
|
||||
db = _FakeDB(provider=provider, keys=[key])
|
||||
|
||||
async def _boom_handler(**kwargs: Any) -> dict[str, Any]:
|
||||
_ = kwargs
|
||||
raise RuntimeError("unit-test boom")
|
||||
|
||||
monkeypatch.setattr(
|
||||
quota_service_module,
|
||||
"_select_refresh_endpoint",
|
||||
lambda provider, provider_type: provider.endpoints[0],
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
quota_service_module, "_resolve_quota_refresh_handler", lambda _: _boom_handler
|
||||
)
|
||||
|
||||
result = await refresh_provider_quota_for_provider(
|
||||
db=cast(Any, db),
|
||||
provider_id="p1",
|
||||
codex_wham_usage_url="https://example.test/wham/usage",
|
||||
)
|
||||
|
||||
assert result["success"] == 0
|
||||
assert result["failed"] == 1
|
||||
assert result["total"] == 1
|
||||
assert result["results"][0]["status"] == "error"
|
||||
assert "unit-test boom" in result["results"][0]["message"]
|
||||
Reference in New Issue
Block a user