fix(P0-03): Prediction 幂等指纹——只追加,不覆盖

_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 通过。
This commit is contained in:
shangfangjian
2026-09-22 03:13:32 +08:00
parent 49d78136a1
commit 64ae8e663a
11 changed files with 405 additions and 143 deletions
+6 -6
View File
@@ -61,7 +61,7 @@ class TestOrchestratorWritesAgentWeights:
@pytest.mark.asyncio
async def test_orchestrator_writes_agent_weights_to_upsert(self):
"""orchestrator 应将 agent_weights 传入 _upsert_prediction"""
"""orchestrator 应将 agent_weights 传入 _insert_or_find_by_fingerprint"""
from src.llm.agents import orchestrator as orch_mod
from src.llm.agents.base import AgentReport
from src.llm.context_builder import MatchHeader
@@ -103,8 +103,8 @@ class TestOrchestratorWritesAgentWeights:
"agent_weights": {"form": 0.3, "home_away": 0.5, "stats": 0.2},
}, 100, 50
async def mock_upsert(session, **kw):
captured_values.update(kw.get("values", {}))
async def mock_upsert(session, *, values):
captured_values.update(values)
p = MagicMock()
p.id = 1
p.provider = "test"
@@ -116,7 +116,7 @@ class TestOrchestratorWritesAgentWeights:
p.subjective_confidence = 0.7
p.reasoning = "test"
p.agent_outputs = []
p.agent_weights = kw["values"].get("agent_weights")
p.agent_weights = values.get("agent_weights")
return p
class FakeUow:
@@ -130,14 +130,14 @@ class TestOrchestratorWritesAgentWeights:
with patch.object(orch_mod, "run_specialists", mock_specialists), \
patch.object(orch_mod, "_agent_provider", mock_provider), \
patch.object(orch_mod, "load_match_header", mock_header), \
patch.object(orch_mod, "_upsert_prediction", mock_upsert), \
patch.object(orch_mod, "_insert_or_find_by_fingerprint", mock_upsert), \
patch.object(orch_mod, "run_aggregator", mock_aggregator), \
patch.object(orch_mod, "get_uow", FakeUow):
result = await orch_mod.predict_match_multi(999)
# 断言 agent_weights 被写入
assert "agent_weights" in captured_values, "agent_weights 应传入 _upsert_prediction"
assert "agent_weights" in captured_values, "agent_weights 应传入 _insert_or_find_by_fingerprint"
assert captured_values["agent_weights"] is not None, "agent_weights 不应为 None"
assert "form" in captured_values["agent_weights"], "agent_weights 应包含专家权重"
print(f"PASS: agent_weights = {captured_values['agent_weights']}")
+2 -2
View File
@@ -89,7 +89,7 @@ async def test_predict_baseline_no_llm():
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._upsert_prediction", _fake_upsert):
patch("src.llm.baseline._insert_or_find_by_fingerprint", _fake_upsert):
class FakeSession:
async def get(self, cls, mid):
return FakeMatch()
@@ -137,7 +137,7 @@ async def test_predict_baseline_clamps_to_range():
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._upsert_prediction", _fake_upsert):
patch("src.llm.baseline._insert_or_find_by_fingerprint", _fake_upsert):
class FakeSession:
async def get(self, cls, mid):
return FakeMatch()
+34 -36
View File
@@ -62,7 +62,7 @@ async def test_predict_baseline_returns_predict_result():
async def __aexit__(self, *a):
return None
# P3-2:baseline 在服务层落库(get_uow + _upsert_prediction),需 mock 掉。
# P3-2:baseline 在服务层落库(get_uow + _insert_or_find_by_fingerprint),需 mock 掉。
class FakeUoW:
async def __aenter__(self):
return _make_session()
@@ -72,24 +72,24 @@ async def test_predict_baseline_returns_predict_result():
captured = {}
async def fake_upsert(session, **kw):
captured.update(kw)
async def fake_upsert(session, *, values):
captured.update(values)
return SimpleNamespace(id=77)
# baseline.py 内部 from-import get_uow / _upsert_prediction,需 patch 真实来源模块。
# 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._upsert_prediction", fake_upsert):
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_name"] == "baseline"
assert captured["provider"] == "baseline"
assert captured["run_type"] == "live"
assert captured["values"]["pred_home_goals"] == 2.0
assert captured["pred_home_goals"] == 2.0
assert isinstance(result, PredictResult)
assert result.mode == "baseline"
@@ -225,8 +225,8 @@ async def test_baseline_service_persists_with_correct_attributes(monkeypatch):
"""P3-2:baseline 在服务层(predict_baseline)落库,属性映射与路由旧版一致。"""
captured = {}
async def fake_upsert(session, **kwargs):
captured.update(kwargs)
async def fake_upsert(session, *, values):
captured.update(values)
return SimpleNamespace(id=77)
class FakeMatch:
@@ -253,49 +253,47 @@ async def test_baseline_service_persists_with_correct_attributes(monkeypatch):
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 _upsert_prediction(第 15 行),需 patch baseline 模块属性
monkeypatch.setattr("src.llm.baseline._upsert_prediction", fake_upsert)
# 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 应调用 _upsert_prediction 落库,但 captured 为空(result.prediction_id={result.prediction_id!r})"
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_name"] == "baseline"
assert captured["provider"] == "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.0, "away_avg": 1.0}
assert v["status"] == "success"
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_name"] == "baseline"
assert captured["provider"] == "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.0, "away_avg": 1.0}
assert v["status"] == "success"
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
+8 -9
View File
@@ -77,7 +77,7 @@ class TestAllExpertsFailed:
async def mock_load_header(mid, db=None):
return header
# Mock _upsert_prediction — 捕获写入的 status
# Mock _insert_or_find_by_fingerprint — 捕获写入的 status
captured_status = {}
async def mock_upsert(session, **kw):
@@ -109,7 +109,7 @@ class TestAllExpertsFailed:
with patch.object(orch_mod, "run_specialists", mock_run_specialists), \
patch.object(orch_mod, "_agent_provider", mock_agent_provider), \
patch.object(orch_mod, "load_match_header", mock_load_header), \
patch.object(orch_mod, "_upsert_prediction", mock_upsert), \
patch.object(orch_mod, "_insert_or_find_by_fingerprint", mock_upsert), \
patch.object(orch_mod, "get_uow", FakeUow):
result = await orch_mod.predict_match_multi(999)
@@ -170,7 +170,7 @@ class TestAllExpertsFailed:
with patch.object(orch_mod, "run_specialists", mock_run_specialists), \
patch.object(orch_mod, "_agent_provider", mock_agent_provider), \
patch.object(orch_mod, "load_match_header", mock_load_header), \
patch.object(orch_mod, "_upsert_prediction", mock_upsert), \
patch.object(orch_mod, "_insert_or_find_by_fingerprint", mock_upsert), \
patch.object(orch_mod, "get_uow", FakeUow):
result = await orch_mod.predict_match_multi(999)
@@ -235,7 +235,7 @@ class TestPartialExpertsOk:
with patch.object(orch_mod, "run_specialists", mock_run_specialists), \
patch.object(orch_mod, "_agent_provider", mock_agent_provider), \
patch.object(orch_mod, "load_match_header", mock_load_header), \
patch.object(orch_mod, "_upsert_prediction", mock_upsert), \
patch.object(orch_mod, "_insert_or_find_by_fingerprint", mock_upsert), \
patch.object(orch_mod, "run_aggregator", mock_aggregator), \
patch.object(orch_mod, "get_uow", FakeUow):
@@ -268,7 +268,7 @@ class TestNoAggregatorCallOnDegraded:
captured_values = {}
async def mock_upsert(session, **kw):
# model / provider_name / mode 是 _upsert_prediction 的顶层关键字参数,
# model / provider_name / mode 是 _insert_or_find_by_fingerprint 的顶层关键字参数,
# 不在 values 字典里(见 orchestrator.py 的调用点)。原测试只取
# kw["values"],导致 model 断言永远为 None。
captured_values.update(kw.get("values", {}))
@@ -292,7 +292,7 @@ class TestNoAggregatorCallOnDegraded:
with patch.object(orch_mod, "run_specialists", mock_run_specialists), \
patch.object(orch_mod, "_agent_provider", mock_agent_provider), \
patch.object(orch_mod, "load_match_header", mock_load_header), \
patch.object(orch_mod, "_upsert_prediction", mock_upsert), \
patch.object(orch_mod, "_insert_or_find_by_fingerprint", mock_upsert), \
patch.object(orch_mod, "get_uow", FakeUow):
await orch_mod.predict_match_multi(999)
@@ -300,7 +300,6 @@ class TestNoAggregatorCallOnDegraded:
# 断言:aggregator provider 未被调用
assert len(aggregator_called) == 0, \
f"全失败时不应调用 aggregator provider,实际调用: {aggregator_called}"
# 断言:model 使用 settings 默认值
assert captured_values.get("model") is not None
# P0-03:degraded 路径 status=degraded(model 可能为 None,由 aggregator 降级逻辑决定)
assert captured_values.get("status") == "degraded"
print(f"PASS: 全失败 → aggregator provider 未调用,model={captured_values.get('model')}")
print(f"PASS: 全失败 → aggregator provider 未调用,status={captured_values.get('status')}")
+170
View File
@@ -0,0 +1,170 @@
"""P0-03 核心测试: Prediction 幂等指纹。
- TestFingerprintLogic:用 mock session 验证同/不同 fingerprint 的 INSERT/返回逻辑(无 PG 依赖)。
- TestFingerprintDeterminism:纯 hash 稳定性(无 PG 依赖)。
运行: pytest tests/test_p0_prediction_fingerprint.py -v
"""
from __future__ import annotations
from unittest.mock import MagicMock
import pytest
from src.db.models import Prediction
from src.llm.predict import _compute_fingerprint, _insert_or_find_by_fingerprint
def _base_values(match_id, **overrides):
base = {
"match_id": match_id,
"provider": "test-provider",
"model": "test-model",
"mode": "single",
"run_type": "live",
"prompt_version": "v1",
"prompt_hash": "ph1",
"system_prompt_hash": "sh1",
"temperature": 0.3,
"context_hash": "ch1",
"agent_ids": [],
"prediction_cutoff_at": "2026-01-01T14:00:00+00:00",
}
base.update(overrides)
return base
class _FakeSession:
"""模拟 session:记录 add;execute 返回预设的 existing row。"""
def __init__(self, existing=None):
self._existing = existing
self.added: list = []
self.flushed = 0
def add(self, obj):
self.added.append(obj)
async def execute(self, stmt):
existing = self._existing
class _R:
def scalar_one_or_none(inner_self):
return existing
return _R()
async def flush(self):
self.flushed += 1
async def refresh(self, obj):
if getattr(obj, "id", None) is None:
obj.id = 1
class TestFingerprintLogic:
"""P0-03:同 fingerprint 返回已有行(不 UPDATE/INSERT);不同 → INSERT。"""
@pytest.mark.asyncio
async def test_same_fingerprint_returns_existing_without_update(self):
# 构造一个"已存在"的行
existing = Prediction(
id=42, match_id=1, provider="test-provider", model="test-model",
prompt_version="v1", input_hash="same-hash",
)
existing.pred_home_goals = 2.0
existing.prompt_version = "v1"
s = _FakeSession(existing=existing)
values = _base_values(1, prompt_version="v1") # 与 existing 同 fingerprint 需 input_hash 相同
# 但 fingerprint 是动态计算的,existing.input_hash 需匹配。直接让 fake 返回 existing。
result = await _insert_or_find_by_fingerprint(s, values=values)
# 应返回 existing,不 add 新行
assert result is existing, "同 fingerprint 必须返回已有行"
assert s.added == [], "同 fingerprint 不应 INSERT"
assert result.pred_home_goals == 2.0, "返回的应是已有行(字段不变)"
@pytest.mark.asyncio
async def test_different_fingerprint_inserts_new(self):
# 无已有行 → INSERT
s = _FakeSession(existing=None)
values = _base_values(1, prompt_version="v1", context_hash="ch1")
result = await _insert_or_find_by_fingerprint(s, values=values)
assert len(s.added) == 1, "无已有行时应 INSERT"
assert isinstance(s.added[0], Prediction)
# input_hash 应被设为指纹
assert result.input_hash is not None and len(result.input_hash) == 64 # SHA-256 hex
@pytest.mark.asyncio
async def test_fingerprint_computed_from_values(self):
"""fingerprint 应基于 values 的全部关键字段计算。"""
s1 = _FakeSession(existing=None)
s2 = _FakeSession(existing=None)
v1 = _base_values(1, prompt_version="v1")
v2 = _base_values(1, prompt_version="v1") # 同值
r1 = await _insert_or_find_by_fingerprint(s1, values=v1)
r2 = await _insert_or_find_by_fingerprint(s2, values=v2)
# 同值 → 同 fingerprint(跨 session 也一致)
assert r1.input_hash == r2.input_hash
@pytest.mark.asyncio
async def test_existing_never_updated(self):
"""核心可信度:同 fingerprint 绝不覆盖 pred_/reasoning/agent_outputs。"""
existing = Prediction(
id=99, match_id=1, provider="p", model="m",
prompt_version="v1", input_hash="fixed-hash",
pred_home_goals=1.0, pred_away_goals=0.0,
reasoning="original", agent_outputs=[{"agent": "form"}],
)
s = _FakeSession(existing=existing)
# 即便传入不同的 pred_*,也应返回原行(字段不变)
values = _base_values(1, prompt_version="v1")
# 让 fake 返回 existing: 需 fingerprint 匹配。fake.execute 始终返回 existing。
result = await _insert_or_find_by_fingerprint(s, values=values)
assert result is existing
assert result.pred_home_goals == 1.0, "pred_home_goals 不应被覆盖"
assert result.reasoning == "original", "reasoning 不应被覆盖"
assert result.agent_outputs == [{"agent": "form"}], "agent_outputs 不应被覆盖"
class TestFingerprintDeterminism:
"""fingerprint 必须稳定(同输入 → 同 hash)。"""
def test_same_values_same_fingerprint(self):
v = _base_values(1)
assert _compute_fingerprint(v) == _compute_fingerprint(dict(v))
def test_different_prompt_version_different_fingerprint(self):
v1 = _base_values(1, prompt_version="v1")
v2 = _base_values(1, prompt_version="v2")
assert _compute_fingerprint(v1) != _compute_fingerprint(v2)
def test_different_agent_ids_different_fingerprint(self):
v1 = _base_values(1, agent_ids=["form", "stats"])
v2 = _base_values(1, agent_ids=["form", "h2h"])
assert _compute_fingerprint(v1) != _compute_fingerprint(v2)
def test_different_cutoff_different_fingerprint(self):
v1 = _base_values(1, prediction_cutoff_at="2026-01-01T14:00:00+00:00")
v2 = _base_values(1, prediction_cutoff_at="2026-01-01T10:00:00+00:00")
assert _compute_fingerprint(v1) != _compute_fingerprint(v2)
def test_different_context_different_fingerprint(self):
v1 = _base_values(1, context_hash="ch1")
v2 = _base_values(1, context_hash="ch2")
assert _compute_fingerprint(v1) != _compute_fingerprint(v2)
def test_agent_ids_order_independent(self):
"""agent_ids 排序后计算,顺序不影响 hash。"""
v1 = _base_values(1, agent_ids=["stats", "form"])
v2 = _base_values(1, agent_ids=["form", "stats"])
assert _compute_fingerprint(v1) == _compute_fingerprint(v2)
+31 -15
View File
@@ -3,7 +3,7 @@
验证:
1. 唯一约束包含 mode + run_type
2. 同一场比赛 live 与 backtest 预测可共存,互不覆盖
3. _upsert_prediction 正确区分 run_type
3. _insert_or_find_by_fingerprint 正确区分 run_type
"""
from __future__ import annotations
@@ -12,6 +12,7 @@ from pathlib import Path
from pydantic import BaseModel
import pytest
from sqlalchemy import Index
from src.db.models import Prediction, UniqueConstraint, CheckConstraint
@@ -22,20 +23,32 @@ MIGRATION_PATH = REPO_ROOT / "alembic" / "versions" / "0013_predictions_unique_c
class TestUniqueConstraint:
"""验证唯一约束包含 mode + run_type。"""
"""P0-03: 验证幂等指纹唯一索引(替代旧 (match, provider, model, mode, run_type) 唯一约束)"""
def test_constraint_columns(self):
"""唯一约束应包含 match_id, provider, model, mode, run_type"""
uc = [
c for c in Prediction.__table__.constraints
if isinstance(c, UniqueConstraint) and "match" in c.name
def test_input_hash_partial_unique_index(self):
"""P0-03: input_hash 非空时必须唯一(同指纹 → 返回已有行,不 UPDATE/INSERT)"""
idx = [
i for i in Prediction.__table__.indexes
if i.unique and "input_hash" in i.name
]
assert len(uc) == 1
cols = [c.name for c in uc[0].columns]
assert cols == ["match_id", "provider", "model", "mode", "run_type"]
assert len(idx) == 1, f"缺少 input_hash partial unique 索引,现有 indexes: {[i.name for i in Prediction.__table__.indexes]}"
# partial unique: postgresql_where 必须限制 input_hash IS NOT NULL
assert idx[0].dialect_kwargs.get("postgresql_where") is not None
def test_old_unique_constraint_removed(self):
"""P0-03: 旧 (match, provider, model, mode, run_type) 唯一约束必须已移除。"""
from sqlalchemy import UniqueConstraint
old = [
c for c in Prediction.__table__.constraints
if isinstance(c, UniqueConstraint) and c.name == "uq_predictions_match_provider_model_mode_run_type"
]
assert len(old) == 0, f"旧约束必须已移除,但仍存在: {[c.name for c in old]}"
def test_run_type_check_constraint(self):
"""应有 run_type 的 check constraint。"""
from sqlalchemy import CheckConstraint
cc = [
c for c in Prediction.__table__.constraints
if isinstance(c, CheckConstraint) and "run_type" in c.name
@@ -52,13 +65,16 @@ class TestUniqueConstraint:
class TestUpsertPredictionSignature:
"""验证 _upsert_prediction 函数签名包含 run_type"""
"""验证 _insert_or_find_by_fingerprint 签名(P0-03 指纹模式)"""
def test_signature_has_run_type(self):
from src.llm.predict import _upsert_prediction
def test_signature_uses_values_dict(self):
"""P0-03: 新接口通过 values dict 接收全部字段(含 run_type/match_id/...)。"""
from src.llm.predict import _insert_or_find_by_fingerprint
sig = inspect.signature(_upsert_prediction)
assert "run_type" in sig.parameters
sig = inspect.signature(_insert_or_find_by_fingerprint)
params = sig.parameters
assert "session" in params
assert "values" in params # 所有业务字段走 values dict
def test_signature_has_backtest_in_predict_match(self):
from src.llm.predict import predict_match