"""D2 工程债回归测试: 统一预测结果类型。 背景: predict_match 三条路径返回类型不一 —— single 返回 PredictResult, multi 返回字段重复定义的 MultiPredictResult dataclass,baseline 返回裸 dict。 后果: (1) 预测路由 PredictOut 映射被迫写 isinstance(result, dict) 双分支; (2) backtest 对 baseline 模式直接 AttributeError(dict 没有 .prediction_id, 潜伏 bug);(3) 字段清单在两处 dataclass 重复维护,加字段必漏一处。 统一方案: 扩展 PredictResult(可选字段)承载全部模式; MultiPredictResult 变为其别名(保留 orchestrator 签名标记,兼容 re-export); baseline 返回 PredictResult;路由单一字段映射。 本测试守护四件事: 1. predict_baseline 返回 PredictResult(属性访问) 2. MultiPredictResult 与 PredictResult 兼容(orchestrator 构造调用的 全字段 kwargs 可直接构造别名) 3. 预测路由不再有 isinstance(result, dict) 分支(源码守卫,仿 R4 范式) 4. _persist_baseline 用属性访问构造 upsert values(baseline 落库语义不变) """ from __future__ import annotations from pathlib import Path from types import SimpleNamespace import pytest from src.llm.baseline import predict_baseline from src.llm.predict import PredictResult from src.llm.agents import MultiPredictResult ROUTE_PATH = Path(__file__).resolve().parents[1] / "src" / "api" / "routes" / "predict.py" # ============================================================ # 1. baseline 返回 PredictResult # ============================================================ @pytest.mark.asyncio async def test_predict_baseline_returns_predict_result(): """基线预测返回 PredictResult 实例,mode=baseline,token/延迟为 0。""" from unittest.mock import patch async def fake_avg(db, *, team_id, side, league_id, before): return 2.4 if side == "home" else 1.6 class FakeMatch: id = 1 home_team_id = 10 away_team_id = 20 league_id = 1 match_status = "scheduled" class FakeSession: async def get(self, cls, mid): return FakeMatch() class FakeCM: async def __aenter__(self): return FakeSession() async def __aexit__(self, *a): return None # P3-2:baseline 在服务层落库(get_uow + _insert_or_find_by_fingerprint),需 mock 掉。 class FakeUoW: async def __aenter__(self): return _make_session() async def __aexit__(self, *a): return None captured = {} async def fake_upsert(session, *, values): captured.update(values) return SimpleNamespace(id=77) # baseline.py 内部 from-import get_uow / _insert_or_find_by_fingerprint,需 patch 真实来源模块。 with patch("src.llm.baseline._avg_goals", fake_avg), \ patch("src.llm.baseline.AsyncSessionLocal") as SLC, \ patch("src.db.unit_of_work.get_uow", FakeUoW), \ patch("src.llm.baseline._insert_or_find_by_fingerprint", fake_upsert): SLC.return_value = FakeCM() result = await predict_baseline(1) # P3-2:验证服务层落库被调用且属性映射正确 assert captured["match_id"] == 1 assert captured["provider"] == "baseline" assert captured["run_type"] == "live" assert captured["pred_home_goals"] == 2.0 assert isinstance(result, PredictResult) assert result.mode == "baseline" assert result.provider == "baseline" assert result.model == "baseline" assert result.prompt_version == "baseline_v1" assert result.pred_home_goals == 2.0 assert result.pred_away_goals == 2.0 assert result.pred_1x2 == "X" assert result.subjective_confidence == 0.5 assert result.prompt_tokens == 0 assert result.completion_tokens == 0 assert result.latency_ms == 0 assert result.status == "success" assert "非投注建议" in (result.reasoning or "") # baseline 不构建 LLM 上下文,但字段必须存在且可安全序列化 assert result.context == "" # 原始统计快照保留在 raw 中 assert result.raw is not None assert "home_avg" in result.raw # ============================================================ # 2. MultiPredictResult 与扩展后的 PredictResult 兼容 # ============================================================ def test_multi_predict_result_is_predict_result_alias(): """multi 结果不再是重复定义的 dataclass,而是扩展 PredictResult 的别名。""" assert MultiPredictResult is PredictResult def test_multi_result_constructor_kwargs_still_supported(): """orchestrator 现有构造调用的全部字段 kwargs 必须仍可构造(别名完整性)。""" # 与 orchestrator.predict_match_multi 的 return MultiPredictResult(...) 逐一对应 result = MultiPredictResult( prediction_id=1, provider="openai", model="gpt-x", prompt_version="multi_v1", mode="multi", pred_home_goals=2.0, pred_away_goals=1.0, alt_pred_home_goals=None, alt_pred_away_goals=None, pred_1x2="1", subjective_confidence=0.7, reasoning="r", status="success", agent_outputs=[{"agent": "form"}], agent_weights={"form": 0.2}, context="ctx", latency_ms=100, prompt_tokens=10, completion_tokens=5, raw={"final": True}, ) assert result.mode == "multi" assert result.agent_outputs == [{"agent": "form"}] assert result.agent_weights == {"form": 0.2} assert result.prompt_tokens == 10 assert result.completion_tokens == 5 # ============================================================ # 3. 路由去 dict 分支(源码守卫) # ============================================================ def test_predict_route_has_no_dict_branch(): """PredictOut 映射必须统一走属性访问,禁止 isinstance(result, dict) 回潮。""" src = ROUTE_PATH.read_text(encoding="utf-8") assert "isinstance(result, dict)" not in src assert ".get(\"pred_home_goals\")" not in src # ============================================================ # 4. P3-2:baseline 服务层落库属性映射(落库已从路由移到 baseline.py) # ============================================================ class _FakeResult: """支持 .scalar_one_or_none() 的最小假结果集。""" def __init__(self, items): self._items = list(items) def scalars(self): return self def all(self): return self._items def scalar_one_or_none(self): return self._items[0] if self._items else None class _FakeUoW: """替代 get_uow 的最小上下文管理器(session.execute 是 async 的)。""" def __init__(self): self.session = _make_session() async def __aenter__(self): return self.session async def __aexit__(self, *a): return None def __call__(self): return self def _make_session(existing=None): """构造带 async execute / add / flush 的假 session。""" sess = SimpleNamespace() async def execute(*a, **k): return _FakeResult(existing or []) sess.execute = execute sess.add = lambda *a, **k: None async def flush(*a, **k): return None sess.flush = flush return sess @pytest.mark.asyncio async def test_baseline_service_persists_with_correct_attributes(monkeypatch): """P3-2:baseline 在服务层(predict_baseline)落库,属性映射与路由旧版一致。""" captured = {} async def fake_upsert(session, *, values): captured.update(values) return SimpleNamespace(id=77) class FakeMatch: id = 1 home_team_id = 10 away_team_id = 20 league_id = 1 match_status = "scheduled" class FakeSession: async def get(self, cls, mid): return FakeMatch() class FakeSLC: async def __aenter__(self): return FakeSession() async def __aexit__(self, *a): return None async def fake_avg(db, *, team_id, side, league_id, before): return 2.0 if side == "home" else 1.0 monkeypatch.setattr("src.llm.baseline._avg_goals", fake_avg) monkeypatch.setattr("src.llm.baseline.AsyncSessionLocal", FakeSLC) monkeypatch.setattr("src.db.unit_of_work.get_uow", lambda: _FakeUoW()) # baseline.py 模块级 import _insert_or_find_by_fingerprint(第 15 行),需 patch baseline 模块属性 monkeypatch.setattr("src.llm.baseline._insert_or_find_by_fingerprint", fake_upsert) result = await predict_baseline(1) # 落库被调用且属性映射正确 assert captured, f"predict_baseline 应调用 _insert_or_find_by_fingerprint 落库,但 captured 为空(result.prediction_id={result.prediction_id!r})" assert captured["match_id"] == 1 assert captured["provider"] == "baseline" assert captured["model"] == "baseline" assert captured["mode"] == "baseline" assert captured["run_type"] == "live" assert captured["prompt_version"] == "baseline_v1" assert captured["pred_home_goals"] == 2.0 assert captured["pred_away_goals"] == 1.0 assert captured["pred_1x2"] == "1" assert captured["subjective_confidence"] == 0.5 assert captured["prompt_tokens"] == 0 assert captured["completion_tokens"] == 0 assert captured["latency_ms"] == 0 assert captured["raw_response"] == {"home_avg": 2.0, "away_avg": 1.0} assert captured["status"] == "success" # 回填真实 prediction_id(服务层落库后取得) assert result.prediction_id == 77 assert result.pred_1x2 == "1" assert captured["match_id"] == 1 assert captured["provider"] == "baseline" assert captured["model"] == "baseline" assert captured["mode"] == "baseline" assert captured["run_type"] == "live" assert captured["prompt_version"] == "baseline_v1" assert captured["pred_home_goals"] == 2.0 assert captured["pred_away_goals"] == 1.0 assert captured["pred_1x2"] == "1" assert captured["subjective_confidence"] == 0.5 assert captured["prompt_tokens"] == 0 assert captured["completion_tokens"] == 0 assert captured["latency_ms"] == 0 assert captured["raw_response"] == {"home_avg": 2.0, "away_avg": 1.0} assert captured["status"] == "success" # 回填真实 prediction_id(服务层落库后取得) assert result.prediction_id == 77