mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
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:
40
src/main.py
40
src/main.py
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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] = {}
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user