refactor: 优化调度器内存占用,增加 session 全局清理

- main: 启动阶段 Provider 查询改用聚合查询,避免全量加载 ORM 对象
- envelope: Claude Code session 增加全局定时清理,防止 scope_key 无界增长
- health_cache: 增量更新时清理已移除的 stale key 条目
- aware_scheduler: 日志中用 key.name 替代 api_key 后四位,避免泄露凭证片段
- candidate_builder: defer api_key/auth_config 等冷字段,减少调度热路径内存占用
This commit is contained in:
fawney19
2026-03-10 11:40:29 +08:00
parent 3063938a82
commit 1518223de6
5 changed files with 81 additions and 27 deletions

View File

@@ -51,40 +51,38 @@ if TYPE_CHECKING:
async def initialize_providers() -> None:
"""从数据库初始化提供商(仅用于日志记录)"""
from sqlalchemy.orm import Session, selectinload
"""从数据库初始化提供商(仅用于日志记录,使用轻量查询"""
from sqlalchemy import func
from sqlalchemy.orm import Session
from src.database.database import create_session
from src.models.database import Provider
from src.models.database import Provider, ProviderEndpoint
try:
# 创建数据库会话
db: Session = create_session()
try:
# 从数据库加载所有活跃的提供商(使用 selectinload 预加载 endpoints 避免 N+1 查询)
providers = (
db.query(Provider)
.options(selectinload(Provider.endpoints))
# 使用聚合查询代替全量加载 ORM 对象,大幅减少内存占用
results = (
db.query(
Provider.name,
func.count(ProviderEndpoint.id).label("total"),
func.count(func.nullif(ProviderEndpoint.is_active, False)).label("active"),
)
.outerjoin(ProviderEndpoint, Provider.id == ProviderEndpoint.provider_id)
.filter(Provider.is_active.is_(True))
.group_by(Provider.id, Provider.name, Provider.provider_priority)
.order_by(Provider.provider_priority.asc())
.all()
)
if not providers:
if not results:
logger.warning("数据库中未找到活跃的提供商")
return
# 记录提供商信息
logger.info(f"从数据库加载了 {len(providers)} 个活跃提供商")
for provider in providers:
# 统计端点信息
endpoint_count = len(provider.endpoints) if provider.endpoints else 0 # type: ignore[arg-type]
active_endpoints = (
sum(1 for ep in provider.endpoints if ep.is_active) if provider.endpoints else 0 # type: ignore[misc,attr-defined]
)
logger.info(f"提供商: {provider.name} (端点: {active_endpoints}/{endpoint_count})")
logger.info(f"从数据库加载了 {len(results)} 个活跃提供商")
for name, total, active in results:
logger.info(f"提供商: {name} (端点: {active}/{total})")
finally:
db.close()
@@ -317,7 +315,9 @@ async def _start_background_services(state: LifecycleState) -> None:
state.model_fetch_scheduler = None
# 启动号池额度主动探测调度器
pool_quota_probe_scheduler_active = await state.task_coordinator.acquire("pool_quota_probe_scheduler")
pool_quota_probe_scheduler_active = await state.task_coordinator.acquire(
"pool_quota_probe_scheduler"
)
if pool_quota_probe_scheduler_active:
logger.info("启动号池额度主动探测调度器...")
await state.pool_quota_probe_scheduler.start()

View File

@@ -35,6 +35,11 @@ _session_runtime_lock = threading.Lock()
_active_sessions: dict[str, dict[str, float]] = {}
# key: scope_key -> (masked_session_uuid, expire_at_monotonic)
_masked_sessions: dict[str, tuple[str, float]] = {}
# 上限保护scope_key 总数超过此值时触发全局清理
_MAX_SCOPE_KEYS = 5000
# 全局清理间隔monotonic 秒),避免高频请求时每次都做全局扫描
_LAST_GLOBAL_CLEANUP: float = 0.0
_GLOBAL_CLEANUP_INTERVAL = 300.0 # 5 分钟
_REDIS_SESSION_KEY_PREFIX = "claude_code:sessions"
_REDIS_SESSION_RESERVE_LUA = """
local key = KEYS[1]
@@ -289,6 +294,40 @@ def _apply_cache_ttl_override(request_body: dict[str, Any], target: str) -> None
logger.debug("Cache TTL override: {} block(s) -> {}", overridden, target)
def _cleanup_stale_scope_keys(now: float, idle_seconds: int) -> None:
"""清理空 bucket 和过期的 _masked_sessions 条目(调用方须持有锁)。"""
global _LAST_GLOBAL_CLEANUP
if now - _LAST_GLOBAL_CLEANUP < _GLOBAL_CLEANUP_INTERVAL:
return
_LAST_GLOBAL_CLEANUP = now
# 清理所有 scope_key 下的过期 session并删除空 bucket
stale_keys = []
for sk, bucket in _active_sessions.items():
expired = [sid for sid, ts in bucket.items() if now - ts > idle_seconds]
for sid in expired:
bucket.pop(sid, None)
if not bucket:
stale_keys.append(sk)
for sk in stale_keys:
_active_sessions.pop(sk, None)
# 清理过期的 masked session 条目
expired_masked = [sk for sk, (_, exp) in _masked_sessions.items() if exp <= now]
for sk in expired_masked:
_masked_sessions.pop(sk, None)
total = len(_active_sessions) + len(_masked_sessions)
if total > 0 and (stale_keys or expired_masked):
logger.debug(
"Session 全局清理: 移除 {} 个 active scope + {} 个 masked scope, 剩余 {}",
len(stale_keys),
len(expired_masked),
total,
)
def _register_or_reject_session(
*,
scope_key: str,
@@ -300,9 +339,16 @@ def _register_or_reject_session(
idle_seconds = max(60, int(idle_timeout_minutes * 60))
with _session_runtime_lock:
# scope_key 总数超限或达到清理间隔时,执行全局清理
if (
len(_active_sessions) + len(_masked_sessions) > _MAX_SCOPE_KEYS
or now - _LAST_GLOBAL_CLEANUP >= _GLOBAL_CLEANUP_INTERVAL
):
_cleanup_stale_scope_keys(now, idle_seconds)
bucket = _active_sessions.setdefault(scope_key, {})
# 先清理过期会话,避免误判占用。
# 先清理当前 bucket 的过期会话,避免误判占用。
expired = [sid for sid, last_seen in bucket.items() if now - last_seen > idle_seconds]
for sid in expired:
bucket.pop(sid, None)

View File

@@ -62,6 +62,10 @@ def get_health_scores(provider_id: str, keys: list[Any]) -> dict[str, float]:
payload[kid] = aggregate_health_score(
getattr(keys_by_id[kid], "health_by_format", None)
)
# Trim stale key entries that are no longer in the current key set
stale_ids = [sid for sid in payload if sid not in keys_by_id]
for sid in stale_ids:
del payload[sid]
return {kid: payload[kid] for kid in keys_by_id}
fresh: dict[str, float] = {}

View File

@@ -316,11 +316,10 @@ class CacheAwareScheduler:
continue
logger.debug(
" └─ 选择 Provider={}, Endpoint={}..., "
"Key=***{}, 缓存命中={}, 并发状态[{}]",
" └─ 选择 Provider={}, Endpoint={}..., " "Key={}, 缓存命中={}, 并发状态[{}]",
provider.name,
endpoint.id[:8],
key.api_key[-4:],
key.name,
is_cached_user,
snapshot.describe(),
)
@@ -669,13 +668,13 @@ class CacheAwareScheduler:
"检测到缓存亲和性: affinity_key={}..., "
"api_format={}, global_model_id={}..., "
"provider={}, endpoint={}..., "
"provider_key=***{}, 使用次数={}",
"provider_key={}, 使用次数={}",
affinity_key[:8],
api_format_str,
global_model_id[:8],
provider.name,
endpoint.id[:8],
key.api_key[-4:],
key.name,
affinity.request_count,
)
else:

View File

@@ -97,8 +97,13 @@ class CandidateBuilder:
db.query(Provider)
.options(
# 预加载 Provider 级别的 api_keys
# defer 排除仅后台管理/模型获取用的冷字段,热路径字段全部加载
# defer 排除调度热路径不需要的字段,减少内存占用
# 凭证类字段api_key/auth_config在执行阶段才由 get_provider_auth() 按需加载
# 注意adjustment_history/utilization_samples 在并发检查阶段
# (AdaptiveReservationManager) 被访问,不可 defer
selectinload(Provider.api_keys)
.defer(ProviderAPIKey.api_key)
.defer(ProviderAPIKey.auth_config)
.defer(ProviderAPIKey.note)
.defer(ProviderAPIKey.last_error_msg)
.defer(ProviderAPIKey.auto_fetch_models)