Files
Profeto/tests/test_d2_predict_unification.py
WorkBuddy 5c7fdce0a3 debt(D2): 统一预测结果类型为 PredictResult,路由去 dict 分支
- PredictResult 扩展可选字段 mode/agent_outputs/agent_weights/prompt_tokens/completion_tokens
- MultiPredictResult 变为 PredictResult 别名(保留 R4 守卫标记与 re-export)
- baseline 改返回 PredictResult(修复 backtest 对 baseline AttributeError 的潜伏 bug)
- 预测路由单一属性映射,删除全部 isinstance(result, dict) 分支
- _persist_baseline 属性化,baseline upsert 语义不变(prompt_version/token/latency 同前)
- TDD: 5 新测试 + test_baseline.py 属性化;全量 257 passed
2026-09-21 19:23:50 +08:00

217 lines
7.3 KiB
Python

"""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
with patch("src.llm.baseline._avg_goals", fake_avg), \
patch("src.llm.baseline.AsyncSessionLocal") as SLC:
SLC.return_value = FakeCM()
result = await predict_baseline(1)
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. _persist_baseline 属性映射(baseline 落库语义不变)
# ============================================================
class _FakeUoW:
"""替代 get_uow 的最小上下文管理器。"""
def __init__(self):
self.session = SimpleNamespace()
async def __aenter__(self):
return self.session
async def __aexit__(self, *a):
return None
@pytest.mark.asyncio
async def test_persist_baseline_maps_attributes(monkeypatch):
captured = {}
async def fake_upsert(session, **kwargs):
captured.update(kwargs)
return SimpleNamespace(id=77)
monkeypatch.setattr("src.db.unit_of_work.get_uow", lambda: _FakeUoW())
monkeypatch.setattr("src.llm.predict._upsert_prediction", fake_upsert)
from src.api.routes.predict import _persist_baseline
baseline = PredictResult(
prediction_id=0, # baseline 不在服务层落库,由 _persist_baseline 落库后取得真实 id
provider="baseline",
model="baseline",
prompt_version="baseline_v1",
mode="baseline",
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.5,
reasoning="r",
context="",
status="success",
latency_ms=0,
prompt_tokens=0,
completion_tokens=0,
raw={"home_avg": 2.1, "away_avg": 1.4},
)
pid = await _persist_baseline(1, baseline)
assert pid == 77
assert captured["match_id"] == 1
assert captured["provider_name"] == "baseline"
assert captured["model"] == "baseline"
assert captured["mode"] == "baseline"
assert captured["run_type"] == "live"
v = captured["values"]
assert v["prompt_version"] == "baseline_v1"
assert v["pred_home_goals"] == 2.0
assert v["pred_away_goals"] == 1.0
assert v["pred_1x2"] == "1"
assert v["subjective_confidence"] == 0.5
assert v["prompt_tokens"] == 0
assert v["completion_tokens"] == 0
assert v["latency_ms"] == 0
assert v["raw_response"] == {"home_avg": 2.1, "away_avg": 1.4}
assert v["status"] == "success"