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/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..c860b84 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,35 @@ 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() 对所有数据源都报「未知数据源」,把导入错误
+ 伪装成「不存在的名字」—— 这类静默失效极难定位,故在此显式记日志。
+
+ 失败时保持 _loaded=False:后续调用会重新尝试导入。正常应用流程中
+ src.api.app 先导入本模块,此处的模块级预热是在 `src.data.bzzoiro`
+ 尚未初始化时发起的,导入链在本模块内成环,首次必然失败(bzzoiro
+ 仍在加载中),由后续 get_source() / list_sources() 调用完成真正的装载。
+ """
+ global _loaded
+ if _loaded:
+ return
+ try:
+ from src.data.bzzoiro import BzzoiroSource # noqa: F401
+ except Exception:
+ logger.exception(
+ "数据源模块导入失败(通常为初始化中途的环状导入),本次注册表为空;"
+ "下次 get_source()/list_sources() 调用会自动重试"
+ )
+ return # 保持 _loaded=False,下次调用可重试
+ _loaded = True
-# 保持向后兼容:模块加载时尝试加载(但不再强制)
-try:
- _load_sources()
-except Exception:
- pass
+# 模块加载时预热(失败会记录日志,不再静默)
+_load_sources()
diff --git a/src/llm/agents/orchestrator.py b/src/llm/agents/orchestrator.py
index b131c35..f54621b 100644
--- a/src/llm/agents/orchestrator.py
+++ b/src/llm/agents/orchestrator.py
@@ -105,21 +105,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 +132,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 +142,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
@@ -148,13 +155,19 @@ 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 的层级默认。
+
+ 注意:model_override 必须作为形参下传,不能用模块级变量中转。
+ backtest 会 asyncio.gather 并发 8 场预测(见 backtest.py 的 Semaphore(8)),
+ 模块级变量会被并发调用互相覆盖,导致 A 场的预测用上 B 场的模型。
"""
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 +235,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 +262,11 @@ async def predict_match_multi(
prediction_cutoff_at = cutoff
# 2. 并行专家(各自独立配置,使用统一 cutoff)
- reports = await run_specialists(header, version=version, before=cutoff)
+ # 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"]
@@ -279,7 +298,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 +366,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_agent_weights_persist.py b/tests/test_agent_weights_persist.py
index b9a41f9..e2a11cc 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
@@ -82,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_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_bzzoiro_source_registry.py b/tests/test_bzzoiro_source_registry.py
new file mode 100644
index 0000000..e0842fa
--- /dev/null
+++ b/tests/test_bzzoiro_source_registry.py
@@ -0,0 +1,144 @@
+"""回归测试:锁定 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 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` 吞掉导入失败。"""
+
+ 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 吞掉导入异常"
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..e03100a 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):
+ async def mock_run_specialists(header, *, version, before=None, model_override=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):
+ async def mock_run_specialists(header, *, version, before=None, model_override=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):
+ async def mock_run_specialists(header, *, version, before=None, model_override=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..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):
@@ -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..4ff3601 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,130 @@ 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` 之前」整段匹配。
+ # 只取 relationship(...) 调用本身的括号内内容。
+ # 注意: 不能把整段(含注释)做子串匹配 —— 关系声明下方的注释里
+ # 恰好也写着 lazy="selectin",会导致「删掉真实 kwarg 但测试仍绿」
+ # 的假阳性(已用变异测试证实: 移除 league 的 lazy 后断言依旧通过)。
m_rel = re.search(
- rf"^\s*{rel}: Mapped.*?(?=^\s*\w+: Mapped|\Z)", body, re.M | re.S
+ 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):
+ """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 +186,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 回归"
diff --git a/tests/test_review_required_fixes.py b/tests/test_review_required_fixes.py
new file mode 100644
index 0000000..e070e45
--- /dev/null
+++ b/tests/test_review_required_fixes.py
@@ -0,0 +1,506 @@
+"""回归测试: 代码评审确认的 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 re
+
+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_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。
+ """
+ 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, model_override=None):
+ seen["override_during_run"] = 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)
+
+ 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", (
+ "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 可能又走回隐式中转,并发下会串味"
+ )
+
+
+# ============================================================
+# 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)
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)"
+ )