refactor: 继续拆分大型模块并增强模块注册健壮性

后端:
- chat_handler_base 错误处理函数提取到 chat_error_utils 子模块
- CLI mixin 引入 CliHandlerProtocol 协议类改善类型标注
- aware_scheduler 拆分为 _candidate_builder 和 _candidate_sorter 子模块
- usage recording 拆分为 _billing_integration 和 _recording_helpers 子模块
- ModuleRegistry 添加循环依赖检测,将写操作从查询方法分离到 reconcile_module_state
- 修正 plugin manager 入度注释

前端:
- 路由守卫逻辑拆分为独立 guards 模块
- ProviderManagement 拆分为 TableHeader/TableRow/BalanceCell/MobileCard 子组件
- SystemSettings 拆分为多个 Section 子组件和 composables
- 提取 useEndpointStatus/useProviderBalance/useProviderFilters composables

测试适配重构后的子模块结构
This commit is contained in:
fawney19
2026-02-14 20:06:38 +08:00
parent 676e918edc
commit 8a670f5524
51 changed files with 7095 additions and 5359 deletions

591
src/services/cache/_candidate_builder.py vendored Normal file
View File

@@ -0,0 +1,591 @@
"""
候选构建器 (CandidateBuilder)
从 CacheAwareScheduler 拆分出的候选构建逻辑,负责:
- 查询活跃 Provider
- 检查模型支持
- 检查 Key 可用性
- 构建候选列表
"""
from __future__ import annotations
import re
from collections.abc import Sequence
from typing import TYPE_CHECKING
from sqlalchemy.orm import Session, selectinload
from src.core.api_format.conversion.compatibility import is_format_compatible
from src.core.api_format.enums import EndpointKind
from src.core.api_format.signature import make_signature_key, parse_signature_key
from src.core.key_capabilities import check_capability_match
from src.core.logger import logger
from src.core.model_permissions import check_model_allowed_with_mappings
from src.models.database import (
Model,
Provider,
ProviderAPIKey,
ProviderEndpoint,
)
from src.services.cache.quota_skipper import is_key_quota_exhausted
from src.services.health.monitor import health_monitor
from src.services.provider.format import normalize_endpoint_signature
if TYPE_CHECKING:
from src.models.database import GlobalModel
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
from src.services.cache.model_cache import ModelCacheService
def _sort_endpoints_by_family_priority(
eps: Sequence[ProviderEndpoint],
) -> list[ProviderEndpoint]:
"""按 ApiFamily 优先级对端点排序(同分组内使用)。"""
from src.core.api_format.enums import ApiFamily
def sort_key(ep: ProviderEndpoint) -> int:
family_str = str(getattr(ep, "api_family", "") or "").strip().lower()
try:
return ApiFamily(family_str).priority
except ValueError:
return 99
return sorted(eps, key=sort_key)
class CandidateBuilder:
"""候选构建器,负责查询 Provider、检查模型支持和 Key 可用性、构建候选列表。"""
def __init__(self, scheduler: CacheAwareScheduler) -> None:
self._scheduler = scheduler
def _query_providers(
self,
db: Session,
provider_offset: int = 0,
provider_limit: int | None = None,
) -> list[Provider]:
"""
查询活跃的 Providers带预加载
Args:
db: 数据库会话
provider_offset: 分页偏移
provider_limit: 分页限制
Returns:
Provider 列表
"""
provider_query = (
db.query(Provider)
.options(
# 预加载 Provider 级别的 api_keys
selectinload(Provider.api_keys),
# 预加载 endpoints用于按 api_format 选择请求配置)
selectinload(Provider.endpoints),
# 同时加载 models 和 global_model 关系
selectinload(Provider.models).selectinload(Model.global_model),
)
.filter(Provider.is_active.is_(True))
.order_by(Provider.provider_priority.asc())
)
if provider_offset:
provider_query = provider_query.offset(provider_offset)
if provider_limit:
provider_query = provider_query.limit(provider_limit)
return provider_query.all()
async def _check_model_support(
self,
db: Session,
provider: Provider,
model_name: str,
api_format: str | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
) -> tuple[bool, str | None, list[str] | None, set[str] | None]:
"""
检查 Provider 是否支持指定模型(可选检查流式支持和能力需求)
模型能力检查在这里进行(而不是在 Key 级别),因为:
- 模型支持的能力是全局的,与具体的 Key 无关
- 如果模型不支持某能力,整个 Provider 的所有 Key 都应该被跳过
仅支持直接匹配 GlobalModel.name外部请求不接受映射名
Args:
db: 数据库会话
provider: Provider 对象
model_name: 模型名称(必须是 GlobalModel.name
is_stream: 是否是流式请求,如果为 True 则同时检查流式支持
capability_requirements: 能力需求(可选),用于检查模型是否支持所需能力
Returns:
(is_supported, skip_reason, supported_capabilities, provider_model_names)
- is_supported: 是否支持
- skip_reason: 跳过原因
- supported_capabilities: 模型支持的能力列表
- provider_model_names: Provider 侧可用的模型名称集合(主名称 + 映射名称,按 api_format 过滤)
"""
# Avoid holding a DB connection while awaiting cache/Redis inside ModelCacheService.
self._scheduler._release_db_connection_before_await(db)
# 仅接受 GlobalModel.name不允许映射名
normalized_name = model_name.strip() if isinstance(model_name, str) else ""
if not normalized_name:
return False, "模型不存在或名称无效", None, None
global_model = await ModelCacheService.get_global_model_by_name(db, normalized_name)
if not global_model or not global_model.is_active:
return False, "模型不存在或已停用", None, None
# 找到 GlobalModel 后,检查当前 Provider 是否支持
is_supported, skip_reason, caps, provider_model_names = (
await self._check_model_support_for_global_model(
db,
provider,
global_model,
model_name,
api_format,
is_stream,
capability_requirements,
)
)
return is_supported, skip_reason, caps, provider_model_names
async def _check_model_support_for_global_model(
self,
db: Session,
provider: Provider,
global_model: GlobalModel,
model_name: str,
api_format: str | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
) -> tuple[bool, str | None, list[str] | None, set[str] | None]:
"""
检查 Provider 是否支持指定的 GlobalModel
Args:
db: 数据库会话
provider: Provider 对象
global_model: GlobalModel 对象
model_name: 用户请求的模型名称(用于错误消息)
is_stream: 是否是流式请求
capability_requirements: 能力需求
Returns:
(is_supported, skip_reason, supported_capabilities, provider_model_names)
"""
# 确保 global_model 附加到当前 Session
# 注意:从缓存重建的对象是 transient 状态,不能使用 load=False
# 使用 load=True默认允许 SQLAlchemy 正确处理 transient 对象
from sqlalchemy import inspect
insp = inspect(global_model)
if insp.transient or insp.detached:
# transient/detached 对象:使用默认 merge会查询 DB 检查是否存在)
global_model = db.merge(global_model)
else:
# persistent 对象:已经附加到 session无需 merge
pass
# 获取模型支持的能力列表
model_supported_capabilities: list[str] = list(global_model.supported_capabilities or [])
# 查询该 Provider 是否有实现这个 GlobalModel
for model in provider.models:
if model.global_model_id == global_model.id and model.is_active:
# 检查流式支持
if is_stream:
supports_streaming = model.get_effective_supports_streaming()
if not supports_streaming:
return False, f"模型 {model_name} 在此 Provider 不支持流式", None, None
# 检查模型是否支持所需的能力(在 Provider 级别检查,而不是 Key 级别)
# 只有当 model_supported_capabilities 非空时才进行检查
# 空列表意味着模型没有配置能力限制,默认支持所有能力
if capability_requirements and model_supported_capabilities:
for cap_name, is_required in capability_requirements.items():
if is_required and cap_name not in model_supported_capabilities:
return (
False,
f"模型 {model_name} 不支持能力: {cap_name}",
list(model_supported_capabilities),
None,
)
provider_model_names: set[str] = {model.provider_model_name}
raw_mappings = model.provider_model_mappings
if isinstance(raw_mappings, list):
for raw in raw_mappings:
if not isinstance(raw, dict):
continue
name = raw.get("name")
if not isinstance(name, str) or not name.strip():
continue
mapping_api_formats = raw.get("api_formats")
if api_format and mapping_api_formats:
# 新模式endpoint signaturefamily:kind按小写 canonical 比较
if isinstance(mapping_api_formats, list):
target = str(api_format).strip().lower()
allowed = {
str(fmt).strip().lower() for fmt in mapping_api_formats if fmt
}
if target not in allowed:
continue
provider_model_names.add(name.strip())
return True, None, list(model_supported_capabilities), provider_model_names
return False, "Provider 未实现此模型", None, None
def _check_key_availability(
self,
key: ProviderAPIKey,
api_format: str | None,
model_name: str,
capability_requirements: dict[str, bool] | None = None,
model_mappings: list[str] | None = None,
candidate_models: set[str] | None = None,
*,
provider_type: str | None = None,
) -> tuple[bool, str | None, str | None]:
"""
检查 API Key 的可用性
注意:模型能力检查已移到 _check_model_support 中进行Provider 级别),
这里只检查 Key 级别的能力匹配。
Args:
key: API Key 对象
model_name: 模型名称GlobalModel.name
capability_requirements: 能力需求(可选)
model_mappings: GlobalModel 的映射列表(用于通配符匹配)
candidate_models: Provider 侧可用的模型名称集合(用于限制映射匹配范围)
Returns:
(is_available, skip_reason, mapping_matched_model)
- is_available: Key 是否可用
- skip_reason: 不可用时的原因
- mapping_matched_model: 通过映射匹配到的模型名(用于实际请求)
"""
# 检查熔断器状态(使用详细状态方法获取更丰富的跳过原因,按 API 格式)
is_available, circuit_reason = health_monitor.get_circuit_breaker_status(
key, api_format=api_format
)
if not is_available:
return False, circuit_reason or "熔断器已打开", None
# 模型权限检查:使用 allowed_models 白名单
# None = 允许所有模型,[] = 拒绝所有模型,["a","b"] = 只允许指定模型
# 支持通配符映射匹配(通过 model_mappings
try:
is_allowed, mapping_matched_model = check_model_allowed_with_mappings(
model_name=model_name,
allowed_models=key.allowed_models,
model_mappings=model_mappings,
candidate_models=candidate_models,
)
if mapping_matched_model:
logger.debug(
"[Scheduler] Key {}... 模型名匹配: model={} -> {}, allowed_models={}",
key.id[:8],
model_name,
mapping_matched_model,
key.allowed_models,
)
except TimeoutError:
# 正则匹配超时(可能是 ReDoS 攻击或复杂模式)
logger.warning("映射匹配超时: key_id={}, model={}", key.id, model_name)
return False, "映射匹配超时,请简化配置", None
except re.error as e:
# 正则语法错误(配置问题)
logger.warning("映射规则无效: key_id={}, model={}, error={}", key.id, model_name, e)
return False, f"映射规则无效: {str(e)}", None
except Exception as e:
# 其他未知异常
logger.error(
"映射匹配异常: key_id={}, model={}, error={}", key.id, model_name, e, exc_info=True
)
# 异常时保守处理:不允许使用该 Key
return False, "映射匹配失败", None
if not is_allowed:
return (
False,
f"Key 不支持 {model_name}",
None,
)
# Key 级别的能力匹配检查
# 注意:模型级别的能力检查已在 _check_model_support 中完成
# 始终执行检查,即使 capability_requirements 为空
# 因为 check_capability_match 会检查 Key 的 EXCLUSIVE 能力是否被浪费
key_caps: dict[str, bool] = dict(key.capabilities or {})
is_match, skip_reason = check_capability_match(key_caps, capability_requirements)
if not is_match:
return False, skip_reason, None
effective_model_name = mapping_matched_model or model_name
quota_exhausted, quota_reason = is_key_quota_exhausted(
provider_type,
key,
model_name=effective_model_name,
)
if quota_exhausted:
return False, quota_reason, mapping_matched_model
return True, None, mapping_matched_model
async def _build_candidates(
self,
db: Session,
providers: list[Provider],
client_format: str,
model_name: str,
affinity_key: str | None,
model_mappings: list[str] | None = None,
max_candidates: int | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
global_conversion_enabled: bool = True,
) -> "list[ProviderCandidate]":
"""
构建候选列表
Key 直属 Provider通过 api_formats 筛选符合端点格式的 Key。
Args:
db: 数据库会话
providers: Provider 列表
client_format: 客户端请求的 API 格式
model_name: 模型名称GlobalModel.name
affinity_key: 亲和性标识符通常为API Key ID
model_mappings: GlobalModel 的映射列表(用于 Key.allowed_models 通配符匹配)
max_candidates: 最大候选数
is_stream: 是否是流式请求,如果为 True 则过滤不支持流式的 Provider
capability_requirements: 能力需求(可选)
global_conversion_enabled: 格式转换总开关(数据库配置),关闭时禁止任何跨格式转换
Returns:
候选列表
"""
from src.services.cache.aware_scheduler import ProviderCandidate
candidates: list[ProviderCandidate] = []
client_format_str = normalize_endpoint_signature(client_format)
client_sig = parse_signature_key(client_format_str)
client_family, client_kind = client_sig.api_family, client_sig.endpoint_kind
# chat/cli 互相可回退用于同协议族下的端点变体video/image 等不跨类回退
if client_kind in {EndpointKind.CHAT, EndpointKind.CLI}:
allowed_kinds = {EndpointKind.CHAT, EndpointKind.CLI}
else:
allowed_kinds = {client_kind}
for provider in providers:
logger.debug(
"[Scheduler] Checking provider: {}, endpoints={}",
provider.name,
len(provider.endpoints) if provider.endpoints else 0,
)
# 按端点格式分别判断兼容性与模型/Key 可用性:
# - 同格式端点优先needs_conversion=False
# - 跨格式端点次之needs_conversion=True
model_support_cache: dict[
str, tuple[bool, str | None, list[str] | None, set[str] | None]
] = {}
exact_candidates: list[ProviderCandidate] = []
convertible_candidates: list[ProviderCandidate] = []
# 使用新架构字段 (api_family, endpoint_kind) 进行预过滤与排序:
# - family/kind 匹配的 endpoint 排在前面(但不做硬过滤,避免破坏格式转换路径)
# - chat/cli 请求允许互相回退(优先同 kind
# - video 等请求只允许同 kind
endpoints = list(provider.endpoints or [])
allowed_kind_values = {k.value for k in allowed_kinds}
preferred: list[ProviderEndpoint] = []
preferred_other_family: list[ProviderEndpoint] = []
fallback: list[ProviderEndpoint] = []
fallback_other_family: list[ProviderEndpoint] = []
for ep in endpoints:
if not getattr(ep, "is_active", False):
continue
raw_family = getattr(ep, "api_family", None)
raw_kind = getattr(ep, "endpoint_kind", None)
if not isinstance(raw_family, str) or not raw_family.strip():
continue
if not isinstance(raw_kind, str) or not raw_kind.strip():
continue
ep_family = raw_family.strip().lower()
ep_kind = raw_kind.strip().lower()
if allowed_kind_values and ep_kind not in allowed_kind_values:
continue
same_family = ep_family == client_family.value
same_kind = ep_kind == client_kind.value
if same_kind and same_family:
preferred.append(ep)
elif same_kind:
preferred_other_family.append(ep)
elif same_family:
fallback.append(ep)
else:
fallback_other_family.append(ep)
endpoints = (
_sort_endpoints_by_family_priority(preferred)
+ _sort_endpoints_by_family_priority(preferred_other_family)
+ _sort_endpoints_by_family_priority(fallback)
+ _sort_endpoints_by_family_priority(fallback_other_family)
)
for endpoint in endpoints:
logger.debug(
"[Scheduler] Checking endpoint: family={}, kind={}, is_active={}, base_url={}",
getattr(endpoint, "api_family", None),
getattr(endpoint, "endpoint_kind", None),
getattr(endpoint, "is_active", None),
(endpoint.base_url[:50] if endpoint.base_url else "N/A"),
)
if not endpoint.is_active:
logger.debug("[Scheduler] Endpoint skipped: not active")
continue
endpoint_format_str = make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
# 计算格式转换开关状态(三层优先级)
#
# 1) 全局开关(数据库配置)关闭 -> 禁止任何跨格式转换
# 2) 全局开关开启 -> 允许跨格式转换
# 3) 提供商覆盖Provider.enable_format_conversion开启 -> 强制允许(跳过端点检查)
# 4) 否则 -> 由端点配置 format_acceptance_config 决定是否允许
provider_allows_conversion = getattr(provider, "enable_format_conversion", True)
skip_endpoint_check = global_conversion_enabled or provider_allows_conversion
is_compatible, needs_conversion, _compat_reason = is_format_compatible(
client_format_str,
endpoint_format_str,
getattr(endpoint, "format_acceptance_config", None),
is_stream,
global_conversion_enabled,
skip_endpoint_check=skip_endpoint_check,
)
logger.debug(
"[Scheduler] Format compatibility: client={}, endpoint={}, compatible={}, "
"global={}, provider={}, skip_endpoint={}, reason={}",
client_format_str,
endpoint_format_str,
is_compatible,
global_conversion_enabled,
provider_allows_conversion,
skip_endpoint_check,
_compat_reason,
)
if not is_compatible:
continue
# 检查模型支持(按端点格式过滤 provider_model_mappings
if endpoint_format_str not in model_support_cache:
model_support_cache[endpoint_format_str] = await self._check_model_support(
db,
provider,
model_name,
api_format=endpoint_format_str,
is_stream=is_stream,
capability_requirements=capability_requirements,
)
supports_model, skip_reason, _model_caps, provider_model_names = (
model_support_cache[endpoint_format_str]
)
logger.debug(
"[Scheduler] Model support: provider={}, model={}, supports={}, reason={}",
provider.name,
model_name,
supports_model,
skip_reason,
)
if not supports_model:
logger.debug(
f"Provider {provider.name} 端点 {endpoint_format_str} "
f"不支持模型 {model_name}: {skip_reason}"
)
continue
# Key 直属 Provider通过 api_formats 按端点格式筛选
# api_formats=None 视为"全支持"(兼容历史数据)
active_keys = [
key
for key in provider.api_keys
if key.is_active
and (key.api_formats is None or endpoint_format_str in key.api_formats)
]
if not active_keys:
continue
# 检查是否所有 Key 都是 TTL=0轮换模式
use_random = all((key.cache_ttl_minutes or 0) == 0 for key in active_keys)
if use_random and len(active_keys) > 1:
logger.debug(
f" Provider {provider.name} 启用 Key 轮换模式 "
f"(endpoint_format={endpoint_format_str}, {len(active_keys)} keys)"
)
keys = self._scheduler._candidate_sorter._shuffle_keys_by_internal_priority(
active_keys, affinity_key, use_random
)
for key in keys:
# Key 级别检查(健康度/熔断按 provider_format bucket
# 传入 provider_model_names 作为 candidate_models
# 用于检查 Key 的 allowed_models 是否支持 Provider 定义的模型名称
is_available, key_skip_reason, mapping_matched_model = (
self._check_key_availability(
key,
endpoint_format_str,
model_name,
capability_requirements,
model_mappings=model_mappings,
candidate_models=provider_model_names,
provider_type=getattr(provider, "provider_type", None),
)
)
candidate = ProviderCandidate(
provider=provider,
endpoint=endpoint,
key=key,
is_skipped=not is_available,
skip_reason=key_skip_reason,
mapping_matched_model=mapping_matched_model,
needs_conversion=needs_conversion,
provider_api_format=str(endpoint_format_str or ""),
)
if needs_conversion:
convertible_candidates.append(candidate)
else:
exact_candidates.append(candidate)
candidates.extend(exact_candidates)
candidates.extend(convertible_candidates)
# max_candidates 截断应在所有候选收集完成后统一处理,确保优先级排序正确
if max_candidates and len(candidates) > max_candidates:
candidates = candidates[:max_candidates]
return candidates

268
src/services/cache/_candidate_sorter.py vendored Normal file
View File

@@ -0,0 +1,268 @@
"""
候选排序器 (CandidateSorter)
从 CacheAwareScheduler 拆分出的候选排序逻辑,负责:
- 优先级模式排序provider / global_key
- 负载均衡模式排序
- Key 内部按优先级分组打乱
"""
from __future__ import annotations
import random
from collections import defaultdict
from typing import TYPE_CHECKING
from src.core.logger import logger
from src.services.system.config import SystemConfigService
if TYPE_CHECKING:
from sqlalchemy.orm import Session
from src.models.database import ProviderAPIKey
from src.services.cache.aware_scheduler import CacheAwareScheduler, ProviderCandidate
class CandidateSorter:
"""候选排序器,负责优先级模式排序、负载均衡排序和 Key 内部打乱。"""
def __init__(self, scheduler: CacheAwareScheduler) -> None:
self._scheduler = scheduler
def _apply_priority_mode_sort(
self,
candidates: list[ProviderCandidate],
db: Session,
affinity_key: str | None = None,
api_format: str | None = None,
) -> list[ProviderCandidate]:
"""
根据优先级模式对候选列表排序(数字越小越优先)
排序规则(受 keep_priority_on_conversion 配置影响):
1. 如果全局配置 keep_priority_on_conversion=True所有候选保持原优先级
2. 否则,按 needs_conversion 和 provider.keep_priority_on_conversion 分组:
- 保持优先级的候选exact 或 provider.keep_priority_on_conversion=True按原优先级排序
- 需要降级的候选convertible 且 provider.keep_priority_on_conversion=False整体排在后面
3. 在同一组内,按优先级模式排序:
- provider: 按 Provider.provider_priority -> Key.internal_priority 排序
- global_key: 按 Key.global_priority_by_format 排序
"""
if not candidates:
return candidates
s = self._scheduler
# 全局配置:如果开启,所有候选保持原优先级
global_keep_priority = SystemConfigService.is_keep_priority_on_conversion(db)
if global_keep_priority:
# 全局开启:不分组,直接按优先级模式排序
if s.priority_mode == s.PRIORITY_MODE_GLOBAL_KEY:
return self._sort_by_global_priority_with_hash(candidates, affinity_key, api_format)
# 提供商优先模式:保持构建时的顺序(已按 provider_priority 排序)
return candidates
# 全局未开启:按是否需要降级分组
# - 不需要降级exact 候选 或 provider.keep_priority_on_conversion=True 的 convertible 候选
# - 需要降级convertible 且 provider.keep_priority_on_conversion=False
keep_priority_candidates: list[ProviderCandidate] = []
demote_candidates: list[ProviderCandidate] = []
for c in candidates:
if not c.needs_conversion:
# exact 候选:不需要降级
keep_priority_candidates.append(c)
elif getattr(c.provider, "keep_priority_on_conversion", False):
# convertible 但提供商配置了保持优先级
keep_priority_candidates.append(c)
else:
# convertible 且未配置保持优先级:降级
demote_candidates.append(c)
if s.priority_mode == s.PRIORITY_MODE_GLOBAL_KEY:
# 全局 Key 优先模式:分别对两组排序后合并
sorted_keep = self._sort_by_global_priority_with_hash(
keep_priority_candidates, affinity_key, api_format
)
sorted_demote = self._sort_by_global_priority_with_hash(
demote_candidates, affinity_key, api_format
)
return sorted_keep + sorted_demote
# 提供商优先模式:保持优先级的在前,降级的在后(各组内部顺序已由构建时保证)
return keep_priority_candidates + demote_candidates
def _sort_by_global_priority_with_hash(
self,
candidates: list[ProviderCandidate],
affinity_key: str | None = None,
api_format: str | None = None,
) -> list[ProviderCandidate]:
"""
按 global_priority_by_format 分组排序,同优先级内通过哈希分散实现负载均衡
排序逻辑:
1. 按 global_priority_by_format[api_format] 分组数字小的优先NULL 排后面)
2. 同优先级组内,使用 affinity_key 哈希分散
3. 确保同一用户请求稳定选择同一个 Key缓存亲和性
"""
def get_priority(candidate: ProviderCandidate) -> int:
"""获取候选的优先级"""
if not candidate.key:
return 999999
priority_by_format = candidate.key.global_priority_by_format or {}
if api_format and api_format in priority_by_format:
return priority_by_format[api_format]
return 999999 # NULL 排在后面
# 按优先级分组
priority_groups: dict[int, list[ProviderCandidate]] = defaultdict(list)
for candidate in candidates:
priority = get_priority(candidate)
priority_groups[priority].append(candidate)
result = []
for priority in sorted(priority_groups.keys()): # 数字小的优先级高
group = priority_groups[priority]
if len(group) > 1 and affinity_key:
# 同优先级内哈希分散负载均衡
scored_candidates = []
for candidate in group:
key_id = candidate.key.id if candidate.key else ""
hash_value = self._scheduler._affinity_hash(affinity_key, key_id)
scored_candidates.append((hash_value, candidate))
# 按哈希值排序
sorted_group = [c for _, c in sorted(scored_candidates, key=lambda x: x[0])]
result.extend(sorted_group)
else:
# 单个候选或没有 affinity_key按次要排序条件排序
def secondary_sort(c: ProviderCandidate) -> tuple[int, int, str]:
pp = c.provider.provider_priority
ip = c.key.internal_priority if c.key else None
return (
pp if pp is not None else 999999,
ip if ip is not None else 999999,
c.key.id if c.key else "",
)
result.extend(sorted(group, key=secondary_sort))
return result
def _apply_load_balance(
self, candidates: list[ProviderCandidate], api_format: str | None = None
) -> list[ProviderCandidate]:
"""
负载均衡模式:同优先级内随机轮换
排序逻辑:
1. 按优先级分组provider_priority, internal_priority 或 global_priority_by_format
2. 同优先级组内随机打乱
3. 不考虑缓存亲和性
"""
if not candidates:
return candidates
s = self._scheduler
priority_groups: dict[tuple, list[ProviderCandidate]] = defaultdict(list)
# 根据优先级模式选择分组方式
if s.priority_mode == s.PRIORITY_MODE_GLOBAL_KEY:
# 全局 Key 优先模式:按格式特定优先级分组
for candidate in candidates:
priority = 999999
if candidate.key:
priority_by_format = candidate.key.global_priority_by_format or {}
if api_format and api_format in priority_by_format:
priority = priority_by_format[api_format]
priority_groups[(priority,)].append(candidate)
else:
# 提供商优先模式:按 (provider_priority, internal_priority) 分组
for candidate in candidates:
pp = candidate.provider.provider_priority
ip = candidate.key.internal_priority if candidate.key else None
key = (
pp if pp is not None else 999999,
ip if ip is not None else 999999,
)
priority_groups[key].append(candidate)
result: list[ProviderCandidate] = []
for priority in sorted(priority_groups.keys()):
group = priority_groups[priority]
if len(group) > 1:
# 同优先级内随机打乱
shuffled = list(group)
random.shuffle(shuffled)
result.extend(shuffled)
else:
result.extend(group)
return result
def _shuffle_keys_by_internal_priority(
self,
keys: list[ProviderAPIKey],
affinity_key: str | None = None,
use_random: bool = False,
) -> list[ProviderAPIKey]:
"""
对 API Key 按 internal_priority 分组,同优先级内部基于 affinity_key 进行确定性打乱
目的:
- 数字越小越优先使用
- 同优先级 Key 之间实现负载均衡
- 使用 affinity_key 哈希确保同一请求 Key 的请求稳定(避免破坏缓存亲和性)
- 当 use_random=True 时,使用随机排序实现轮换(用于 TTL=0 的场景)
Args:
keys: API Key 列表
affinity_key: 亲和性标识符(通常为 API Key ID用于确定性打乱
use_random: 是否使用随机排序TTL=0 时为 True
Returns:
排序后的 Key 列表
"""
if not keys:
return []
# 按 internal_priority 分组
priority_groups: dict[int, list[ProviderAPIKey]] = defaultdict(list)
for key in keys:
priority = key.internal_priority if key.internal_priority is not None else 999999
priority_groups[priority].append(key)
# 对每个优先级组内的 Key 进行打乱
result = []
for priority in sorted(priority_groups.keys()): # 数字小的优先级高,排前面
group_keys = priority_groups[priority]
if len(group_keys) > 1:
if use_random:
# TTL=0 模式:使用随机排序实现 Key 轮换
shuffled = list(group_keys)
random.shuffle(shuffled)
result.extend(shuffled)
elif affinity_key:
# 正常模式:使用哈希确定性打乱(保持缓存亲和性)
key_scores = []
for key in group_keys:
hash_value = self._scheduler._affinity_hash(affinity_key, key.id)
key_scores.append((hash_value, key))
# 按哈希值排序
sorted_group = [key for _, key in sorted(key_scores, key=lambda x: x[0])]
result.extend(sorted_group)
else:
# 没有 affinity_key 时按 ID 排序保持稳定性
result.extend(sorted(group_keys, key=lambda k: k.id))
else:
# 单个 Key 直接添加
result.extend(group_keys)
return result

View File

@@ -32,46 +32,37 @@ from __future__ import annotations
import hashlib
import math
import random
import re
import time
from collections import defaultdict
from collections.abc import Sequence
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
from typing import Any
from sqlalchemy.orm import Session, selectinload
from sqlalchemy.orm import Session
from src.core.api_format.conversion.compatibility import is_format_compatible
from src.core.api_format.enums import ApiFamily, EndpointKind
from src.core.api_format.signature import make_signature_key, parse_signature_key
from src.core.exceptions import ModelNotSupportedException, ProviderNotAvailableException
from src.core.key_capabilities import check_capability_match
from src.core.logger import logger
from src.core.model_permissions import (
check_model_allowed,
check_model_allowed_with_mappings,
get_allowed_models_preview,
merge_allowed_models,
)
from src.models.database import (
ApiKey,
Model,
Provider,
ProviderAPIKey,
ProviderEndpoint,
)
from src.services.cache.quota_skipper import is_key_quota_exhausted
if TYPE_CHECKING:
from src.models.database import GlobalModel
from src.services.cache._candidate_builder import (
CandidateBuilder,
)
from src.services.cache._candidate_builder import (
_sort_endpoints_by_family_priority as _sort_endpoints_by_family_priority,
)
from src.services.cache._candidate_sorter import CandidateSorter
from src.services.cache.affinity_manager import (
CacheAffinityManager,
get_affinity_manager,
)
from src.services.cache.model_cache import ModelCacheService
from src.services.health.monitor import health_monitor
from src.services.provider.format import normalize_endpoint_signature
from src.services.rate_limit.adaptive_reservation import (
AdaptiveReservationManager,
@@ -154,21 +145,6 @@ class ConcurrencySnapshot:
)
def _sort_endpoints_by_family_priority(
eps: Sequence[ProviderEndpoint],
) -> list[ProviderEndpoint]:
"""按 ApiFamily 优先级对端点排序(同分组内使用)。"""
def sort_key(ep: ProviderEndpoint) -> int:
family_str = str(getattr(ep, "api_family", "") or "").strip().lower()
try:
return ApiFamily(family_str).priority
except ValueError:
return 99
return sorted(eps, key=sort_key)
class CacheAwareScheduler:
"""
缓存感知调度器
@@ -248,6 +224,10 @@ class CacheAwareScheduler:
"last_reservation_result": None,
}
# 初始化拆分出的子模块
self._candidate_builder = CandidateBuilder(self)
self._candidate_sorter = CandidateSorter(self)
@staticmethod
def _release_db_connection_before_await(db: Session) -> None:
"""
@@ -436,7 +416,7 @@ class CacheAwareScheduler:
- 总槽位: 有效 RPM 限制(固定值或学习到的值)
- 预留比例: 由 AdaptiveReservationManager 根据置信度和负载动态计算
- 缓存用户可用: 全部槽位
- 新用户可用: 总槽位 × (1 - 动态预留比例)
- 新用户可用: 总槽位 x (1 - 动态预留比例)
Args:
key: ProviderAPIKey对象
@@ -625,8 +605,8 @@ class CacheAwareScheduler:
预先获取所有可用的 Provider/Endpoint/Key 组合
重构后的方法将逻辑拆分为:
1. _query_providers: 数据库查询逻辑
2. _build_candidates: 候选构建逻辑
1. _query_providers: 数据库查询逻辑(委托给 CandidateBuilder
2. _build_candidates: 候选构建逻辑(委托给 CandidateBuilder
3. _apply_cache_affinity: 缓存亲和性处理
Args:
@@ -714,8 +694,8 @@ class CacheAwareScheduler:
)
return [], global_model_id, queried_provider_count
# 1. 查询 Providers
providers = self._query_providers(
# 1. 查询 Providers(委托给 CandidateBuilder
providers = self._candidate_builder._query_providers(
db=db,
provider_offset=provider_offset,
provider_limit=provider_limit,
@@ -755,12 +735,12 @@ class CacheAwareScheduler:
if not providers:
return [], global_model_id, queried_provider_count
# 2. 构建候选列表(传入 is_stream 和 capability_requirements 用于过滤
# 2. 构建候选列表(委托给 CandidateBuilder
# 格式转换总开关(数据库配置):关闭时禁止任何跨格式候选进入队列
global_conversion_enabled = SystemConfigService.is_format_conversion_enabled(db)
candidates = await self._build_candidates(
candidates = await self._candidate_builder._build_candidates(
db=db,
providers=providers,
client_format=target_format,
@@ -818,8 +798,10 @@ class CacheAwareScheduler:
if not candidates:
return candidates
# 1. 优先级模式排序
candidates = self._apply_priority_mode_sort(candidates, db, affinity_key, api_format)
# 1. 优先级模式排序(委托给 CandidateSorter
candidates = self._candidate_sorter._apply_priority_mode_sort(
candidates, db, affinity_key, api_format
)
# 2. 调度模式排序
if self.scheduling_mode == self.SCHEDULING_MODE_CACHE_AFFINITY:
@@ -832,7 +814,7 @@ class CacheAwareScheduler:
global_model_id=global_model_id,
)
elif self.scheduling_mode == self.SCHEDULING_MODE_LOAD_BALANCE:
candidates = self._apply_load_balance(candidates, api_format)
candidates = self._candidate_sorter._apply_load_balance(candidates, api_format)
for candidate in candidates:
candidate.is_cached = False
else:
@@ -841,532 +823,6 @@ class CacheAwareScheduler:
return candidates
def _query_providers(
self,
db: Session,
provider_offset: int = 0,
provider_limit: int | None = None,
) -> list[Provider]:
"""
查询活跃的 Providers带预加载
Args:
db: 数据库会话
provider_offset: 分页偏移
provider_limit: 分页限制
Returns:
Provider 列表
"""
provider_query = (
db.query(Provider)
.options(
# 预加载 Provider 级别的 api_keys
selectinload(Provider.api_keys),
# 预加载 endpoints用于按 api_format 选择请求配置)
selectinload(Provider.endpoints),
# 同时加载 models 和 global_model 关系
selectinload(Provider.models).selectinload(Model.global_model),
)
.filter(Provider.is_active == True)
.order_by(Provider.provider_priority.asc())
)
if provider_offset:
provider_query = provider_query.offset(provider_offset)
if provider_limit:
provider_query = provider_query.limit(provider_limit)
return provider_query.all()
async def _check_model_support(
self,
db: Session,
provider: Provider,
model_name: str,
api_format: str | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
) -> tuple[bool, str | None, list[str] | None, set[str] | None]:
"""
检查 Provider 是否支持指定模型(可选检查流式支持和能力需求)
模型能力检查在这里进行(而不是在 Key 级别),因为:
- 模型支持的能力是全局的,与具体的 Key 无关
- 如果模型不支持某能力,整个 Provider 的所有 Key 都应该被跳过
仅支持直接匹配 GlobalModel.name外部请求不接受映射名
Args:
db: 数据库会话
provider: Provider 对象
model_name: 模型名称(必须是 GlobalModel.name
is_stream: 是否是流式请求,如果为 True 则同时检查流式支持
capability_requirements: 能力需求(可选),用于检查模型是否支持所需能力
Returns:
(is_supported, skip_reason, supported_capabilities, provider_model_names)
- is_supported: 是否支持
- skip_reason: 跳过原因
- supported_capabilities: 模型支持的能力列表
- provider_model_names: Provider 侧可用的模型名称集合(主名称 + 映射名称,按 api_format 过滤)
"""
# Avoid holding a DB connection while awaiting cache/Redis inside ModelCacheService.
self._release_db_connection_before_await(db)
# 仅接受 GlobalModel.name不允许映射名
normalized_name = model_name.strip() if isinstance(model_name, str) else ""
if not normalized_name:
return False, "模型不存在或名称无效", None, None
global_model = await ModelCacheService.get_global_model_by_name(db, normalized_name)
if not global_model or not global_model.is_active:
return False, "模型不存在或已停用", None, None
# 找到 GlobalModel 后,检查当前 Provider 是否支持
is_supported, skip_reason, caps, provider_model_names = (
await self._check_model_support_for_global_model(
db,
provider,
global_model,
model_name,
api_format,
is_stream,
capability_requirements,
)
)
return is_supported, skip_reason, caps, provider_model_names
async def _check_model_support_for_global_model(
self,
db: Session,
provider: Provider,
global_model: GlobalModel,
model_name: str,
api_format: str | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
) -> tuple[bool, str | None, list[str] | None, set[str] | None]:
"""
检查 Provider 是否支持指定的 GlobalModel
Args:
db: 数据库会话
provider: Provider 对象
global_model: GlobalModel 对象
model_name: 用户请求的模型名称(用于错误消息)
is_stream: 是否是流式请求
capability_requirements: 能力需求
Returns:
(is_supported, skip_reason, supported_capabilities, provider_model_names)
"""
# 确保 global_model 附加到当前 Session
# 注意:从缓存重建的对象是 transient 状态,不能使用 load=False
# 使用 load=True默认允许 SQLAlchemy 正确处理 transient 对象
from sqlalchemy import inspect
insp = inspect(global_model)
if insp.transient or insp.detached:
# transient/detached 对象:使用默认 merge会查询 DB 检查是否存在)
global_model = db.merge(global_model)
else:
# persistent 对象:已经附加到 session无需 merge
pass
# 获取模型支持的能力列表
model_supported_capabilities: list[str] = list(global_model.supported_capabilities or [])
# 查询该 Provider 是否有实现这个 GlobalModel
for model in provider.models:
if model.global_model_id == global_model.id and model.is_active:
# 检查流式支持
if is_stream:
supports_streaming = model.get_effective_supports_streaming()
if not supports_streaming:
return False, f"模型 {model_name} 在此 Provider 不支持流式", None, None
# 检查模型是否支持所需的能力(在 Provider 级别检查,而不是 Key 级别)
# 只有当 model_supported_capabilities 非空时才进行检查
# 空列表意味着模型没有配置能力限制,默认支持所有能力
if capability_requirements and model_supported_capabilities:
for cap_name, is_required in capability_requirements.items():
if is_required and cap_name not in model_supported_capabilities:
return (
False,
f"模型 {model_name} 不支持能力: {cap_name}",
list(model_supported_capabilities),
None,
)
provider_model_names: set[str] = {model.provider_model_name}
raw_mappings = model.provider_model_mappings
if isinstance(raw_mappings, list):
for raw in raw_mappings:
if not isinstance(raw, dict):
continue
name = raw.get("name")
if not isinstance(name, str) or not name.strip():
continue
mapping_api_formats = raw.get("api_formats")
if api_format and mapping_api_formats:
# 新模式endpoint signaturefamily:kind按小写 canonical 比较
if isinstance(mapping_api_formats, list):
target = str(api_format).strip().lower()
allowed = {
str(fmt).strip().lower() for fmt in mapping_api_formats if fmt
}
if target not in allowed:
continue
provider_model_names.add(name.strip())
return True, None, list(model_supported_capabilities), provider_model_names
return False, "Provider 未实现此模型", None, None
def _check_key_availability(
self,
key: ProviderAPIKey,
api_format: str | None,
model_name: str,
capability_requirements: dict[str, bool] | None = None,
model_mappings: list[str] | None = None,
candidate_models: set[str] | None = None,
*,
provider_type: str | None = None,
) -> tuple[bool, str | None, str | None]:
"""
检查 API Key 的可用性
注意:模型能力检查已移到 _check_model_support 中进行Provider 级别),
这里只检查 Key 级别的能力匹配。
Args:
key: API Key 对象
model_name: 模型名称GlobalModel.name
capability_requirements: 能力需求(可选)
model_mappings: GlobalModel 的映射列表(用于通配符匹配)
candidate_models: Provider 侧可用的模型名称集合(用于限制映射匹配范围)
Returns:
(is_available, skip_reason, mapping_matched_model)
- is_available: Key 是否可用
- skip_reason: 不可用时的原因
- mapping_matched_model: 通过映射匹配到的模型名(用于实际请求)
"""
# 检查熔断器状态(使用详细状态方法获取更丰富的跳过原因,按 API 格式)
is_available, circuit_reason = health_monitor.get_circuit_breaker_status(
key, api_format=api_format
)
if not is_available:
return False, circuit_reason or "熔断器已打开", None
# 模型权限检查:使用 allowed_models 白名单
# None = 允许所有模型,[] = 拒绝所有模型,["a","b"] = 只允许指定模型
# 支持通配符映射匹配(通过 model_mappings
try:
is_allowed, mapping_matched_model = check_model_allowed_with_mappings(
model_name=model_name,
allowed_models=key.allowed_models,
model_mappings=model_mappings,
candidate_models=candidate_models,
)
if mapping_matched_model:
logger.debug(
"[Scheduler] Key {}... 模型名匹配: model={} -> {}, allowed_models={}",
key.id[:8],
model_name,
mapping_matched_model,
key.allowed_models,
)
except TimeoutError:
# 正则匹配超时(可能是 ReDoS 攻击或复杂模式)
logger.warning("映射匹配超时: key_id={}, model={}", key.id, model_name)
return False, "映射匹配超时,请简化配置", None
except re.error as e:
# 正则语法错误(配置问题)
logger.warning("映射规则无效: key_id={}, model={}, error={}", key.id, model_name, e)
return False, f"映射规则无效: {str(e)}", None
except Exception as e:
# 其他未知异常
logger.error(
"映射匹配异常: key_id={}, model={}, error={}", key.id, model_name, e, exc_info=True
)
# 异常时保守处理:不允许使用该 Key
return False, "映射匹配失败", None
if not is_allowed:
return (
False,
f"Key 不支持 {model_name}",
None,
)
# Key 级别的能力匹配检查
# 注意:模型级别的能力检查已在 _check_model_support 中完成
# 始终执行检查,即使 capability_requirements 为空
# 因为 check_capability_match 会检查 Key 的 EXCLUSIVE 能力是否被浪费
key_caps: dict[str, bool] = dict(key.capabilities or {})
is_match, skip_reason = check_capability_match(key_caps, capability_requirements)
if not is_match:
return False, skip_reason, None
effective_model_name = mapping_matched_model or model_name
quota_exhausted, quota_reason = is_key_quota_exhausted(
provider_type,
key,
model_name=effective_model_name,
)
if quota_exhausted:
return False, quota_reason, mapping_matched_model
return True, None, mapping_matched_model
async def _build_candidates(
self,
db: Session,
providers: list[Provider],
client_format: str,
model_name: str,
affinity_key: str | None,
model_mappings: list[str] | None = None,
max_candidates: int | None = None,
is_stream: bool = False,
capability_requirements: dict[str, bool] | None = None,
global_conversion_enabled: bool = True,
) -> list[ProviderCandidate]:
"""
构建候选列表
Key 直属 Provider通过 api_formats 筛选符合端点格式的 Key。
Args:
db: 数据库会话
providers: Provider 列表
client_format: 客户端请求的 API 格式
model_name: 模型名称GlobalModel.name
affinity_key: 亲和性标识符通常为API Key ID
model_mappings: GlobalModel 的映射列表(用于 Key.allowed_models 通配符匹配)
max_candidates: 最大候选数
is_stream: 是否是流式请求,如果为 True 则过滤不支持流式的 Provider
capability_requirements: 能力需求(可选)
global_conversion_enabled: 格式转换总开关(数据库配置),关闭时禁止任何跨格式转换
Returns:
候选列表
"""
candidates: list[ProviderCandidate] = []
client_format_str = normalize_endpoint_signature(client_format)
client_sig = parse_signature_key(client_format_str)
client_family, client_kind = client_sig.api_family, client_sig.endpoint_kind
# chat/cli 互相可回退用于同协议族下的端点变体video/image 等不跨类回退
if client_kind in {EndpointKind.CHAT, EndpointKind.CLI}:
allowed_kinds = {EndpointKind.CHAT, EndpointKind.CLI}
else:
allowed_kinds = {client_kind}
for provider in providers:
logger.debug(
"[Scheduler] Checking provider: {}, endpoints={}",
provider.name,
len(provider.endpoints) if provider.endpoints else 0,
)
# 按端点格式分别判断兼容性与模型/Key 可用性:
# - 同格式端点优先needs_conversion=False
# - 跨格式端点次之needs_conversion=True
model_support_cache: dict[
str, tuple[bool, str | None, list[str] | None, set[str] | None]
] = {}
exact_candidates: list[ProviderCandidate] = []
convertible_candidates: list[ProviderCandidate] = []
# 使用新架构字段 (api_family, endpoint_kind) 进行预过滤与排序:
# - family/kind 匹配的 endpoint 排在前面(但不做硬过滤,避免破坏格式转换路径)
# - chat/cli 请求允许互相回退(优先同 kind
# - video 等请求只允许同 kind
endpoints = list(provider.endpoints or [])
allowed_kind_values = {k.value for k in allowed_kinds}
preferred: list[ProviderEndpoint] = []
preferred_other_family: list[ProviderEndpoint] = []
fallback: list[ProviderEndpoint] = []
fallback_other_family: list[ProviderEndpoint] = []
for ep in endpoints:
if not getattr(ep, "is_active", False):
continue
raw_family = getattr(ep, "api_family", None)
raw_kind = getattr(ep, "endpoint_kind", None)
if not isinstance(raw_family, str) or not raw_family.strip():
continue
if not isinstance(raw_kind, str) or not raw_kind.strip():
continue
ep_family = raw_family.strip().lower()
ep_kind = raw_kind.strip().lower()
if allowed_kind_values and ep_kind not in allowed_kind_values:
continue
same_family = ep_family == client_family.value
same_kind = ep_kind == client_kind.value
if same_kind and same_family:
preferred.append(ep)
elif same_kind:
preferred_other_family.append(ep)
elif same_family:
fallback.append(ep)
else:
fallback_other_family.append(ep)
endpoints = (
_sort_endpoints_by_family_priority(preferred)
+ _sort_endpoints_by_family_priority(preferred_other_family)
+ _sort_endpoints_by_family_priority(fallback)
+ _sort_endpoints_by_family_priority(fallback_other_family)
)
for endpoint in endpoints:
logger.debug(
"[Scheduler] Checking endpoint: family={}, kind={}, is_active={}, base_url={}",
getattr(endpoint, "api_family", None),
getattr(endpoint, "endpoint_kind", None),
getattr(endpoint, "is_active", None),
(endpoint.base_url[:50] if endpoint.base_url else "N/A"),
)
if not endpoint.is_active:
logger.debug("[Scheduler] Endpoint skipped: not active")
continue
endpoint_format_str = make_signature_key(
str(getattr(endpoint, "api_family", "")).strip().lower(),
str(getattr(endpoint, "endpoint_kind", "")).strip().lower(),
)
# 计算格式转换开关状态(三层优先级)
#
# 1) 全局开关(数据库配置)关闭 -> 禁止任何跨格式转换
# 2) 全局开关开启 -> 允许跨格式转换
# 3) 提供商覆盖Provider.enable_format_conversion开启 -> 强制允许(跳过端点检查)
# 4) 否则 -> 由端点配置 format_acceptance_config 决定是否允许
provider_allows_conversion = getattr(provider, "enable_format_conversion", True)
skip_endpoint_check = global_conversion_enabled or provider_allows_conversion
is_compatible, needs_conversion, _compat_reason = is_format_compatible(
client_format_str,
endpoint_format_str,
getattr(endpoint, "format_acceptance_config", None),
is_stream,
global_conversion_enabled,
skip_endpoint_check=skip_endpoint_check,
)
logger.debug(
"[Scheduler] Format compatibility: client={}, endpoint={}, compatible={}, "
"global={}, provider={}, skip_endpoint={}, reason={}",
client_format_str,
endpoint_format_str,
is_compatible,
global_conversion_enabled,
provider_allows_conversion,
skip_endpoint_check,
_compat_reason,
)
if not is_compatible:
continue
# 检查模型支持(按端点格式过滤 provider_model_mappings
if endpoint_format_str not in model_support_cache:
model_support_cache[endpoint_format_str] = await self._check_model_support(
db,
provider,
model_name,
api_format=endpoint_format_str,
is_stream=is_stream,
capability_requirements=capability_requirements,
)
supports_model, skip_reason, _model_caps, provider_model_names = (
model_support_cache[endpoint_format_str]
)
logger.debug(
"[Scheduler] Model support: provider={}, model={}, supports={}, reason={}",
provider.name,
model_name,
supports_model,
skip_reason,
)
if not supports_model:
logger.debug(
f"Provider {provider.name} 端点 {endpoint_format_str} 不支持模型 {model_name}: {skip_reason}"
)
continue
# Key 直属 Provider通过 api_formats 按端点格式筛选
# api_formats=None 视为"全支持"(兼容历史数据)
active_keys = [
key
for key in provider.api_keys
if key.is_active
and (key.api_formats is None or endpoint_format_str in key.api_formats)
]
if not active_keys:
continue
# 检查是否所有 Key 都是 TTL=0轮换模式
use_random = all((key.cache_ttl_minutes or 0) == 0 for key in active_keys)
if use_random and len(active_keys) > 1:
logger.debug(
f" Provider {provider.name} 启用 Key 轮换模式 "
f"(endpoint_format={endpoint_format_str}, {len(active_keys)} keys)"
)
keys = self._shuffle_keys_by_internal_priority(
active_keys, affinity_key, use_random
)
for key in keys:
# Key 级别检查(健康度/熔断按 provider_format bucket
# 传入 provider_model_names 作为 candidate_models
# 用于检查 Key 的 allowed_models 是否支持 Provider 定义的模型名称
is_available, key_skip_reason, mapping_matched_model = (
self._check_key_availability(
key,
endpoint_format_str,
model_name,
capability_requirements,
model_mappings=model_mappings,
candidate_models=provider_model_names,
provider_type=getattr(provider, "provider_type", None),
)
)
candidate = ProviderCandidate(
provider=provider,
endpoint=endpoint,
key=key,
is_skipped=not is_available,
skip_reason=key_skip_reason,
mapping_matched_model=mapping_matched_model,
needs_conversion=needs_conversion,
provider_api_format=str(endpoint_format_str or ""),
)
if needs_conversion:
convertible_candidates.append(candidate)
else:
exact_candidates.append(candidate)
candidates.extend(exact_candidates)
candidates.extend(convertible_candidates)
# max_candidates 截断应在所有候选收集完成后统一处理,确保优先级排序正确
if max_candidates and len(candidates) > max_candidates:
candidates = candidates[:max_candidates]
return candidates
async def _apply_cache_affinity(
self,
candidates: list[ProviderCandidate],
@@ -1545,241 +1001,6 @@ class CacheAwareScheduler:
self.scheduling_mode = normalized
logger.debug(f"[CacheAwareScheduler] 切换调度模式为: {self.scheduling_mode}")
def _apply_priority_mode_sort(
self,
candidates: list[ProviderCandidate],
db: Session,
affinity_key: str | None = None,
api_format: str | None = None,
) -> list[ProviderCandidate]:
"""
根据优先级模式对候选列表排序(数字越小越优先)
排序规则(受 keep_priority_on_conversion 配置影响):
1. 如果全局配置 keep_priority_on_conversion=True所有候选保持原优先级
2. 否则,按 needs_conversion 和 provider.keep_priority_on_conversion 分组:
- 保持优先级的候选exact 或 provider.keep_priority_on_conversion=True按原优先级排序
- 需要降级的候选convertible 且 provider.keep_priority_on_conversion=False整体排在后面
3. 在同一组内,按优先级模式排序:
- provider: 按 Provider.provider_priority -> Key.internal_priority 排序
- global_key: 按 Key.global_priority_by_format 排序
"""
if not candidates:
return candidates
# 全局配置:如果开启,所有候选保持原优先级
global_keep_priority = SystemConfigService.is_keep_priority_on_conversion(db)
if global_keep_priority:
# 全局开启:不分组,直接按优先级模式排序
if self.priority_mode == self.PRIORITY_MODE_GLOBAL_KEY:
return self._sort_by_global_priority_with_hash(candidates, affinity_key, api_format)
# 提供商优先模式:保持构建时的顺序(已按 provider_priority 排序)
return candidates
# 全局未开启:按是否需要降级分组
# - 不需要降级exact 候选 或 provider.keep_priority_on_conversion=True 的 convertible 候选
# - 需要降级convertible 且 provider.keep_priority_on_conversion=False
keep_priority_candidates: list[ProviderCandidate] = []
demote_candidates: list[ProviderCandidate] = []
for c in candidates:
if not c.needs_conversion:
# exact 候选:不需要降级
keep_priority_candidates.append(c)
elif getattr(c.provider, "keep_priority_on_conversion", False):
# convertible 但提供商配置了保持优先级
keep_priority_candidates.append(c)
else:
# convertible 且未配置保持优先级:降级
demote_candidates.append(c)
if self.priority_mode == self.PRIORITY_MODE_GLOBAL_KEY:
# 全局 Key 优先模式:分别对两组排序后合并
sorted_keep = self._sort_by_global_priority_with_hash(
keep_priority_candidates, affinity_key, api_format
)
sorted_demote = self._sort_by_global_priority_with_hash(
demote_candidates, affinity_key, api_format
)
return sorted_keep + sorted_demote
# 提供商优先模式:保持优先级的在前,降级的在后(各组内部顺序已由构建时保证)
return keep_priority_candidates + demote_candidates
def _sort_by_global_priority_with_hash(
self,
candidates: list[ProviderCandidate],
affinity_key: str | None = None,
api_format: str | None = None,
) -> list[ProviderCandidate]:
"""
按 global_priority_by_format 分组排序,同优先级内通过哈希分散实现负载均衡
排序逻辑:
1. 按 global_priority_by_format[api_format] 分组数字小的优先NULL 排后面)
2. 同优先级组内,使用 affinity_key 哈希分散
3. 确保同一用户请求稳定选择同一个 Key缓存亲和性
"""
def get_priority(candidate: ProviderCandidate) -> int:
"""获取候选的优先级"""
if not candidate.key:
return 999999
priority_by_format = candidate.key.global_priority_by_format or {}
if api_format and api_format in priority_by_format:
return priority_by_format[api_format]
return 999999 # NULL 排在后面
# 按优先级分组
priority_groups: dict[int, list[ProviderCandidate]] = defaultdict(list)
for candidate in candidates:
priority = get_priority(candidate)
priority_groups[priority].append(candidate)
result = []
for priority in sorted(priority_groups.keys()): # 数字小的优先级高
group = priority_groups[priority]
if len(group) > 1 and affinity_key:
# 同优先级内哈希分散负载均衡
scored_candidates = []
for candidate in group:
key_id = candidate.key.id if candidate.key else ""
hash_value = self._affinity_hash(affinity_key, key_id)
scored_candidates.append((hash_value, candidate))
# 按哈希值排序
sorted_group = [c for _, c in sorted(scored_candidates, key=lambda x: x[0])]
result.extend(sorted_group)
else:
# 单个候选或没有 affinity_key按次要排序条件排序
def secondary_sort(c: ProviderCandidate) -> tuple[int, int, str]:
pp = c.provider.provider_priority
ip = c.key.internal_priority if c.key else None
return (
pp if pp is not None else 999999,
ip if ip is not None else 999999,
c.key.id if c.key else "",
)
result.extend(sorted(group, key=secondary_sort))
return result
def _apply_load_balance(
self, candidates: list[ProviderCandidate], api_format: str | None = None
) -> list[ProviderCandidate]:
"""
负载均衡模式:同优先级内随机轮换
排序逻辑:
1. 按优先级分组provider_priority, internal_priority 或 global_priority_by_format
2. 同优先级组内随机打乱
3. 不考虑缓存亲和性
"""
if not candidates:
return candidates
priority_groups: dict[tuple, list[ProviderCandidate]] = defaultdict(list)
# 根据优先级模式选择分组方式
if self.priority_mode == self.PRIORITY_MODE_GLOBAL_KEY:
# 全局 Key 优先模式:按格式特定优先级分组
for candidate in candidates:
priority = 999999
if candidate.key:
priority_by_format = candidate.key.global_priority_by_format or {}
if api_format and api_format in priority_by_format:
priority = priority_by_format[api_format]
priority_groups[(priority,)].append(candidate)
else:
# 提供商优先模式:按 (provider_priority, internal_priority) 分组
for candidate in candidates:
pp = candidate.provider.provider_priority
ip = candidate.key.internal_priority if candidate.key else None
key = (
pp if pp is not None else 999999,
ip if ip is not None else 999999,
)
priority_groups[key].append(candidate)
result: list[ProviderCandidate] = []
for priority in sorted(priority_groups.keys()):
group = priority_groups[priority]
if len(group) > 1:
# 同优先级内随机打乱
shuffled = list(group)
random.shuffle(shuffled)
result.extend(shuffled)
else:
result.extend(group)
return result
def _shuffle_keys_by_internal_priority(
self,
keys: list[ProviderAPIKey],
affinity_key: str | None = None,
use_random: bool = False,
) -> list[ProviderAPIKey]:
"""
对 API Key 按 internal_priority 分组,同优先级内部基于 affinity_key 进行确定性打乱
目的:
- 数字越小越优先使用
- 同优先级 Key 之间实现负载均衡
- 使用 affinity_key 哈希确保同一请求 Key 的请求稳定(避免破坏缓存亲和性)
- 当 use_random=True 时,使用随机排序实现轮换(用于 TTL=0 的场景)
Args:
keys: API Key 列表
affinity_key: 亲和性标识符(通常为 API Key ID用于确定性打乱
use_random: 是否使用随机排序TTL=0 时为 True
Returns:
排序后的 Key 列表
"""
if not keys:
return []
# 按 internal_priority 分组
priority_groups: dict[int, list[ProviderAPIKey]] = defaultdict(list)
for key in keys:
priority = key.internal_priority if key.internal_priority is not None else 999999
priority_groups[priority].append(key)
# 对每个优先级组内的 Key 进行打乱
result = []
for priority in sorted(priority_groups.keys()): # 数字小的优先级高,排前面
group_keys = priority_groups[priority]
if len(group_keys) > 1:
if use_random:
# TTL=0 模式:使用随机排序实现 Key 轮换
shuffled = list(group_keys)
random.shuffle(shuffled)
result.extend(shuffled)
elif affinity_key:
# 正常模式:使用哈希确定性打乱(保持缓存亲和性)
key_scores = []
for key in group_keys:
hash_value = self._affinity_hash(affinity_key, key.id)
key_scores.append((hash_value, key))
# 按哈希值排序
sorted_group = [key for _, key in sorted(key_scores, key=lambda x: x[0])]
result.extend(sorted_group)
else:
# 没有 affinity_key 时按 ID 排序保持稳定性
result.extend(sorted(group_keys, key=lambda k: k.id))
else:
# 单个 Key 直接添加
result.extend(group_keys)
return result
async def invalidate_cache(
self,
affinity_key: str,

View File

@@ -0,0 +1,216 @@
from __future__ import annotations
from typing import Any
from src.core.api_format.signature import normalize_signature_key
from src.services.billing.token_normalization import normalize_input_tokens_for_billing
from src.services.usage._recording_helpers import (
build_usage_params,
sanitize_request_metadata,
)
from src.services.usage._types import UsageCostInfo, UsageRecordParams
class UsageBillingIntegrationMixin:
"""计费集成方法 -- 准备用量记录的共享逻辑"""
@classmethod
async def _prepare_usage_record(
cls,
params: UsageRecordParams,
) -> tuple[dict[str, Any], float]:
"""准备用量记录的共享逻辑
此方法提取了 record_usage 和 record_usage_async 的公共处理逻辑:
- 获取费率倍数
- 计算成本
- 构建 Usage 参数
Args:
params: 用量记录参数数据类
Returns:
(usage_params 字典, total_cost 总成本)
"""
# 计费口径以 Provider 为准(优先 endpoint_api_format
billing_api_format: str | None = None
if params.endpoint_api_format:
try:
billing_api_format = normalize_signature_key(str(params.endpoint_api_format))
except Exception:
billing_api_format = None
if billing_api_format is None and params.api_format:
try:
billing_api_format = normalize_signature_key(str(params.api_format))
except Exception:
billing_api_format = None
input_tokens_for_billing = normalize_input_tokens_for_billing(
billing_api_format,
params.input_tokens,
params.cache_read_input_tokens,
)
# 获取费率倍数和是否免费套餐(传递 api_format 支持按格式配置的倍率)
actual_rate_multiplier, is_free_tier = await cls._get_rate_multiplier_and_free_tier(
params.db, params.provider_api_key_id, params.provider_id, billing_api_format
)
metadata = dict(params.metadata or {})
is_failed_request = params.status_code >= 400 or params.error_message is not None
# Helper: compute billing task_type (billing domain)
billing_task_type = (params.request_type or "").lower()
if billing_task_type not in {"chat", "cli", "video", "image", "audio"}:
billing_task_type = "chat"
# 使用新计费系统计算费用
from src.services.billing.service import BillingService
request_count = 0 if is_failed_request else 1
dims: dict[str, Any] = {
"input_tokens": input_tokens_for_billing,
"output_tokens": params.output_tokens,
"cache_creation_input_tokens": params.cache_creation_input_tokens,
"cache_read_input_tokens": params.cache_read_input_tokens,
"request_count": request_count,
}
if params.cache_ttl_minutes is not None:
dims["cache_ttl_minutes"] = params.cache_ttl_minutes
# If tiered pricing is disabled, force first tier by using tier-key=0.
if not params.use_tiered_pricing:
dims["total_input_context"] = 0
billing = BillingService(params.db)
result = billing.calculate(
task_type=billing_task_type,
model=params.model,
provider_id=params.provider_id or "",
dimensions=dims,
strict_mode=None,
)
snap = result.snapshot
breakdown = snap.cost_breakdown or {}
input_cost = float(breakdown.get("input_cost", 0.0))
output_cost = float(breakdown.get("output_cost", 0.0))
cache_creation_cost = float(breakdown.get("cache_creation_cost", 0.0))
cache_read_cost = float(breakdown.get("cache_read_cost", 0.0))
request_cost = float(breakdown.get("request_cost", 0.0))
cache_cost = cache_creation_cost + cache_read_cost
total_cost = float(snap.total_cost or 0.0)
rv = snap.resolved_variables or {}
def _as_float(v: Any, d: float | None) -> float | None:
try:
if v is None:
return d
return float(v)
except Exception:
return d
input_price = _as_float(rv.get("input_price_per_1m"), 0.0) or 0.0
output_price = _as_float(rv.get("output_price_per_1m"), 0.0) or 0.0
cache_creation_price = _as_float(rv.get("cache_creation_price_per_1m"), None)
cache_read_price = _as_float(rv.get("cache_read_price_per_1m"), None)
request_price = _as_float(rv.get("price_per_request"), None)
# Audit snapshot (pruned later by sanitize_request_metadata)
metadata["billing_snapshot"] = snap.to_dict()
# Best-effort prune metadata to reduce DB/memory pressure.
metadata = sanitize_request_metadata(metadata)
# 构建 Usage 参数
usage_params = build_usage_params(
db=params.db,
user=params.user,
api_key=params.api_key,
provider=params.provider,
model=params.model,
input_tokens=input_tokens_for_billing,
output_tokens=params.output_tokens,
cache_creation_input_tokens=params.cache_creation_input_tokens,
cache_read_input_tokens=params.cache_read_input_tokens,
request_type=params.request_type,
api_format=params.api_format,
endpoint_api_format=params.endpoint_api_format,
has_format_conversion=params.has_format_conversion,
is_stream=params.is_stream,
response_time_ms=params.response_time_ms,
first_byte_time_ms=params.first_byte_time_ms,
status_code=params.status_code,
error_message=params.error_message,
metadata=metadata,
request_headers=params.request_headers,
request_body=params.request_body,
provider_request_headers=params.provider_request_headers,
response_headers=params.response_headers,
client_response_headers=params.client_response_headers,
response_body=params.response_body,
request_id=params.request_id,
provider_id=params.provider_id,
provider_endpoint_id=params.provider_endpoint_id,
provider_api_key_id=params.provider_api_key_id,
status=params.status,
target_model=params.target_model,
cost=UsageCostInfo(
input_cost=input_cost,
output_cost=output_cost,
cache_creation_cost=cache_creation_cost,
cache_read_cost=cache_read_cost,
cache_cost=cache_cost,
request_cost=request_cost,
total_cost=total_cost,
input_price=input_price,
output_price=output_price,
cache_creation_price=cache_creation_price,
cache_read_price=cache_read_price,
request_price=request_price,
actual_rate_multiplier=actual_rate_multiplier,
is_free_tier=is_free_tier,
),
)
return usage_params, total_cost
@classmethod
async def _prepare_usage_records_batch(
cls,
params_list: list[UsageRecordParams],
) -> list[tuple[dict[str, Any], float, Exception | None]]:
"""批量并行准备用量记录(性能优化)
并行调用 _prepare_usage_record提高批量处理效率。
Args:
params_list: 用量记录参数列表
Returns:
列表,每项为 (usage_params, total_cost, exception)
如果处理成功exception 为 None
"""
import asyncio
async def prepare_single(
params: UsageRecordParams,
) -> tuple[dict[str, Any], float, Exception | None]:
try:
usage_params, total_cost = await cls._prepare_usage_record(params)
return (usage_params, total_cost, None)
except Exception as e:
return ({}, 0.0, e)
if not params_list:
return []
# 避免一次性创建过多 task并且 _prepare_usage_record 内部也可能包含并行调用)
# 这里采用分批 gather 来限制并发量。
chunk_size = 50
results: list[tuple[dict[str, Any], float, Exception | None]] = []
for i in range(0, len(params_list), chunk_size):
chunk = params_list[i : i + chunk_size]
chunk_results = await asyncio.gather(*(prepare_single(p) for p in chunk))
results.extend(chunk_results)
return results

View File

@@ -0,0 +1,310 @@
from __future__ import annotations
import json
from typing import Any
from sqlalchemy.orm import Session
from src.models.database import ApiKey, Usage, User
from src.services.system.config import SystemConfigService
from src.services.usage._types import UsageCostInfo
from src.services.usage.error_classifier import classify_error
# Metadata pruning configuration (ordered by priority - drop first to last)
METADATA_PRUNE_KEYS: tuple[str, ...] = (
"raw_response_ref",
"poll_raw_response",
"trace",
"debug",
"dimensions",
"provider_response_headers",
"client_response_headers",
)
# Keys to preserve even under aggressive pruning
METADATA_KEEP_KEYS: frozenset[str] = frozenset(
{
"billing_snapshot",
"billing_updated_at",
"perf",
"_metadata_truncated",
}
)
def build_usage_params(
*,
db: Session,
user: User | None,
api_key: ApiKey | None,
provider: str,
model: str,
input_tokens: int,
output_tokens: int,
cache_creation_input_tokens: int,
cache_read_input_tokens: int,
request_type: str,
api_format: str | None,
endpoint_api_format: str | None,
has_format_conversion: bool,
is_stream: bool,
response_time_ms: int | None,
first_byte_time_ms: int | None,
status_code: int,
error_message: str | None,
metadata: dict[str, Any] | None,
request_headers: dict[str, Any] | None,
request_body: Any | None,
provider_request_headers: dict[str, Any] | None,
response_headers: dict[str, Any] | None,
client_response_headers: dict[str, Any] | None,
response_body: Any | None,
request_id: str,
provider_id: str | None,
provider_endpoint_id: str | None,
provider_api_key_id: str | None,
status: str,
target_model: str | None,
cost: UsageCostInfo,
) -> dict[str, Any]:
"""构建 Usage 记录的参数字典(内部方法,避免代码重复)"""
# 展开成本信息
input_cost = cost.input_cost
output_cost = cost.output_cost
cache_creation_cost = cost.cache_creation_cost
cache_read_cost = cost.cache_read_cost
cache_cost = cost.cache_cost
request_cost = cost.request_cost
total_cost = cost.total_cost
input_price = cost.input_price
output_price = cost.output_price
cache_creation_price = cost.cache_creation_price
cache_read_price = cost.cache_read_price
request_price = cost.request_price
actual_rate_multiplier = cost.actual_rate_multiplier
is_free_tier = cost.is_free_tier
# 根据配置决定是否记录请求详情
should_log_headers = SystemConfigService.should_log_headers(db)
should_log_body = SystemConfigService.should_log_body(db)
# 处理请求头(可能需要脱敏)
processed_request_headers = None
if should_log_headers and request_headers:
processed_request_headers = SystemConfigService.mask_sensitive_headers(db, request_headers)
# 处理提供商请求头(可能需要脱敏)
processed_provider_request_headers = None
if should_log_headers and provider_request_headers:
processed_provider_request_headers = SystemConfigService.mask_sensitive_headers(
db, provider_request_headers
)
# 处理请求体和响应体(可能需要截断)
processed_request_body = None
processed_response_body = None
if should_log_body:
if request_body:
processed_request_body = SystemConfigService.truncate_body(
db, request_body, is_request=True
)
if response_body:
processed_response_body = SystemConfigService.truncate_body(
db, response_body, is_request=False
)
# 处理响应头
processed_response_headers = None
if should_log_headers and response_headers:
processed_response_headers = SystemConfigService.mask_sensitive_headers(
db, response_headers
)
# 处理返回给客户端的响应头
processed_client_response_headers = None
if should_log_headers and client_response_headers:
processed_client_response_headers = SystemConfigService.mask_sensitive_headers(
db, client_response_headers
)
# 计算真实成本(表面成本 * 倍率),免费套餐实际费用为 0
if is_free_tier:
actual_input_cost = 0.0
actual_output_cost = 0.0
actual_cache_creation_cost = 0.0
actual_cache_read_cost = 0.0
actual_request_cost = 0.0
actual_total_cost = 0.0
else:
actual_input_cost = input_cost * actual_rate_multiplier
actual_output_cost = output_cost * actual_rate_multiplier
actual_cache_creation_cost = cache_creation_cost * actual_rate_multiplier
actual_cache_read_cost = cache_read_cost * actual_rate_multiplier
actual_request_cost = request_cost * actual_rate_multiplier
actual_total_cost = total_cost * actual_rate_multiplier
error_category = None
if status_code >= 400 or error_message or status in {"failed", "cancelled"}:
error_category = classify_error(status_code, error_message, status).value
return {
"user_id": user.id if user else None,
"api_key_id": api_key.id if api_key else None,
"request_id": request_id,
"provider_name": provider,
"model": model,
"target_model": target_model,
"provider_id": provider_id,
"provider_endpoint_id": provider_endpoint_id,
"provider_api_key_id": provider_api_key_id,
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"total_tokens": input_tokens + output_tokens,
"cache_creation_input_tokens": cache_creation_input_tokens,
"cache_read_input_tokens": cache_read_input_tokens,
"input_cost_usd": input_cost,
"output_cost_usd": output_cost,
"cache_cost_usd": cache_cost,
"cache_creation_cost_usd": cache_creation_cost,
"cache_read_cost_usd": cache_read_cost,
"request_cost_usd": request_cost,
"total_cost_usd": total_cost,
"actual_input_cost_usd": actual_input_cost,
"actual_output_cost_usd": actual_output_cost,
"actual_cache_creation_cost_usd": actual_cache_creation_cost,
"actual_cache_read_cost_usd": actual_cache_read_cost,
"actual_request_cost_usd": actual_request_cost,
"actual_total_cost_usd": actual_total_cost,
"rate_multiplier": actual_rate_multiplier,
"input_price_per_1m": input_price,
"output_price_per_1m": output_price,
"cache_creation_price_per_1m": cache_creation_price,
"cache_read_price_per_1m": cache_read_price,
"price_per_request": request_price,
"request_type": request_type,
"api_format": api_format,
"endpoint_api_format": endpoint_api_format,
"has_format_conversion": has_format_conversion,
"is_stream": is_stream,
"status_code": status_code,
"error_message": error_message,
"error_category": error_category,
"response_time_ms": response_time_ms,
"first_byte_time_ms": first_byte_time_ms,
"status": status,
"request_metadata": metadata,
"request_headers": processed_request_headers,
"request_body": processed_request_body,
"provider_request_headers": processed_provider_request_headers,
"response_headers": processed_response_headers,
"client_response_headers": processed_client_response_headers,
"response_body": processed_response_body,
}
def update_existing_usage(
existing_usage: Usage,
usage_params: dict[str, Any],
target_model: str | None,
) -> None:
"""更新已存在的 Usage 记录(内部方法)"""
# 更新关键字段
existing_usage.provider_name = usage_params["provider_name"]
existing_usage.model = usage_params["model"]
existing_usage.request_type = usage_params["request_type"]
existing_usage.api_format = usage_params["api_format"]
existing_usage.endpoint_api_format = usage_params["endpoint_api_format"]
existing_usage.has_format_conversion = usage_params["has_format_conversion"]
existing_usage.is_stream = usage_params["is_stream"]
existing_usage.status = usage_params["status"]
existing_usage.status_code = usage_params["status_code"]
existing_usage.error_message = usage_params["error_message"]
existing_usage.error_category = usage_params.get("error_category")
existing_usage.response_time_ms = usage_params["response_time_ms"]
existing_usage.first_byte_time_ms = usage_params["first_byte_time_ms"]
# 更新请求头和请求体(如果有新值)
if usage_params["request_headers"] is not None:
existing_usage.request_headers = usage_params["request_headers"]
if usage_params["request_body"] is not None:
existing_usage.request_body = usage_params["request_body"]
if usage_params["provider_request_headers"] is not None:
existing_usage.provider_request_headers = usage_params["provider_request_headers"]
existing_usage.response_body = usage_params["response_body"]
existing_usage.response_headers = usage_params["response_headers"]
existing_usage.client_response_headers = usage_params["client_response_headers"]
# 更新 token 和费用信息
existing_usage.input_tokens = usage_params["input_tokens"]
existing_usage.output_tokens = usage_params["output_tokens"]
existing_usage.total_tokens = usage_params["total_tokens"]
existing_usage.cache_creation_input_tokens = usage_params["cache_creation_input_tokens"]
existing_usage.cache_read_input_tokens = usage_params["cache_read_input_tokens"]
existing_usage.input_cost_usd = usage_params["input_cost_usd"]
existing_usage.output_cost_usd = usage_params["output_cost_usd"]
existing_usage.cache_cost_usd = usage_params["cache_cost_usd"]
existing_usage.cache_creation_cost_usd = usage_params["cache_creation_cost_usd"]
existing_usage.cache_read_cost_usd = usage_params["cache_read_cost_usd"]
existing_usage.request_cost_usd = usage_params["request_cost_usd"]
existing_usage.total_cost_usd = usage_params["total_cost_usd"]
existing_usage.actual_input_cost_usd = usage_params["actual_input_cost_usd"]
existing_usage.actual_output_cost_usd = usage_params["actual_output_cost_usd"]
existing_usage.actual_cache_creation_cost_usd = usage_params["actual_cache_creation_cost_usd"]
existing_usage.actual_cache_read_cost_usd = usage_params["actual_cache_read_cost_usd"]
existing_usage.actual_request_cost_usd = usage_params["actual_request_cost_usd"]
existing_usage.actual_total_cost_usd = usage_params["actual_total_cost_usd"]
existing_usage.rate_multiplier = usage_params["rate_multiplier"]
# 更新 Provider 侧追踪信息
existing_usage.provider_id = usage_params["provider_id"]
existing_usage.provider_endpoint_id = usage_params["provider_endpoint_id"]
existing_usage.provider_api_key_id = usage_params["provider_api_key_id"]
# 更新元数据(如 billing_snapshot/dimensions 等)
if usage_params.get("request_metadata") is not None:
existing_usage.request_metadata = usage_params["request_metadata"]
# 更新模型映射信息
if target_model is not None:
existing_usage.target_model = target_model
def sanitize_request_metadata(metadata: dict[str, Any]) -> dict[str, Any]:
"""
Best-effort metadata pruning to reduce DB/CPU/memory pressure.
This is called right before persisting Usage rows (or updating request_metadata).
Pruning order is defined by `METADATA_PRUNE_KEYS` (first key is dropped first).
"""
if not isinstance(metadata, dict) or not metadata:
return {}
from src.config.settings import config
# Enforce global metadata size limit (best-effort)
max_bytes = int(getattr(config, "usage_metadata_max_bytes", 0) or 0)
if max_bytes <= 0:
return metadata
def _size(d: dict[str, Any]) -> int:
try:
return len(json.dumps(d, ensure_ascii=False, default=str))
except Exception:
return len(str(d))
if _size(metadata) <= max_bytes:
return metadata
# Progressive pruning (configurable order)
metadata["_metadata_truncated"] = True
for k in METADATA_PRUNE_KEYS:
if k in metadata:
metadata.pop(k, None)
if _size(metadata) <= max_bytes:
return metadata
# Fallback: keep only billing-related metadata
reduced = {k: metadata.get(k) for k in METADATA_KEEP_KEYS if k in metadata}
return reduced

View File

@@ -1,217 +1,39 @@
from __future__ import annotations
import json
import uuid
from datetime import datetime, timezone
from typing import Any
from sqlalchemy.orm import Session
from src.core.api_format.signature import normalize_signature_key
from src.core.logger import logger
from src.models.database import ApiKey, Provider, Usage, User
from src.services.billing.token_normalization import normalize_input_tokens_for_billing
from src.services.system.config import SystemConfigService
from src.services.usage._billing_integration import UsageBillingIntegrationMixin
from src.services.usage._recording_helpers import (
METADATA_KEEP_KEYS,
METADATA_PRUNE_KEYS,
build_usage_params,
sanitize_request_metadata,
update_existing_usage,
)
from src.services.usage._types import UsageCostInfo, UsageRecordParams
from src.services.usage.error_classifier import classify_error
class UsageRecordingMixin:
class UsageRecordingMixin(UsageBillingIntegrationMixin):
"""记录用量相关方法"""
# Metadata pruning configuration (ordered by priority - drop first to last)
_METADATA_PRUNE_KEYS: tuple[str, ...] = (
"raw_response_ref",
"poll_raw_response",
"trace",
"debug",
"dimensions",
"provider_response_headers",
"client_response_headers",
)
# Metadata pruning configuration -- re-export from helpers for backward compatibility
_METADATA_PRUNE_KEYS: tuple[str, ...] = METADATA_PRUNE_KEYS
_METADATA_KEEP_KEYS: frozenset[str] = METADATA_KEEP_KEYS
# Keys to preserve even under aggressive pruning
_METADATA_KEEP_KEYS: frozenset[str] = frozenset(
{
"billing_snapshot",
"billing_updated_at",
"perf",
"_metadata_truncated",
}
)
# ------------------------------------------------------------------
# Backward-compatible thin wrappers
# ------------------------------------------------------------------
@staticmethod
def _build_usage_params(
*,
db: Session,
user: User | None,
api_key: ApiKey | None,
provider: str,
model: str,
input_tokens: int,
output_tokens: int,
cache_creation_input_tokens: int,
cache_read_input_tokens: int,
request_type: str,
api_format: str | None,
endpoint_api_format: str | None,
has_format_conversion: bool,
is_stream: bool,
response_time_ms: int | None,
first_byte_time_ms: int | None,
status_code: int,
error_message: str | None,
metadata: dict[str, Any] | None,
request_headers: dict[str, Any] | None,
request_body: Any | None,
provider_request_headers: dict[str, Any] | None,
response_headers: dict[str, Any] | None,
client_response_headers: dict[str, Any] | None,
response_body: Any | None,
request_id: str,
provider_id: str | None,
provider_endpoint_id: str | None,
provider_api_key_id: str | None,
status: str,
target_model: str | None,
cost: UsageCostInfo,
) -> dict[str, Any]:
"""构建 Usage 记录的参数字典(内部方法,避免代码重复)"""
# 展开成本信息
input_cost = cost.input_cost
output_cost = cost.output_cost
cache_creation_cost = cost.cache_creation_cost
cache_read_cost = cost.cache_read_cost
cache_cost = cost.cache_cost
request_cost = cost.request_cost
total_cost = cost.total_cost
input_price = cost.input_price
output_price = cost.output_price
cache_creation_price = cost.cache_creation_price
cache_read_price = cost.cache_read_price
request_price = cost.request_price
actual_rate_multiplier = cost.actual_rate_multiplier
is_free_tier = cost.is_free_tier
# 根据配置决定是否记录请求详情
should_log_headers = SystemConfigService.should_log_headers(db)
should_log_body = SystemConfigService.should_log_body(db)
# 处理请求头(可能需要脱敏)
processed_request_headers = None
if should_log_headers and request_headers:
processed_request_headers = SystemConfigService.mask_sensitive_headers(
db, request_headers
)
# 处理提供商请求头(可能需要脱敏)
processed_provider_request_headers = None
if should_log_headers and provider_request_headers:
processed_provider_request_headers = SystemConfigService.mask_sensitive_headers(
db, provider_request_headers
)
# 处理请求体和响应体(可能需要截断)
processed_request_body = None
processed_response_body = None
if should_log_body:
if request_body:
processed_request_body = SystemConfigService.truncate_body(
db, request_body, is_request=True
)
if response_body:
processed_response_body = SystemConfigService.truncate_body(
db, response_body, is_request=False
)
# 处理响应头
processed_response_headers = None
if should_log_headers and response_headers:
processed_response_headers = SystemConfigService.mask_sensitive_headers(
db, response_headers
)
# 处理返回给客户端的响应头
processed_client_response_headers = None
if should_log_headers and client_response_headers:
processed_client_response_headers = SystemConfigService.mask_sensitive_headers(
db, client_response_headers
)
# 计算真实成本(表面成本 * 倍率),免费套餐实际费用为 0
if is_free_tier:
actual_input_cost = 0.0
actual_output_cost = 0.0
actual_cache_creation_cost = 0.0
actual_cache_read_cost = 0.0
actual_request_cost = 0.0
actual_total_cost = 0.0
else:
actual_input_cost = input_cost * actual_rate_multiplier
actual_output_cost = output_cost * actual_rate_multiplier
actual_cache_creation_cost = cache_creation_cost * actual_rate_multiplier
actual_cache_read_cost = cache_read_cost * actual_rate_multiplier
actual_request_cost = request_cost * actual_rate_multiplier
actual_total_cost = total_cost * actual_rate_multiplier
error_category = None
if status_code >= 400 or error_message or status in {"failed", "cancelled"}:
error_category = classify_error(status_code, error_message, status).value
return {
"user_id": user.id if user else None,
"api_key_id": api_key.id if api_key else None,
"request_id": request_id,
"provider_name": provider,
"model": model,
"target_model": target_model,
"provider_id": provider_id,
"provider_endpoint_id": provider_endpoint_id,
"provider_api_key_id": provider_api_key_id,
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"total_tokens": input_tokens + output_tokens,
"cache_creation_input_tokens": cache_creation_input_tokens,
"cache_read_input_tokens": cache_read_input_tokens,
"input_cost_usd": input_cost,
"output_cost_usd": output_cost,
"cache_cost_usd": cache_cost,
"cache_creation_cost_usd": cache_creation_cost,
"cache_read_cost_usd": cache_read_cost,
"request_cost_usd": request_cost,
"total_cost_usd": total_cost,
"actual_input_cost_usd": actual_input_cost,
"actual_output_cost_usd": actual_output_cost,
"actual_cache_creation_cost_usd": actual_cache_creation_cost,
"actual_cache_read_cost_usd": actual_cache_read_cost,
"actual_request_cost_usd": actual_request_cost,
"actual_total_cost_usd": actual_total_cost,
"rate_multiplier": actual_rate_multiplier,
"input_price_per_1m": input_price,
"output_price_per_1m": output_price,
"cache_creation_price_per_1m": cache_creation_price,
"cache_read_price_per_1m": cache_read_price,
"price_per_request": request_price,
"request_type": request_type,
"api_format": api_format,
"endpoint_api_format": endpoint_api_format,
"has_format_conversion": has_format_conversion,
"is_stream": is_stream,
"status_code": status_code,
"error_message": error_message,
"error_category": error_category,
"response_time_ms": response_time_ms,
"first_byte_time_ms": first_byte_time_ms,
"status": status,
"request_metadata": metadata,
"request_headers": processed_request_headers,
"request_body": processed_request_body,
"provider_request_headers": processed_provider_request_headers,
"response_headers": processed_response_headers,
"client_response_headers": processed_client_response_headers,
"response_body": processed_response_body,
}
def _build_usage_params(**kwargs: Any) -> dict[str, Any]:
"""构建 Usage 记录的参数字典(委托到模块级函数)"""
return build_usage_params(**kwargs)
@staticmethod
def _update_existing_usage(
@@ -219,309 +41,17 @@ class UsageRecordingMixin:
usage_params: dict[str, Any],
target_model: str | None,
) -> None:
"""更新已存在的 Usage 记录(内部方法"""
# 更新关键字段
existing_usage.provider_name = usage_params["provider_name"]
existing_usage.model = usage_params["model"]
existing_usage.request_type = usage_params["request_type"]
existing_usage.api_format = usage_params["api_format"]
existing_usage.endpoint_api_format = usage_params["endpoint_api_format"]
existing_usage.has_format_conversion = usage_params["has_format_conversion"]
existing_usage.is_stream = usage_params["is_stream"]
existing_usage.status = usage_params["status"]
existing_usage.status_code = usage_params["status_code"]
existing_usage.error_message = usage_params["error_message"]
existing_usage.error_category = usage_params.get("error_category")
existing_usage.response_time_ms = usage_params["response_time_ms"]
existing_usage.first_byte_time_ms = usage_params["first_byte_time_ms"]
# 更新请求头和请求体(如果有新值)
if usage_params["request_headers"] is not None:
existing_usage.request_headers = usage_params["request_headers"]
if usage_params["request_body"] is not None:
existing_usage.request_body = usage_params["request_body"]
if usage_params["provider_request_headers"] is not None:
existing_usage.provider_request_headers = usage_params["provider_request_headers"]
existing_usage.response_body = usage_params["response_body"]
existing_usage.response_headers = usage_params["response_headers"]
existing_usage.client_response_headers = usage_params["client_response_headers"]
# 更新 token 和费用信息
existing_usage.input_tokens = usage_params["input_tokens"]
existing_usage.output_tokens = usage_params["output_tokens"]
existing_usage.total_tokens = usage_params["total_tokens"]
existing_usage.cache_creation_input_tokens = usage_params["cache_creation_input_tokens"]
existing_usage.cache_read_input_tokens = usage_params["cache_read_input_tokens"]
existing_usage.input_cost_usd = usage_params["input_cost_usd"]
existing_usage.output_cost_usd = usage_params["output_cost_usd"]
existing_usage.cache_cost_usd = usage_params["cache_cost_usd"]
existing_usage.cache_creation_cost_usd = usage_params["cache_creation_cost_usd"]
existing_usage.cache_read_cost_usd = usage_params["cache_read_cost_usd"]
existing_usage.request_cost_usd = usage_params["request_cost_usd"]
existing_usage.total_cost_usd = usage_params["total_cost_usd"]
existing_usage.actual_input_cost_usd = usage_params["actual_input_cost_usd"]
existing_usage.actual_output_cost_usd = usage_params["actual_output_cost_usd"]
existing_usage.actual_cache_creation_cost_usd = usage_params[
"actual_cache_creation_cost_usd"
]
existing_usage.actual_cache_read_cost_usd = usage_params["actual_cache_read_cost_usd"]
existing_usage.actual_request_cost_usd = usage_params["actual_request_cost_usd"]
existing_usage.actual_total_cost_usd = usage_params["actual_total_cost_usd"]
existing_usage.rate_multiplier = usage_params["rate_multiplier"]
# 更新 Provider 侧追踪信息
existing_usage.provider_id = usage_params["provider_id"]
existing_usage.provider_endpoint_id = usage_params["provider_endpoint_id"]
existing_usage.provider_api_key_id = usage_params["provider_api_key_id"]
# 更新元数据(如 billing_snapshot/dimensions 等)
if usage_params.get("request_metadata") is not None:
existing_usage.request_metadata = usage_params["request_metadata"]
# 更新模型映射信息
if target_model is not None:
existing_usage.target_model = target_model
"""更新已存在的 Usage 记录(委托到模块级函数"""
update_existing_usage(existing_usage, usage_params, target_model)
@classmethod
def _sanitize_request_metadata(cls, metadata: dict[str, Any]) -> dict[str, Any]:
"""
Best-effort metadata pruning to reduce DB/CPU/memory pressure.
"""元数据清理(委托到模块级函数)"""
return sanitize_request_metadata(metadata)
This is called right before persisting Usage rows (or updating request_metadata).
Pruning order is defined by `_METADATA_PRUNE_KEYS` (first key is dropped first).
"""
if not isinstance(metadata, dict) or not metadata:
return {}
from src.config.settings import config
# Enforce global metadata size limit (best-effort)
max_bytes = int(getattr(config, "usage_metadata_max_bytes", 0) or 0)
if max_bytes <= 0:
return metadata
def _size(d: dict[str, Any]) -> int:
try:
return len(json.dumps(d, ensure_ascii=False, default=str))
except Exception:
return len(str(d))
if _size(metadata) <= max_bytes:
return metadata
# Progressive pruning (configurable order)
metadata["_metadata_truncated"] = True
for k in cls._METADATA_PRUNE_KEYS:
if k in metadata:
metadata.pop(k, None)
if _size(metadata) <= max_bytes:
return metadata
# Fallback: keep only billing-related metadata
reduced = {k: metadata.get(k) for k in cls._METADATA_KEEP_KEYS if k in metadata}
return reduced
@classmethod
async def _prepare_usage_record(
cls,
params: UsageRecordParams,
) -> tuple[dict[str, Any], float]:
"""准备用量记录的共享逻辑
此方法提取了 record_usage 和 record_usage_async 的公共处理逻辑:
- 获取费率倍数
- 计算成本
- 构建 Usage 参数
Args:
params: 用量记录参数数据类
Returns:
(usage_params 字典, total_cost 总成本)
"""
# 计费口径以 Provider 为准(优先 endpoint_api_format
billing_api_format: str | None = None
if params.endpoint_api_format:
try:
billing_api_format = normalize_signature_key(str(params.endpoint_api_format))
except Exception:
billing_api_format = None
if billing_api_format is None and params.api_format:
try:
billing_api_format = normalize_signature_key(str(params.api_format))
except Exception:
billing_api_format = None
input_tokens_for_billing = normalize_input_tokens_for_billing(
billing_api_format,
params.input_tokens,
params.cache_read_input_tokens,
)
# 获取费率倍数和是否免费套餐(传递 api_format 支持按格式配置的倍率)
actual_rate_multiplier, is_free_tier = await cls._get_rate_multiplier_and_free_tier(
params.db, params.provider_api_key_id, params.provider_id, billing_api_format
)
metadata = dict(params.metadata or {})
is_failed_request = params.status_code >= 400 or params.error_message is not None
# Helper: compute billing task_type (billing domain)
billing_task_type = (params.request_type or "").lower()
if billing_task_type not in {"chat", "cli", "video", "image", "audio"}:
billing_task_type = "chat"
# 使用新计费系统计算费用
from src.services.billing.service import BillingService
request_count = 0 if is_failed_request else 1
dims: dict[str, Any] = {
"input_tokens": input_tokens_for_billing,
"output_tokens": params.output_tokens,
"cache_creation_input_tokens": params.cache_creation_input_tokens,
"cache_read_input_tokens": params.cache_read_input_tokens,
"request_count": request_count,
}
if params.cache_ttl_minutes is not None:
dims["cache_ttl_minutes"] = params.cache_ttl_minutes
# If tiered pricing is disabled, force first tier by using tier-key=0.
if not params.use_tiered_pricing:
dims["total_input_context"] = 0
billing = BillingService(params.db)
result = billing.calculate(
task_type=billing_task_type,
model=params.model,
provider_id=params.provider_id or "",
dimensions=dims,
strict_mode=None,
)
snap = result.snapshot
breakdown = snap.cost_breakdown or {}
input_cost = float(breakdown.get("input_cost", 0.0))
output_cost = float(breakdown.get("output_cost", 0.0))
cache_creation_cost = float(breakdown.get("cache_creation_cost", 0.0))
cache_read_cost = float(breakdown.get("cache_read_cost", 0.0))
request_cost = float(breakdown.get("request_cost", 0.0))
cache_cost = cache_creation_cost + cache_read_cost
total_cost = float(snap.total_cost or 0.0)
rv = snap.resolved_variables or {}
def _as_float(v: Any, d: float | None) -> float | None:
try:
if v is None:
return d
return float(v)
except Exception:
return d
input_price = _as_float(rv.get("input_price_per_1m"), 0.0) or 0.0
output_price = _as_float(rv.get("output_price_per_1m"), 0.0) or 0.0
cache_creation_price = _as_float(rv.get("cache_creation_price_per_1m"), None)
cache_read_price = _as_float(rv.get("cache_read_price_per_1m"), None)
request_price = _as_float(rv.get("price_per_request"), None)
# Audit snapshot (pruned later by _sanitize_request_metadata)
metadata["billing_snapshot"] = snap.to_dict()
# Best-effort prune metadata to reduce DB/memory pressure.
metadata = cls._sanitize_request_metadata(metadata)
# 构建 Usage 参数
usage_params = cls._build_usage_params(
db=params.db,
user=params.user,
api_key=params.api_key,
provider=params.provider,
model=params.model,
input_tokens=input_tokens_for_billing,
output_tokens=params.output_tokens,
cache_creation_input_tokens=params.cache_creation_input_tokens,
cache_read_input_tokens=params.cache_read_input_tokens,
request_type=params.request_type,
api_format=params.api_format,
endpoint_api_format=params.endpoint_api_format,
has_format_conversion=params.has_format_conversion,
is_stream=params.is_stream,
response_time_ms=params.response_time_ms,
first_byte_time_ms=params.first_byte_time_ms,
status_code=params.status_code,
error_message=params.error_message,
metadata=metadata,
request_headers=params.request_headers,
request_body=params.request_body,
provider_request_headers=params.provider_request_headers,
response_headers=params.response_headers,
client_response_headers=params.client_response_headers,
response_body=params.response_body,
request_id=params.request_id,
provider_id=params.provider_id,
provider_endpoint_id=params.provider_endpoint_id,
provider_api_key_id=params.provider_api_key_id,
status=params.status,
target_model=params.target_model,
cost=UsageCostInfo(
input_cost=input_cost,
output_cost=output_cost,
cache_creation_cost=cache_creation_cost,
cache_read_cost=cache_read_cost,
cache_cost=cache_cost,
request_cost=request_cost,
total_cost=total_cost,
input_price=input_price,
output_price=output_price,
cache_creation_price=cache_creation_price,
cache_read_price=cache_read_price,
request_price=request_price,
actual_rate_multiplier=actual_rate_multiplier,
is_free_tier=is_free_tier,
),
)
return usage_params, total_cost
@classmethod
async def _prepare_usage_records_batch(
cls,
params_list: list[UsageRecordParams],
) -> list[tuple[dict[str, Any], float, Exception | None]]:
"""批量并行准备用量记录(性能优化)
并行调用 _prepare_usage_record提高批量处理效率。
Args:
params_list: 用量记录参数列表
Returns:
列表,每项为 (usage_params, total_cost, exception)
如果处理成功exception 为 None
"""
import asyncio
async def prepare_single(
params: UsageRecordParams,
) -> tuple[dict[str, Any], float, Exception | None]:
try:
usage_params, total_cost = await cls._prepare_usage_record(params)
return (usage_params, total_cost, None)
except Exception as e:
return ({}, 0.0, e)
if not params_list:
return []
# 避免一次性创建过多 task并且 _prepare_usage_record 内部也可能包含并行调用)
# 这里采用分批 gather 来限制并发量。
chunk_size = 50
results: list[tuple[dict[str, Any], float, Exception | None]] = []
for i in range(0, len(params_list), chunk_size):
chunk = params_list[i : i + chunk_size]
chunk_results = await asyncio.gather(*(prepare_single(p) for p in chunk))
results.extend(chunk_results)
return results
# ------------------------------------------------------------------
# Recording methods
# ------------------------------------------------------------------
@classmethod
async def record_usage_async(
@@ -891,7 +421,7 @@ class UsageRecordingMixin:
)
total_cost = float(total_cost_usd)
usage_params = cls._build_usage_params(
usage_params = build_usage_params(
db=db,
user=user,
api_key=api_key,