refactor: 优化调度器并发与内存占用,修复循环依赖

- fetch_scheduler: 模型获取从串行改为 Semaphore 限并发并行
- pool_quota_probe_scheduler: 取消启动时立即探测避免阻塞,
  改为逐 provider 独立查询 key 避免一次性加载全部到内存,
  移除未使用的 _ProviderProbeTask 数据类
- proxy_node/__init__: 移除 health_scheduler 导入避免循环依赖
This commit is contained in:
fawney19
2026-03-10 11:22:11 +08:00
parent 86449cae52
commit 3063938a82
3 changed files with 118 additions and 115 deletions

View File

@@ -19,7 +19,7 @@ from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any
from sqlalchemy.orm import Session, joinedload
from sqlalchemy.orm import joinedload
from src.core.cache_service import CacheService
from src.core.crypto import crypto_service
@@ -243,30 +243,37 @@ class ModelFetchScheduler:
logger.info(f"找到 {len(key_ids)} 个启用自动获取模型的 Key")
# 逐个处理每个 Key每个 Key 使用独立的数据库会话
for key_id in key_ids:
if not self._running:
logger.info("调度器已停止,中断模型获取任务")
break
# 使用 Semaphore 限制并发数,避免大号池场景下同时发起过多 HTTP 请求
semaphore = asyncio.Semaphore(MAX_CONCURRENT_REQUESTS)
try:
# 添加超时保护
result = await asyncio.wait_for(
self._fetch_models_for_key_by_id(key_id),
timeout=KEY_FETCH_TIMEOUT_SECONDS,
)
if result == "success":
success_count += 1
elif result == "skip":
skip_count += 1
else:
error_count += 1
except TimeoutError:
logger.error(f"处理 Key {key_id} 超时({KEY_FETCH_TIMEOUT_SECONDS}s")
self._update_key_error(key_id, f"Timeout after {KEY_FETCH_TIMEOUT_SECONDS}s")
error_count += 1
except Exception:
logger.exception(f"处理 Key {key_id} 时出错")
async def _fetch_with_limit(key_id: str) -> str:
if not self._running:
return "skip"
async with semaphore:
if not self._running:
return "skip"
try:
return await asyncio.wait_for(
self._fetch_models_for_key_by_id(key_id),
timeout=KEY_FETCH_TIMEOUT_SECONDS,
)
except TimeoutError:
logger.error(f"处理 Key {key_id} 超时({KEY_FETCH_TIMEOUT_SECONDS}s")
self._update_key_error(key_id, f"Timeout after {KEY_FETCH_TIMEOUT_SECONDS}s")
return "error"
except Exception as exc:
logger.exception(f"处理 Key {key_id} 时出错")
self._update_key_error(key_id, str(exc))
return "error"
results = await asyncio.gather(*[_fetch_with_limit(kid) for kid in key_ids])
for result in results:
if result == "success":
success_count += 1
elif result == "skip":
skip_count += 1
else:
error_count += 1
logger.info(

View File

@@ -11,7 +11,6 @@ from __future__ import annotations
import os
import time
from dataclasses import dataclass
from datetime import datetime, timezone
from typing import Any
@@ -139,14 +138,6 @@ def _select_probe_key_ids(
return [key_id for _, key_id in stale]
@dataclass(frozen=True, slots=True)
class _ProviderProbeTask:
provider_id: str
provider_type: str
probe_key_ids: list[str]
interval_seconds: int
class PoolQuotaProbeScheduler:
"""按号池高级配置执行额度主动探测。"""
@@ -185,9 +176,8 @@ class PoolQuotaProbeScheduler:
job_id="pool_quota_probe_check",
name="号池额度主动探测检查",
)
# 启动时立即执行一次,避免首次等待一个轮询周期
await self._run_probe_cycle()
# 不在启动时立即探测:大号池场景下会阻塞启动并占用大量内存。
# 首次探测由定时调度器在 scan_interval_seconds 后自动触发。
async def stop(self) -> Any:
if not self.running:
@@ -248,14 +238,11 @@ class PoolQuotaProbeScheduler:
now_ts = int(time.time())
redis_client = await get_redis_client(require_redis=False)
# 第一阶段:用一个短生命周期 session 查出需要探测的 provider / key 信息
probe_tasks: list[_ProviderProbeTask] = []
# 第一阶段:查出符合条件的 provider 列表(轻量查询)
eligible_providers: list[tuple[str, str, int]] = [] # (id, type, interval_seconds)
db = create_session()
try:
providers = db.query(Provider).filter(Provider.is_active == True).all() # noqa: E712
# 先筛选出符合条件的 provider收集其 ID 和配置
eligible_providers: list[tuple[str, str, int]] = [] # (id, type, interval_seconds)
for provider in providers:
provider_id = str(getattr(provider, "id", "") or "")
provider_type = normalize_provider_type(getattr(provider, "provider_type", ""))
@@ -270,91 +257,50 @@ class PoolQuotaProbeScheduler:
pool_cfg.probing_interval_minutes
)
eligible_providers.append((provider_id, provider_type, interval_minutes * 60))
# 批量查询所有符合条件的 provider 的活跃 keys避免 N+1
eligible_ids = [p[0] for p in eligible_providers]
all_keys: list[ProviderAPIKey] = []
if eligible_ids:
all_keys = (
db.query(ProviderAPIKey)
.options(
load_only(
ProviderAPIKey.id,
ProviderAPIKey.provider_id,
ProviderAPIKey.last_used_at,
ProviderAPIKey.upstream_metadata,
)
)
.filter(
ProviderAPIKey.provider_id.in_(eligible_ids),
ProviderAPIKey.is_active == True, # noqa: E712
)
.all()
)
# 按 provider_id 分组
keys_by_provider: dict[str, list[ProviderAPIKey]] = {}
for key in all_keys:
pid = str(key.provider_id)
keys_by_provider.setdefault(pid, []).append(key)
for provider_id, provider_type, interval_seconds in eligible_providers:
keys = keys_by_provider.get(provider_id, [])
if not keys:
continue
key_ids = [str(key.id) for key in keys if getattr(key, "id", None)]
probe_stamps = await self._load_probe_timestamps(
redis_client=redis_client,
provider_id=provider_id,
key_ids=key_ids,
)
probe_key_ids = _select_probe_key_ids(
keys=keys,
provider_type=provider_type,
now_ts=now_ts,
interval_seconds=interval_seconds,
last_probe_timestamps=probe_stamps,
limit=self.max_keys_per_provider,
)
if not probe_key_ids:
continue
probe_tasks.append(
_ProviderProbeTask(
provider_id=provider_id,
provider_type=provider_type,
probe_key_ids=probe_key_ids,
interval_seconds=interval_seconds,
)
)
finally:
db.close()
# 第二阶段:每个 provider 使用独立 session 执行探测
for task in probe_tasks:
if not eligible_providers:
return
# 第二阶段:逐个 provider 查询 key 并筛选探测目标
# 避免一次性加载所有 provider 的全部 key 到内存
for provider_id, provider_type, interval_seconds in eligible_providers:
if not self.running:
break
probe_key_ids = await self._select_keys_for_provider(
provider_id=provider_id,
provider_type=provider_type,
interval_seconds=interval_seconds,
now_ts=now_ts,
redis_client=redis_client,
)
if not probe_key_ids:
continue
# 先写探测节流时间戳,避免异常时高频重入
await self._mark_probe_timestamps(
redis_client=redis_client,
provider_id=task.provider_id,
key_ids=task.probe_key_ids,
provider_id=provider_id,
key_ids=probe_key_ids,
now_ts=now_ts,
interval_seconds=task.interval_seconds,
interval_seconds=interval_seconds,
)
probe_db = create_session()
try:
result = await refresh_provider_quota_for_provider(
db=probe_db,
provider_id=task.provider_id,
provider_id=provider_id,
codex_wham_usage_url=CODEX_WHAM_USAGE_URL,
key_ids=task.probe_key_ids,
key_ids=probe_key_ids,
)
logger.info(
"[POOL_PROBE] Provider {} ({}) 静默探测完成: selected={}, success={}, failed={}",
task.provider_id[:8],
task.provider_type,
len(task.probe_key_ids),
provider_id[:8],
provider_type,
len(probe_key_ids),
int(result.get("success") or 0),
int(result.get("failed") or 0),
)
@@ -365,13 +311,61 @@ class PoolQuotaProbeScheduler:
pass
logger.warning(
"[POOL_PROBE] Provider {} ({}) 静默探测失败: {}",
task.provider_id[:8],
task.provider_type,
provider_id[:8],
provider_type,
exc,
)
finally:
probe_db.close()
async def _select_keys_for_provider(
self,
*,
provider_id: str,
provider_type: str,
interval_seconds: int,
now_ts: int,
redis_client: Any,
) -> list[str]:
"""为单个 provider 筛选需要探测的 key使用独立短生命周期 session。"""
db = create_session()
try:
keys = (
db.query(ProviderAPIKey)
.options(
load_only(
ProviderAPIKey.id,
ProviderAPIKey.provider_id,
ProviderAPIKey.last_used_at,
ProviderAPIKey.upstream_metadata,
)
)
.filter(
ProviderAPIKey.provider_id == provider_id,
ProviderAPIKey.is_active == True, # noqa: E712
)
.all()
)
if not keys:
return []
key_ids = [str(key.id) for key in keys if getattr(key, "id", None)]
probe_stamps = await self._load_probe_timestamps(
redis_client=redis_client,
provider_id=provider_id,
key_ids=key_ids,
)
return _select_probe_key_ids(
keys=keys,
provider_type=provider_type,
now_ts=now_ts,
interval_seconds=interval_seconds,
last_probe_timestamps=probe_stamps,
limit=self.max_keys_per_provider,
)
finally:
db.close()
_pool_quota_probe_scheduler: PoolQuotaProbeScheduler | None = None

View File

@@ -1,6 +1,10 @@
"""代理节点服务"""
"""代理节点服务
注意: health_scheduler 不在此处导入,避免触发 system.scheduler -> system.__init__
-> maintenance_scheduler -> ... -> cache_service 的循环依赖。
使用方应直接 from src.services.proxy_node.health_scheduler import ... 导入。
"""
from .health_scheduler import ProxyNodeHealthScheduler, get_proxy_node_health_scheduler
from .resolver import (
build_post_kwargs,
build_proxy_url,
@@ -21,8 +25,6 @@ from .resolver import (
from .service import ProxyNodeService, node_to_dict
__all__ = [
"ProxyNodeHealthScheduler",
"get_proxy_node_health_scheduler",
"ProxyNodeService",
"node_to_dict",
"build_post_kwargs",