_upsert_prediction 改为 _insert_or_find_by_fingerprint: - 同 input_hash → 返回已有行(绝不 UPDATE pred_/reasoning/agent_outputs) - 不同 input_hash → INSERT 新行 input_hash 升级为规范 JSON SHA-256,捕获:match_id, cutoff, prompt_version, prompt_hash, system_prompt_hash, provider, model, mode, run_type, temperature, context_hash, agent_ids。移除旧 (match, provider, model, mode, run_type) 唯一约束, 改为 partial unique index(WHERE input_hash IS NOT NULL,兼容旧 NULL 数据)。 三条路径(single/multi/baseline)统一传足指纹字段。 迁移 0024 + 测试 test_p0_prediction_fingerprint(10/10);全量 295 通过。
300 lines
10 KiB
Python
300 lines
10 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
|
|
|
|
# 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
|