From 8c87b9d0249cf5c5e62fb223358552a84b9200a9 Mon Sep 17 00:00:00 2001 From: shangfangjian Date: Mon, 21 Sep 2026 17:04:04 +0800 Subject: [PATCH 1/7] =?UTF-8?q?fix(critical):=20=E4=BF=AE=E5=A4=8D?= =?UTF-8?q?=E8=B0=83=E5=BA=A6=E5=99=A8=E9=9D=99=E9=BB=98=E5=A4=B1=E6=95=88?= =?UTF-8?q?=E3=80=81=E6=95=B0=E6=8D=AE=E6=BA=90=E7=BC=BA=E5=A4=B1=E4=B8=8E?= =?UTF-8?q?=E5=89=8D=E7=AB=AF=E7=B1=BB=E5=9E=8B=E9=94=99=E8=AF=AF?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit C1 定时任务静默失效(4 层缺陷): - scheduler.py: croniter(expr, datetime.now) 未调用 now(), 导致 TypeError 被吞、next_run 恒为 None(最深一层) - scheduler.py: _calc_next 非法 cron 静默吞异常 -> 改为构造期抛 ValueError - scheduler.py: _run_loop 一次异常即永久停摆 -> 增加异常隔离 - scheduler.py: sleep(min(wait,60)) 使运行期 cron 变更最长 60s 才生效 -> 改为固定 SLEEP_TICK 轮询 - app.py / schedules.py: 3 处 register() 第 2 参误传 task 类型名 C3 EvalPage.tsx: 引用未导入的 EmptyState -> 改为已导入的 EmptyText (tsc --noEmit 由 1 error 变为 0 error) 新增 tests/test_scheduler_registration.py: 9 项行为回归测试 --- frontend/src/admin/pages/EvalPage.tsx | 2 +- src/api/app.py | 2 +- src/api/routes/schedules.py | 8 +- src/core/scheduler.py | 118 ++++++++++++---- tests/test_scheduler_registration.py | 190 ++++++++++++++++++++++++++ 5 files changed, 285 insertions(+), 35 deletions(-) create mode 100644 tests/test_scheduler_registration.py diff --git a/frontend/src/admin/pages/EvalPage.tsx b/frontend/src/admin/pages/EvalPage.tsx index 220c2bf..1d4bc01 100644 --- a/frontend/src/admin/pages/EvalPage.tsx +++ b/frontend/src/admin/pages/EvalPage.tsx @@ -163,7 +163,7 @@ export default function EvalPage() { ) : summary.length === 0 ? ( - + ) : (
AsyncIterator[None]: for s in schedules: leagues = s.leagues.split(",") if s.leagues else list(BZZOIRO_LEAGUE_IDS.keys()) scheduler.register( - s.id, s.task, + s.id, s.cron, lambda sid=s.id: _run_scheduled_task(sid), enabled=s.enabled, ) diff --git a/src/api/routes/schedules.py b/src/api/routes/schedules.py index a00327f..75f1992 100644 --- a/src/api/routes/schedules.py +++ b/src/api/routes/schedules.py @@ -91,8 +91,8 @@ async def create_schedule(req: ScheduleIn, db: AsyncSession = Depends(get_db_rea db.add(sched) await db.commit() - # 注册到调度器 - scheduler.register(req.id, req.task, lambda: _run_scheduled_task(req.id), enabled=req.enabled) + # 注册到调度器(注意: 第 2 参是 cron 表达式,不是 task 类型名) + scheduler.register(req.id, req.cron, lambda: _run_scheduled_task(req.id), enabled=req.enabled) return {"ok": True, "id": req.id} @@ -115,8 +115,8 @@ async def update_schedule(schedule_id: str, req: ScheduleUpdate, db: AsyncSessio sched.enabled = req.enabled await db.commit() - # 更新调度器(使用最终值) - scheduler.register(schedule_id, sched.task, lambda: _run_scheduled_task(schedule_id), enabled=sched.enabled) + # 更新调度器(第 2 参传 cron 表达式;使用最终值) + scheduler.register(schedule_id, sched.cron, lambda: _run_scheduled_task(schedule_id), enabled=sched.enabled) return {"ok": True} diff --git a/src/core/scheduler.py b/src/core/scheduler.py index ac057e5..214bc30 100644 --- a/src/core/scheduler.py +++ b/src/core/scheduler.py @@ -14,6 +14,9 @@ logger = logging.getLogger(__name__) class ScheduledTask: """一个定时任务。""" + #: 运行循环的轮询粒度(秒)。同时决定运行期 cron/next_run 变更的生效延迟上限。 + SLEEP_TICK: float = 1.0 + def __init__( self, task_id: str, @@ -28,13 +31,29 @@ class ScheduledTask: self.last_run: datetime | None = None self.next_run: datetime | None = None self._task: asyncio.Task | None = None - self._calc_next() + # 构造时立即校验 cron:非法表达式直接报错,不留给 _run_loop 静默吞掉。 + # (全量审查 C1: 传入任务类型字符串时 croniter 抛异常被吞 → 任务永不触发) + self._calc_next(raise_on_error=True) - def _calc_next(self) -> None: + def _calc_next(self, *, raise_on_error: bool = False) -> None: + """计算下次运行时间。 + + raise_on_error=True 时非法 cron 抛 ValueError(构造期用); + 否则仅告警并置 next_run=None(运行期容错)。 + """ try: - self.next_run = croniter(self.cron, datetime.now).get_next(datetime) - except Exception: + # 必须调用 datetime.now():传方法对象本身会让 croniter 在 + # (start_time or now) 的算术里抛 TypeError,导致 next_run 永远为 None。 + self.next_run = croniter(self.cron, datetime.now()).get_next(datetime) + except Exception as e: self.next_run = None + msg = ( + f"任务 {self.task_id!r} 的 cron 表达式非法: {self.cron!r} ({e})。" + "注意: 此处应为 cron 表达式(如 '0 8 * * *'),不是任务类型名。" + ) + if raise_on_error: + raise ValueError(msg) from e + logger.error(msg) def update(self, cron: str | None = None, enabled: bool | None = None) -> None: if cron is not None: @@ -44,27 +63,45 @@ class ScheduledTask: self._calc_next() async def _run_loop(self) -> None: + """任务运行循环。 + + 睡眠策略: 使用固定的短 tick(SLEEP_TICK 秒)轮询 next_run,而不是 + 一次性 sleep 到 next_run。原因是运行期可通过 API 更新 cron / 手动 + 调整 next_run;若按 wait_seconds 长时间沉睡,变更最长要等 + wait_seconds 才生效(实测可达 60s),表现为「改了不生效」。 + """ while True: - if not self.enabled or not self.next_run: - await asyncio.sleep(60) - self._calc_next() - continue - - now = datetime.now() - wait_seconds = (self.next_run - now).total_seconds() - if wait_seconds > 0: - await asyncio.sleep(min(wait_seconds, 60)) - continue - - # 执行任务 - self.last_run = datetime.now() - self._calc_next() try: - logger.info("定时任务触发: %s (cron=%s)", self.task_id, self.cron) - await self.fn() - logger.info("定时任务完成: %s", self.task_id) + if not self.enabled or not self.next_run: + await asyncio.sleep(self.SLEEP_TICK) + self._calc_next() + continue + + now = datetime.now() + wait_seconds = (self.next_run - now).total_seconds() + if wait_seconds > 0: + # 短 tick 轮询,保证 next_run/cron 变更能及时被感知 + await asyncio.sleep(min(wait_seconds, self.SLEEP_TICK)) + continue + + # 执行任务:先推进 next_run 再执行,避免任务耗时导致重复触发 + self.last_run = datetime.now() + self._calc_next() + try: + logger.info("定时任务触发: %s (cron=%s)", self.task_id, self.cron) + await self.fn() + logger.info("定时任务完成: %s", self.task_id) + except asyncio.CancelledError: + raise + except Exception: + # 单个任务失败不能终止循环,否则一次异常即永久停摆 + logger.exception("定时任务失败: %s", self.task_id) + except asyncio.CancelledError: + raise except Exception: - logger.exception("定时任务失败: %s", self.task_id) + # 循环体自身的意外异常(如 _calc_next)也不能终止调度 + logger.exception("调度循环异常: %s", self.task_id) + await asyncio.sleep(self.SLEEP_TICK) class DataQualityScheduler: @@ -152,6 +189,7 @@ class Scheduler: def __init__(self) -> None: self._tasks: dict[str, ScheduledTask] = {} + self._running = False def register( self, @@ -160,13 +198,31 @@ class Scheduler: fn: Callable[[], Coroutine], enabled: bool = True, ) -> ScheduledTask: - if task_id in self._tasks: - self._tasks[task_id].update(cron=cron, enabled=enabled) - return self._tasks[task_id] - task = ScheduledTask(task_id, cron, fn, enabled) - self._tasks[task_id] = task + """注册(或更新)一个定时任务。 + + 若调度器已 start,新任务会立即启动其运行循环 —— + 否则运行期通过 API 新建的任务永远不会被执行(全量审查 C1 缺陷 2)。 + """ + existing = self._tasks.get(task_id) + if existing is not None: + existing.update(cron=cron, enabled=enabled) + # 更新 cron 后重新校验:改坏了要立刻报错,而不是静默失活 + existing._calc_next(raise_on_error=True) + existing.fn = fn + task = existing + else: + task = ScheduledTask(task_id, cron, fn, enabled) + self._tasks[task_id] = task + + if self._running: + self._ensure_loop(task) return task + def _ensure_loop(self, task: ScheduledTask) -> None: + """为任务启动运行循环(幂等:已在运行则跳过)。""" + if task._task is None or task._task.done(): + task._task = asyncio.create_task(task._run_loop()) + def get(self, task_id: str) -> ScheduledTask | None: return self._tasks.get(task_id) @@ -174,14 +230,18 @@ class Scheduler: return list(self._tasks.values()) def remove(self, task_id: str) -> None: - self._tasks.pop(task_id, None) + task = self._tasks.pop(task_id, None) + if task is not None and task._task is not None: + task._task.cancel() async def start(self) -> None: + self._running = True for task in self._tasks.values(): - task._task = asyncio.create_task(task._run_loop()) + self._ensure_loop(task) logger.info("定时调度器已启动, 共 %d 个任务", len(self._tasks)) async def stop(self) -> None: + self._running = False for task in self._tasks.values(): if task._task: task._task.cancel() diff --git a/tests/test_scheduler_registration.py b/tests/test_scheduler_registration.py new file mode 100644 index 0000000..e7604c9 --- /dev/null +++ b/tests/test_scheduler_registration.py @@ -0,0 +1,190 @@ +"""C1 回归测试: 定时任务调度器「静默失效」缺陷。 + +原始缺陷(全量审查 C1): + 1. `Scheduler.register()` 的第 2 参是 cron 表达式,但 app.py:52 / + schedules.py:95,119 三处调用传的都是任务类型字符串("events" 等)。 + `croniter("events")` 抛异常被 `_calc_next` 的 except 吞掉 → + next_run 永远为 None → 所有定时任务从不触发且无任何报错。 + 2. 运行期通过 API 新建的任务只进了 _tasks 字典,从未 + asyncio.create_task 启动 _run_loop(只有 Scheduler.start() 会建循环)。 + +这些用例不依赖数据库,直接对调度器类做行为断言。 +""" +from __future__ import annotations + +import asyncio + +import pytest + +from src.core.scheduler import ScheduledTask, Scheduler + + +class TestCronValidation: + """非法 cron 必须显式报错,不能静默变成永不触发。""" + + def test_nonevent_cron_raises(self): + """把任务类型字符串当 cron 传(原缺陷) → 构造时立即抛 ValueError。""" + async def _noop() -> None: ... + with pytest.raises(ValueError, match="cron"): + ScheduledTask("daily-events", "events", _noop) + + def test_valid_cron_accepted(self): + """合法 cron 正常构造,且 next_run 被计算出来(非 None)。""" + async def _noop() -> None: ... + task = ScheduledTask("daily-events", "0 8 * * *", _noop) + assert task.next_run is not None + + def test_next_run_is_in_future(self): + """next_run 必须落在未来,否则 _run_loop 会立刻误触发。""" + from datetime import datetime + async def _noop() -> None: ... + task = ScheduledTask("t", "*/30 * * * *", _noop) + assert task.next_run > datetime.now() + + def test_croniter_receives_called_datetime(self): + """回归: `croniter(expr, datetime.now)`(未调用)会抛 + TypeError: 'builtin_function_or_method' object cannot be interpreted + as an integer —— 这是比「传错参数」更深一层的静默失效根源。 + 断言 next_run 是真实 datetime,而非 None。""" + from datetime import datetime + async def _noop() -> None: ... + task = ScheduledTask("t", "0 8 * * *", _noop) + assert isinstance(task.next_run, datetime), ( + f"next_run 应为 datetime,实际 {task.next_run!r} —— " + "croniter 的 start_time 未正确传入" + ) + + +class TestRunLoopActuallyFires: + """注册后任务必须真的在到期时被执行(端到端行为,不 mock 内部)。""" + + async def test_register_then_start_runs_task(self): + """cron 到点后任务函数被调用一次。用 1 秒粒度的 * * * * * 加速验证。""" + calls: list[str] = [] + + async def _job() -> None: + calls.append("ran") + + sched = Scheduler() + sched.register("t1", "* * * * *", _job, enabled=True) + await sched.start() + try: + # _run_loop 用 min(wait, 60) 分段睡;下一分钟边界最多 60 秒。 + # 为让测试可跑,直接把 next_run 拨到过去,触发一次执行。 + task = sched.get("t1") + from datetime import datetime, timedelta + task.next_run = datetime.now() - timedelta(seconds=1) + for _ in range(40): + if calls: + break + await asyncio.sleep(0.05) + finally: + await sched.stop() + assert calls == ["ran"], "注册并 start 后,到期任务未被执行" + + async def test_register_after_start_also_runs(self): + """运行期(已 start 之后)新 register 的任务也必须被启动(原缺陷 2)。""" + calls: list[str] = [] + + async def _job() -> None: + calls.append("late") + + sched = Scheduler() + await sched.start() + try: + sched.register("late-task", "* * * * *", _job, enabled=True) + from datetime import datetime, timedelta + sched.get("late-task").next_run = datetime.now() - timedelta(seconds=1) + for _ in range(40): + if calls: + break + await asyncio.sleep(0.05) + finally: + await sched.stop() + assert calls == ["late"], "start 之后注册的任务未被启动循环" + + async def test_disabled_task_does_not_run(self): + """enabled=False 的任务不执行。""" + calls: list[str] = [] + + async def _job() -> None: + calls.append("nope") + + sched = Scheduler() + sched.register("off", "* * * * *", _job, enabled=False) + await sched.start() + try: + from datetime import datetime, timedelta + sched.get("off").next_run = datetime.now() - timedelta(seconds=1) + await asyncio.sleep(0.3) + finally: + await sched.stop() + assert calls == [] + + +class TestTaskFailureIsolation: + """单个任务抛异常不能杀掉循环(否则一次失败永久停摆)。""" + + async def test_exception_does_not_kill_loop(self): + """任务抛异常后,循环仍存活并能在下一次到期时继续执行。""" + runs: list[int] = [] + + async def _boom() -> None: + runs.append(1) + if len(runs) == 1: + raise RuntimeError("boom") + + sched = Scheduler() + sched.register("flaky", "* * * * *", _boom, enabled=True) + await sched.start() + try: + from datetime import datetime, timedelta + task = sched.get("flaky") + for _ in range(40): + if len(runs) >= 2: + break + task.next_run = datetime.now() - timedelta(seconds=1) + await asyncio.sleep(0.05) + assert len(runs) >= 2, "任务首次抛异常后循环未继续" + assert task._task is not None and not task._task.done(), ( + "循环在任务抛异常后已终止" + ) + finally: + await sched.stop() + + async def test_next_run_change_takes_effect_promptly(self): + """运行期把 next_run 改到过去,循环须在 SLEEP_TICK 内感知并触发。 + + 回归: 原实现 sleep(min(wait_seconds, 60)),当 wait_seconds<60 时 + 会一次性睡满 wait_seconds,导致运行期通过 API 更新 cron 后 + 最长 60s 不生效。 + """ + calls: list[str] = [] + + async def _job() -> None: + calls.append("ran") + + import src.core.scheduler as sched_mod + old_tick = sched_mod.ScheduledTask.SLEEP_TICK + sched_mod.ScheduledTask.SLEEP_TICK = 0.05 + try: + sched = Scheduler() + sched.register("tick", "* * * * *", _job, enabled=True) + await sched.start() + try: + from datetime import datetime, timedelta + # 先等循环进入 sleep(wait 接近 60s 的最坏情况) + await asyncio.sleep(0.15) + sched.get("tick").next_run = datetime.now() - timedelta(seconds=1) + # 短 tick 应在 ~0.15s 内感知;给 1s 容差 + for _ in range(40): + if calls: + break + await asyncio.sleep(0.05) + finally: + await sched.stop() + finally: + sched_mod.ScheduledTask.SLEEP_TICK = old_tick + assert calls == ["ran"], ( + "next_run 变更后循环未在 tick 内响应(疑似一次性睡满 wait_seconds)" + ) -- 2.39.5 From 9ccda4bf51c5ae534f8212736bb74e11ff838d98 Mon Sep 17 00:00:00 2001 From: shangfangjian Date: Mon, 21 Sep 2026 17:10:17 +0800 Subject: [PATCH 2/7] =?UTF-8?q?fix(critical):=20=E6=81=A2=E5=A4=8D=20bzzoi?= =?UTF-8?q?ro=20events=20=E9=87=87=E9=9B=86=E7=AE=A1=E7=BA=BF,=E6=B6=88?= =?UTF-8?q?=E9=99=A4=E6=95=B0=E6=8D=AE=E6=BA=90=E6=B3=A8=E5=86=8C=E8=A1=A8?= =?UTF-8?q?=E9=9D=99=E9=BB=98=E5=A4=B1=E6=95=88?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit C2 数据丢失修复。重构提交 6a49940 删除了 src/data/bzzoiro.py 中的 fetch_bzzoiro_events 与 BzzoiroSource 后未做替换,同时 sources.py 用 `try/except Exception: pass` 吞掉了 ImportError,导致: - _SOURCES 注册表恒为空 - get_source("bzzoiro") 恒抛 ValueError("未知数据源: bzzoiro") - src/api/routes/ingest.py 与 schedules.py 运行时全线失效, 且日志中无任何导入错误痕迹 修复内容: 1. 从 f05dc1a 恢复并移植 events 管线: - fetch_bzzoiro_events():分页抓取 /events/(纯异步,复用现有 _fetch_json_async) - @register class BzzoiroSource + ingest():保留原签名 (db, *, leagues, date_from, date_to, status),保持「不 commit, 事务由调用方 UnitOfWork 控制」的契约,并保留单联赛抓取失败的 错误隔离(记录 error 后 continue) - 复用现有 _to_date / _to_int_or_none / _match_key,未重复定义 - 血缘字段继续使用与 nm 配对的 raw(避免 P0-3 回归) 2. sources.py:导入失败改为 logger.exception 记录,新增 _loaded 标记 保证 _load_sources() 至多执行一次;失败时不置位以便后续重试。 get_source() 对真正未知的名字仍抛 ValueError。 验证: - 新增 tests/test_bzzoiro_source_registry.py(8 项):注册表非空、 get_source/list_sources 可用、ingest 签名契约、不再静默吞异常 - test_regressions.py:3 failed -> 2 failed - 全量:14 failed -> 13 failed,通过数 190 -> 199 注意事项:未改动任何 caller(routes/ingest.py、routes/schedules.py 原样 可用,证明签名保持正确)。 --- src/data/bzzoiro.py | 293 +++++++++++++++++++++++++- src/data/sources.py | 38 +++- tests/test_bzzoiro_source_registry.py | 92 ++++++++ 3 files changed, 410 insertions(+), 13 deletions(-) create mode 100644 tests/test_bzzoiro_source_registry.py diff --git a/src/data/bzzoiro.py b/src/data/bzzoiro.py index 4fb345e..0278828 100644 --- a/src/data/bzzoiro.py +++ b/src/data/bzzoiro.py @@ -23,7 +23,7 @@ import httpx from src.core.runtime_config import get_runtime_value from src.core.http_client import get_client from src.data.config import BZZOIRO_LEAGUE_IDS, LEAGUE_COUNTRIES, LEAGUE_NAMES, REQUEST_INTERVAL -from src.data.key_ring import get_key_ring +from src.data.key_ring import _mask, get_key_ring from src.data.normalize import normalize_bzzoiro from src.data.team_names_zh import zh_name from src.data.sources import register @@ -101,7 +101,7 @@ async def _fetch_json_async(path: str, params: dict | None = None, max_retries: # 限流:标记当前 key 冷却,切换到下一个 new_key = ring.report_rate_limited(key) if new_key and new_key != key: - logger.info("bzzoiro 429 → 切换 key: %s → %s,立即重试", _km(key), _km(new_key)) + logger.info("bzzoiro 429 → 切换 key: %s → %s,立即重试", _mask(key), _mask(new_key)) key = new_key continue # 立即重试,不等待 # 单 key 或全部冷却:等待最早恢复的 key @@ -130,6 +130,209 @@ async def _fetch_json_async(path: str, params: dict | None = None, max_retries: raise RuntimeError(f"bzzoiro request failed after {max_retries} attempts: {last_exc}") +async def fetch_bzzoiro_events( + league_code: str, + *, + status: str = "finished", + date_from: str | None = None, + date_to: str | None = None, + limit: int = 200, +) -> list[dict]: + """抓取 bzzoiro 原始事件(纯异步,无需 run_in_executor)。""" + league_id = BZZOIRO_LEAGUE_IDS.get(league_code) + if league_id is None: + raise ValueError(f"未知联赛代码: {league_code}") + + rows: list[dict] = [] + offset = 0 + while True: + params: dict = { + "league_id": league_id, + "status": status, + "limit": limit, + "offset": offset, + } + if date_from: + params["date_from"] = str(date_from)[:10] + if date_to: + params["date_to"] = str(date_to)[:10] + payload = await _fetch_json_async("/events/", params) + batch = payload.get("results") or [] + if not batch: + break + rows.extend(batch) + total = payload.get("total") + offset += limit + if total is not None and offset >= total: + break + if len(batch) < limit: + break + await asyncio.sleep(REQUEST_INTERVAL) + return rows + + +@register +class BzzoiroSource: + """bzzoiro 数据源(实现 DataSource 协议)。""" + + name = "bzzoiro" + + async def ingest( + self, + db, + *, + leagues: Iterable[str], + date_from: str | None = None, + date_to: str | None = None, + status: str = "finished", + ) -> dict: + """采集 bzzoiro → 入库。返回统计。 + + 注意: 本方法不控制事务(commit/rollback),由调用方通过 UnitOfWork 控制。 + """ + result: dict = {"leagues": {}, "total_inserted": 0, "total_updated": 0, "errors": []} + + for code in leagues: + league_r: dict = {"inserted": 0, "updated": 0, "errors": []} + try: + raw_events = await fetch_bzzoiro_events(code, status=status, date_from=date_from, date_to=date_to) + except Exception as e: + # 单联赛抓取失败隔离:记录错误后继续其余联赛,不拖垮整批 + logger.exception("bzzoiro fetch failed for %s", code) + league_r["errors"].append(f"fetch failed: {e}") + result["leagues"][code] = league_r + continue + + # 获取或创建联赛 + stmt = select(League).where(League.code == code) + league = (await db.execute(stmt)).scalar_one_or_none() + if league is None: + league = League(code=code, name=LEAGUE_NAMES.get(code, code), country=LEAGUE_COUNTRIES.get(code)) + db.add(league) + await db.flush() + + # === 批量优化: 预加载球队和已有比赛到内存 === + team_name_to_id: dict[str, int] = {} + existing_matches: dict[tuple[int, int, str], Match] = {} # 完整对象,避免重复查询 + # (NormalizedMatch, 原始 event) 成对保存:后续写 source_event_id 时 + # 必须用配对的那条 event,不能依赖外层循环变量残留值。 + normalized_matches: list[tuple] = [] + + if raw_events: + # 一次遍历: 收集球队名 + 规范化 + all_team_names = set() + for raw in raw_events: + nm = normalize_bzzoiro(raw, code) + if nm is not None: + try: + nm.validate() + except Exception as e: + # P1-3: 统一使用 warning,不追加到 errors(仅运行时错误入 errors) + logger.warning("normalize skip: %s", e) + continue + normalized_matches.append((nm, raw)) + all_team_names.add(nm.home_team) + all_team_names.add(nm.away_team) + + if all_team_names: + stmt = select(Team).where(Team.name.in_(all_team_names)) + teams = (await db.execute(stmt)).scalars().all() + team_name_to_id = {t.name: t.id for t in teams} + + # P1-2: 按需加载,只加载 raw_events 涉及日期范围的比赛(加 30 天缓冲) + # 避免加载联赛全部历史比赛到内存(多赛季采集时内存溢出) + if normalized_matches: + # normalized_matches 存的是 (nm, raw) 元组,遍历需解包 + dates = [nm.date for nm, _raw in normalized_matches if nm.date is not None] + if dates: + min_dt = min(dates) - timedelta(days=30) + max_dt = max(dates) + timedelta(days=30) + stmt = ( + select(Match) + .where(Match.league_id == league.id) + .where(Match.match_date >= min_dt) + .where(Match.match_date <= max_dt) + ) + existing_matches = { + _match_key(m.home_team_id, m.away_team_id, m.match_date_date): m + for m in (await db.execute(stmt)).scalars() + } + # else: existing_matches 保持空 dict(全量新比赛) + + for nm, raw in normalized_matches: + # 球队: 内存查找 + 按需创建 + home_team_id = team_name_to_id.get(nm.home_team) + if home_team_id is None: + home = Team(name=nm.home_team, name_zh=zh_name(nm.home_team)) + db.add(home) + await db.flush() + home_team_id = home.id + team_name_to_id[nm.home_team] = home_team_id + + away_team_id = team_name_to_id.get(nm.away_team) + if away_team_id is None: + away = Team(name=nm.away_team, name_zh=zh_name(nm.away_team)) + db.add(away) + await db.flush() + away_team_id = away.id + team_name_to_id[nm.away_team] = away_team_id + + # 查找已有比赛: 内存查找 + match_key = _match_key(home_team_id, away_team_id, nm.date) + existing_match = existing_matches.get(match_key) + + if existing_match is None: + m = Match( + league_id=league.id, + season=nm.season_label or None, + home_team_id=home_team_id, + away_team_id=away_team_id, + match_date=nm.date, + match_date_date=_to_date(nm.date), + match_status=nm.match_status, + home_goals=nm.home_goals, + away_goals=nm.away_goals, + home_ht_goals=nm.home_ht_goals, + away_ht_goals=nm.away_ht_goals, + match_stage=nm.match_stage, + source_event_id=_to_int_or_none(raw.get("id")), + ) + db.add(m) + await db.flush() + existing_matches[match_key] = m # 防止同批重复 + # 统计字段不在 /events/ 载荷中(单独由 stats 管线回填), + # 此处不再创建 MatchStats。 + league_r["inserted"] += 1 + else: + # 已有比赛: 直接从内存获取对象更新(无需再查询) + changed = False + if existing_match.match_status != nm.match_status and nm.match_status == "finished": + existing_match.match_status = nm.match_status + changed = True + if existing_match.home_goals is None and nm.home_goals is not None: + existing_match.home_goals = nm.home_goals + existing_match.away_goals = nm.away_goals + existing_match.home_ht_goals = nm.home_ht_goals + existing_match.away_ht_goals = nm.away_ht_goals + changed = True + if existing_match.match_stage is None and nm.match_stage: + existing_match.match_stage = nm.match_stage + changed = True + if existing_match.source_event_id is None: + eid = _to_int_or_none(raw.get("id")) + if eid is not None: + existing_match.source_event_id = eid + changed = True + if changed: + league_r["updated"] += 1 + + # 注意: 不在此处 commit,由调用方 UnitOfWork 控制事务 + result["leagues"][code] = league_r + result["total_inserted"] += league_r["inserted"] + result["total_updated"] += league_r["updated"] + return result + + # ============================================================ # 管线基础设施:RawEvent / IngestFailure / DataLineage # ============================================================ @@ -226,6 +429,92 @@ async def ingest_bzzoiro_standings(db, *, leagues: Iterable[str], season: str | continue rows = payload.get("standings") or [] + if not rows: + result["leagues"][code] = {"error": "无积分榜数据(赛季未开始或未提供)"} + result["errors"].append(f"{code}: 无积分榜数据") + continue + + # 联赛(get-or-create) + stmt = select(League).where(League.code == code) + league = (await db.execute(stmt)).scalar_one_or_none() + if league is None: + league = League(code=code, name=LEAGUE_NAMES.get(code, code), country=LEAGUE_COUNTRIES.get(code)) + db.add(league) + await db.flush() + + # 赛季标签:优先用返回的 season 对象推导 + season_obj = payload.get("season") or {} + season_label = _season_label_from_dates( + season_obj.get("start_date"), season_obj.get("end_date") + ) + if season_label == "?": + season_label = season or "" + + # 批量预载球队(与 events 管线使用同一 normalize 规则,保证 Team 匹配) + names = {normalize_name(str(r.get("team_name", ""))) for r in rows} + names.discard("") + team_map: dict[str, Team] = {} + if names: + stmt = select(Team).where(Team.name.in_(names)) + for t in (await db.execute(stmt)).scalars(): + team_map[t.name] = t + + now = datetime.now(timezone.utc) + for r in rows: + team_name = normalize_name(str(r.get("team_name", ""))) + if not team_name: + continue + team = team_map.get(team_name) + if team is None: + team = Team(name=team_name, name_zh=zh_name(team_name)) + db.add(team) + await db.flush() + team_map[team_name] = team + league_r["teams_created"] += 1 + + zone = r.get("zone") or {} + values = dict( + position=_to_int_or_none(r.get("position")) or 0, + played=_to_int_or_none(r.get("played")) or 0, + won=_to_int_or_none(r.get("won")) or 0, + drawn=_to_int_or_none(r.get("drawn")) or 0, + lost=_to_int_or_none(r.get("lost")) or 0, + goals_for=_to_int_or_none(r.get("gf")) or 0, + goals_against=_to_int_or_none(r.get("ga")) or 0, + goal_diff=_to_int_or_none(r.get("gd")) or 0, + points=_to_int_or_none(r.get("pts")) or 0, + xg_for=_to_float_or_none(r.get("xgf")), + xg_against=_to_float_or_none(r.get("xga")), + form=r.get("form") or None, + zone=zone.get("label") or zone.get("key") or None, + updated_at=now, + retrieved_at=now, + ) + + # 同一联赛同一赛季只保留最新快照:按 (league, season, team) upsert + stmt = select(Standing).where( + Standing.league_id == league.id, + Standing.season == season_label, + Standing.team_id == team.id, + ) + standing = (await db.execute(stmt)).scalar_one_or_none() + if standing is None: + standing = Standing( + league_id=league.id, season=season_label, team_id=team.id, **values + ) + db.add(standing) + else: + for k, v in values.items(): + setattr(standing, k, v) + league_r["upserted"] += 1 + + league_r["rows"] = len(rows) + result["leagues"][code] = league_r + result["total_upserted"] += league_r["upserted"] + logger.info( + "bzzoiro standings 采集完成: %s 赛季 %s, upsert %d/%d", + code, season_label, league_r["upserted"], league_r["rows"], + ) return result diff --git a/src/data/sources.py b/src/data/sources.py index 6466323..f4f363a 100644 --- a/src/data/sources.py +++ b/src/data/sources.py @@ -8,10 +8,13 @@ """ from __future__ import annotations +import logging from typing import Protocol from src.db.base import AsyncSession +logger = logging.getLogger(__name__) + class DataSource(Protocol): """比赛数据源契约:抓取 → 规范化 → 入库。""" @@ -29,6 +32,9 @@ class DataSource(Protocol): # ── 注册表 ── _SOURCES: dict[str, DataSource] = {} +# 延迟加载标记:保证 _load_sources() 最多执行一次,重复调用廉价 +_loaded = False + def register(source): """装饰器:将数据源注册到全局注册表。 @@ -42,8 +48,7 @@ def register(source): def get_source(name: str) -> DataSource: """按名获取数据源。""" - if not _SOURCES: - _load_sources() + _load_sources() if name not in _SOURCES: raise ValueError(f"未知数据源: {name}") return _SOURCES[name] @@ -51,18 +56,29 @@ def get_source(name: str) -> DataSource: def list_sources() -> list[str]: """列出所有已注册数据源名。""" - if not _SOURCES: - _load_sources() + _load_sources() return list(_SOURCES.keys()) def _load_sources() -> None: - """延迟导入数据源触发 @register(避免循环导入)。""" - from src.data.bzzoiro import BzzoiroSource # noqa: F811 + """延迟导入数据源触发 @register(避免循环导入)。 + + 导入失败必须留下痕迹:静默吞掉 ImportError 会让注册表恒为空, + 导致 get_source() 对所有数据源都报「未知数据源」,把导入错误 + 伪装成「不存在的名字」—— 这类静默失效极难定位,故在此显式记日志。 + """ + global _loaded + if _loaded: + return + try: + from src.data.bzzoiro import BzzoiroSource # noqa: F401 + except Exception: + logger.exception( + "数据源模块导入失败,注册表将为空 —— get_source() 会对所有名字报「未知数据源」" + ) + return # 保持 _loaded=False,下次调用可重试 + _loaded = True -# 保持向后兼容:模块加载时尝试加载(但不再强制) -try: - _load_sources() -except Exception: - pass +# 模块加载时预热(失败会记录日志,不再静默) +_load_sources() diff --git a/tests/test_bzzoiro_source_registry.py b/tests/test_bzzoiro_source_registry.py new file mode 100644 index 0000000..e38aab2 --- /dev/null +++ b/tests/test_bzzoiro_source_registry.py @@ -0,0 +1,92 @@ +"""回归测试:锁定 bzzoiro events 管线被删除 + 注册表静默失效的缺陷不再复发。 + +背景(真实数据丢失事故): + 重构提交 6a49940 删除了 src/data/bzzoiro.py 中的 fetch_bzzoiro_events 与 + BzzoiroSource,但 src/data/sources.py 的模块级预热把 ImportError 用 + `try/except Exception: pass` 吞掉了。后果: + - _SOURCES 注册表恒为空 + - get_source("bzzoiro") 恒抛 ValueError("未知数据源: bzzoiro") + - src/api/routes/ingest.py 与 schedules.py 在运行时全线失效, + 且日志中看不到任何导入错误的痕迹 —— 这才是它长期漏网的原因。 + +本测试用静态结构断言 + 真实导入来锁死这两点,不 mock 网络: + - 注册表必须非空(直接守卫「静默吞异常」回归) + - get_source("bzzoiro") 必须返回真实实例 + - ingest() 的签名必须保持 caller 依赖的 keyword-only 参数 +""" +from __future__ import annotations + +import inspect +from pathlib import Path + +import pytest + +from src.data.bzzoiro import BzzoiroSource, fetch_bzzoiro_events +from src.data.sources import get_source, list_sources + +SRC = Path(__file__).resolve().parent.parent / "src" + + +class TestSourceRegistry: + """P0: 注册表必须真的装载到 bzzoiro,而不是静默为空。""" + + def test_registry_is_not_empty(self): + """注册表恒非空 —— 直接守卫 `try/except: pass` 静默吞 ImportError。""" + assert list_sources(), ( + "数据源注册表为空 —— _load_sources() 的导入失败了," + "且(修复前)异常被静默吞掉。任何导入错误都必须被 logger.exception 记录。" + ) + + def test_get_source_returns_bzzoiro(self): + """get_source('bzzoiro') 必须返回实例,而不是抛 ValueError。""" + source = get_source("bzzoiro") + assert source.name == "bzzoiro" + + def test_list_sources_contains_bzzoiro(self): + assert "bzzoiro" in list_sources() + + def test_register_decorator_still_exports_the_class(self): + """@register 必须仍然返回原类(不能再被改成返回实例)。""" + assert isinstance(BzzoiroSource, type), "@register 不应把类替换成实例" + + +class TestIngestContract: + """caller 依赖 ingest() 的签名,routes 传的就是这些参数。""" + + def test_ingest_exists_and_is_async(self): + assert hasattr(BzzoiroSource, "ingest"), "BzzoiroSource.ingest 丢失(events 管线被删)" + assert inspect.iscoroutinefunction(BzzoiroSource.ingest), "ingest 必须是 async" + + def test_ingest_accepts_caller_keyword_args(self): + """ingest.py:65 与 schedules.py:41 依赖的 keyword-only 参数必须齐全。""" + params = inspect.signature(BzzoiroSource.ingest).parameters + for name in ("leagues", "date_from", "date_to", "status"): + assert name in params, f"ingest() 缺少 keyword 参数 {name!r} —— caller 会 TypeError" + # leagues 必须 keyword-only(src/api/routes/ingest.py 用 leagues=[code] 传) + assert params["leagues"].kind is inspect.Parameter.KEYWORD_ONLY + # 默认值契约:schedules.py 只传 leagues + status + assert params["date_from"].default is None + assert params["date_to"].default is None + assert params["status"].default == "finished" + + def test_fetch_bzzoiro_events_signature(self): + """抓取函数也必须存在,且 leagues 抓取走 keyword-only 的 status/日期。""" + assert callable(fetch_bzzoiro_events) + params = inspect.signature(fetch_bzzoiro_events).parameters + assert "league_code" in params + for name in ("status", "date_from", "date_to"): + assert name in params, f"fetch_bzzoiro_events() 缺少 {name!r}" + + +class TestNoSilentImportSwallow: + """sources.py 不许再用 `except Exception: pass` 吞掉导入失败。""" + + def test_load_sources_logs_failure(self): + src = (SRC / "data" / "sources.py").read_text(encoding="utf-8") + assert "logger" in src, "sources.py 缺少模块 logger,无法记录导入失败" + body = src[src.index("def _load_sources"):] + assert "logger.exception" in body, ( + "_load_sources() 失败时未记日志 —— 静默吞异常会让注册表恒为空," + "把 ImportError 伪装成「未知数据源」(P0 事故根因)" + ) + assert "pass" not in body.split("def ")[0], "不得再用裸 pass 吞掉导入异常" -- 2.39.5 From b9d7b6423e05436c2fd073220dfeeb5b9903871f Mon Sep 17 00:00:00 2001 From: shangfangjian Date: Mon, 21 Sep 2026 17:13:57 +0800 Subject: [PATCH 3/7] =?UTF-8?q?test(c2):=20=E8=A1=A5=E9=BD=90=E5=AF=BC?= =?UTF-8?q?=E5=85=A5=E6=AC=A1=E5=BA=8F=E5=9B=9E=E5=BD=92,=E6=B3=A8?= =?UTF-8?q?=E6=98=8E=E7=8E=AF=E7=8A=B6=E5=AF=BC=E5=85=A5=E4=B8=BA=E9=A2=84?= =?UTF-8?q?=E6=9C=9F=E4=B8=94=E5=8F=AF=E8=87=AA=E6=84=88?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit test_bzzoiro_source_registry.py 增至 11 项: - 新增 TestImportOrderSelfHeal:用独立子进程验证两种真实导入次序 (sources 先 / bzzoiro 先)下注册表均可用。该缺陷只在真实模块初始化 时序下成立,mock 照不出来,故不 mock。 - 新增未知名字仍抛 ValueError 的契约守卫。 src/data/sources.py:补全 _load_sources 文档,说明首次装载失败属正常的 环状导入(bzzoiro 仍在初始化),失败时保持 _loaded=False 使后续 get_source()/list_sources() 调用可重试装载(自愈)。 验证: - tests/test_bzzoiro_source_registry.py:11 passed - 两种导入次序实测均输出 bzzoiro ['bzzoiro'] --- src/data/sources.py | 8 ++++- tests/test_bzzoiro_source_registry.py | 52 +++++++++++++++++++++++++++ 2 files changed, 59 insertions(+), 1 deletion(-) diff --git a/src/data/sources.py b/src/data/sources.py index f4f363a..c860b84 100644 --- a/src/data/sources.py +++ b/src/data/sources.py @@ -66,6 +66,11 @@ def _load_sources() -> None: 导入失败必须留下痕迹:静默吞掉 ImportError 会让注册表恒为空, 导致 get_source() 对所有数据源都报「未知数据源」,把导入错误 伪装成「不存在的名字」—— 这类静默失效极难定位,故在此显式记日志。 + + 失败时保持 _loaded=False:后续调用会重新尝试导入。正常应用流程中 + src.api.app 先导入本模块,此处的模块级预热是在 `src.data.bzzoiro` + 尚未初始化时发起的,导入链在本模块内成环,首次必然失败(bzzoiro + 仍在加载中),由后续 get_source() / list_sources() 调用完成真正的装载。 """ global _loaded if _loaded: @@ -74,7 +79,8 @@ def _load_sources() -> None: from src.data.bzzoiro import BzzoiroSource # noqa: F401 except Exception: logger.exception( - "数据源模块导入失败,注册表将为空 —— get_source() 会对所有名字报「未知数据源」" + "数据源模块导入失败(通常为初始化中途的环状导入),本次注册表为空;" + "下次 get_source()/list_sources() 调用会自动重试" ) return # 保持 _loaded=False,下次调用可重试 _loaded = True diff --git a/tests/test_bzzoiro_source_registry.py b/tests/test_bzzoiro_source_registry.py index e38aab2..e0842fa 100644 --- a/tests/test_bzzoiro_source_registry.py +++ b/tests/test_bzzoiro_source_registry.py @@ -78,6 +78,58 @@ class TestIngestContract: assert name in params, f"fetch_bzzoiro_events() 缺少 {name!r}" +class TestImportOrderSelfHeal: + """导入次序不得影响注册表 —— 这是本缺陷的第二个隐藏面。 + + 背景:src.api.app 先 `import src.data.sources`,其模块级预热在 + `src.data.bzzoiro` 尚未初始化时发起,导入链在 sources.py 内成环, + 首次导入必然失败(bzzoiro 仍在加载中)。修复前 `except Exception: pass` + 把这次失败同时变得「无声」且「不可恢复」,注册表就永久空掉了。 + + 这里用独立子进程验证真实导入次序,不 mock —— 因为该缺陷只在 + 真实的模块初始化时序下才成立,单元测试里的 mock 反而照不出来。 + """ + + @staticmethod + def _run(import_lines: list[str]) -> str: + """在全新解释器中按指定次序导入,返回 get_source/list_sources 结果。""" + import subprocess + import sys + + root = Path(__file__).resolve().parent.parent + code = ( + "import sys; sys.path.insert(0, r'%s')\n" % root + + "\n".join(import_lines) + + "\nfrom src.data.sources import get_source, list_sources\n" + "print('RESULT', get_source('bzzoiro').name, list_sources())\n" + ) + proc = subprocess.run( + [sys.executable, "-c", code], + capture_output=True, text=True, timeout=120, cwd=str(root), + ) + assert proc.returncode == 0, ( + f"导入次序 {import_lines} 下 get_source('bzzoiro') 失败:\n{proc.stderr[-1500:]}" + ) + return proc.stdout.strip() + + def test_sources_first_then_bzzoiro(self): + """次序 B:sources 先导入(正常应用路径)。""" + out = self._run(["import src.data.sources"]) + assert out == "RESULT bzzoiro ['bzzoiro']", out + + def test_bzzoiro_first_then_sources(self): + """次序 A:bzzoiro 先导入 —— 注册表必须仍然可用(自愈)。""" + out = self._run(["import src.data.bzzoiro"]) + assert out == "RESULT bzzoiro ['bzzoiro']", ( + out + " —— 注册表为空说明首次装载失败后没有再重试(静默失效回归)" + ) + + def test_unknown_name_still_raises(self): + """真正未知的名字仍须抛 ValueError —— 修复不应放宽这一契约。""" + with pytest.raises(ValueError, match="未知数据源"): + get_source("no_such_source_xyz") + + class TestNoSilentImportSwallow: """sources.py 不许再用 `except Exception: pass` 吞掉导入失败。""" -- 2.39.5 From 11da4d3631b2f9cfac0129568d8c02ba0ba906f3 Mon Sep 17 00:00:00 2001 From: shangfangjian Date: Mon, 21 Sep 2026 17:19:24 +0800 Subject: [PATCH 4/7] =?UTF-8?q?fix(critical):=20=E4=BF=AE=E5=A4=8D?= =?UTF-8?q?=E8=AF=84=E5=AE=A1=E7=A1=AE=E8=AE=A4=E7=9A=84=20R1-R5=20?= =?UTF-8?q?=E4=BA=94=E5=A4=84=E7=BC=BA=E9=99=B7?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit R1 429 key 轮换路径调用不存在的 _km → NameError(且位于凭证脱敏日志行): 改用 src/data/key_ring._mask,全树不再有 _km 调用。 [注: 该改动已随另一工作流提交 9ccda4b 一并入库] R2 ingest_bzzoiro_standings 被截断(return 前无 upsert 逻辑,standings 表 永不写入、total_upserted 恒为 0):从 f05dc1a 移植完整实现,按 (league_id, season, team_id) upsert,保留逐联赛错误隔离,不加 db.commit()。 [注: 同上,已随 9ccda4b 入库] R3 src/llm/agents/orchestrator.py 完成日志把 list 喂给 %d → logging TypeError: 改为 len(ok_reports),与同文件降级日志行写法一致。 R4 mode=multi 静默丢弃调用方传入的 model:predict_match_multi 新增 model 关键字参数,经 _ACTIVE_MODEL_OVERRIDE 下传至 _agent_provider, 显式 model 优先级最高;override 生效时跳过 provider 缓存读写以免串味, 并在 predict.py 派发点透传。 R5 回测把字符串日期直接与 timestamptz 列比较:新增 _parse_date_bound 助手, 支持 YYYY-MM-DD / 完整 ISO / datetime / None,裸日期按 UTC 锚定, 结束日取当天末刻(闭区间,避免最后一天被静默排除),非法输入抛 ValueError。 新增 tests/test_review_required_fixes.py 覆盖 R1-R5(R2/R4 为行为测试), 20 项全通过。 --- src/llm/agents/orchestrator.py | 56 +++- src/llm/backtest.py | 41 ++- src/llm/predict.py | 3 +- tests/test_review_required_fixes.py | 452 ++++++++++++++++++++++++++++ 4 files changed, 531 insertions(+), 21 deletions(-) create mode 100644 tests/test_review_required_fixes.py diff --git a/src/llm/agents/orchestrator.py b/src/llm/agents/orchestrator.py index b131c35..35fc16d 100644 --- a/src/llm/agents/orchestrator.py +++ b/src/llm/agents/orchestrator.py @@ -34,6 +34,11 @@ logger = logging.getLogger(__name__) _AGENT_PROVIDER_CACHE: dict[str, tuple[float, LLMProvider]] = {} _AGENT_PROVIDER_CACHE_TTL = 60.0 +# 当前预测的模型覆盖(调用方显式指定)。用模块级变量而非新增形参,因为 +# run_specialists 会被既有测试 mock,加形参会破坏那些测试的调用签名。 +# 由 predict_match_multi 在进入时 set、退出时 reset。 +_ACTIVE_MODEL_OVERRIDE: str | None = None + # ── 5 个专家 agent 定义 ── # A=近期状态 B=攻防数据 C=主客因素 D=联赛排名 E=历史交锋 @@ -105,21 +110,24 @@ class MultiPredictResult: raw: dict | None = None -async def _agent_provider(agent_id: str, *, tier: str) -> LLMProvider: +async def _agent_provider(agent_id: str, *, tier: str, model_override: str | None = None) -> LLMProvider: """构造某 agent 专属 provider。 覆盖优先级: - 模型: AGENT_MODEL_{ID}(运行时) → 层级默认(LLM_SPECIALIST/AGGREGATOR_MODEL) → 全局 LLM_MODEL + 模型: model_override(调用方显式指定) → AGENT_MODEL_{ID}(运行时) → 层级默认(LLM_SPECIALIST/AGGREGATOR_MODEL) → 全局 LLM_MODEL 地址/密钥: AGENT_BASE_URL_{ID} / AGENT_API_KEY_{ID}(运行时) → 全局 LLM_BASE_URL / LLM_API_KEY P3-2: 结果缓存 60 秒,避免每次预测都多次查询运行时配置 DB。 + 注意:model_override 生效时跳过缓存读写 —— 否则带 override 的结果会泄漏给 + 不带 override 的调用(反之亦然),导致跨调用的模型串味。 """ cache_key = f"{agent_id}:{tier}" - cached = _AGENT_PROVIDER_CACHE.get(cache_key) - if cached is not None: - ts, provider = cached - if time.time() - ts < _AGENT_PROVIDER_CACHE_TTL: - return provider + if model_override is None: + cached = _AGENT_PROVIDER_CACHE.get(cache_key) + if cached is not None: + ts, provider = cached + if time.time() - ts < _AGENT_PROVIDER_CACHE_TTL: + return provider pfx = f"AGENT_{agent_id.upper()}_" p = await get_default_provider() @@ -129,6 +137,9 @@ async def _agent_provider(agent_id: str, *, tier: str) -> LLMProvider: model = await get_runtime_value(f"{pfx}MODEL") if model: p.model = model + # 调用方显式传入的 model 优先级最高,高于 agent 级与层级默认 + if model_override: + p.model = model_override base = await get_runtime_value(f"{pfx}BASE_URL") if base: p.base_url = base @@ -136,10 +147,11 @@ async def _agent_provider(agent_id: str, *, tier: str) -> LLMProvider: if key: p.api_key = key - _AGENT_PROVIDER_CACHE[cache_key] = (time.time(), p) - # 简单淘汰:超过 20 条时清空(60s TTL 下不会累积太多) - if len(_AGENT_PROVIDER_CACHE) > 20: - _AGENT_PROVIDER_CACHE.clear() + if model_override is None: + _AGENT_PROVIDER_CACHE[cache_key] = (time.time(), p) + # 简单淘汰:超过 20 条时清空(60s TTL 下不会累积太多) + if len(_AGENT_PROVIDER_CACHE) > 20: + _AGENT_PROVIDER_CACHE.clear() return p @@ -152,9 +164,14 @@ async def run_specialists( """并行执行 5 个专家 agent。fail-open: 单个失败不影响其他。 before: 数据截止时间(回测防泄漏)。None 表示不限制。 + + 模型覆盖通过 _ACTIVE_MODEL_OVERRIDE 传递,而不是新增形参:既有测试 + 会 mock 本函数(见 tests/test_agent_weights_persist.py),加形参会破坏 + 它们的调用签名。predict_match_multi 在调用前后 set/reset 该变量。 """ + model_override = _ACTIVE_MODEL_OVERRIDE tasks = [ - _run_one(spec, header, await _agent_provider(spec.name, tier="specialist"), version=version, before=before) + _run_one(spec, header, await _agent_provider(spec.name, tier="specialist", model_override=model_override), version=version, before=before) for spec in SPECIALIST_SPECS ] results = await asyncio.gather(*tasks, return_exceptions=True) @@ -222,11 +239,13 @@ async def predict_match_multi( version: str = "v1", backtest: bool = False, cutoff_at=None, + model: str | None = None, ) -> MultiPredictResult: """多 agent 端到端预测: 切片 → 并行专家 → 终裁 → 存库。 backtest: 回测模式。True 时 cutoff 自动设为 match_dt - 1 天。 cutoff_at: 显式截止时间(优先于 backtest 自动计算)。 + model: 显式指定模型,优先于 agent 级/层级默认配置(single 模式语义一致)。 """ start = time.perf_counter() @@ -247,7 +266,14 @@ async def predict_match_multi( prediction_cutoff_at = cutoff # 2. 并行专家(各自独立配置,使用统一 cutoff) - reports = await run_specialists(header, version=version, before=cutoff) + # 显式 model 覆盖通过模块级变量下传,避免改动 run_specialists 的签名 + global _ACTIVE_MODEL_OVERRIDE + _prev_override = _ACTIVE_MODEL_OVERRIDE + _ACTIVE_MODEL_OVERRIDE = model + try: + reports = await run_specialists(header, version=version, before=cutoff) + finally: + _ACTIVE_MODEL_OVERRIDE = _prev_override # 2.5 统计有效专家报告数量 ok_reports = [r for r in reports if r.status == "ok"] @@ -279,7 +305,7 @@ async def predict_match_multi( } agg_prompt_tokens = 0 agg_completion_tokens = 0 - aggregator_model = settings.LLM_MODEL # 占位,无实际 LLM 调用 + aggregator_model = model or settings.LLM_MODEL # 占位,无实际 LLM 调用 latency_ms = int((time.perf_counter() - start) * 1000) @@ -347,7 +373,7 @@ async def predict_match_multi( "预测完成 match=%s mode=%s status=%s pred=%s:%s (%s) latency=%sms, experts=%d/%d, prediction_id=%s", match_id, "multi", pred_status, pred.pred_home_goals, pred.pred_away_goals, pred.pred_1x2, - latency_ms, ok_reports, len(reports), pred.id, + latency_ms, len(ok_reports), len(reports), pred.id, ) return MultiPredictResult( diff --git a/src/llm/backtest.py b/src/llm/backtest.py index 258f180..a7c7272 100644 --- a/src/llm/backtest.py +++ b/src/llm/backtest.py @@ -10,7 +10,7 @@ from __future__ import annotations import asyncio import logging from dataclasses import dataclass, field -from datetime import datetime +from datetime import datetime, timezone from sqlalchemy import select from sqlalchemy.orm import selectinload @@ -79,6 +79,35 @@ class BacktestSummary: results: list[BacktestMatchResult] = field(default_factory=list) +def _parse_date_bound(value, *, end_of_day: bool) -> datetime | None: + """把日期入参解析成可与 timestamptz 列比较的 aware datetime。 + + 支持 "YYYY-MM-DD"、完整 ISO 串(可带偏移)以及 datetime 对象;None 原样返回。 + 裸日期按 UTC 锚定 —— Match.match_date 是 timestamptz,naive datetime 与之 + 比较会因时区不同而偏移;start 取当天 00:00,end 取当天 23:59:59.999999 + (闭区间,否则最后一天会被静默排除)。 + + 解析失败抛 ValueError(不静默吞掉):fromisoformat 对非法输入统一抛 ValueError, + 这里包一层以带上原始值,便于定位是哪个参数写错了。 + """ + if value is None: + return None + if isinstance(value, datetime): + dt = value + else: + try: + dt = datetime.fromisoformat(str(value)) + except ValueError as e: + raise ValueError(f"无法解析日期: {value!r}(应为 YYYY-MM-DD 或 ISO 格式)") from e + + if dt.tzinfo is None: + dt = dt.replace(tzinfo=timezone.utc) + # 闭区间上界:日期串解析出来是 00:00,取当天末刻才能让最后一天参与回测 + if end_of_day: + dt = dt.replace(hour=23, minute=59, second=59, microsecond=999999) + return dt + + async def _get_historical_matches( db, *, @@ -106,10 +135,12 @@ async def _get_historical_matches( ) if league_id is not None: stmt = stmt.where(Match.league_id == league_id) - if date_from: - stmt = stmt.where(Match.match_date >= date_from) - if date_to: - stmt = stmt.where(Match.match_date <= date_to) + dt_from = _parse_date_bound(date_from, end_of_day=False) + if dt_from is not None: + stmt = stmt.where(Match.match_date >= dt_from) + dt_to = _parse_date_bound(date_to, end_of_day=True) + if dt_to is not None: + stmt = stmt.where(Match.match_date <= dt_to) stmt = stmt.order_by(Match.match_date.desc()).limit(limit) result = await db.execute(stmt) diff --git a/src/llm/predict.py b/src/llm/predict.py index 2f384ac..383b79f 100644 --- a/src/llm/predict.py +++ b/src/llm/predict.py @@ -187,13 +187,14 @@ async def predict_match( ) from src.llm.agents.orchestrator import predict_match_multi - # 回测参数完整传递到 multi-agent 路径 + # 回测参数 + 模型覆盖完整传递到 multi-agent 路径 return await predict_match_multi( match_id, provider=provider, version=(prompt_version or "v1").removeprefix("multi_"), backtest=backtest, cutoff_at=cutoff_at, + model=model, ) diff --git a/tests/test_review_required_fixes.py b/tests/test_review_required_fixes.py new file mode 100644 index 0000000..4eaa07f --- /dev/null +++ b/tests/test_review_required_fixes.py @@ -0,0 +1,452 @@ +"""回归测试: 代码评审确认的 5 个缺陷修复(R1-R5)。 + +R1 429 key 轮换路径调用不存在的 _km → NameError(且是凭证脱敏点) +R2 ingest_bzzoiro_standings 被截断,永不写 standings 表 +R3 orchestrator 完成日志把 list 喂给 %d → logging TypeError +R4 mode="multi" 静默丢弃调用方传入的 model +R5 回测把字符串日期直接与 timestamptz 列比较 + +R2 采用行为测试(假 db + monkeypatch 抓取函数),其余为单元/结构断言。 +""" +from __future__ import annotations + +import inspect +import logging +import pathlib + +import pytest + +from src.data.key_ring import _mask +from src.llm import backtest as bt_mod +from src.llm.agents import orchestrator as orch_mod + +_REPO_ROOT = pathlib.Path(__file__).resolve().parents[1] +_BZZOIRO_SRC = _REPO_ROOT / "src" / "data" / "bzzoiro.py" + +# 在导入期就抓取真实的 _agent_provider。 +# 原因: tests/test_multi_agent_cutoff.py:52/87/117 会直接 +# orch._agent_provider = lambda agent_id, **kw: MagicMock(model="test") +# 且不做清理(既有测试,本次任务不允许改动),导致模块属性在整套测试跑完后 +# 被永久替换成同步 lambda。导入期快照可以规避这种跨测试污染。 +_REAL_AGENT_PROVIDER = orch_mod._agent_provider + + +def _bzzoiro_source() -> str: + return _BZZOIRO_SRC.read_text(encoding="utf-8") + + +# ============================================================ +# R1 — _km → _mask +# ============================================================ + +class TestR1KeyMasking: + def test_mask_is_importable_from_bzzoiro(self): + """修复点:bzzoiro 通过 key_ring 复用它,而不是本地重实现。""" + import src.data.bzzoiro as bz + + assert bz._mask("abcd1234efgh5678") == "abcd...5678" + + def test_mask_long_key_shows_head_and_tail(self): + assert _mask("abcd1234efgh5678") == "abcd...5678" + + def test_mask_short_key_hides_middle(self): + assert _mask("short") == "sh***" + + def test_no_km_call_remains_in_bzzoiro_source(self): + """源码守卫: _km 在整个代码库不存在,这行一旦执行必抛 NameError。 + + 这是 NameError 类缺陷(静态即可判定),且位于凭证脱敏日志行上, + 因此用源码文本守卫是恰当的,而不是只测运行时路径。 + """ + assert "_km(" not in _bzzoiro_source() + + +# ============================================================ +# R2 — 积分榜 upsert(行为测试) +# ============================================================ + +class _FakeScalars: + def __init__(self, items): + self._items = list(items) + + def all(self): + return list(self._items) + + def __iter__(self): + return iter(self._items) + + def scalar_one_or_none(self): + return self._items[0] if self._items else None + + +class _FakeResult: + def __init__(self, items): + self._items = list(items) + + def scalars(self): + return _FakeScalars(self._items) + + def scalar_one_or_none(self): + return self._items[0] if self._items else None + + +class _FakeDb: + """极简假 db: 记录 add/flush,按队列返回预置查询结果。 + + 只实现 ingest_bzzoiro_standings 真正用到的部分: + - execute(...) → 依次弹出 _results 里的结果 + - add(obj) → 记录 + - flush() → 给尚无 id 的对象补一个自增 id(模拟 DB 回填主键) + """ + + def __init__(self, results=None): + self._results = list(results or []) + self.added: list = [] + self.flush_count = 0 + self._next_id = 1000 + + async def execute(self, _stmt): + if self._results: + return self._results.pop(0) + return _FakeResult([]) + + def add(self, obj): + self.added.append(obj) + + async def flush(self): + self.flush_count += 1 + for obj in self.added: + if getattr(obj, "id", None) is None: + self._next_id += 1 + obj.id = self._next_id + + +async def test_r2_standings_actually_upserts(monkeypatch): + """行为测试: 喂一份积分榜 payload,断言真的构造了 Standing 且计数 > 0。""" + import src.data.bzzoiro as bz + from src.db.models import League, Standing, Team + + payload = { + "season": {"start_date": "2025-08-01", "end_date": "2026-05-31"}, + "standings": [ + { + "position": 1, "team_name": "Arsenal FC", + "played": 10, "won": 8, "drawn": 1, "lost": 1, + "gf": 22, "ga": 8, "gd": 14, "pts": 25, + "xgf": 18.5, "xga": 9.1, "form": "WWDLW", + "zone": {"key": "champions_league", "label": "Champions League"}, + }, + { + "position": 2, "team_name": "Chelsea FC", + "played": 10, "won": 6, "drawn": 2, "lost": 2, + "gf": 18, "ga": 12, "gd": 6, "pts": 20, + }, + ], + } + + async def _fake_fetch(league_code, season=None): + return payload + + monkeypatch.setattr(bz, "fetch_bzzoiro_standings", _fake_fetch, raising=True) + + # 查询顺序: League 命中(避免建联赛) → Team 预载(空) → 每行 Standing(未命中) + league = League(code="EPL", name="Premier League", country="England") + league.id = 42 + db = _FakeDb(results=[_FakeResult([league]), _FakeResult([])]) + + result = await bz.ingest_bzzoiro_standings(db, leagues=["EPL"]) + + assert result["errors"] == [] + assert result["total_upserted"] == 2, result + assert result["leagues"]["EPL"]["rows"] == 2 + assert result["leagues"]["EPL"]["upserted"] == 2 + # 两支球队都是新建的 + assert result["leagues"]["EPL"]["teams_created"] == 2 + + standings = [o for o in db.added if isinstance(o, Standing)] + assert len(standings) == 2, "应真的构造 Standing 行" + assert all(isinstance(o, (Standing, Team)) for o in db.added) + + first = standings[0] + assert first.league_id == 42 + assert first.season == "2025-2026" # 8 月起 → 跨年标签 + assert first.position == 1 + assert first.points == 25 + assert first.xg_for == 18.5 + assert first.zone == "Champions League" # 优先取 label + + +async def test_r2_standings_upsert_updates_existing(monkeypatch): + """行为测试: 已存在同 (league, season, team) 时应就地更新而非新增。""" + import src.data.bzzoiro as bz + from src.db.models import League, Standing + + payload = { + "season": {"start_date": "2025-08-01", "end_date": "2026-05-31"}, + "standings": [{"position": 1, "team_name": "Arsenal FC", "pts": 30}], + } + + async def _fake_fetch(league_code, season=None): + return payload + + monkeypatch.setattr(bz, "fetch_bzzoiro_standings", _fake_fetch, raising=True) + + league = League(code="EPL", name="Premier League", country="England") + league.id = 42 + team = __import__("src.db.models", fromlist=["Team"]).Team(name="Arsenal FC", name_zh="阿森纳") + team.id = 7 + + existing = Standing(league_id=42, season="2025-2026", team_id=7, position=9) + existing.points = 1 + + # 查询顺序: League → Team 预载(命中) → Standing 查询(命中已有行) + db = _FakeDb(results=[_FakeResult([league]), _FakeResult([team]), _FakeResult([existing])]) + + result = await bz.ingest_bzzoiro_standings(db, leagues=["EPL"]) + + assert result["total_upserted"] == 1 + assert existing.points == 30, "已有行应被就地更新" + assert result["leagues"]["EPL"]["teams_created"] == 0 + # 不应新增 Standing(只有 league/team 层面的 add) + assert not [o for o in db.added if isinstance(o, Standing)] + + +def test_r2_source_contains_real_upsert_loop(): + """结构断言(行为测试之外的兜底): 确认函数未被截断。""" + import src.data.bzzoiro as bz + + src = inspect.getsource(bz.ingest_bzzoiro_standings) + assert "total_upserted" in src + assert 'result["total_upserted"] +=' in src, "total_upserted 必须真的被累加" + assert "Standing(" in src, "必须真的构造 Standing" + assert "select(Standing)" in src, "必须查询已有快照以决定 insert/update" + + +# ============================================================ +# R3 — logging 参数类型 +# ============================================================ + +def test_r3_completion_log_record_formats_without_raising(): + """%d 占位符拿到 list 时 logging 会抛 TypeError;修复后应为计数。""" + fmt = ( + "预测完成 match=%s mode=%s status=%s pred=%s:%s (%s) latency=%sms, " + "experts=%d/%d, prediction_id=%s" + ) + ok_reports = ["a", "b", "c"] + reports = ["a", "b", "c", "d", "e"] + + record = logging.LogRecord( + "src.llm.agents.orchestrator", logging.INFO, __file__, 1, fmt, + (999, "multi", "success", 2, 1, "1", 100, len(ok_reports), len(reports), 7), + None, + ) + msg = record.getMessage() # 修复前此处抛 TypeError + assert "experts=3/5" in msg + + # 反证: 原缺陷写法(直接传 list)确实会炸,确保这条测试真的有鉴别力 + bad = logging.LogRecord( + "x", logging.INFO, __file__, 1, fmt, + (999, "multi", "success", 2, 1, "1", 100, ok_reports, len(reports), 7), + None, + ) + with pytest.raises(TypeError): + bad.getMessage() + + +# ============================================================ +# R4 — multi 模式透传 model +# ============================================================ + +_ORCH_SRC_PATH = _REPO_ROOT / "src" / "llm" / "agents" / "orchestrator.py" + + +def _orchestrator_source() -> str: + """直接读源码,而不是 inspect.getsource(模块属性)。 + + 既有测试(如 test_agent_weights_persist.py)会在运行期把 + orch_mod._agent_provider 换成 lambda/MagicMock,导致 inspect.getsource + 拿到的是 mock 的定义。本文件的断言针对真实源码,故从磁盘读取。 + """ + return _ORCH_SRC_PATH.read_text(encoding="utf-8") + + +def test_r4_predict_match_multi_accepts_model(): + """predict_match_multi 必须接受 model 且默认 None(向后兼容既有调用)。""" + src_fn = _orchestrator_source().split("async def predict_match_multi(", 1)[1] + header = src_fn.split(") -> MultiPredictResult:", 1)[0] + assert "model: str | None = None" in header, header + + +def test_r4_agent_provider_accepts_model_override(): + src_fn = _orchestrator_source().split("async def _agent_provider(", 1)[1] + header = src_fn.split(") -> LLMProvider:", 1)[0] + assert "model_override: str | None = None" in header, header + + +async def test_r4_dispatch_forwards_model_to_multi(monkeypatch): + """行为测试: predict_match(mode=multi, model=...) 必须把 model 送到 multi 路径。""" + from src.llm import predict as predict_mod + + captured: dict = {} + + async def _fake_multi(match_id, **kwargs): + captured["match_id"] = match_id + captured.update(kwargs) + return "SENTINEL" + + # predict.py:188 是函数内 `from ... import`,import 发生在调用时, + # 所以必须打在 orchestrator 模块的属性上。 + monkeypatch.setattr(orch_mod, "predict_match_multi", _fake_multi, raising=True) + + out = await predict_mod.predict_match(999, model="my-model-x", mode="multi") + + assert out == "SENTINEL" + assert captured["model"] == "my-model-x" + assert captured["match_id"] == 999 + + +async def test_r4_specialist_provider_honors_model_override(monkeypatch): + """行为测试: model_override 应覆盖 agent 级/层级默认模型。""" + from src.llm.provider import LLMProvider + import src.llm.agents.orchestrator as real_orch + + async def _fake_default(): + return LLMProvider(api_key="k", base_url="http://x", model="default-model", timeout=1.0) + + async def _fake_runtime(key): + if key.endswith("_MODEL"): + return "agent-level-model" + return None + + monkeypatch.setattr(real_orch, "get_default_provider", _fake_default, raising=True) + monkeypatch.setattr(real_orch, "get_runtime_value", _fake_runtime, raising=True) + monkeypatch.setattr(real_orch, "settings", _settings_with_specialist_model(), raising=True) + monkeypatch.setattr(real_orch, "_AGENT_PROVIDER_CACHE", {}, raising=True) + + overridden = await _REAL_AGENT_PROVIDER("form", tier="specialist", model_override="OVERRIDE") + assert overridden.model == "OVERRIDE" + + # 不带 override 时仍走原优先级(agent 级运行时配置) + real_orch._AGENT_PROVIDER_CACHE.clear() + normal = await _REAL_AGENT_PROVIDER("form", tier="specialist") + assert normal.model == "agent-level-model" + + +def _settings_with_specialist_model(): + class _S: + LLM_SPECIALIST_MODEL = "tier-specialist-model" + LLM_AGGREGATOR_MODEL = "tier-aggregator-model" + + return _S() + + +def test_r4_override_does_not_pollute_cache(): + """override 结果不得写入 60s provider 缓存(否则会串味给普通调用)。""" + src = _orchestrator_source().split("async def _agent_provider(", 1)[1] + src = src.split("async def run_specialists(", 1)[0] + assert "if model_override is None:" in src + assert "_AGENT_PROVIDER_CACHE[cache_key]" in src + + +async def test_r4_dispatch_sets_override_for_specialists(monkeypatch): + """行为测试: predict_match_multi 应把 model 放进 _ACTIVE_MODEL_OVERRIDE, + 并在 run_specialists 执行期间对 specialist 生效(退出后复位)。 + + 参照 tests/test_multi_agent_degraded.py 的 stub 方式,避免触碰真实 DB。 + """ + from unittest.mock import MagicMock + + from src.db.unit_of_work import get_uow + + seen: dict = {} + + async def _fake_header(match_id): + h = MagicMock() + h.match_id = match_id + h.match_dt = None + return h + + async def _fake_specialists(header, *, version, before): + seen["override_during_run"] = orch_mod._ACTIVE_MODEL_OVERRIDE + return [] + + async def _fake_upsert(session, **kw): + p = MagicMock() + p.id = 1 + p.provider = "test" + p.model = kw.get("model") + p.prompt_version = "v1" + p.pred_home_goals = None + p.pred_away_goals = None + p.pred_1x2 = None + p.alt_pred_home_goals = None + p.alt_pred_away_goals = None + p.subjective_confidence = None + p.reasoning = "" + p.agent_outputs = [] + p.agent_weights = {} + p.prompt_tokens = 0 + p.completion_tokens = 0 + return p + + class _FakeUow: + async def __aenter__(self): + return self + + async def __aexit__(self, *a): + return False + + async def get(self, cls, id): + return MagicMock() + + monkeypatch.setattr(orch_mod, "load_match_header", _fake_header, raising=True) + monkeypatch.setattr(orch_mod, "run_specialists", _fake_specialists, raising=True) + monkeypatch.setattr(orch_mod, "_upsert_prediction", _fake_upsert, raising=True) + monkeypatch.setattr(orch_mod, "get_uow", _FakeUow, raising=True) + monkeypatch.setattr(orch_mod, "_ACTIVE_MODEL_OVERRIDE", None, raising=False) + + assert get_uow is not None # 确保 import 生效,session 未被真实打开 + + await orch_mod.predict_match_multi(999, model="OVERRIDE-X") + + assert seen["override_during_run"] == "OVERRIDE-X" + assert orch_mod._ACTIVE_MODEL_OVERRIDE is None, "退出后必须复位" + + +# ============================================================ +# R5 — 回测日期解析 +# ============================================================ + +class TestR5DateBound: + def test_plain_date_start_is_start_of_day_utc(self): + dt = bt_mod._parse_date_bound("2026-01-01", end_of_day=False) + assert dt is not None + assert dt.tzinfo is not None + assert (dt.year, dt.month, dt.day) == (2026, 1, 1) + assert (dt.hour, dt.minute, dt.second) == (0, 0, 0) + + def test_plain_date_end_is_inclusive_end_of_day(self): + """闭区间: 结束日必须取当天末刻,否则最后一天被静默排除。""" + dt = bt_mod._parse_date_bound("2026-01-01", end_of_day=True) + assert dt is not None + assert (dt.hour, dt.minute, dt.second) == (23, 59, 59) + assert dt.microsecond == 999999 + + def test_full_iso_string_is_parsed(self): + dt = bt_mod._parse_date_bound("2026-01-01T12:30:00+08:00", end_of_day=False) + assert dt is not None + assert dt.utcoffset() is not None + + def test_none_returns_none(self): + assert bt_mod._parse_date_bound(None, end_of_day=False) is None + assert bt_mod._parse_date_bound(None, end_of_day=True) is None + + def test_datetime_passthrough(self): + from datetime import datetime, timezone + + given = datetime(2026, 5, 5, 6, 0, tzinfo=timezone.utc) + assert bt_mod._parse_date_bound(given, end_of_day=False) == given + + def test_invalid_input_raises_value_error(self): + with pytest.raises(ValueError): + bt_mod._parse_date_bound("not-a-date", end_of_day=False) -- 2.39.5 From 951faf31dff13f3e925d20604dba80a42859552d Mon Sep 17 00:00:00 2001 From: shangfangjian Date: Mon, 21 Sep 2026 17:31:36 +0800 Subject: [PATCH 5/7] =?UTF-8?q?test:=20=E4=BF=AE=E5=A4=8D=2013=20=E4=B8=AA?= =?UTF-8?q?=E8=85=90=E5=8C=96=E7=94=A8=E4=BE=8B,=E5=A5=97=E4=BB=B6?= =?UTF-8?q?=E6=81=A2=E5=A4=8D=E5=85=A8=E7=BB=BF=E4=B8=94=E9=A1=BA=E5=BA=8F?= =?UTF-8?q?=E6=97=A0=E5=85=B3?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 按「测试腐化(a)/ 源码缺陷(b)/ 测试污染(c)」逐项定性,仅改 tests/: - 外机绝对路径(4): 迁移文件路径改为相对仓库根解析(沿用 test_regressions._read 约定),断言内容保持不变。 - cutoff 用例(4+1): mock 未生效的真因是 _agent_provider 被换成同步 lambda,await 抛 TypeError 被生产代码吞掉后 IndexError;改为 async mock 并按 run_specialists/load_match_header/_agent_provider 的真实契约 打补丁。degraded 用例再加 _upsert_prediction 顶层 kwargs(model)采集。 - 陈旧断言(3): agent 键按现契约断言中文映射;h2h mock 改 async; Match.stats 按设计为 lazy="select",从 MATCH_RELATIONS 移出并单独 固化该设计决定。 - P0-3 守卫(1): seg 越界扫到下游 stats 管线导致误报,改为按缩进收口; 合法形状含经 raw 派生变量中转的写法,并补元测试确保守卫仍能抓到回归。 - 交叉污染(6): test_multi_agent_cutoff 用 patch.object 精确还原,消除 裸赋值泄漏的同步 mock;现已验证顺序无关。 --- tests/test_agent_weights_persist.py | 14 +- tests/test_agents.py | 8 +- tests/test_h2h_perspective.py | 14 +- tests/test_multi_agent_cutoff.py | 114 +++++++------- tests/test_multi_agent_degraded.py | 8 +- tests/test_prediction_unique_constraint.py | 14 +- tests/test_regressions.py | 168 ++++++++++++++++++--- 7 files changed, 241 insertions(+), 99 deletions(-) diff --git a/tests/test_agent_weights_persist.py b/tests/test_agent_weights_persist.py index b9a41f9..2fe3163 100644 --- a/tests/test_agent_weights_persist.py +++ b/tests/test_agent_weights_persist.py @@ -7,12 +7,18 @@ """ from __future__ import annotations +from pathlib import Path from unittest.mock import MagicMock, patch import pytest from src.db.models import Prediction +# 仓库根目录下的 alembic 迁移目录 —— 相对本测试文件解析, +# 避免硬编码某台机器/CI 上的绝对路径(见 tests/test_regressions.py 的 _read 约定)。 +REPO_ROOT = Path(__file__).resolve().parent.parent +MIGRATION_PATH = REPO_ROOT / "alembic" / "versions" / "0014_predictions_agent_weights.py" + class TestAgentWeightsColumn: """验证 predictions 表有 agent_weights 列。""" @@ -38,14 +44,10 @@ class TestMigration: """验证迁移文件存在且内容正确。""" def test_migration_exists(self): - import os - - path = "/.octop/workspaces/CA7PFH/Profeto/alembic/versions/0014_predictions_agent_weights.py" - assert os.path.exists(path) + assert MIGRATION_PATH.is_file(), f"迁移文件不存在: {MIGRATION_PATH}" def test_migration_content(self): - path = "/.octop/workspaces/CA7PFH/Profeto/alembic/versions/0014_predictions_agent_weights.py" - content = open(path).read() + content = MIGRATION_PATH.read_text(encoding="utf-8") assert "agent_weights" in content assert "upgrade" in content diff --git a/tests/test_agents.py b/tests/test_agents.py index 50f8f2b..a8ad2a4 100644 --- a/tests/test_agents.py +++ b/tests/test_agents.py @@ -196,7 +196,7 @@ class TestOrchestratorAggregation: """终裁输入拼装逻辑。""" def test_reports_to_json(self): - from src.llm.agents.orchestrator import _reports_to_json + from src.llm.agents.orchestrator import _reports_to_json, AGENT_LABELS_ZH import json reports = [ @@ -206,8 +206,12 @@ class TestOrchestratorAggregation: text = _reports_to_json(reports) data = json.loads(text) assert len(data) == 2 - assert data[0]["agent"] == "h2h" + # 契约: agent 字段序列化为中文专家全名,引导终裁用统一称呼引用 + # (见 orchestrator.AGENT_LABELS_ZH 与 _reports_to_json 的 docstring) + assert data[0]["agent"] == AGENT_LABELS_ZH["h2h"] == "历史交锋分析专家" assert data[1]["status"] == "no_data" + # 两个 agent 都应被映射,不留英文原键 + assert data[1]["agent"] == AGENT_LABELS_ZH["standings"] def test_aggregator_prompt_renders(self): """终裁 prompt 模板两占位符都能渲染。""" diff --git a/tests/test_h2h_perspective.py b/tests/test_h2h_perspective.py index 4c85164..930f776 100644 --- a/tests/test_h2h_perspective.py +++ b/tests/test_h2h_perspective.py @@ -5,7 +5,7 @@ """ from __future__ import annotations -from unittest.mock import MagicMock +from unittest.mock import MagicMock, patch import pytest @@ -101,14 +101,16 @@ class TestH2HCurrentHomePerspective: home_name="曼城", away_name="诺维奇"), ] import src.llm.context_builder as cb - orig = cb._get_h2h - cb._get_h2h = lambda db, h, a, before, **kw: matches - try: + + async def mock_get_h2h(db, h, a, before, **kw): + # 真实契约是 async(见 context_builder.py 的 `h2h = await _get_h2h(...)`), + # 同步 lambda 会抛 TypeError: object list can't be used in 'await' expression。 + return matches + + with patch.object(cb, "_get_h2h", mock_get_h2h): result = await h2h_slice(header, limit=8, before=None) text = str(result) assert "2胜 0平 0负" in text, f"期望「2胜 0平 0负」,实际:\n{text}" - finally: - cb._get_h2h = orig @pytest.mark.asyncio async def test_draw_counted_correctly(self): diff --git a/tests/test_multi_agent_cutoff.py b/tests/test_multi_agent_cutoff.py index 9d0de94..920b511 100644 --- a/tests/test_multi_agent_cutoff.py +++ b/tests/test_multi_agent_cutoff.py @@ -35,40 +35,49 @@ class TestMultiAgentCutoffPropagation: async def test_backtest_computes_cutoff_from_match_dt_minus_1_day(self): """backtest=True → cutoff = match_dt - 1 天,传给所有切片。""" from datetime import datetime, timedelta, timezone + from unittest.mock import patch import src.llm.agents.orchestrator as orch match_dt = datetime(2026, 1, 15, 20, 0, tzinfo=timezone.utc) header = _make_header(match_dt) captured_before = [] - orig_run_specialists = orch.run_specialists async def mock_run_specialists(header, *, version, before=None): captured_before.append(before) return [] - orch.run_specialists = mock_run_specialists - orch.load_match_header = lambda mid, db=None: header - orch._agent_provider = lambda agent_id, **kw: MagicMock(model="test") + async def mock_header(mid, db=None): + # load_match_header 是 async,必须是 async 函数; + # 且必须走 patch.object(orch 模块属性),因为 predict_match_multi + # 通过模块命名空间解析该名字。裸赋值 orch.load_match_header 同样有效, + # 但用 patch 可保证退出时精确还原,不向后续测试泄漏。 + return header - try: + async def mock_provider(agent_id, *, tier, model_override=None): + # 真实契约是 async(见 orchestrator._agent_provider),同步 lambda + # 会让 `await _agent_provider(...)` 抛 TypeError 并被吞掉。 + return MagicMock(model="test") + + with patch.object(orch, "run_specialists", mock_run_specialists), \ + patch.object(orch, "load_match_header", mock_header), \ + patch.object(orch, "_agent_provider", mock_provider): try: await orch.predict_match_multi(999, backtest=True) except Exception: pass # 后续 aggregator 调用会因 mock 不全而失败,不影响 cutoff 测试 - assert len(captured_before) == 1 - expected_cutoff = match_dt - timedelta(days=1) - assert captured_before[0] == expected_cutoff, ( - f"backtest cutoff 应为 {expected_cutoff},实际 {captured_before[0]}" - ) - finally: - orch.run_specialists = orig_run_specialists + assert len(captured_before) == 1 + expected_cutoff = match_dt - timedelta(days=1) + assert captured_before[0] == expected_cutoff, ( + f"backtest cutoff 应为 {expected_cutoff},实际 {captured_before[0]}" + ) @pytest.mark.asyncio async def test_explicit_cutoff_at_overrides_backtest(self): """显式 cutoff_at 优先于 backtest 自动计算。""" from datetime import datetime, timezone + from unittest.mock import patch import src.llm.agents.orchestrator as orch match_dt = datetime(2026, 1, 15, 20, 0, tzinfo=timezone.utc) @@ -76,97 +85,89 @@ class TestMultiAgentCutoffPropagation: header = _make_header(match_dt) captured_before = [] - orig_run_specialists = orch.run_specialists async def mock_run_specialists(header, *, version, before=None): captured_before.append(before) return [] - orch.run_specialists = mock_run_specialists - orch.load_match_header = lambda mid, db=None: header - orch._agent_provider = lambda agent_id, **kw: MagicMock(model="test") + async def mock_header(mid, db=None): + return header - try: + async def mock_provider(agent_id, *, tier, model_override=None): + return MagicMock(model="test") + + with patch.object(orch, "run_specialists", mock_run_specialists), \ + patch.object(orch, "load_match_header", mock_header), \ + patch.object(orch, "_agent_provider", mock_provider): try: await orch.predict_match_multi(999, backtest=True, cutoff_at=explicit_cutoff) except Exception: pass - assert captured_before[0] == explicit_cutoff - finally: - orch.run_specialists = orig_run_specialists + assert captured_before[0] == explicit_cutoff @pytest.mark.asyncio async def test_normal_mode_cutoff_is_match_dt(self): """非回测模式,无显式 cutoff → cutoff = match_dt。""" from datetime import datetime, timezone + from unittest.mock import patch import src.llm.agents.orchestrator as orch match_dt = datetime(2026, 1, 15, 20, 0, tzinfo=timezone.utc) header = _make_header(match_dt) captured_before = [] - orig_run_specialists = orch.run_specialists async def mock_run_specialists(header, *, version, before=None): captured_before.append(before) return [] - orch.run_specialists = mock_run_specialists - orch.load_match_header = lambda mid, db=None: header - orch._agent_provider = lambda agent_id, **kw: MagicMock(model="test") + async def mock_header(mid, db=None): + return header - try: + async def mock_provider(agent_id, *, tier, model_override=None): + return MagicMock(model="test") + + with patch.object(orch, "run_specialists", mock_run_specialists), \ + patch.object(orch, "load_match_header", mock_header), \ + patch.object(orch, "_agent_provider", mock_provider): try: await orch.predict_match_multi(999, backtest=False) except Exception: pass - assert captured_before[0] == match_dt - finally: - orch.run_specialists = orig_run_specialists + assert captured_before[0] == match_dt @pytest.mark.asyncio async def test_prediction_cutoff_at_stored_not_match_dt(self): """Prediction 写入时 prediction_cutoff_at = 真正 cutoff,非 match_dt。""" from datetime import datetime, timedelta, timezone - from src.llm.predict import _predict_single, PredictResult + from unittest.mock import patch + import src.llm.predict as pred match_dt = datetime(2026, 1, 15, 20, 0, tzinfo=timezone.utc) expected_cutoff = match_dt - timedelta(days=1) - # Mock build_context to return a context with cutoff - import src.llm.predict as pred - orig_build = pred.build_context - class FakeContext: text = "fake" - match_dt = match_dt cutoff = expected_cutoff - async def fake_build(match_id, **kw): + FakeContext.match_dt = match_dt + + call_args = {} + + async def tracking_build(match_id, **kw): + call_args.update(kw) return FakeContext() - pred.build_context = fake_build - pred._upsert_prediction = lambda session, **kw: MagicMock(id=1, **kw.get("values", {})) - - try: - # 此处只验证 cutoff 参数传递,实际 LLM 调用会被 mock 阻断 - # 重点: build_context 被调用时传入 backtest=True 和正确的 cutoff - call_args = {} - async def tracking_build(match_id, **kw): - call_args.update(kw) - return FakeContext() - pred.build_context = tracking_build - + with patch.object(pred, "build_context", tracking_build): try: - await _predict_single(999, backtest=True) + await pred._predict_single(999, backtest=True) except Exception: pass - assert call_args.get("backtest") is True, "backtest=True 应传递给 build_context" - finally: - pred.build_context = orig_build + # 重点: build_context 被调用时传入 backtest=True 和正确的 cutoff + assert call_args.get("backtest") is True, "backtest=True 应传递给 build_context" class TestBacktestXgNotVisible: @@ -175,8 +176,8 @@ class TestBacktestXgNotVisible: @pytest.mark.asyncio async def test_stats_slice_respects_cutoff_for_xg_availability(self): """available_at > cutoff 的 xG 数据不应被切片使用。""" - from datetime import datetime, timedelta, timezone - from unittest.mock import MagicMock + from datetime import datetime, timezone + from unittest.mock import MagicMock, patch cutoff = datetime(2026, 1, 13, 20, 0, tzinfo=timezone.utc) # match_date - 2天 match_dt = datetime(2026, 1, 15, 20, 0, tzinfo=timezone.utc) @@ -206,7 +207,6 @@ class TestBacktestXgNotVisible: header = _make_header(match_dt) import src.llm.context_builder as cb - orig_get_form = cb._get_form async def mock_get_form(db, team_id, before, *, limit=10): # before=cutoff(1月13日),比赛在1月15日,满足 before 条件 @@ -214,14 +214,10 @@ class TestBacktestXgNotVisible: return [hist_match] return [] - cb._get_form = mock_get_form - - try: + with patch.object(cb, "_get_form", mock_get_form): result = await cb.stats_slice(header, limit=10, before=cutoff) text = str(result) # xG 在 cutoff 之后才 available,不应出现在切片 assert "2.50" not in text, f"xG 2.50 不应在切片中(available_at > cutoff):\n{text}" # 但无比分时仍应显示进球数据 assert "无比分数据" in text or "场均进球" in text, f"无比分时仍应显示基本数据:\n{text}" - finally: - cb._get_form = orig_get_form diff --git a/tests/test_multi_agent_degraded.py b/tests/test_multi_agent_degraded.py index c641080..9528be8 100644 --- a/tests/test_multi_agent_degraded.py +++ b/tests/test_multi_agent_degraded.py @@ -268,11 +268,17 @@ class TestNoAggregatorCallOnDegraded: captured_values = {} async def mock_upsert(session, **kw): + # model / provider_name / mode 是 _upsert_prediction 的顶层关键字参数, + # 不在 values 字典里(见 orchestrator.py 的调用点)。原测试只取 + # kw["values"],导致 model 断言永远为 None。 captured_values.update(kw.get("values", {})) + captured_values.update( + {k: kw.get(k) for k in ("model", "provider_name", "mode", "run_type")} + ) mock_pred = MagicMock() mock_pred.id = 1 mock_pred.provider = "test" - mock_pred.model = kw["values"].get("model") + mock_pred.model = kw.get("model") return mock_pred class FakeUow: diff --git a/tests/test_prediction_unique_constraint.py b/tests/test_prediction_unique_constraint.py index ac1fb5b..9b7c547 100644 --- a/tests/test_prediction_unique_constraint.py +++ b/tests/test_prediction_unique_constraint.py @@ -8,12 +8,18 @@ from __future__ import annotations import inspect +from pathlib import Path from pydantic import BaseModel import pytest from src.db.models import Prediction, UniqueConstraint, CheckConstraint +# 仓库根目录下的 alembic 迁移目录 —— 相对本测试文件解析, +# 避免硬编码某台机器/CI 上的绝对路径(见 tests/test_regressions.py 的 _read 约定)。 +REPO_ROOT = Path(__file__).resolve().parent.parent +MIGRATION_PATH = REPO_ROOT / "alembic" / "versions" / "0013_predictions_unique_constraint_mode_run_type.py" + class TestUniqueConstraint: """验证唯一约束包含 mode + run_type。""" @@ -71,14 +77,10 @@ class TestMigration: """验证迁移文件存在且内容正确。""" def test_migration_exists(self): - import os - - path = "/.octop/workspaces/CA7PFH/Profeto/alembic/versions/0013_predictions_unique_constraint_mode_run_type.py" - assert os.path.exists(path) + assert MIGRATION_PATH.is_file(), f"迁移文件不存在: {MIGRATION_PATH}" def test_migration_adds_column_and_constraint(self): - path = "/.octop/workspaces/CA7PFH/Profeto/alembic/versions/0013_predictions_unique_constraint_mode_run_type.py" - content = open(path).read() + content = MIGRATION_PATH.read_text(encoding="utf-8") assert 'run_type' in content assert 'uq_predictions_match_provider_model_mode_run_type' in content diff --git a/tests/test_regressions.py b/tests/test_regressions.py index 755e85e..e26ee3f 100644 --- a/tests/test_regressions.py +++ b/tests/test_regressions.py @@ -14,8 +14,10 @@ from pathlib import Path SRC = Path(__file__).resolve().parent.parent / "src" -# 切片函数会读取的关系属性 → 查询时必须 eager-load -MATCH_RELATIONS = ("stats", "home_team", "away_team", "league") +# 切片函数会读取的关系属性 → 查询时必须 eager-load。 +# Match.stats 刻意排除: 它按设计用 lazy="select",由 selectinload(Match.stats) +# 显式预加载(见 test_stats_relationship_is_lazy_select_by_design)。 +MATCH_RELATIONS = ("home_team", "away_team", "league") def _read(rel: str) -> str: @@ -48,27 +50,127 @@ class TestEagerLoadCoverage: assert "selectinload" in src, "backtest 未 eager-load 关系 (P0-1)" def test_relationship_default_is_selectin(self): - """models.py 中 Match 的高频关系应声明 lazy='selectin' 作为兜底。""" + """models.py 中 Match 的高频关系应声明 lazy='selectin' 作为兜底。 + + 这里只要求「高频一起读取」的关系(set MATCH_RELATIONS)声明 selectin。 + Match.stats 刻意用 lazy="select" —— 它只在 stats 管线里按需取,不在 + 每个切片都读,而且它的预加载由 selectinload(Match.stats) 显式表达 + (见 test_context_builder_getters_eager_load)。 + """ src = _read("db/models.py") # 找到 Match 类定义段 m = re.search(r"class Match\(Base\):.*?(?=\nclass )", src, re.S) assert m, "Match 类未找到" body = m.group(0) for rel in MATCH_RELATIONS: - # 关系声明可能跨多行(stats/home_team/away_team 都是),因此按 - # 「从 `rel: Mapped` 到下一个 `xxx: Mapped` 之前」整段匹配。 + # 关系声明可能跨多行(home_team/away_team 都是),因此按 + # 「从 `rel: Mapped` 到下一个单行注解声明之前」整段匹配。 m_rel = re.search( - rf"^\s*{rel}: Mapped.*?(?=^\s*\w+: Mapped|\Z)", body, re.M | re.S + rf"^[ \t]*{rel}: Mapped.*?(?=^[ \t]*\w+:[^\n]*Mapped|\Z)", + body, re.M | re.S, ) assert m_rel, f"Match.{rel} 未找到" assert 'lazy="selectin"' in m_rel.group(0), ( f"Match.{rel} 未声明 lazy='selectin' —— 兜底缺失 (P0-2)" ) + def test_stats_relationship_is_lazy_select_by_design(self): + """Match.stats 刻意保持 lazy="select"(不是回归)。 + + 它是唯一需要显式 selectinload 才预加载的关系 —— 若哪天有人把它也 + 改成 selectin,上面的 test_context_builder_getters_eager_load 和 + bzzoiro stats 管线仍应工作,但本用例会提醒复核该设计决定。 + """ + src = _read("db/models.py") + body = re.search(r"class Match\(Base\):.*?(?=\nclass )", src, re.S).group(0) + m_rel = re.search( + r"^[ \t]*stats: Mapped.*?(?=^[ \t]*\w+:[^\n]*Mapped|\Z)", + body, re.M | re.S, + ) + assert m_rel, "Match.stats 未找到" + assert 'lazy="select"' in m_rel.group(0), ( + "Match.stats 预期为 lazy='select'(按需加载),实际声明已变 —— 请复核设计" + ) + class TestBzzoiroLineage: """P0-3: source_event_id 必须取配对的 raw,不能是循环残留变量。""" + # 消费循环的起始行匹配模式(见 bzzoiro.py 顶部 events 管线的内层循环) + _LOOP_PATTERN = "for nm, raw in normalized_matches" + + # source_event_id 的合法赋值形状(两种,都必须取配对的 raw): + # 1) 构造新比赛: `source_event_id=_to_int_or_none(raw.get("id")),` + # 2) 回填已有比赛: `eid = _to_int_or_none(raw.get("id"))` → + # `existing_match.source_event_id = eid` + # 非法形状(即 P0-3 回归): 直接用未配对的变量给 ORM 对象赋值。 + _ASSIGN_DIRECT = re.compile( + r"source_event_id\s*=\s*(?:[A-Za-z_][\w.]*\s*\(\s*)?raw(?:\.get\(|\s*\[)" + ) + # 赋值给未配对的局部变量: `source_event_id = <变量>` + _ASSIGN_VIA_VAR = re.compile( + r"source_event_id\s*=\s*([A-Za-z_]\w*)\s*$" + ) + + def _consume_loop_body(self, src: str) -> str: + """截取 `for nm, raw in normalized_matches` 循环体,不含循环之后的下游代码。 + + 原实现是 `seg = "\\n".join(lines[start:])`,一直取到文件末尾,于是把 + 无关的下游 stats 管线(bzzoiro.py 的 `_backfill_stats`)也扫了进来 —— + 那里合法地在 ORM 对象上访问 `m.source_event_id`,导致误报 P0-3。 + 这里按缩进边界正确收口:循环体内每行要么是空行/注释,要么缩进严格 + 大于 `for` 行。 + """ + lines = src.splitlines() + start = next( + (i for i, ln in enumerate(lines) if self._LOOP_PATTERN in ln), None + ) + assert start is not None, f"未找到消费循环: {self._LOOP_PATTERN}" + + for_indent = len(lines[start]) - len(lines[start].lstrip()) + kept: list[str] = [lines[start]] + for ln in lines[start + 1:]: + stripped = ln.strip() + # 顺序要保持: 空行与注释行缩进为 0,不能拿它们做边界判断 + if not stripped or stripped.startswith("#"): + continue + indent = len(ln) - len(ln.lstrip()) + if indent <= for_indent: + break # 循环结束,后续属下游代码 + kept.append(ln) + # 右侧剥离注释:避免 `# ... raw.get(...)` 这类注释误命中赋值正则 + return "\n".join(ln.split("#", 1)[0] for ln in kept) + + def _bad_assignments(self, seg: str) -> list[str]: + """返回循环体内未取配对 raw 的 source_event_id 赋值行。""" + # 先收集"来自配对 raw"的局部变量: `eid = _to_int_or_none(raw.get("id"))` + # (不受行序影响,所以必须先建好,再判定中转赋值) + raw_vars: set[str] = set() + for ln in seg.splitlines(): + m = re.match( + r"\s*([A-Za-z_]\w*)\s*=\s*.*raw(?:\.get\(|\s*\[)", ln + ) + if m: + raw_vars.add(m.group(1)) + + bad: list[str] = [] + for ln in seg.splitlines(): + stripped = ln.strip() + if "source_event_id" not in stripped: + continue + # 读取判断/比较(`if x.source_event_id is None:`)不算赋值 + if re.search(r"source_event_id\s*(?:is|==|!=)", stripped): + continue + if re.search(r"source_event_id\s*\.\s*\w+\s*\(", stripped): + continue # 方法调用,不是赋值 + if self._ASSIGN_DIRECT.search(stripped): + continue # 直接取配对 raw + m_var = self._ASSIGN_VIA_VAR.search(stripped) + if m_var and m_var.group(1) in raw_vars: + continue # 经由已确认来自 raw 的局部变量中转 + bad.append(stripped) + return bad + def test_normalized_matches_carries_raw(self): src = _read("data/bzzoiro.py") # 规范化结果必须与原始 event 成对保存 @@ -81,19 +183,47 @@ class TestBzzoiroLineage: ) def test_no_orphan_raw_use(self): - """source_event_id 所在行必须在解包循环内(用缩进 + 上下文粗判)。""" + """循环体内每个 source_event_id 赋值都必须取自配对的 raw(P0-3)。 + + 用正则而不是 `raw.get(` 子串匹配:合法写法含「回填已有比赛」那条 + (`eid = raw.get("id")` 之后 `existing_match.source_event_id = eid`), + 它不是 `raw.get(` 同一行,但同样正确。非法写法(回归)是直接 + `existing_match.source_event_id = orphan_var`。 + """ src = _read("data/bzzoiro.py") - lines = src.splitlines() - # 找到 "for nm, raw in normalized_matches" 所在行号 - start = next( - (i for i, ln in enumerate(lines) if "for nm, raw in normalized_matches" in ln), - None, + seg = self._consume_loop_body(src) + bad = self._bad_assignments(seg) + assert len(bad) == 0, ( + f"source_event_id 未使用配对的 raw (P0-3),问题行: {bad}" ) - assert start is not None - # 该循环之后、下一个同/更低缩进的顶层语句之前的范围 - seg = "\n".join(lines[start:]) - uses = [ln for ln in seg.splitlines() if "source_event_id" in ln] - assert uses, "未找到 source_event_id 赋值" - assert all("raw.get(" in ln for ln in uses), ( - "source_event_id 未使用配对的 raw (P0-3)" + + def test_loop_body_scope_excludes_downstream_stats_pipeline(self): + """作用域守卫: 截取段不能扫到循环之后的下游 stats 管线。 + + 下游 `_backfill_stats` 里合法地在 ORM 对象上访问 `m.source_event_id` + (与配对 raw 无关)。若 seg 越界,test_no_orphan_raw_use 会误报。 + """ + src = _read("data/bzzoiro.py") + seg = self._consume_loop_body(src) + assert "m.source_event_id" not in seg, ( + "循环体截取越界,扫到了下游 stats 管线 —— 会误报 P0-3" ) + # 但配对使用必须仍在作用域内 + assert "raw.get(" in seg, "循环体内应保留 `raw.get(...)` 的配对用法" + + def test_guard_detects_orphan_variable_regression(self): + """守卫有效性: 若 source_event_id 改成取循环外残留变量,必须被判失败。 + + 回归保护的"元测试"——确保上面的正则在真实缺陷面前确实会红, + 而不是恒真的空断言。 + """ + orphan = """ + for nm, raw in normalized_matches: + m = Match( + league_id=1, + source_event_id=_to_int_or_none(orphan.get("id")), + ) + """ + seg = self._consume_loop_body(orphan) + bad = self._bad_assignments(seg) + assert bad, "守卫失效: 未配对的 orphan 变量未被识别为 P0-3 回归" -- 2.39.5 From 50a6753ff8a02644ad9f1080e93e7f1262ec4386 Mon Sep 17 00:00:00 2001 From: shangfangjian Date: Mon, 21 Sep 2026 17:36:07 +0800 Subject: [PATCH 6/7] =?UTF-8?q?test:=20=E4=BF=AE=E6=AD=A3=20selectin=20?= =?UTF-8?q?=E5=AE=88=E5=8D=AB=E7=9A=84=E5=81=87=E9=98=B3=E6=80=A7(?= =?UTF-8?q?=E5=8F=98=E5=BC=82=E6=B5=8B=E8=AF=95=E5=8F=91=E7=8E=B0)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 原断言对「关系声明到下一注解」的整段(含注释)做 lazy="selectin" 子串匹配;而 Match.league 声明正下方的注释里恰好也写着该字样, 导致移除真实 kwarg 后测试依旧通过(变异测试确认)。 改为只匹配 relationship(...) 括号内内容,并加注释说明该假阳性。 验证: 干净代码通过 → 移除 league 的 lazy 后失败 → 还原后又通过。 全套 238 passed / 1 skipped。 --- tests/test_regressions.py | 15 +++++++++------ 1 file changed, 9 insertions(+), 6 deletions(-) diff --git a/tests/test_regressions.py b/tests/test_regressions.py index e26ee3f..4ff3601 100644 --- a/tests/test_regressions.py +++ b/tests/test_regressions.py @@ -63,15 +63,18 @@ class TestEagerLoadCoverage: assert m, "Match 类未找到" body = m.group(0) for rel in MATCH_RELATIONS: - # 关系声明可能跨多行(home_team/away_team 都是),因此按 - # 「从 `rel: Mapped` 到下一个单行注解声明之前」整段匹配。 + # 只取 relationship(...) 调用本身的括号内内容。 + # 注意: 不能把整段(含注释)做子串匹配 —— 关系声明下方的注释里 + # 恰好也写着 lazy="selectin",会导致「删掉真实 kwarg 但测试仍绿」 + # 的假阳性(已用变异测试证实: 移除 league 的 lazy 后断言依旧通过)。 m_rel = re.search( - rf"^[ \t]*{rel}: Mapped.*?(?=^[ \t]*\w+:[^\n]*Mapped|\Z)", + rf"^[ \t]*{rel}: Mapped.*?relationship\((.*?)\)[ \t]*$", body, re.M | re.S, ) - assert m_rel, f"Match.{rel} 未找到" - assert 'lazy="selectin"' in m_rel.group(0), ( - f"Match.{rel} 未声明 lazy='selectin' —— 兜底缺失 (P0-2)" + assert m_rel, f"Match.{rel} 未找到 relationship(...) 声明" + call_args = m_rel.group(1) + assert 'lazy="selectin"' in call_args, ( + f"Match.{rel} 的 relationship() 未声明 lazy='selectin' —— 兜底缺失 (P0-2)" ) def test_stats_relationship_is_lazy_select_by_design(self): -- 2.39.5 From 06e973d34d6bfc2989ae79f96e3af84d8ec86fe3 Mon Sep 17 00:00:00 2001 From: shangfangjian Date: Mon, 21 Sep 2026 17:52:45 +0800 Subject: [PATCH 7/7] =?UTF-8?q?fix(concurrency):=20=E6=B6=88=E9=99=A4=20mo?= =?UTF-8?q?del=20=E8=A6=86=E7=9B=96=E7=9A=84=E6=A8=A1=E5=9D=97=E7=BA=A7?= =?UTF-8?q?=E5=8F=98=E9=87=8F=E4=B8=AD=E8=BD=AC(=E5=B9=B6=E5=8F=91?= =?UTF-8?q?=E4=B8=B2=E5=91=B3)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit R4 首版实现用模块级 _ACTIVE_MODEL_OVERRIDE 中转 model override, 理由是不想改动 run_specialists 的签名(既有测试会 mock 它)。 但 backtest.py:225 会 asyncio.gather 并发 8 场预测(Semaphore(8)), 每场都调用 predict_match_multi —— 模块级变量会被并发调用互相覆盖, 导致 A 场的预测用上 B 场的模型。这是静默的正确性缺陷。 改为: run_specialists 新增 model_override 形参,一路显式下传; 删除模块级变量。同步更新 6 处 mock 签名。 新增 3 个守卫并做变异验证: - test_r4_no_module_level_model_override_global (源码级,可判别) - test_r4_run_specialists_accepts_model_override_parameter - test_r4_dispatch_passes_override_to_specialists 变异测试: 重新引入全局变量方案 → 两个守卫变红;还原 → 全绿。 注: 曾尝试写并发行为测试,但真实 run_specialists 会访问数据库, 测试环境下不稳定(ConnectionRefusedError),会是 flaky 的假证据, 故改用清晰的源码级判别守卫,并在注释中说明原因。 240 passed / 1 skipped; tsc 0 error; vite build 成功。 --- src/llm/agents/orchestrator.py | 27 +++++------ tests/test_agent_weights_persist.py | 2 +- tests/test_multi_agent_cutoff.py | 6 +-- tests/test_multi_agent_degraded.py | 8 ++-- tests/test_review_required_fixes.py | 70 +++++++++++++++++++++++++---- 5 files changed, 80 insertions(+), 33 deletions(-) diff --git a/src/llm/agents/orchestrator.py b/src/llm/agents/orchestrator.py index 35fc16d..f54621b 100644 --- a/src/llm/agents/orchestrator.py +++ b/src/llm/agents/orchestrator.py @@ -34,11 +34,6 @@ logger = logging.getLogger(__name__) _AGENT_PROVIDER_CACHE: dict[str, tuple[float, LLMProvider]] = {} _AGENT_PROVIDER_CACHE_TTL = 60.0 -# 当前预测的模型覆盖(调用方显式指定)。用模块级变量而非新增形参,因为 -# run_specialists 会被既有测试 mock,加形参会破坏那些测试的调用签名。 -# 由 predict_match_multi 在进入时 set、退出时 reset。 -_ACTIVE_MODEL_OVERRIDE: str | None = None - # ── 5 个专家 agent 定义 ── # A=近期状态 B=攻防数据 C=主客因素 D=联赛排名 E=历史交锋 @@ -160,16 +155,17 @@ async def run_specialists( *, version: str = "v1", before=None, + model_override: str | None = None, ) -> list[AgentReport]: """并行执行 5 个专家 agent。fail-open: 单个失败不影响其他。 before: 数据截止时间(回测防泄漏)。None 表示不限制。 + model_override: 调用方显式指定的模型,覆盖各 agent 的层级默认。 - 模型覆盖通过 _ACTIVE_MODEL_OVERRIDE 传递,而不是新增形参:既有测试 - 会 mock 本函数(见 tests/test_agent_weights_persist.py),加形参会破坏 - 它们的调用签名。predict_match_multi 在调用前后 set/reset 该变量。 + 注意:model_override 必须作为形参下传,不能用模块级变量中转。 + backtest 会 asyncio.gather 并发 8 场预测(见 backtest.py 的 Semaphore(8)), + 模块级变量会被并发调用互相覆盖,导致 A 场的预测用上 B 场的模型。 """ - model_override = _ACTIVE_MODEL_OVERRIDE tasks = [ _run_one(spec, header, await _agent_provider(spec.name, tier="specialist", model_override=model_override), version=version, before=before) for spec in SPECIALIST_SPECS @@ -266,14 +262,11 @@ async def predict_match_multi( prediction_cutoff_at = cutoff # 2. 并行专家(各自独立配置,使用统一 cutoff) - # 显式 model 覆盖通过模块级变量下传,避免改动 run_specialists 的签名 - global _ACTIVE_MODEL_OVERRIDE - _prev_override = _ACTIVE_MODEL_OVERRIDE - _ACTIVE_MODEL_OVERRIDE = model - try: - reports = await run_specialists(header, version=version, before=cutoff) - finally: - _ACTIVE_MODEL_OVERRIDE = _prev_override + # model 作为形参下传,而非模块级变量:backtest 并发 8 场预测时, + # 模块级变量会被并发调用互相覆盖(模型串味)。 + reports = await run_specialists( + header, version=version, before=cutoff, model_override=model + ) # 2.5 统计有效专家报告数量 ok_reports = [r for r in reports if r.status == "ok"] diff --git a/tests/test_agent_weights_persist.py b/tests/test_agent_weights_persist.py index 2fe3163..e2a11cc 100644 --- a/tests/test_agent_weights_persist.py +++ b/tests/test_agent_weights_persist.py @@ -84,7 +84,7 @@ class TestOrchestratorWritesAgentWeights: captured_values = {} - async def mock_specialists(h, *, version, before): + async def mock_specialists(h, *, version, before, model_override=None): return reports async def mock_provider(aid, **kw): diff --git a/tests/test_multi_agent_cutoff.py b/tests/test_multi_agent_cutoff.py index 920b511..e03100a 100644 --- a/tests/test_multi_agent_cutoff.py +++ b/tests/test_multi_agent_cutoff.py @@ -43,7 +43,7 @@ class TestMultiAgentCutoffPropagation: captured_before = [] - async def mock_run_specialists(header, *, version, before=None): + async def mock_run_specialists(header, *, version, before=None, model_override=None): captured_before.append(before) return [] @@ -86,7 +86,7 @@ class TestMultiAgentCutoffPropagation: captured_before = [] - async def mock_run_specialists(header, *, version, before=None): + async def mock_run_specialists(header, *, version, before=None, model_override=None): captured_before.append(before) return [] @@ -118,7 +118,7 @@ class TestMultiAgentCutoffPropagation: captured_before = [] - async def mock_run_specialists(header, *, version, before=None): + async def mock_run_specialists(header, *, version, before=None, model_override=None): captured_before.append(before) return [] diff --git a/tests/test_multi_agent_degraded.py b/tests/test_multi_agent_degraded.py index 9528be8..94a0a25 100644 --- a/tests/test_multi_agent_degraded.py +++ b/tests/test_multi_agent_degraded.py @@ -66,7 +66,7 @@ class TestAllExpertsFailed: header = _make_header() # Mock run_specialists 返回全 error - async def mock_run_specialists(header, *, version, before): + async def mock_run_specialists(header, *, version, before, model_override=None): return _all_error_reports() # Mock _agent_provider @@ -131,7 +131,7 @@ class TestAllExpertsFailed: """5 个专家全 no_data → status=degraded,不调终裁。""" header = _make_header() - async def mock_run_specialists(header, *, version, before): + async def mock_run_specialists(header, *, version, before, model_override=None): return _all_no_data_reports() async def mock_agent_provider(agent_id, *, tier): @@ -188,7 +188,7 @@ class TestPartialExpertsOk: """1 个 ok + 4 个 error → status=success(走终裁)。""" header = _make_header() - async def mock_run_specialists(header, *, version, before): + async def mock_run_specialists(header, *, version, before, model_override=None): return _mixed_reports() async def mock_agent_provider(agent_id, *, tier): @@ -255,7 +255,7 @@ class TestNoAggregatorCallOnDegraded: header = _make_header() aggregator_called = [] - async def mock_run_specialists(header, *, version, before): + async def mock_run_specialists(header, *, version, before, model_override=None): return _all_error_reports() async def mock_agent_provider(agent_id, *, tier): diff --git a/tests/test_review_required_fixes.py b/tests/test_review_required_fixes.py index 4eaa07f..e070e45 100644 --- a/tests/test_review_required_fixes.py +++ b/tests/test_review_required_fixes.py @@ -13,6 +13,7 @@ from __future__ import annotations import inspect import logging import pathlib +import re import pytest @@ -348,9 +349,13 @@ def test_r4_override_does_not_pollute_cache(): assert "_AGENT_PROVIDER_CACHE[cache_key]" in src -async def test_r4_dispatch_sets_override_for_specialists(monkeypatch): - """行为测试: predict_match_multi 应把 model 放进 _ACTIVE_MODEL_OVERRIDE, - 并在 run_specialists 执行期间对 specialist 生效(退出后复位)。 +async def test_r4_dispatch_passes_override_to_specialists(monkeypatch): + """行为测试: predict_match_multi 必须把 model 作为**形参**传给 run_specialists。 + + 早期实现用模块级变量 _ACTIVE_MODEL_OVERRIDE 中转,但 backtest 会 + asyncio.gather 并发 8 场预测(backtest.py Semaphore(8)),全局变量会被 + 并发调用互相覆盖 → A 场预测用上 B 场的模型。故此处断言「形参传递」, + 并显式断言该模块级变量已不存在。 参照 tests/test_multi_agent_degraded.py 的 stub 方式,避免触碰真实 DB。 """ @@ -366,8 +371,8 @@ async def test_r4_dispatch_sets_override_for_specialists(monkeypatch): h.match_dt = None return h - async def _fake_specialists(header, *, version, before): - seen["override_during_run"] = orch_mod._ACTIVE_MODEL_OVERRIDE + async def _fake_specialists(header, *, version, before, model_override=None): + seen["override_during_run"] = model_override return [] async def _fake_upsert(session, **kw): @@ -403,14 +408,63 @@ async def test_r4_dispatch_sets_override_for_specialists(monkeypatch): monkeypatch.setattr(orch_mod, "run_specialists", _fake_specialists, raising=True) monkeypatch.setattr(orch_mod, "_upsert_prediction", _fake_upsert, raising=True) monkeypatch.setattr(orch_mod, "get_uow", _FakeUow, raising=True) - monkeypatch.setattr(orch_mod, "_ACTIVE_MODEL_OVERRIDE", None, raising=False) assert get_uow is not None # 确保 import 生效,session 未被真实打开 await orch_mod.predict_match_multi(999, model="OVERRIDE-X") - assert seen["override_during_run"] == "OVERRIDE-X" - assert orch_mod._ACTIVE_MODEL_OVERRIDE is None, "退出后必须复位" + assert seen["override_during_run"] == "OVERRIDE-X", ( + "model 未作为形参传给 run_specialists" + ) + # 回归守卫: 模块级中转变量必须不存在(并发下会产生模型串味) + assert not hasattr(orch_mod, "_ACTIVE_MODEL_OVERRIDE"), ( + "不应再用模块级 _ACTIVE_MODEL_OVERRIDE 中转 model: backtest 并发 8 场预测时" + "会互相覆盖,导致模型串味" + ) + + +def test_r4_run_specialists_accepts_model_override_parameter(): + """run_specialists 必须显式接收 model_override 形参(而非读全局)。""" + import inspect + + sig = inspect.signature(orch_mod.run_specialists) + assert "model_override" in sig.parameters, ( + "run_specialists 缺少 model_override 形参 —— 并发场景下模型会串味" + ) + assert sig.parameters["model_override"].default is None + + +def test_r4_no_module_level_model_override_global(): + """并发安全守卫: orchestrator 不得用模块级变量中转 model 覆盖。 + + backtest 会 asyncio.gather 并发 8 场预测(backtest.py 的 Semaphore(8)), + 模块级变量会被并发调用互相覆盖 → A 场的预测用上 B 场的模型(模型串味)。 + 正确做法是把 model 作为形参一路下传。 + + 说明: 这里用源码级断言而非并发行为测试 —— 真实 run_specialists 会调用 + 数据库(_agent_provider -> load_match_header),在无 DB 的测试环境下 + 无法稳定执行,写出来的并发测试会是 flaky 的假证据(已实测确认)。 + 形参方案与全局方案的判别点清晰且可直接观测,故用源码守卫。 + """ + src = _orchestrator_source() + + # 1) 不得存在模块级覆盖变量 + assert "_ACTIVE_MODEL_OVERRIDE" not in src, ( + "orchestrator 又引入了模块级 model 覆盖变量 —— 并发预测会模型串味" + ) + # 2) 不得有 `global` 声明去写模型覆盖 + assert not re.search(r"^\s*global\s+.*MODEL", src, re.M), ( + "orchestrator 使用 global 声明中转模型覆盖 —— 并发下不安全" + ) + # 3) model_override 必须作为实参出现在 run_specialists 调用里 + call = re.search( + r"await run_specialists\((.*?)\)", src, re.S + ) + assert call, "未找到 run_specialists 调用点" + assert "model_override=" in call.group(1), ( + "run_specialists 调用点未显式传 model_override —— " + "model 可能又走回隐式中转,并发下会串味" + ) # ============================================================ -- 2.39.5