fix: 修复全量审查确认的 3 Critical + 5 Required,并修复 13 个腐化用例 #8

Merged
shangfangjian merged 7 commits from fix-review-report-critical into main 2026-09-21 18:04:54 +08:00
19 changed files with 1586 additions and 179 deletions
+1 -1
View File
@@ -163,7 +163,7 @@ export default function EvalPage() {
<Spinner /> <Spinner />
</div> </div>
) : summary.length === 0 ? ( ) : summary.length === 0 ? (
<EmptyState text="暂无评估数据,请调整筛选条件或先完成预测与结算" /> <EmptyText text="暂无评估数据,请调整筛选条件或先完成预测与结算" />
) : ( ) : (
<div className="overflow-x-auto"> <div className="overflow-x-auto">
<DataTable <DataTable
+1 -1
View File
@@ -50,7 +50,7 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]:
for s in schedules: for s in schedules:
leagues = s.leagues.split(",") if s.leagues else list(BZZOIRO_LEAGUE_IDS.keys()) leagues = s.leagues.split(",") if s.leagues else list(BZZOIRO_LEAGUE_IDS.keys())
scheduler.register( scheduler.register(
s.id, s.task, s.id, s.cron,
lambda sid=s.id: _run_scheduled_task(sid), lambda sid=s.id: _run_scheduled_task(sid),
enabled=s.enabled, enabled=s.enabled,
) )
+4 -4
View File
@@ -91,8 +91,8 @@ async def create_schedule(req: ScheduleIn, db: AsyncSession = Depends(get_db_rea
db.add(sched) db.add(sched)
await db.commit() await db.commit()
# 注册到调度器 # 注册到调度器(注意: 第 2 参是 cron 表达式,不是 task 类型名)
scheduler.register(req.id, req.task, lambda: _run_scheduled_task(req.id), enabled=req.enabled) scheduler.register(req.id, req.cron, lambda: _run_scheduled_task(req.id), enabled=req.enabled)
return {"ok": True, "id": req.id} 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 sched.enabled = req.enabled
await db.commit() await db.commit()
# 更新调度器(使用最终值) # 更新调度器(第 2 参传 cron 表达式;使用最终值)
scheduler.register(schedule_id, sched.task, lambda: _run_scheduled_task(schedule_id), enabled=sched.enabled) scheduler.register(schedule_id, sched.cron, lambda: _run_scheduled_task(schedule_id), enabled=sched.enabled)
return {"ok": True} return {"ok": True}
+89 -29
View File
@@ -14,6 +14,9 @@ logger = logging.getLogger(__name__)
class ScheduledTask: class ScheduledTask:
"""一个定时任务。""" """一个定时任务。"""
#: 运行循环的轮询粒度(秒)。同时决定运行期 cron/next_run 变更的生效延迟上限。
SLEEP_TICK: float = 1.0
def __init__( def __init__(
self, self,
task_id: str, task_id: str,
@@ -28,13 +31,29 @@ class ScheduledTask:
self.last_run: datetime | None = None self.last_run: datetime | None = None
self.next_run: datetime | None = None self.next_run: datetime | None = None
self._task: asyncio.Task | 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: try:
self.next_run = croniter(self.cron, datetime.now).get_next(datetime) # 必须调用 datetime.now():传方法对象本身会让 croniter 在
except Exception: # (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 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: def update(self, cron: str | None = None, enabled: bool | None = None) -> None:
if cron is not None: if cron is not None:
@@ -44,27 +63,45 @@ class ScheduledTask:
self._calc_next() self._calc_next()
async def _run_loop(self) -> None: 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: 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: try:
logger.info("定时任务触发: %s (cron=%s)", self.task_id, self.cron) if not self.enabled or not self.next_run:
await self.fn() await asyncio.sleep(self.SLEEP_TICK)
logger.info("定时任务完成: %s", self.task_id) 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: except Exception:
logger.exception("定时任务失败: %s", self.task_id) # 循环体自身的意外异常(如 _calc_next)也不能终止调度
logger.exception("调度循环异常: %s", self.task_id)
await asyncio.sleep(self.SLEEP_TICK)
class DataQualityScheduler: class DataQualityScheduler:
@@ -152,6 +189,7 @@ class Scheduler:
def __init__(self) -> None: def __init__(self) -> None:
self._tasks: dict[str, ScheduledTask] = {} self._tasks: dict[str, ScheduledTask] = {}
self._running = False
def register( def register(
self, self,
@@ -160,13 +198,31 @@ class Scheduler:
fn: Callable[[], Coroutine], fn: Callable[[], Coroutine],
enabled: bool = True, enabled: bool = True,
) -> ScheduledTask: ) -> ScheduledTask:
if task_id in self._tasks: """注册(或更新)一个定时任务。
self._tasks[task_id].update(cron=cron, enabled=enabled)
return self._tasks[task_id] 若调度器已 start,新任务会立即启动其运行循环 ——
task = ScheduledTask(task_id, cron, fn, enabled) 否则运行期通过 API 新建的任务永远不会被执行(全量审查 C1 缺陷 2)。
self._tasks[task_id] = task """
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 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: def get(self, task_id: str) -> ScheduledTask | None:
return self._tasks.get(task_id) return self._tasks.get(task_id)
@@ -174,14 +230,18 @@ class Scheduler:
return list(self._tasks.values()) return list(self._tasks.values())
def remove(self, task_id: str) -> None: 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: async def start(self) -> None:
self._running = True
for task in self._tasks.values(): for task in self._tasks.values():
task._task = asyncio.create_task(task._run_loop()) self._ensure_loop(task)
logger.info("定时调度器已启动, 共 %d 个任务", len(self._tasks)) logger.info("定时调度器已启动, 共 %d 个任务", len(self._tasks))
async def stop(self) -> None: async def stop(self) -> None:
self._running = False
for task in self._tasks.values(): for task in self._tasks.values():
if task._task: if task._task:
task._task.cancel() task._task.cancel()
+291 -2
View File
@@ -23,7 +23,7 @@ import httpx
from src.core.runtime_config import get_runtime_value from src.core.runtime_config import get_runtime_value
from src.core.http_client import get_client 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.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.normalize import normalize_bzzoiro
from src.data.team_names_zh import zh_name from src.data.team_names_zh import zh_name
from src.data.sources import register from src.data.sources import register
@@ -101,7 +101,7 @@ async def _fetch_json_async(path: str, params: dict | None = None, max_retries:
# 限流:标记当前 key 冷却,切换到下一个 # 限流:标记当前 key 冷却,切换到下一个
new_key = ring.report_rate_limited(key) new_key = ring.report_rate_limited(key)
if new_key and new_key != 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 key = new_key
continue # 立即重试,不等待 continue # 立即重试,不等待
# 单 key 或全部冷却:等待最早恢复的 key # 单 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}") 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 # 管线基础设施:RawEvent / IngestFailure / DataLineage
# ============================================================ # ============================================================
@@ -226,6 +429,92 @@ async def ingest_bzzoiro_standings(db, *, leagues: Iterable[str], season: str |
continue continue
rows = payload.get("standings") or [] 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 return result
+33 -11
View File
@@ -8,10 +8,13 @@
""" """
from __future__ import annotations from __future__ import annotations
import logging
from typing import Protocol from typing import Protocol
from src.db.base import AsyncSession from src.db.base import AsyncSession
logger = logging.getLogger(__name__)
class DataSource(Protocol): class DataSource(Protocol):
"""比赛数据源契约:抓取 → 规范化 → 入库。""" """比赛数据源契约:抓取 → 规范化 → 入库。"""
@@ -29,6 +32,9 @@ class DataSource(Protocol):
# ── 注册表 ── # ── 注册表 ──
_SOURCES: dict[str, DataSource] = {} _SOURCES: dict[str, DataSource] = {}
# 延迟加载标记:保证 _load_sources() 最多执行一次,重复调用廉价
_loaded = False
def register(source): def register(source):
"""装饰器:将数据源注册到全局注册表。 """装饰器:将数据源注册到全局注册表。
@@ -42,8 +48,7 @@ def register(source):
def get_source(name: str) -> DataSource: def get_source(name: str) -> DataSource:
"""按名获取数据源。""" """按名获取数据源。"""
if not _SOURCES: _load_sources()
_load_sources()
if name not in _SOURCES: if name not in _SOURCES:
raise ValueError(f"未知数据源: {name}") raise ValueError(f"未知数据源: {name}")
return _SOURCES[name] return _SOURCES[name]
@@ -51,18 +56,35 @@ def get_source(name: str) -> DataSource:
def list_sources() -> list[str]: def list_sources() -> list[str]:
"""列出所有已注册数据源名。""" """列出所有已注册数据源名。"""
if not _SOURCES: _load_sources()
_load_sources()
return list(_SOURCES.keys()) return list(_SOURCES.keys())
def _load_sources() -> None: def _load_sources() -> None:
"""延迟导入数据源触发 @register(避免循环导入)。""" """延迟导入数据源触发 @register(避免循环导入)。
from src.data.bzzoiro import BzzoiroSource # noqa: F811
导入失败必须留下痕迹:静默吞掉 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()
_load_sources()
except Exception:
pass
+34 -15
View File
@@ -105,21 +105,24 @@ class MultiPredictResult:
raw: dict | None = None 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 专属 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 地址/密钥: AGENT_BASE_URL_{ID} / AGENT_API_KEY_{ID}(运行时) → 全局 LLM_BASE_URL / LLM_API_KEY
P3-2: 结果缓存 60 秒,避免每次预测都多次查询运行时配置 DB。 P3-2: 结果缓存 60 秒,避免每次预测都多次查询运行时配置 DB。
注意:model_override 生效时跳过缓存读写 —— 否则带 override 的结果会泄漏给
不带 override 的调用(反之亦然),导致跨调用的模型串味。
""" """
cache_key = f"{agent_id}:{tier}" cache_key = f"{agent_id}:{tier}"
cached = _AGENT_PROVIDER_CACHE.get(cache_key) if model_override is None:
if cached is not None: cached = _AGENT_PROVIDER_CACHE.get(cache_key)
ts, provider = cached if cached is not None:
if time.time() - ts < _AGENT_PROVIDER_CACHE_TTL: ts, provider = cached
return provider if time.time() - ts < _AGENT_PROVIDER_CACHE_TTL:
return provider
pfx = f"AGENT_{agent_id.upper()}_" pfx = f"AGENT_{agent_id.upper()}_"
p = await get_default_provider() 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") model = await get_runtime_value(f"{pfx}MODEL")
if model: if model:
p.model = model p.model = model
# 调用方显式传入的 model 优先级最高,高于 agent 级与层级默认
if model_override:
p.model = model_override
base = await get_runtime_value(f"{pfx}BASE_URL") base = await get_runtime_value(f"{pfx}BASE_URL")
if base: if base:
p.base_url = base p.base_url = base
@@ -136,10 +142,11 @@ async def _agent_provider(agent_id: str, *, tier: str) -> LLMProvider:
if key: if key:
p.api_key = key p.api_key = key
_AGENT_PROVIDER_CACHE[cache_key] = (time.time(), p) if model_override is None:
# 简单淘汰:超过 20 条时清空(60s TTL 下不会累积太多) _AGENT_PROVIDER_CACHE[cache_key] = (time.time(), p)
if len(_AGENT_PROVIDER_CACHE) > 20: # 简单淘汰:超过 20 条时清空(60s TTL 下不会累积太多)
_AGENT_PROVIDER_CACHE.clear() if len(_AGENT_PROVIDER_CACHE) > 20:
_AGENT_PROVIDER_CACHE.clear()
return p return p
@@ -148,13 +155,19 @@ async def run_specialists(
*, *,
version: str = "v1", version: str = "v1",
before=None, before=None,
model_override: str | None = None,
) -> list[AgentReport]: ) -> list[AgentReport]:
"""并行执行 5 个专家 agent。fail-open: 单个失败不影响其他。 """并行执行 5 个专家 agent。fail-open: 单个失败不影响其他。
before: 数据截止时间(回测防泄漏)。None 表示不限制。 before: 数据截止时间(回测防泄漏)。None 表示不限制。
model_override: 调用方显式指定的模型,覆盖各 agent 的层级默认。
注意:model_override 必须作为形参下传,不能用模块级变量中转。
backtest 会 asyncio.gather 并发 8 场预测(见 backtest.py 的 Semaphore(8)),
模块级变量会被并发调用互相覆盖,导致 A 场的预测用上 B 场的模型。
""" """
tasks = [ 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 for spec in SPECIALIST_SPECS
] ]
results = await asyncio.gather(*tasks, return_exceptions=True) results = await asyncio.gather(*tasks, return_exceptions=True)
@@ -222,11 +235,13 @@ async def predict_match_multi(
version: str = "v1", version: str = "v1",
backtest: bool = False, backtest: bool = False,
cutoff_at=None, cutoff_at=None,
model: str | None = None,
) -> MultiPredictResult: ) -> MultiPredictResult:
"""多 agent 端到端预测: 切片 → 并行专家 → 终裁 → 存库。 """多 agent 端到端预测: 切片 → 并行专家 → 终裁 → 存库。
backtest: 回测模式。True 时 cutoff 自动设为 match_dt - 1 天。 backtest: 回测模式。True 时 cutoff 自动设为 match_dt - 1 天。
cutoff_at: 显式截止时间(优先于 backtest 自动计算)。 cutoff_at: 显式截止时间(优先于 backtest 自动计算)。
model: 显式指定模型,优先于 agent 级/层级默认配置(single 模式语义一致)。
""" """
start = time.perf_counter() start = time.perf_counter()
@@ -247,7 +262,11 @@ async def predict_match_multi(
prediction_cutoff_at = cutoff prediction_cutoff_at = cutoff
# 2. 并行专家(各自独立配置,使用统一 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 统计有效专家报告数量 # 2.5 统计有效专家报告数量
ok_reports = [r for r in reports if r.status == "ok"] 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_prompt_tokens = 0
agg_completion_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) 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=%s mode=%s status=%s pred=%s:%s (%s) latency=%sms, experts=%d/%d, prediction_id=%s",
match_id, "multi", pred_status, match_id, "multi", pred_status,
pred.pred_home_goals, pred.pred_away_goals, pred.pred_1x2, 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( return MultiPredictResult(
+36 -5
View File
@@ -10,7 +10,7 @@ from __future__ import annotations
import asyncio import asyncio
import logging import logging
from dataclasses import dataclass, field from dataclasses import dataclass, field
from datetime import datetime from datetime import datetime, timezone
from sqlalchemy import select from sqlalchemy import select
from sqlalchemy.orm import selectinload from sqlalchemy.orm import selectinload
@@ -79,6 +79,35 @@ class BacktestSummary:
results: list[BacktestMatchResult] = field(default_factory=list) 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( async def _get_historical_matches(
db, db,
*, *,
@@ -106,10 +135,12 @@ async def _get_historical_matches(
) )
if league_id is not None: if league_id is not None:
stmt = stmt.where(Match.league_id == league_id) stmt = stmt.where(Match.league_id == league_id)
if date_from: dt_from = _parse_date_bound(date_from, end_of_day=False)
stmt = stmt.where(Match.match_date >= date_from) if dt_from is not None:
if date_to: stmt = stmt.where(Match.match_date >= dt_from)
stmt = stmt.where(Match.match_date <= date_to) 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) stmt = stmt.order_by(Match.match_date.desc()).limit(limit)
result = await db.execute(stmt) result = await db.execute(stmt)
+2 -1
View File
@@ -187,13 +187,14 @@ async def predict_match(
) )
from src.llm.agents.orchestrator import predict_match_multi from src.llm.agents.orchestrator import predict_match_multi
# 回测参数完整传递到 multi-agent 路径 # 回测参数 + 模型覆盖完整传递到 multi-agent 路径
return await predict_match_multi( return await predict_match_multi(
match_id, match_id,
provider=provider, provider=provider,
version=(prompt_version or "v1").removeprefix("multi_"), version=(prompt_version or "v1").removeprefix("multi_"),
backtest=backtest, backtest=backtest,
cutoff_at=cutoff_at, cutoff_at=cutoff_at,
model=model,
) )
+9 -7
View File
@@ -7,12 +7,18 @@
""" """
from __future__ import annotations from __future__ import annotations
from pathlib import Path
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
import pytest import pytest
from src.db.models import Prediction 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: class TestAgentWeightsColumn:
"""验证 predictions 表有 agent_weights 列。""" """验证 predictions 表有 agent_weights 列。"""
@@ -38,14 +44,10 @@ class TestMigration:
"""验证迁移文件存在且内容正确。""" """验证迁移文件存在且内容正确。"""
def test_migration_exists(self): def test_migration_exists(self):
import os assert MIGRATION_PATH.is_file(), f"迁移文件不存在: {MIGRATION_PATH}"
path = "/.octop/workspaces/CA7PFH/Profeto/alembic/versions/0014_predictions_agent_weights.py"
assert os.path.exists(path)
def test_migration_content(self): def test_migration_content(self):
path = "/.octop/workspaces/CA7PFH/Profeto/alembic/versions/0014_predictions_agent_weights.py" content = MIGRATION_PATH.read_text(encoding="utf-8")
content = open(path).read()
assert "agent_weights" in content assert "agent_weights" in content
assert "upgrade" in content assert "upgrade" in content
@@ -82,7 +84,7 @@ class TestOrchestratorWritesAgentWeights:
captured_values = {} captured_values = {}
async def mock_specialists(h, *, version, before): async def mock_specialists(h, *, version, before, model_override=None):
return reports return reports
async def mock_provider(aid, **kw): async def mock_provider(aid, **kw):
+6 -2
View File
@@ -196,7 +196,7 @@ class TestOrchestratorAggregation:
"""终裁输入拼装逻辑。""" """终裁输入拼装逻辑。"""
def test_reports_to_json(self): 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 import json
reports = [ reports = [
@@ -206,8 +206,12 @@ class TestOrchestratorAggregation:
text = _reports_to_json(reports) text = _reports_to_json(reports)
data = json.loads(text) data = json.loads(text)
assert len(data) == 2 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" assert data[1]["status"] == "no_data"
# 两个 agent 都应被映射,不留英文原键
assert data[1]["agent"] == AGENT_LABELS_ZH["standings"]
def test_aggregator_prompt_renders(self): def test_aggregator_prompt_renders(self):
"""终裁 prompt 模板两占位符都能渲染。""" """终裁 prompt 模板两占位符都能渲染。"""
+144
View File
@@ -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 吞掉导入异常"
+8 -6
View File
@@ -5,7 +5,7 @@
""" """
from __future__ import annotations from __future__ import annotations
from unittest.mock import MagicMock from unittest.mock import MagicMock, patch
import pytest import pytest
@@ -101,14 +101,16 @@ class TestH2HCurrentHomePerspective:
home_name="曼城", away_name="诺维奇"), home_name="曼城", away_name="诺维奇"),
] ]
import src.llm.context_builder as cb import src.llm.context_builder as cb
orig = cb._get_h2h
cb._get_h2h = lambda db, h, a, before, **kw: matches async def mock_get_h2h(db, h, a, before, **kw):
try: # 真实契约是 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) result = await h2h_slice(header, limit=8, before=None)
text = str(result) text = str(result)
assert "2胜 0平 0负" in text, f"期望「2胜 0平 0负」,实际:\n{text}" assert "2胜 0平 0负" in text, f"期望「2胜 0平 0负」,实际:\n{text}"
finally:
cb._get_h2h = orig
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_draw_counted_correctly(self): async def test_draw_counted_correctly(self):
+58 -62
View File
@@ -35,40 +35,49 @@ class TestMultiAgentCutoffPropagation:
async def test_backtest_computes_cutoff_from_match_dt_minus_1_day(self): async def test_backtest_computes_cutoff_from_match_dt_minus_1_day(self):
"""backtest=True → cutoff = match_dt - 1 天,传给所有切片。""" """backtest=True → cutoff = match_dt - 1 天,传给所有切片。"""
from datetime import datetime, timedelta, timezone from datetime import datetime, timedelta, timezone
from unittest.mock import patch
import src.llm.agents.orchestrator as orch import src.llm.agents.orchestrator as orch
match_dt = datetime(2026, 1, 15, 20, 0, tzinfo=timezone.utc) match_dt = datetime(2026, 1, 15, 20, 0, tzinfo=timezone.utc)
header = _make_header(match_dt) header = _make_header(match_dt)
captured_before = [] 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) captured_before.append(before)
return [] return []
orch.run_specialists = mock_run_specialists async def mock_header(mid, db=None):
orch.load_match_header = lambda mid, db=None: header # load_match_header 是 async,必须是 async 函数;
orch._agent_provider = lambda agent_id, **kw: MagicMock(model="test") # 且必须走 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: try:
await orch.predict_match_multi(999, backtest=True) await orch.predict_match_multi(999, backtest=True)
except Exception: except Exception:
pass # 后续 aggregator 调用会因 mock 不全而失败,不影响 cutoff 测试 pass # 后续 aggregator 调用会因 mock 不全而失败,不影响 cutoff 测试
assert len(captured_before) == 1 assert len(captured_before) == 1
expected_cutoff = match_dt - timedelta(days=1) expected_cutoff = match_dt - timedelta(days=1)
assert captured_before[0] == expected_cutoff, ( assert captured_before[0] == expected_cutoff, (
f"backtest cutoff 应为 {expected_cutoff},实际 {captured_before[0]}" f"backtest cutoff 应为 {expected_cutoff},实际 {captured_before[0]}"
) )
finally:
orch.run_specialists = orig_run_specialists
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_explicit_cutoff_at_overrides_backtest(self): async def test_explicit_cutoff_at_overrides_backtest(self):
"""显式 cutoff_at 优先于 backtest 自动计算。""" """显式 cutoff_at 优先于 backtest 自动计算。"""
from datetime import datetime, timezone from datetime import datetime, timezone
from unittest.mock import patch
import src.llm.agents.orchestrator as orch import src.llm.agents.orchestrator as orch
match_dt = datetime(2026, 1, 15, 20, 0, tzinfo=timezone.utc) match_dt = datetime(2026, 1, 15, 20, 0, tzinfo=timezone.utc)
@@ -76,97 +85,89 @@ class TestMultiAgentCutoffPropagation:
header = _make_header(match_dt) header = _make_header(match_dt)
captured_before = [] 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) captured_before.append(before)
return [] return []
orch.run_specialists = mock_run_specialists async def mock_header(mid, db=None):
orch.load_match_header = lambda mid, db=None: header return header
orch._agent_provider = lambda agent_id, **kw: MagicMock(model="test")
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: try:
await orch.predict_match_multi(999, backtest=True, cutoff_at=explicit_cutoff) await orch.predict_match_multi(999, backtest=True, cutoff_at=explicit_cutoff)
except Exception: except Exception:
pass pass
assert captured_before[0] == explicit_cutoff assert captured_before[0] == explicit_cutoff
finally:
orch.run_specialists = orig_run_specialists
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_normal_mode_cutoff_is_match_dt(self): async def test_normal_mode_cutoff_is_match_dt(self):
"""非回测模式,无显式 cutoff → cutoff = match_dt。""" """非回测模式,无显式 cutoff → cutoff = match_dt。"""
from datetime import datetime, timezone from datetime import datetime, timezone
from unittest.mock import patch
import src.llm.agents.orchestrator as orch import src.llm.agents.orchestrator as orch
match_dt = datetime(2026, 1, 15, 20, 0, tzinfo=timezone.utc) match_dt = datetime(2026, 1, 15, 20, 0, tzinfo=timezone.utc)
header = _make_header(match_dt) header = _make_header(match_dt)
captured_before = [] 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) captured_before.append(before)
return [] return []
orch.run_specialists = mock_run_specialists async def mock_header(mid, db=None):
orch.load_match_header = lambda mid, db=None: header return header
orch._agent_provider = lambda agent_id, **kw: MagicMock(model="test")
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: try:
await orch.predict_match_multi(999, backtest=False) await orch.predict_match_multi(999, backtest=False)
except Exception: except Exception:
pass pass
assert captured_before[0] == match_dt assert captured_before[0] == match_dt
finally:
orch.run_specialists = orig_run_specialists
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_prediction_cutoff_at_stored_not_match_dt(self): async def test_prediction_cutoff_at_stored_not_match_dt(self):
"""Prediction 写入时 prediction_cutoff_at = 真正 cutoff,非 match_dt。""" """Prediction 写入时 prediction_cutoff_at = 真正 cutoff,非 match_dt。"""
from datetime import datetime, timedelta, timezone 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) match_dt = datetime(2026, 1, 15, 20, 0, tzinfo=timezone.utc)
expected_cutoff = match_dt - timedelta(days=1) 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: class FakeContext:
text = "fake" text = "fake"
match_dt = match_dt
cutoff = expected_cutoff 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() return FakeContext()
pred.build_context = fake_build with patch.object(pred, "build_context", tracking_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
try: try:
await _predict_single(999, backtest=True) await pred._predict_single(999, backtest=True)
except Exception: except Exception:
pass pass
assert call_args.get("backtest") is True, "backtest=True 应传递给 build_context" # 重点: build_context 被调用时传入 backtest=True 和正确的 cutoff
finally: assert call_args.get("backtest") is True, "backtest=True 应传递给 build_context"
pred.build_context = orig_build
class TestBacktestXgNotVisible: class TestBacktestXgNotVisible:
@@ -175,8 +176,8 @@ class TestBacktestXgNotVisible:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_stats_slice_respects_cutoff_for_xg_availability(self): async def test_stats_slice_respects_cutoff_for_xg_availability(self):
"""available_at > cutoff 的 xG 数据不应被切片使用。""" """available_at > cutoff 的 xG 数据不应被切片使用。"""
from datetime import datetime, timedelta, timezone from datetime import datetime, timezone
from unittest.mock import MagicMock from unittest.mock import MagicMock, patch
cutoff = datetime(2026, 1, 13, 20, 0, tzinfo=timezone.utc) # match_date - 2天 cutoff = datetime(2026, 1, 13, 20, 0, tzinfo=timezone.utc) # match_date - 2天
match_dt = datetime(2026, 1, 15, 20, 0, tzinfo=timezone.utc) match_dt = datetime(2026, 1, 15, 20, 0, tzinfo=timezone.utc)
@@ -206,7 +207,6 @@ class TestBacktestXgNotVisible:
header = _make_header(match_dt) header = _make_header(match_dt)
import src.llm.context_builder as cb import src.llm.context_builder as cb
orig_get_form = cb._get_form
async def mock_get_form(db, team_id, before, *, limit=10): async def mock_get_form(db, team_id, before, *, limit=10):
# before=cutoff(1月13日),比赛在1月15日,满足 before 条件 # before=cutoff(1月13日),比赛在1月15日,满足 before 条件
@@ -214,14 +214,10 @@ class TestBacktestXgNotVisible:
return [hist_match] return [hist_match]
return [] return []
cb._get_form = mock_get_form with patch.object(cb, "_get_form", mock_get_form):
try:
result = await cb.stats_slice(header, limit=10, before=cutoff) result = await cb.stats_slice(header, limit=10, before=cutoff)
text = str(result) text = str(result)
# xG 在 cutoff 之后才 available,不应出现在切片 # xG 在 cutoff 之后才 available,不应出现在切片
assert "2.50" not in text, f"xG 2.50 不应在切片中(available_at > cutoff):\n{text}" assert "2.50" not in text, f"xG 2.50 不应在切片中(available_at > cutoff):\n{text}"
# 但无比分时仍应显示进球数据 # 但无比分时仍应显示进球数据
assert "无比分数据" in text or "场均进球" in text, f"无比分时仍应显示基本数据:\n{text}" assert "无比分数据" in text or "场均进球" in text, f"无比分时仍应显示基本数据:\n{text}"
finally:
cb._get_form = orig_get_form
+11 -5
View File
@@ -66,7 +66,7 @@ class TestAllExpertsFailed:
header = _make_header() header = _make_header()
# Mock run_specialists 返回全 error # 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() return _all_error_reports()
# Mock _agent_provider # Mock _agent_provider
@@ -131,7 +131,7 @@ class TestAllExpertsFailed:
"""5 个专家全 no_data → status=degraded,不调终裁。""" """5 个专家全 no_data → status=degraded,不调终裁。"""
header = _make_header() 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() return _all_no_data_reports()
async def mock_agent_provider(agent_id, *, tier): async def mock_agent_provider(agent_id, *, tier):
@@ -188,7 +188,7 @@ class TestPartialExpertsOk:
"""1 个 ok + 4 个 error → status=success(走终裁)。""" """1 个 ok + 4 个 error → status=success(走终裁)。"""
header = _make_header() 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() return _mixed_reports()
async def mock_agent_provider(agent_id, *, tier): async def mock_agent_provider(agent_id, *, tier):
@@ -255,7 +255,7 @@ class TestNoAggregatorCallOnDegraded:
header = _make_header() header = _make_header()
aggregator_called = [] 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() return _all_error_reports()
async def mock_agent_provider(agent_id, *, tier): async def mock_agent_provider(agent_id, *, tier):
@@ -268,11 +268,17 @@ class TestNoAggregatorCallOnDegraded:
captured_values = {} captured_values = {}
async def mock_upsert(session, **kw): 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(kw.get("values", {}))
captured_values.update(
{k: kw.get(k) for k in ("model", "provider_name", "mode", "run_type")}
)
mock_pred = MagicMock() mock_pred = MagicMock()
mock_pred.id = 1 mock_pred.id = 1
mock_pred.provider = "test" mock_pred.provider = "test"
mock_pred.model = kw["values"].get("model") mock_pred.model = kw.get("model")
return mock_pred return mock_pred
class FakeUow: class FakeUow:
+8 -6
View File
@@ -8,12 +8,18 @@
from __future__ import annotations from __future__ import annotations
import inspect import inspect
from pathlib import Path
from pydantic import BaseModel from pydantic import BaseModel
import pytest import pytest
from src.db.models import Prediction, UniqueConstraint, CheckConstraint 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: class TestUniqueConstraint:
"""验证唯一约束包含 mode + run_type。""" """验证唯一约束包含 mode + run_type。"""
@@ -71,14 +77,10 @@ class TestMigration:
"""验证迁移文件存在且内容正确。""" """验证迁移文件存在且内容正确。"""
def test_migration_exists(self): def test_migration_exists(self):
import os assert MIGRATION_PATH.is_file(), f"迁移文件不存在: {MIGRATION_PATH}"
path = "/.octop/workspaces/CA7PFH/Profeto/alembic/versions/0013_predictions_unique_constraint_mode_run_type.py"
assert os.path.exists(path)
def test_migration_adds_column_and_constraint(self): def test_migration_adds_column_and_constraint(self):
path = "/.octop/workspaces/CA7PFH/Profeto/alembic/versions/0013_predictions_unique_constraint_mode_run_type.py" content = MIGRATION_PATH.read_text(encoding="utf-8")
content = open(path).read()
assert 'run_type' in content assert 'run_type' in content
assert 'uq_predictions_match_provider_model_mode_run_type' in content assert 'uq_predictions_match_provider_model_mode_run_type' in content
+155 -22
View File
@@ -14,8 +14,10 @@ from pathlib import Path
SRC = Path(__file__).resolve().parent.parent / "src" SRC = Path(__file__).resolve().parent.parent / "src"
# 切片函数会读取的关系属性 → 查询时必须 eager-load # 切片函数会读取的关系属性 → 查询时必须 eager-load
MATCH_RELATIONS = ("stats", "home_team", "away_team", "league") # 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: def _read(rel: str) -> str:
@@ -48,27 +50,130 @@ class TestEagerLoadCoverage:
assert "selectinload" in src, "backtest 未 eager-load 关系 (P0-1)" assert "selectinload" in src, "backtest 未 eager-load 关系 (P0-1)"
def test_relationship_default_is_selectin(self): 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") src = _read("db/models.py")
# 找到 Match 类定义段 # 找到 Match 类定义段
m = re.search(r"class Match\(Base\):.*?(?=\nclass )", src, re.S) m = re.search(r"class Match\(Base\):.*?(?=\nclass )", src, re.S)
assert m, "Match 类未找到" assert m, "Match 类未找到"
body = m.group(0) body = m.group(0)
for rel in MATCH_RELATIONS: for rel in MATCH_RELATIONS:
# 关系声明可能跨多行(stats/home_team/away_team 都是),因此按 # 只取 relationship(...) 调用本身的括号内内容。
# 「从 `rel: Mapped` 到下一个 `xxx: Mapped` 之前」整段匹配。 # 注意: 不能把整段(含注释)做子串匹配 —— 关系声明下方的注释里
# 恰好也写着 lazy="selectin",会导致「删掉真实 kwarg 但测试仍绿」
# 的假阳性(已用变异测试证实: 移除 league 的 lazy 后断言依旧通过)。
m_rel = re.search( 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 m_rel, f"Match.{rel} 未找到 relationship(...) 声明"
assert 'lazy="selectin"' in m_rel.group(0), ( call_args = m_rel.group(1)
f"Match.{rel} 未声明 lazy='selectin' —— 兜底缺失 (P0-2)" 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: class TestBzzoiroLineage:
"""P0-3: source_event_id 必须取配对的 raw,不能是循环残留变量。""" """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): def test_normalized_matches_carries_raw(self):
src = _read("data/bzzoiro.py") src = _read("data/bzzoiro.py")
# 规范化结果必须与原始 event 成对保存 # 规范化结果必须与原始 event 成对保存
@@ -81,19 +186,47 @@ class TestBzzoiroLineage:
) )
def test_no_orphan_raw_use(self): 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") src = _read("data/bzzoiro.py")
lines = src.splitlines() seg = self._consume_loop_body(src)
# 找到 "for nm, raw in normalized_matches" 所在行号 bad = self._bad_assignments(seg)
start = next( assert len(bad) == 0, (
(i for i, ln in enumerate(lines) if "for nm, raw in normalized_matches" in ln), f"source_event_id 未使用配对的 raw (P0-3),问题行: {bad}"
None,
) )
assert start is not None
# 该循环之后、下一个同/更低缩进的顶层语句之前的范围 def test_loop_body_scope_excludes_downstream_stats_pipeline(self):
seg = "\n".join(lines[start:]) """作用域守卫: 截取段不能扫到循环之后的下游 stats 管线。
uses = [ln for ln in seg.splitlines() if "source_event_id" in ln]
assert uses, "未找到 source_event_id 赋值" 下游 `_backfill_stats` 里合法地在 ORM 对象上访问 `m.source_event_id`
assert all("raw.get(" in ln for ln in uses), ( (与配对 raw 无关)。若 seg 越界,test_no_orphan_raw_use 会误报。
"source_event_id 未使用配对的 raw (P0-3)" """
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 回归"
+506
View File
@@ -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)
+190
View File
@@ -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)"
)