refactor: 优化模型获取调度器并发模型和缓存格式

- 用固定 worker 数的队列消费模式替换 Semaphore+gather,避免大号池一次性创建大量协程
- 分批扫描 Key ID(keyset pagination),避免一次性加载全量 ID
- 上游模型缓存从 per-api_format 去重改为 per-model-id 聚合 api_formats 列表,减少 Redis 占用
- 启动阶段支持环境变量控制(MODEL_FETCH_STARTUP_ENABLED / MODEL_FETCH_STARTUP_DELAY_SECONDS)
- 新增单元测试覆盖聚合逻辑和并发上限验证
This commit is contained in:
fawney19
2026-03-10 12:47:50 +08:00
parent 9ee27308db
commit 2b21a75982
2 changed files with 307 additions and 61 deletions

View File

@@ -15,6 +15,7 @@ import asyncio
import fnmatch import fnmatch
import json import json
import os import os
from collections.abc import Awaitable, Callable
from dataclasses import dataclass from dataclasses import dataclass
from datetime import datetime, timezone from datetime import datetime, timezone
from typing import Any from typing import Any
@@ -52,6 +53,18 @@ KEY_FETCH_TIMEOUT_SECONDS = 120
# 使用较短的超时10秒避免不支持 /models 端点的提供商长时间阻塞 # 使用较短的超时10秒避免不支持 /models 端点的提供商长时间阻塞
MODEL_FETCH_HTTP_TIMEOUT = 10.0 MODEL_FETCH_HTTP_TIMEOUT = 10.0
# 启动时首次自动获取开关与延迟
MODEL_FETCH_STARTUP_ENABLED = (
os.getenv("MODEL_FETCH_STARTUP_ENABLED", "true").lower() == "true"
)
MODEL_FETCH_STARTUP_DELAY_SECONDS = max(
0,
int(os.getenv("MODEL_FETCH_STARTUP_DELAY_SECONDS", "10")),
)
# 单批扫描的 Key 数量(仅拉取 ID避免一次性扫描整个大号池
AUTO_FETCH_KEY_BATCH_SIZE = max(MAX_CONCURRENT_REQUESTS, 100)
# 上游模型缓存 TTL与定时任务间隔保持一致 # 上游模型缓存 TTL与定时任务间隔保持一致
UPSTREAM_MODELS_CACHE_TTL_SECONDS = MODEL_FETCH_INTERVAL_MINUTES * 60 UPSTREAM_MODELS_CACHE_TTL_SECONDS = MODEL_FETCH_INTERVAL_MINUTES * 60
@@ -155,6 +168,122 @@ async def set_upstream_models_to_cache(
logger.debug(f"上游模型已缓存: {cache_key}, 数量={len(models)}") logger.debug(f"上游模型已缓存: {cache_key}, 数量={len(models)}")
def _aggregate_models_for_cache(models: list[dict]) -> list[dict]:
"""聚合缓存模型,按 model id 合并 api_formats减少 Redis 占用。"""
aggregated: dict[str, dict[str, Any]] = {}
ordered_ids: list[str] = []
for model in models:
if not isinstance(model, dict):
continue
model_id = str(model.get("id") or "").strip()
if not model_id:
continue
api_format = str(model.get("api_format") or "").strip()
existing = aggregated.get(model_id)
if existing is None:
payload = dict(model)
payload.pop("api_format", None)
payload["api_formats"] = [api_format] if api_format else []
aggregated[model_id] = payload
ordered_ids.append(model_id)
continue
api_formats = existing.setdefault("api_formats", [])
if not isinstance(api_formats, list):
api_formats = []
existing["api_formats"] = api_formats
if api_format and api_format not in api_formats:
api_formats.append(api_format)
for key, value in model.items():
if key not in existing and key != "api_format":
existing[key] = value
for model_id in ordered_ids:
api_formats = aggregated[model_id].get("api_formats")
if isinstance(api_formats, list):
aggregated[model_id]["api_formats"] = sorted(
{str(fmt) for fmt in api_formats if str(fmt).strip()}
)
return [aggregated[model_id] for model_id in ordered_ids]
async def _run_key_fetch_workers(
key_ids: list[str],
*,
max_concurrent: int,
timeout_seconds: float,
running_predicate: Callable[[], bool],
fetch_one: Callable[[str], Awaitable[str]],
on_timeout: Callable[[str], None],
on_error: Callable[[str, str], None],
) -> tuple[int, int, int]:
"""用固定 worker 数处理 key避免一次性创建大量协程导致内存峰值过高。"""
if not key_ids:
return 0, 0, 0
worker_count = max(1, min(max_concurrent, len(key_ids)))
key_queue: asyncio.Queue[str] = asyncio.Queue()
for key_id in key_ids:
key_queue.put_nowait(key_id)
async def _worker() -> tuple[int, int, int]:
success_count = 0
error_count = 0
skip_count = 0
while True:
if not running_predicate():
break
try:
key_id = key_queue.get_nowait()
except asyncio.QueueEmpty:
break
result = "error"
try:
if not running_predicate():
result = "skip"
else:
result = await asyncio.wait_for(fetch_one(key_id), timeout=timeout_seconds)
except TimeoutError:
logger.error(f"处理 Key {key_id} 超时({timeout_seconds}s")
on_timeout(key_id)
result = "error"
except Exception as exc:
logger.exception(f"处理 Key {key_id} 时出错")
on_error(key_id, str(exc))
result = "error"
finally:
key_queue.task_done()
if result == "success":
success_count += 1
elif result == "skip":
skip_count += 1
else:
error_count += 1
return success_count, error_count, skip_count
results = await asyncio.gather(*[asyncio.create_task(_worker()) for _ in range(worker_count)])
success_count = sum(success for success, _, _ in results)
error_count = sum(error for _, error, _ in results)
skip_count = sum(skip for _, _, skip in results)
# 停止过程中尚未消费的队列项统一计为 skip。
skip_count += key_queue.qsize()
return success_count, error_count, skip_count
class ModelFetchScheduler: class ModelFetchScheduler:
"""模型自动获取调度器""" """模型自动获取调度器"""
@@ -200,7 +329,13 @@ class ModelFetchScheduler:
async def _run_startup_task(self) -> None: async def _run_startup_task(self) -> None:
"""启动时执行的初始化任务""" """启动时执行的初始化任务"""
try: try:
await asyncio.sleep(10) # 等待系统完全启动 if not MODEL_FETCH_STARTUP_ENABLED:
logger.info("启动时模型自动获取已禁用MODEL_FETCH_STARTUP_ENABLED=false")
return
if MODEL_FETCH_STARTUP_DELAY_SECONDS > 0:
await asyncio.sleep(MODEL_FETCH_STARTUP_DELAY_SECONDS)
if not self._running: if not self._running:
return return
logger.info("启动时执行首次模型获取...") logger.info("启动时执行首次模型获取...")
@@ -216,68 +351,77 @@ class ModelFetchScheduler:
async with self._lock: async with self._lock:
await self._perform_fetch_all_keys() await self._perform_fetch_all_keys()
def _list_auto_fetch_key_id_batch(
self,
*,
after_id: str | None = None,
limit: int = AUTO_FETCH_KEY_BATCH_SIZE,
) -> list[str]:
"""分批返回启用自动获取模型的 Key ID。"""
with create_session() as db:
query = db.query(ProviderAPIKey.id).filter(
ProviderAPIKey.auto_fetch_models == True, # noqa: E712
ProviderAPIKey.is_active == True, # noqa: E712
)
if after_id:
query = query.filter(ProviderAPIKey.id > after_id)
return [row[0] for row in query.order_by(ProviderAPIKey.id.asc()).limit(limit).all()]
async def _perform_fetch_all_keys(self) -> None: async def _perform_fetch_all_keys(self) -> None:
"""获取所有启用自动获取的 Key 并拉取模型""" """获取所有启用自动获取的 Key,并以固定批次/并发节奏拉取。"""
logger.info("开始自动获取模型任务...") logger.info("开始自动获取模型任务...")
# 统计信息
success_count = 0 success_count = 0
error_count = 0 error_count = 0
skip_count = 0 skip_count = 0
total_count = 0
last_id: str | None = None
batch_index = 0
with create_session() as db: while self._running:
# 查询所有启用了 auto_fetch_models 的 Key只获取 ID 列表) key_ids = self._list_auto_fetch_key_id_batch(after_id=last_id)
key_ids = [ if not key_ids:
row[0] break
for row in db.query(ProviderAPIKey.id)
.filter(
ProviderAPIKey.auto_fetch_models == True, # noqa: E712
ProviderAPIKey.is_active == True, # noqa: E712
)
.all()
]
if not key_ids: batch_index += 1
total_count += len(key_ids)
last_id = key_ids[-1]
logger.info(
"自动获取模型任务处理第 {} 批 Key: {}",
batch_index,
len(key_ids),
)
batch_success, batch_error, batch_skip = await _run_key_fetch_workers(
key_ids,
max_concurrent=MAX_CONCURRENT_REQUESTS,
timeout_seconds=KEY_FETCH_TIMEOUT_SECONDS,
running_predicate=lambda: self._running,
fetch_one=self._fetch_models_for_key_by_id,
on_timeout=lambda key_id: self._update_key_error(
key_id, f"Timeout after {KEY_FETCH_TIMEOUT_SECONDS}s"
),
on_error=self._update_key_error,
)
success_count += batch_success
error_count += batch_error
skip_count += batch_skip
if len(key_ids) < AUTO_FETCH_KEY_BATCH_SIZE:
break
await asyncio.sleep(0)
if total_count == 0:
logger.debug("没有启用自动获取模型的 Key") logger.debug("没有启用自动获取模型的 Key")
return return
logger.info(f"找到 {len(key_ids)} 个启用自动获取模型的 Key")
# 使用 Semaphore 限制并发数,避免大号池场景下同时发起过多 HTTP 请求
semaphore = asyncio.Semaphore(MAX_CONCURRENT_REQUESTS)
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( logger.info(
f"自动获取模型任务完成: 成功={success_count}, 失败={error_count}, 跳过={skip_count}" "自动获取模型任务完成: 总数={}, 成功={}, 失败={}, 跳过={}",
total_count,
success_count,
error_count,
skip_count,
) )
def _update_key_error(self, key_id: str, error_msg: str) -> None: def _update_key_error(self, key_id: str, error_msg: str) -> None:
@@ -554,16 +698,8 @@ class ModelFetchScheduler:
f"Provider {provider_name} Key {key.id} 获取到 {len(fetched_model_ids)} 个唯一模型" f"Provider {provider_name} Key {key.id} 获取到 {len(fetched_model_ids)} 个唯一模型"
) )
# 写入上游模型缓存(按 model id + api_format 去重后的完整模型信息 # 写入上游模型缓存(按 model id 聚合 api_formats减少 Redis 内存占用
seen_keys: set[str] = set() unique_models = _aggregate_models_for_cache(all_models)
unique_models: list[dict] = []
for model in all_models:
model_id = model.get("id")
api_format = model.get("api_format", "")
unique_key = f"{model_id}:{api_format}"
if model_id and unique_key not in seen_keys:
seen_keys.add(unique_key)
unique_models.append(model)
await set_upstream_models_to_cache(provider_id, key.id, unique_models) await set_upstream_models_to_cache(provider_id, key.id, unique_models)
# 更新 allowed_models保留 locked_models # 更新 allowed_models保留 locked_models

View File

@@ -0,0 +1,110 @@
from __future__ import annotations
import asyncio
import pytest
import src.services.model.fetch_scheduler as fetch_scheduler_module
from src.services.model.fetch_scheduler import (
ModelFetchScheduler,
_aggregate_models_for_cache,
_run_key_fetch_workers,
)
def test_aggregate_models_for_cache_merges_formats_by_model_id() -> None:
models: list[dict] = [
{"id": "gpt-4.1", "api_format": "openai:chat", "label": "GPT 4.1"},
{"id": "gpt-4.1", "api_format": "openai:cli", "extra": {"tier": "pro"}},
{"id": "claude-sonnet", "api_format": "claude:chat", "label": "Sonnet"},
{"id": "", "api_format": "ignored"},
]
aggregated = _aggregate_models_for_cache(models)
assert len(aggregated) == 2
by_id = {item["id"]: item for item in aggregated}
assert by_id["gpt-4.1"]["api_formats"] == ["openai:chat", "openai:cli"]
assert by_id["gpt-4.1"]["label"] == "GPT 4.1"
assert by_id["gpt-4.1"]["extra"] == {"tier": "pro"}
assert "api_format" not in by_id["gpt-4.1"]
assert by_id["claude-sonnet"]["api_formats"] == ["claude:chat"]
@pytest.mark.asyncio
async def test_run_key_fetch_workers_caps_inflight_tasks() -> None:
inflight = 0
max_inflight = 0
async def fetch_one(key_id: str) -> str:
nonlocal inflight, max_inflight
inflight += 1
max_inflight = max(max_inflight, inflight)
await asyncio.sleep(0.01)
inflight -= 1
return "success"
timed_out: list[str] = []
failed: list[tuple[str, str]] = []
result = await _run_key_fetch_workers(
[f"key-{index}" for index in range(8)],
max_concurrent=3,
timeout_seconds=1.0,
running_predicate=lambda: True,
fetch_one=fetch_one,
on_timeout=timed_out.append,
on_error=lambda key_id, message: failed.append((key_id, message)),
)
assert result == (8, 0, 0)
assert max_inflight <= 3
assert timed_out == []
assert failed == []
@pytest.mark.asyncio
async def test_perform_fetch_all_keys_scans_in_batches(monkeypatch: pytest.MonkeyPatch) -> None:
scheduler = ModelFetchScheduler()
scheduler._running = True
batch_requests: list[str | None] = []
processed_batches: list[list[str]] = []
pages = {
None: ["a", "b"],
"b": ["c"],
"c": [],
}
def fake_list_batch(*, after_id: str | None = None, limit: int = 0) -> list[str]:
batch_requests.append(after_id)
return list(pages.get(after_id, []))
async def fake_run_key_fetch_workers(
key_ids: list[str],
*,
max_concurrent: int,
timeout_seconds: float,
running_predicate,
fetch_one,
on_timeout,
on_error,
) -> tuple[int, int, int]:
processed_batches.append(list(key_ids))
return len(key_ids), 0, 0
monkeypatch.setattr(fetch_scheduler_module, "AUTO_FETCH_KEY_BATCH_SIZE", 2)
monkeypatch.setattr(scheduler, "_list_auto_fetch_key_id_batch", fake_list_batch)
monkeypatch.setattr(
fetch_scheduler_module,
"_run_key_fetch_workers",
fake_run_key_fetch_workers,
)
await scheduler._perform_fetch_all_keys()
assert batch_requests == [None, "b"]
assert processed_batches == [["a", "b"], ["c"]]