refactor: defer ProviderAPIKey 大 JSON 字段,减少调度器和模型获取的内存占用

- candidate_builder: defer adjustment_history/utilization_samples/upstream_metadata,
  号池路径在释放 DB 连接前预计算 _pool_account_state 避免后续 N+1 查询
- pool/manager: 优先读取预计算的 _pool_account_state,fallback 到实时解析
- fetch_scheduler: 三处 ProviderAPIKey 查询改用 defer/load_only 排除无关字段
This commit is contained in:
fawney19
2026-03-10 12:33:15 +08:00
parent 1518223de6
commit 9ee27308db
3 changed files with 95 additions and 24 deletions

View File

@@ -19,7 +19,7 @@ from dataclasses import dataclass
from datetime import datetime, timezone from datetime import datetime, timezone
from typing import Any from typing import Any
from sqlalchemy.orm import joinedload from sqlalchemy.orm import defer, joinedload, load_only
from src.core.cache_service import CacheService from src.core.cache_service import CacheService
from src.core.crypto import crypto_service from src.core.crypto import crypto_service
@@ -284,7 +284,19 @@ class ModelFetchScheduler:
"""更新 Key 的错误信息(独立事务)""" """更新 Key 的错误信息(独立事务)"""
try: try:
with create_session() as db: with create_session() as db:
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first() key = (
db.query(ProviderAPIKey)
.options(
defer(ProviderAPIKey.adjustment_history),
defer(ProviderAPIKey.utilization_samples),
defer(ProviderAPIKey.health_by_format),
defer(ProviderAPIKey.circuit_breaker_by_format),
defer(ProviderAPIKey.upstream_metadata),
defer(ProviderAPIKey.allowed_models),
)
.filter(ProviderAPIKey.id == key_id)
.first()
)
if key: if key:
key.last_models_fetch_at = datetime.now(timezone.utc) key.last_models_fetch_at = datetime.now(timezone.utc)
key.last_models_fetch_error = error_msg key.last_models_fetch_error = error_msg
@@ -384,7 +396,18 @@ class ModelFetchScheduler:
with create_session() as db: with create_session() as db:
key = ( key = (
db.query(ProviderAPIKey) db.query(ProviderAPIKey)
.options(joinedload(ProviderAPIKey.provider)) .options(
load_only(
ProviderAPIKey.id,
ProviderAPIKey.is_active,
ProviderAPIKey.auto_fetch_models,
ProviderAPIKey.api_key,
ProviderAPIKey.auth_type,
ProviderAPIKey.auth_config,
ProviderAPIKey.provider_id,
ProviderAPIKey.proxy,
),
)
.filter(ProviderAPIKey.id == key_id) .filter(ProviderAPIKey.id == key_id)
.first() .first()
) )
@@ -482,7 +505,18 @@ class ModelFetchScheduler:
with create_session() as db: with create_session() as db:
# 重新获取 Key因为之前的连接已关闭 # 重新获取 Key因为之前的连接已关闭
key = db.query(ProviderAPIKey).filter(ProviderAPIKey.id == key_id).first() # defer 不需要的大 JSON 字段,减少内存占用
key = (
db.query(ProviderAPIKey)
.options(
defer(ProviderAPIKey.adjustment_history),
defer(ProviderAPIKey.utilization_samples),
defer(ProviderAPIKey.health_by_format),
defer(ProviderAPIKey.circuit_breaker_by_format),
)
.filter(ProviderAPIKey.id == key_id)
.first()
)
if not key: if not key:
logger.warning(f"Key {key_id} 在更新时不存在") logger.warning(f"Key {key_id} 在更新时不存在")
return "error" return "error"

View File

@@ -221,16 +221,22 @@ class PoolManager:
pass pass
# --- 3. Classify candidates ----------------------------------- # --- 3. Classify candidates -----------------------------------
# Pre-compute account states to avoid repeated dict parsing inside loop. # Use precomputed account states from CandidateBuilder when available
# (upstream_metadata is deferred on pool keys to save memory).
# Fall back to on-the-fly resolution for non-pool candidates.
account_states: dict[str, Any] = {} account_states: dict[str, Any] = {}
for c in candidates: for c in candidates:
kid = str(c.key.id) kid = str(c.key.id)
if kid not in account_states: if kid not in account_states:
account_states[kid] = resolve_pool_account_state( precomputed = getattr(c.key, "_pool_account_state", None)
provider_type=provider_type, if precomputed is not None:
upstream_metadata=getattr(c.key, "upstream_metadata", None), account_states[kid] = precomputed
oauth_invalid_reason=getattr(c.key, "oauth_invalid_reason", None), else:
) account_states[kid] = resolve_pool_account_state(
provider_type=provider_type,
upstream_metadata=getattr(c.key, "upstream_metadata", None),
oauth_invalid_reason=getattr(c.key, "oauth_invalid_reason", None),
)
sticky_candidate: ProviderCandidate | None = None sticky_candidate: ProviderCandidate | None = None
available: list[ProviderCandidate] = [] available: list[ProviderCandidate] = []
@@ -582,16 +588,20 @@ class PoolManager:
pass pass
# --- 3. Classify keys ------------------------------------------------- # --- 3. Classify keys -------------------------------------------------
# Pre-compute account states to avoid repeated dict parsing inside loop. # Use precomputed account states when available.
account_states_sk: dict[str, Any] = {} account_states_sk: dict[str, Any] = {}
for k in keys: for k in keys:
kid = str(k.id) kid = str(k.id)
if kid not in account_states_sk: if kid not in account_states_sk:
account_states_sk[kid] = resolve_pool_account_state( precomputed = getattr(k, "_pool_account_state", None)
provider_type=self.provider_type, if precomputed is not None:
upstream_metadata=getattr(k, "upstream_metadata", None), account_states_sk[kid] = precomputed
oauth_invalid_reason=getattr(k, "oauth_invalid_reason", None), else:
) account_states_sk[kid] = resolve_pool_account_state(
provider_type=self.provider_type,
upstream_metadata=getattr(k, "upstream_metadata", None),
oauth_invalid_reason=getattr(k, "oauth_invalid_reason", None),
)
sticky_key: ProviderAPIKey | None = None sticky_key: ProviderAPIKey | None = None
available: list[ProviderAPIKey] = [] available: list[ProviderAPIKey] = []

View File

@@ -35,6 +35,9 @@ from src.models.database import (
) )
from src.services.health.monitor import health_monitor from src.services.health.monitor import health_monitor
from src.services.provider.format import normalize_endpoint_signature from src.services.provider.format import normalize_endpoint_signature
from src.services.provider.pool.account_state import (
resolve_pool_account_state as _resolve_pool_account_state,
)
from src.services.scheduling.quota_skipper import is_key_quota_exhausted from src.services.scheduling.quota_skipper import is_key_quota_exhausted
from src.services.scheduling.utils import release_db_connection_before_await from src.services.scheduling.utils import release_db_connection_before_await
@@ -97,10 +100,13 @@ class CandidateBuilder:
db.query(Provider) db.query(Provider)
.options( .options(
# 预加载 Provider 级别的 api_keys # 预加载 Provider 级别的 api_keys
# defer 排除调度热路径不需要的字段,减少内存占用 # defer 排除调度热路径不需要的大 JSON 字段,减少号池场景内存占用
# 凭证类字段(api_key/auth_config在执行阶段由 get_provider_auth() 按需加载 # - 凭证类: api_key/auth_config 在执行阶段由 get_provider_auth() 按需加载
# 注意:adjustment_history/utilization_samples 在并发检查阶段 # - adjustment_history/utilization_samples: 仅 AdaptiveReservationManager
# (AdaptiveReservationManager) 被访问,不可 defer # 在并发检查时读取单个 key号池/非号池均可 lazy load
# - upstream_metadata: 号池模式在 _build_candidates 中预计算账号封禁
# 状态并挂到 key._pool_account_state排序阶段不再需要原始 JSON
# 非号池模式 key 少lazy load 可忽略
selectinload(Provider.api_keys) selectinload(Provider.api_keys)
.defer(ProviderAPIKey.api_key) .defer(ProviderAPIKey.api_key)
.defer(ProviderAPIKey.auth_config) .defer(ProviderAPIKey.auth_config)
@@ -113,7 +119,10 @@ class CandidateBuilder:
.defer(ProviderAPIKey.last_models_fetch_at) .defer(ProviderAPIKey.last_models_fetch_at)
.defer(ProviderAPIKey.last_models_fetch_error) .defer(ProviderAPIKey.last_models_fetch_error)
.defer(ProviderAPIKey.max_probe_interval_minutes) .defer(ProviderAPIKey.max_probe_interval_minutes)
.defer(ProviderAPIKey.expires_at), .defer(ProviderAPIKey.expires_at)
.defer(ProviderAPIKey.adjustment_history)
.defer(ProviderAPIKey.utilization_samples)
.defer(ProviderAPIKey.upstream_metadata),
# 预加载 endpoints用于按 api_format 选择请求配置) # 预加载 endpoints用于按 api_format 选择请求配置)
selectinload(Provider.endpoints), selectinload(Provider.endpoints),
# 同时加载 models 和 global_model 关系 # 同时加载 models 和 global_model 关系
@@ -570,15 +579,33 @@ class CandidateBuilder:
if pool_cfg is not None: if pool_cfg is not None:
# 号池优化:跳过逐 key 的 _check_key_availability 检查, # 号池优化:跳过逐 key 的 _check_key_availability 检查,
# 直接收集全部 active key将检查推迟到 PoolManager 排序后分页执行。 # 直接收集全部 active key将检查推迟到 PoolManager 排序后分页执行。
# 在此释放 DB 连接,因为后续的 PoolManager 排序涉及大量 Redis 操作,
# 避免在 Redis I/O 期间长时间占用 DB 连接池。
release_db_connection_before_await(db)
pool_keys = list(keys_to_check) pool_keys = list(keys_to_check)
if not pool_keys: if not pool_keys:
continue continue
# 在释放 DB 连接前预计算账号封禁状态并挂到 Key 对象上。
# upstream_metadata 是 deferred 字段,逐条 lazy load 会产生 N+1 查询;
# 这里集中触发后PoolManager 排序时直接读取 _pool_account_state 即可。
provider_type_str = (
str(getattr(provider, "provider_type", "") or "").strip().lower() or None
)
for pk in pool_keys:
setattr(
pk,
"_pool_account_state",
_resolve_pool_account_state(
provider_type=provider_type_str,
upstream_metadata=getattr(pk, "upstream_metadata", None),
oauth_invalid_reason=getattr(pk, "oauth_invalid_reason", None),
),
)
# 释放 DB 连接,因为后续的 PoolManager 排序涉及大量 Redis 操作,
# 避免在 Redis I/O 期间长时间占用 DB 连接池。
release_db_connection_before_await(db)
provider_priority_raw = getattr(provider, "provider_priority", None) provider_priority_raw = getattr(provider, "provider_priority", None)
try: try:
provider_priority = ( provider_priority = (