fix(critical): 修复评审确认的 R1-R5 五处缺陷
R1 429 key 轮换路径调用不存在的 _km → NameError(且位于凭证脱敏日志行): 改用 src/data/key_ring._mask,全树不再有 _km 调用。 [注: 该改动已随另一工作流提交9ccda4b一并入库] R2 ingest_bzzoiro_standings 被截断(return 前无 upsert 逻辑,standings 表 永不写入、total_upserted 恒为 0):从f05dc1a移植完整实现,按 (league_id, season, team_id) upsert,保留逐联赛错误隔离,不加 db.commit()。 [注: 同上,已随9ccda4b入库] R3 src/llm/agents/orchestrator.py 完成日志把 list 喂给 %d → logging TypeError: 改为 len(ok_reports),与同文件降级日志行写法一致。 R4 mode=multi 静默丢弃调用方传入的 model:predict_match_multi 新增 model 关键字参数,经 _ACTIVE_MODEL_OVERRIDE 下传至 _agent_provider, 显式 model 优先级最高;override 生效时跳过 provider 缓存读写以免串味, 并在 predict.py 派发点透传。 R5 回测把字符串日期直接与 timestamptz 列比较:新增 _parse_date_bound 助手, 支持 YYYY-MM-DD / 完整 ISO / datetime / None,裸日期按 UTC 锚定, 结束日取当天末刻(闭区间,避免最后一天被静默排除),非法输入抛 ValueError。 新增 tests/test_review_required_fixes.py 覆盖 R1-R5(R2/R4 为行为测试), 20 项全通过。
This commit is contained in:
@@ -0,0 +1,452 @@
|
||||
"""回归测试: 代码评审确认的 5 个缺陷修复(R1-R5)。
|
||||
|
||||
R1 429 key 轮换路径调用不存在的 _km → NameError(且是凭证脱敏点)
|
||||
R2 ingest_bzzoiro_standings 被截断,永不写 standings 表
|
||||
R3 orchestrator 完成日志把 list 喂给 %d → logging TypeError
|
||||
R4 mode="multi" 静默丢弃调用方传入的 model
|
||||
R5 回测把字符串日期直接与 timestamptz 列比较
|
||||
|
||||
R2 采用行为测试(假 db + monkeypatch 抓取函数),其余为单元/结构断言。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
import logging
|
||||
import pathlib
|
||||
|
||||
import pytest
|
||||
|
||||
from src.data.key_ring import _mask
|
||||
from src.llm import backtest as bt_mod
|
||||
from src.llm.agents import orchestrator as orch_mod
|
||||
|
||||
_REPO_ROOT = pathlib.Path(__file__).resolve().parents[1]
|
||||
_BZZOIRO_SRC = _REPO_ROOT / "src" / "data" / "bzzoiro.py"
|
||||
|
||||
# 在导入期就抓取真实的 _agent_provider。
|
||||
# 原因: tests/test_multi_agent_cutoff.py:52/87/117 会直接
|
||||
# orch._agent_provider = lambda agent_id, **kw: MagicMock(model="test")
|
||||
# 且不做清理(既有测试,本次任务不允许改动),导致模块属性在整套测试跑完后
|
||||
# 被永久替换成同步 lambda。导入期快照可以规避这种跨测试污染。
|
||||
_REAL_AGENT_PROVIDER = orch_mod._agent_provider
|
||||
|
||||
|
||||
def _bzzoiro_source() -> str:
|
||||
return _BZZOIRO_SRC.read_text(encoding="utf-8")
|
||||
|
||||
|
||||
# ============================================================
|
||||
# R1 — _km → _mask
|
||||
# ============================================================
|
||||
|
||||
class TestR1KeyMasking:
|
||||
def test_mask_is_importable_from_bzzoiro(self):
|
||||
"""修复点:bzzoiro 通过 key_ring 复用它,而不是本地重实现。"""
|
||||
import src.data.bzzoiro as bz
|
||||
|
||||
assert bz._mask("abcd1234efgh5678") == "abcd...5678"
|
||||
|
||||
def test_mask_long_key_shows_head_and_tail(self):
|
||||
assert _mask("abcd1234efgh5678") == "abcd...5678"
|
||||
|
||||
def test_mask_short_key_hides_middle(self):
|
||||
assert _mask("short") == "sh***"
|
||||
|
||||
def test_no_km_call_remains_in_bzzoiro_source(self):
|
||||
"""源码守卫: _km 在整个代码库不存在,这行一旦执行必抛 NameError。
|
||||
|
||||
这是 NameError 类缺陷(静态即可判定),且位于凭证脱敏日志行上,
|
||||
因此用源码文本守卫是恰当的,而不是只测运行时路径。
|
||||
"""
|
||||
assert "_km(" not in _bzzoiro_source()
|
||||
|
||||
|
||||
# ============================================================
|
||||
# R2 — 积分榜 upsert(行为测试)
|
||||
# ============================================================
|
||||
|
||||
class _FakeScalars:
|
||||
def __init__(self, items):
|
||||
self._items = list(items)
|
||||
|
||||
def all(self):
|
||||
return list(self._items)
|
||||
|
||||
def __iter__(self):
|
||||
return iter(self._items)
|
||||
|
||||
def scalar_one_or_none(self):
|
||||
return self._items[0] if self._items else None
|
||||
|
||||
|
||||
class _FakeResult:
|
||||
def __init__(self, items):
|
||||
self._items = list(items)
|
||||
|
||||
def scalars(self):
|
||||
return _FakeScalars(self._items)
|
||||
|
||||
def scalar_one_or_none(self):
|
||||
return self._items[0] if self._items else None
|
||||
|
||||
|
||||
class _FakeDb:
|
||||
"""极简假 db: 记录 add/flush,按队列返回预置查询结果。
|
||||
|
||||
只实现 ingest_bzzoiro_standings 真正用到的部分:
|
||||
- execute(...) → 依次弹出 _results 里的结果
|
||||
- add(obj) → 记录
|
||||
- flush() → 给尚无 id 的对象补一个自增 id(模拟 DB 回填主键)
|
||||
"""
|
||||
|
||||
def __init__(self, results=None):
|
||||
self._results = list(results or [])
|
||||
self.added: list = []
|
||||
self.flush_count = 0
|
||||
self._next_id = 1000
|
||||
|
||||
async def execute(self, _stmt):
|
||||
if self._results:
|
||||
return self._results.pop(0)
|
||||
return _FakeResult([])
|
||||
|
||||
def add(self, obj):
|
||||
self.added.append(obj)
|
||||
|
||||
async def flush(self):
|
||||
self.flush_count += 1
|
||||
for obj in self.added:
|
||||
if getattr(obj, "id", None) is None:
|
||||
self._next_id += 1
|
||||
obj.id = self._next_id
|
||||
|
||||
|
||||
async def test_r2_standings_actually_upserts(monkeypatch):
|
||||
"""行为测试: 喂一份积分榜 payload,断言真的构造了 Standing 且计数 > 0。"""
|
||||
import src.data.bzzoiro as bz
|
||||
from src.db.models import League, Standing, Team
|
||||
|
||||
payload = {
|
||||
"season": {"start_date": "2025-08-01", "end_date": "2026-05-31"},
|
||||
"standings": [
|
||||
{
|
||||
"position": 1, "team_name": "Arsenal FC",
|
||||
"played": 10, "won": 8, "drawn": 1, "lost": 1,
|
||||
"gf": 22, "ga": 8, "gd": 14, "pts": 25,
|
||||
"xgf": 18.5, "xga": 9.1, "form": "WWDLW",
|
||||
"zone": {"key": "champions_league", "label": "Champions League"},
|
||||
},
|
||||
{
|
||||
"position": 2, "team_name": "Chelsea FC",
|
||||
"played": 10, "won": 6, "drawn": 2, "lost": 2,
|
||||
"gf": 18, "ga": 12, "gd": 6, "pts": 20,
|
||||
},
|
||||
],
|
||||
}
|
||||
|
||||
async def _fake_fetch(league_code, season=None):
|
||||
return payload
|
||||
|
||||
monkeypatch.setattr(bz, "fetch_bzzoiro_standings", _fake_fetch, raising=True)
|
||||
|
||||
# 查询顺序: League 命中(避免建联赛) → Team 预载(空) → 每行 Standing(未命中)
|
||||
league = League(code="EPL", name="Premier League", country="England")
|
||||
league.id = 42
|
||||
db = _FakeDb(results=[_FakeResult([league]), _FakeResult([])])
|
||||
|
||||
result = await bz.ingest_bzzoiro_standings(db, leagues=["EPL"])
|
||||
|
||||
assert result["errors"] == []
|
||||
assert result["total_upserted"] == 2, result
|
||||
assert result["leagues"]["EPL"]["rows"] == 2
|
||||
assert result["leagues"]["EPL"]["upserted"] == 2
|
||||
# 两支球队都是新建的
|
||||
assert result["leagues"]["EPL"]["teams_created"] == 2
|
||||
|
||||
standings = [o for o in db.added if isinstance(o, Standing)]
|
||||
assert len(standings) == 2, "应真的构造 Standing 行"
|
||||
assert all(isinstance(o, (Standing, Team)) for o in db.added)
|
||||
|
||||
first = standings[0]
|
||||
assert first.league_id == 42
|
||||
assert first.season == "2025-2026" # 8 月起 → 跨年标签
|
||||
assert first.position == 1
|
||||
assert first.points == 25
|
||||
assert first.xg_for == 18.5
|
||||
assert first.zone == "Champions League" # 优先取 label
|
||||
|
||||
|
||||
async def test_r2_standings_upsert_updates_existing(monkeypatch):
|
||||
"""行为测试: 已存在同 (league, season, team) 时应就地更新而非新增。"""
|
||||
import src.data.bzzoiro as bz
|
||||
from src.db.models import League, Standing
|
||||
|
||||
payload = {
|
||||
"season": {"start_date": "2025-08-01", "end_date": "2026-05-31"},
|
||||
"standings": [{"position": 1, "team_name": "Arsenal FC", "pts": 30}],
|
||||
}
|
||||
|
||||
async def _fake_fetch(league_code, season=None):
|
||||
return payload
|
||||
|
||||
monkeypatch.setattr(bz, "fetch_bzzoiro_standings", _fake_fetch, raising=True)
|
||||
|
||||
league = League(code="EPL", name="Premier League", country="England")
|
||||
league.id = 42
|
||||
team = __import__("src.db.models", fromlist=["Team"]).Team(name="Arsenal FC", name_zh="阿森纳")
|
||||
team.id = 7
|
||||
|
||||
existing = Standing(league_id=42, season="2025-2026", team_id=7, position=9)
|
||||
existing.points = 1
|
||||
|
||||
# 查询顺序: League → Team 预载(命中) → Standing 查询(命中已有行)
|
||||
db = _FakeDb(results=[_FakeResult([league]), _FakeResult([team]), _FakeResult([existing])])
|
||||
|
||||
result = await bz.ingest_bzzoiro_standings(db, leagues=["EPL"])
|
||||
|
||||
assert result["total_upserted"] == 1
|
||||
assert existing.points == 30, "已有行应被就地更新"
|
||||
assert result["leagues"]["EPL"]["teams_created"] == 0
|
||||
# 不应新增 Standing(只有 league/team 层面的 add)
|
||||
assert not [o for o in db.added if isinstance(o, Standing)]
|
||||
|
||||
|
||||
def test_r2_source_contains_real_upsert_loop():
|
||||
"""结构断言(行为测试之外的兜底): 确认函数未被截断。"""
|
||||
import src.data.bzzoiro as bz
|
||||
|
||||
src = inspect.getsource(bz.ingest_bzzoiro_standings)
|
||||
assert "total_upserted" in src
|
||||
assert 'result["total_upserted"] +=' in src, "total_upserted 必须真的被累加"
|
||||
assert "Standing(" in src, "必须真的构造 Standing"
|
||||
assert "select(Standing)" in src, "必须查询已有快照以决定 insert/update"
|
||||
|
||||
|
||||
# ============================================================
|
||||
# R3 — logging 参数类型
|
||||
# ============================================================
|
||||
|
||||
def test_r3_completion_log_record_formats_without_raising():
|
||||
"""%d 占位符拿到 list 时 logging 会抛 TypeError;修复后应为计数。"""
|
||||
fmt = (
|
||||
"预测完成 match=%s mode=%s status=%s pred=%s:%s (%s) latency=%sms, "
|
||||
"experts=%d/%d, prediction_id=%s"
|
||||
)
|
||||
ok_reports = ["a", "b", "c"]
|
||||
reports = ["a", "b", "c", "d", "e"]
|
||||
|
||||
record = logging.LogRecord(
|
||||
"src.llm.agents.orchestrator", logging.INFO, __file__, 1, fmt,
|
||||
(999, "multi", "success", 2, 1, "1", 100, len(ok_reports), len(reports), 7),
|
||||
None,
|
||||
)
|
||||
msg = record.getMessage() # 修复前此处抛 TypeError
|
||||
assert "experts=3/5" in msg
|
||||
|
||||
# 反证: 原缺陷写法(直接传 list)确实会炸,确保这条测试真的有鉴别力
|
||||
bad = logging.LogRecord(
|
||||
"x", logging.INFO, __file__, 1, fmt,
|
||||
(999, "multi", "success", 2, 1, "1", 100, ok_reports, len(reports), 7),
|
||||
None,
|
||||
)
|
||||
with pytest.raises(TypeError):
|
||||
bad.getMessage()
|
||||
|
||||
|
||||
# ============================================================
|
||||
# R4 — multi 模式透传 model
|
||||
# ============================================================
|
||||
|
||||
_ORCH_SRC_PATH = _REPO_ROOT / "src" / "llm" / "agents" / "orchestrator.py"
|
||||
|
||||
|
||||
def _orchestrator_source() -> str:
|
||||
"""直接读源码,而不是 inspect.getsource(模块属性)。
|
||||
|
||||
既有测试(如 test_agent_weights_persist.py)会在运行期把
|
||||
orch_mod._agent_provider 换成 lambda/MagicMock,导致 inspect.getsource
|
||||
拿到的是 mock 的定义。本文件的断言针对真实源码,故从磁盘读取。
|
||||
"""
|
||||
return _ORCH_SRC_PATH.read_text(encoding="utf-8")
|
||||
|
||||
|
||||
def test_r4_predict_match_multi_accepts_model():
|
||||
"""predict_match_multi 必须接受 model 且默认 None(向后兼容既有调用)。"""
|
||||
src_fn = _orchestrator_source().split("async def predict_match_multi(", 1)[1]
|
||||
header = src_fn.split(") -> MultiPredictResult:", 1)[0]
|
||||
assert "model: str | None = None" in header, header
|
||||
|
||||
|
||||
def test_r4_agent_provider_accepts_model_override():
|
||||
src_fn = _orchestrator_source().split("async def _agent_provider(", 1)[1]
|
||||
header = src_fn.split(") -> LLMProvider:", 1)[0]
|
||||
assert "model_override: str | None = None" in header, header
|
||||
|
||||
|
||||
async def test_r4_dispatch_forwards_model_to_multi(monkeypatch):
|
||||
"""行为测试: predict_match(mode=multi, model=...) 必须把 model 送到 multi 路径。"""
|
||||
from src.llm import predict as predict_mod
|
||||
|
||||
captured: dict = {}
|
||||
|
||||
async def _fake_multi(match_id, **kwargs):
|
||||
captured["match_id"] = match_id
|
||||
captured.update(kwargs)
|
||||
return "SENTINEL"
|
||||
|
||||
# predict.py:188 是函数内 `from ... import`,import 发生在调用时,
|
||||
# 所以必须打在 orchestrator 模块的属性上。
|
||||
monkeypatch.setattr(orch_mod, "predict_match_multi", _fake_multi, raising=True)
|
||||
|
||||
out = await predict_mod.predict_match(999, model="my-model-x", mode="multi")
|
||||
|
||||
assert out == "SENTINEL"
|
||||
assert captured["model"] == "my-model-x"
|
||||
assert captured["match_id"] == 999
|
||||
|
||||
|
||||
async def test_r4_specialist_provider_honors_model_override(monkeypatch):
|
||||
"""行为测试: model_override 应覆盖 agent 级/层级默认模型。"""
|
||||
from src.llm.provider import LLMProvider
|
||||
import src.llm.agents.orchestrator as real_orch
|
||||
|
||||
async def _fake_default():
|
||||
return LLMProvider(api_key="k", base_url="http://x", model="default-model", timeout=1.0)
|
||||
|
||||
async def _fake_runtime(key):
|
||||
if key.endswith("_MODEL"):
|
||||
return "agent-level-model"
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(real_orch, "get_default_provider", _fake_default, raising=True)
|
||||
monkeypatch.setattr(real_orch, "get_runtime_value", _fake_runtime, raising=True)
|
||||
monkeypatch.setattr(real_orch, "settings", _settings_with_specialist_model(), raising=True)
|
||||
monkeypatch.setattr(real_orch, "_AGENT_PROVIDER_CACHE", {}, raising=True)
|
||||
|
||||
overridden = await _REAL_AGENT_PROVIDER("form", tier="specialist", model_override="OVERRIDE")
|
||||
assert overridden.model == "OVERRIDE"
|
||||
|
||||
# 不带 override 时仍走原优先级(agent 级运行时配置)
|
||||
real_orch._AGENT_PROVIDER_CACHE.clear()
|
||||
normal = await _REAL_AGENT_PROVIDER("form", tier="specialist")
|
||||
assert normal.model == "agent-level-model"
|
||||
|
||||
|
||||
def _settings_with_specialist_model():
|
||||
class _S:
|
||||
LLM_SPECIALIST_MODEL = "tier-specialist-model"
|
||||
LLM_AGGREGATOR_MODEL = "tier-aggregator-model"
|
||||
|
||||
return _S()
|
||||
|
||||
|
||||
def test_r4_override_does_not_pollute_cache():
|
||||
"""override 结果不得写入 60s provider 缓存(否则会串味给普通调用)。"""
|
||||
src = _orchestrator_source().split("async def _agent_provider(", 1)[1]
|
||||
src = src.split("async def run_specialists(", 1)[0]
|
||||
assert "if model_override is None:" in src
|
||||
assert "_AGENT_PROVIDER_CACHE[cache_key]" in src
|
||||
|
||||
|
||||
async def test_r4_dispatch_sets_override_for_specialists(monkeypatch):
|
||||
"""行为测试: predict_match_multi 应把 model 放进 _ACTIVE_MODEL_OVERRIDE,
|
||||
并在 run_specialists 执行期间对 specialist 生效(退出后复位)。
|
||||
|
||||
参照 tests/test_multi_agent_degraded.py 的 stub 方式,避免触碰真实 DB。
|
||||
"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from src.db.unit_of_work import get_uow
|
||||
|
||||
seen: dict = {}
|
||||
|
||||
async def _fake_header(match_id):
|
||||
h = MagicMock()
|
||||
h.match_id = match_id
|
||||
h.match_dt = None
|
||||
return h
|
||||
|
||||
async def _fake_specialists(header, *, version, before):
|
||||
seen["override_during_run"] = orch_mod._ACTIVE_MODEL_OVERRIDE
|
||||
return []
|
||||
|
||||
async def _fake_upsert(session, **kw):
|
||||
p = MagicMock()
|
||||
p.id = 1
|
||||
p.provider = "test"
|
||||
p.model = kw.get("model")
|
||||
p.prompt_version = "v1"
|
||||
p.pred_home_goals = None
|
||||
p.pred_away_goals = None
|
||||
p.pred_1x2 = None
|
||||
p.alt_pred_home_goals = None
|
||||
p.alt_pred_away_goals = None
|
||||
p.subjective_confidence = None
|
||||
p.reasoning = ""
|
||||
p.agent_outputs = []
|
||||
p.agent_weights = {}
|
||||
p.prompt_tokens = 0
|
||||
p.completion_tokens = 0
|
||||
return p
|
||||
|
||||
class _FakeUow:
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, *a):
|
||||
return False
|
||||
|
||||
async def get(self, cls, id):
|
||||
return MagicMock()
|
||||
|
||||
monkeypatch.setattr(orch_mod, "load_match_header", _fake_header, raising=True)
|
||||
monkeypatch.setattr(orch_mod, "run_specialists", _fake_specialists, raising=True)
|
||||
monkeypatch.setattr(orch_mod, "_upsert_prediction", _fake_upsert, raising=True)
|
||||
monkeypatch.setattr(orch_mod, "get_uow", _FakeUow, raising=True)
|
||||
monkeypatch.setattr(orch_mod, "_ACTIVE_MODEL_OVERRIDE", None, raising=False)
|
||||
|
||||
assert get_uow is not None # 确保 import 生效,session 未被真实打开
|
||||
|
||||
await orch_mod.predict_match_multi(999, model="OVERRIDE-X")
|
||||
|
||||
assert seen["override_during_run"] == "OVERRIDE-X"
|
||||
assert orch_mod._ACTIVE_MODEL_OVERRIDE is None, "退出后必须复位"
|
||||
|
||||
|
||||
# ============================================================
|
||||
# R5 — 回测日期解析
|
||||
# ============================================================
|
||||
|
||||
class TestR5DateBound:
|
||||
def test_plain_date_start_is_start_of_day_utc(self):
|
||||
dt = bt_mod._parse_date_bound("2026-01-01", end_of_day=False)
|
||||
assert dt is not None
|
||||
assert dt.tzinfo is not None
|
||||
assert (dt.year, dt.month, dt.day) == (2026, 1, 1)
|
||||
assert (dt.hour, dt.minute, dt.second) == (0, 0, 0)
|
||||
|
||||
def test_plain_date_end_is_inclusive_end_of_day(self):
|
||||
"""闭区间: 结束日必须取当天末刻,否则最后一天被静默排除。"""
|
||||
dt = bt_mod._parse_date_bound("2026-01-01", end_of_day=True)
|
||||
assert dt is not None
|
||||
assert (dt.hour, dt.minute, dt.second) == (23, 59, 59)
|
||||
assert dt.microsecond == 999999
|
||||
|
||||
def test_full_iso_string_is_parsed(self):
|
||||
dt = bt_mod._parse_date_bound("2026-01-01T12:30:00+08:00", end_of_day=False)
|
||||
assert dt is not None
|
||||
assert dt.utcoffset() is not None
|
||||
|
||||
def test_none_returns_none(self):
|
||||
assert bt_mod._parse_date_bound(None, end_of_day=False) is None
|
||||
assert bt_mod._parse_date_bound(None, end_of_day=True) is None
|
||||
|
||||
def test_datetime_passthrough(self):
|
||||
from datetime import datetime, timezone
|
||||
|
||||
given = datetime(2026, 5, 5, 6, 0, tzinfo=timezone.utc)
|
||||
assert bt_mod._parse_date_bound(given, end_of_day=False) == given
|
||||
|
||||
def test_invalid_input_raises_value_error(self):
|
||||
with pytest.raises(ValueError):
|
||||
bt_mod._parse_date_bound("not-a-date", end_of_day=False)
|
||||
Reference in New Issue
Block a user