mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
refactor: 优化调度器并发与内存占用,修复循环依赖
- fetch_scheduler: 模型获取从串行改为 Semaphore 限并发并行 - pool_quota_probe_scheduler: 取消启动时立即探测避免阻塞, 改为逐 provider 独立查询 key 避免一次性加载全部到内存, 移除未使用的 _ProviderProbeTask 数据类 - proxy_node/__init__: 移除 health_scheduler 导入避免循环依赖
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user