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 />
</div>
) : summary.length === 0 ? (
<EmptyState text="暂无评估数据,请调整筛选条件或先完成预测与结算" />
<EmptyText text="暂无评估数据,请调整筛选条件或先完成预测与结算" />
) : (
<div className="overflow-x-auto">
<DataTable
+1 -1
View File
@@ -50,7 +50,7 @@ async def lifespan(app: FastAPI) -> 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,
)
+4 -4
View File
@@ -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}
+89 -29
View File
@@ -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()
+291 -2
View File
@@ -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
+33 -11
View File
@@ -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()
+34 -15
View File
@@ -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(
+36 -5
View File
@@ -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)
+2 -1
View File
@@ -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,
)
+9 -7
View File
@@ -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):
+6 -2
View File
@@ -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 模板两占位符都能渲染。"""
+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 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):
+58 -62
View File
@@ -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
+11 -5
View File
@@ -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:
+8 -6
View File
@@ -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
+155 -22
View File
@@ -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 回归"
+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)"
)