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
@@ -0,0 +1,41 @@
"""P0-03: Prediction 幂等指纹——移除旧唯一约束,改为 partial unique on input_hash
input_hash 非空时唯一(同指纹返回已有行,不 UPDATE/INSERT);
兼容旧数据 NULL input_hash(不强制回填)。
Revision ID: 0024_prediction_idempotent_fingerprint
Revises: 0023_standings_append_only
Create Date: 2026-09-22
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
revision: str = '0024_prediction_idempotent_fingerprint'
down_revision: Union[str, None] = '0023_standings_append_only'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
# 移除旧唯一约束(match, provider, model, mode, run_type)
op.drop_constraint(
'uq_predictions_match_provider_model_mode_run_type',
'predictions', type_unique=True,
)
# P0-03: partial unique on input_hash(非空时唯一)
op.create_index(
'ix_predictions_input_hash_unique', 'predictions', ['input_hash'], unique=True,
postgresql_where=sa.text('input_hash IS NOT NULL'),
)
def downgrade() -> None:
op.drop_index('ix_predictions_input_hash_unique', table_name='predictions')
op.create_unique_constraint(
'uq_predictions_match_provider_model_mode_run_type',
'predictions',
['match_id', 'provider', 'model', 'mode', 'run_type'],
)
+6 -6
View File
@@ -19,6 +19,7 @@ from sqlalchemy import (
String, String,
Text, Text,
UniqueConstraint, UniqueConstraint,
column,
func, func,
) )
from sqlalchemy.dialects.postgresql import JSONB from sqlalchemy.dialects.postgresql import JSONB
@@ -289,14 +290,13 @@ class Prediction(Base):
match: Mapped[Match] = relationship(back_populates="predictions") match: Mapped[Match] = relationship(back_populates="predictions")
__table_args__ = ( __table_args__ = (
# Fix: 唯一约束增加 mode + run_type,允许 live 与 backtest 共存 # P0-03: 幂等指纹——input_hash 非空时唯一(同指纹→返回已有行,不 UPDATE/INSERT);
# 防止回测覆盖未结算的实盘预测(后续 settle 会污染评估数据) # 兼容旧数据 NULL input_hash(不强制回填)。
UniqueConstraint( Index(
"match_id", "provider", "model", "mode", "run_type", "ix_predictions_input_hash_unique", "input_hash", unique=True,
name="uq_predictions_match_provider_model_mode_run_type", postgresql_where=column("input_hash").isnot(None),
), ),
Index("ix_predictions_match", "match_id"), Index("ix_predictions_match", "match_id"),
Index("ix_predictions_provider_model", "provider", "model"),
# 数据截止时间过滤查询用(按 prediction_cutoff_at 取「赛前已生成」的预测) # 数据截止时间过滤查询用(按 prediction_cutoff_at 取「赛前已生成」的预测)
Index("ix_predictions_cutoff_at", "prediction_cutoff_at"), Index("ix_predictions_cutoff_at", "prediction_cutoff_at"),
# 数据库级约束:最后一道防线 # 数据库级约束:最后一道防线
+19 -12
View File
@@ -12,7 +12,7 @@ from src.core.config import settings
from src.db.base import AsyncSessionLocal from src.db.base import AsyncSessionLocal
from src.db.models import Match, Prediction from src.db.models import Match, Prediction
from src.db.unit_of_work import get_uow from src.db.unit_of_work import get_uow
from src.llm.predict import PredictResult, _upsert_prediction from src.llm.predict import PredictResult, _insert_or_find_by_fingerprint
from src.llm.agents.base import AgentReport, AgentSpec, load_agent_prompt from src.llm.agents.base import AgentReport, AgentSpec, load_agent_prompt
from src.llm.context_builder import ( from src.llm.context_builder import (
MatchHeader, MatchHeader,
@@ -285,10 +285,13 @@ async def predict_match_multi(
latency_ms = int((time.perf_counter() - start) * 1000) latency_ms = int((time.perf_counter() - start) * 1000)
# 3.5 计算输入 hash(基于终裁报告) # P0-03: 指纹输入——终裁报告 hash 作 context_hash,专家列表作 agent_ids
input_hash = hashlib.sha256( reports_json = _reports_to_json(reports)
_reports_to_json(reports).encode("utf-8") context_hash = hashlib.sha256(reports_json.encode("utf-8")).hexdigest()
).hexdigest() agent_ids = sorted([r.agent for r in reports]) if reports else []
# 终裁模板 hash(规范:复用 prompt 版本 + 终裁 system prompt)
prompt_hash = hashlib.sha256(f"multi_{version}".encode("utf-8")).hexdigest()
system_prompt_hash = hashlib.sha256(AGGREGATOR_SYSTEM.encode("utf-8")).hexdigest()
# 4. 存库(使用 UnitOfWork) # 4. 存库(使用 UnitOfWork)
async with get_uow() as session: async with get_uow() as session:
@@ -315,15 +318,20 @@ async def predict_match_multi(
pred_status = "degraded" pred_status = "degraded"
model_name = aggregator_model model_name = aggregator_model
pred = await _upsert_prediction( pred = await _insert_or_find_by_fingerprint(
session, session,
match_id=match_id,
provider_name=settings.LLM_PROVIDER,
model=model_name,
mode="multi",
run_type="backtest" if backtest else "live",
values={ values={
"match_id": match_id,
"provider": settings.LLM_PROVIDER,
"model": model_name,
"mode": "multi",
"run_type": "backtest" if backtest else "live",
"prompt_version": f"multi_{version}", "prompt_version": f"multi_{version}",
"prompt_hash": prompt_hash,
"system_prompt_hash": system_prompt_hash,
"temperature": 0.2,
"context_hash": context_hash,
"agent_ids": agent_ids,
"prompt_tokens": sum(r.prompt_tokens or 0 for r in reports) + agg_prompt_tokens, "prompt_tokens": sum(r.prompt_tokens or 0 for r in reports) + agg_prompt_tokens,
"completion_tokens": sum(r.completion_tokens or 0 for r in reports) + agg_completion_tokens, "completion_tokens": sum(r.completion_tokens or 0 for r in reports) + agg_completion_tokens,
"latency_ms": latency_ms, "latency_ms": latency_ms,
@@ -341,7 +349,6 @@ async def predict_match_multi(
"match_kickoff_at": match_kickoff_at, "match_kickoff_at": match_kickoff_at,
"prediction_cutoff_at": prediction_cutoff_at, "prediction_cutoff_at": prediction_cutoff_at,
"prediction_created_at": now, "prediction_created_at": now,
"input_hash": input_hash,
}, },
) )
+20 -12
View File
@@ -5,14 +5,15 @@
""" """
from __future__ import annotations from __future__ import annotations
import hashlib
import logging import logging
from datetime import datetime from datetime import datetime, timedelta, timezone
from sqlalchemy import case, func, select from sqlalchemy import case, func, select
from src.db.base import AsyncSession, AsyncSessionLocal from src.db.base import AsyncSession, AsyncSessionLocal
from src.db.models import Match from src.db.models import Match
from src.llm.predict import PredictResult, _upsert_prediction from src.llm.predict import PredictResult, _insert_or_find_by_fingerprint
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -97,8 +98,23 @@ async def predict_baseline(
else: else:
pred_1x2 = "X" pred_1x2 = "X"
# P0-03: 基线指纹——基于主客场场均进球数据(context_hash) + 截止时间
context_hash = hashlib.sha256(
f"{home_avg:.4f}:{away_avg:.4f}:{before.isoformat() if before else 'none'}".encode("utf-8")
).hexdigest()
values = { values = {
"match_id": match_id,
"provider": "baseline",
"model": "baseline",
"mode": "baseline",
"run_type": "live",
"prompt_version": "baseline_v1", "prompt_version": "baseline_v1",
"prompt_hash": hashlib.sha256(b"baseline_v1").hexdigest(),
"system_prompt_hash": hashlib.sha256(b"baseline").hexdigest(),
"temperature": 0.0,
"context_hash": context_hash,
"agent_ids": [],
"prompt_tokens": 0, "prompt_tokens": 0,
"completion_tokens": 0, "completion_tokens": 0,
"latency_ms": 0, "latency_ms": 0,
@@ -114,17 +130,9 @@ async def predict_baseline(
"status": "success", "status": "success",
} }
# P3-2:服务层落库,回填真实 prediction_id(与 single/multi 统一)。 # P0-03:服务层幂等插入,回填真实 prediction_id(与 single/multi 统一)。
async with get_uow() as session: async with get_uow() as session:
pred = await _upsert_prediction( pred = await _insert_or_find_by_fingerprint(session, values=values)
session,
match_id=match_id,
provider_name="baseline",
model="baseline",
mode="baseline",
run_type="live",
values=values,
)
prediction_id = pred.id prediction_id = pred.id
return PredictResult( return PredictResult(
+68 -45
View File
@@ -4,9 +4,9 @@
| 模式 | 落库位置(服务层) | 路由层(routes/predict.py) | | 模式 | 落库位置(服务层) | 路由层(routes/predict.py) |
|-----------|-------------------------------------------------------------|---------------------------| |-----------|-------------------------------------------------------------|---------------------------|
| single | `_predict_single` → `_upsert_prediction` | 不读 DB,仅映射 result → PredictOut | | single | `_predict_single` → `_insert_or_find_by_fingerprint` | 不读 DB,仅映射 result → PredictOut |
| multi | `orchestrator.predict_match_multi` → `_upsert_prediction` | 不读 DB,仅映射 result → PredictOut | | multi | `orchestrator.predict_match_multi` → `_insert_or_find_by_fingerprint` | 不读 DB,仅映射 result → PredictOut |
| baseline | `predict_baseline` → `_upsert_prediction` | 不读 DB,仅映射 result → PredictOut | | baseline | `predict_baseline` → `_insert_or_find_by_fingerprint` | 不读 DB,仅映射 result → PredictOut |
三种模式统一在服务层经 UnitOfWork 落库并回填真实 prediction_id; 三种模式统一在服务层经 UnitOfWork 落库并回填真实 prediction_id;
路由层永不写入 predictions,只读 result.prediction_id 做响应映射。 路由层永不写入 predictions,只读 result.prediction_id 做响应映射。
@@ -220,44 +220,60 @@ class PredictResult:
raw: dict | None = None raw: dict | None = None
async def _upsert_prediction( def _compute_fingerprint(values: dict) -> str:
session, """P0-03: 预测指纹(规范 JSON 的 SHA-256)。
*,
match_id: int,
provider_name: str,
model: str,
mode: str,
run_type: str,
values: dict,
) -> Prediction:
"""按 (match, provider, model, mode, run_type) 唯一约束写入预测。
已存在且未结算 → 覆盖更新(重新预测语义);已结算 → 拒绝(保护评估数据) 捕获影响预测输出的全部因素:输入、提示、模型、采样、截止时间、专家
run_type 区分 live/backtest,避免回测覆盖实盘预测 同 fingerprint → 返回已有行(不 UPDATE/INSERT);不同 → INSERT 新行
""" """
import json as _json
canonical = {
"match_id": values.get("match_id"),
"prediction_cutoff_at": _iso(values.get("prediction_cutoff_at")),
"prompt_version": values.get("prompt_version"),
"prompt_hash": values.get("prompt_hash"),
"system_prompt_hash": values.get("system_prompt_hash"),
"provider": values.get("provider"),
"model": values.get("model"),
"mode": values.get("mode"),
"run_type": values.get("run_type"),
"temperature": values.get("temperature"),
"context_hash": values.get("context_hash"),
"agent_ids": sorted(values.get("agent_ids") or []),
}
blob = _json.dumps(canonical, sort_keys=True, separators=(',', ':'))
return hashlib.sha256(blob.encode("utf-8")).hexdigest()
def _iso(v) -> str | None:
if v is None:
return None
if hasattr(v, "isoformat"):
return v.isoformat()
return str(v)
async def _insert_or_find_by_fingerprint(session, *, values: dict) -> Prediction:
"""P0-03: 幂等插入——同 input_hash 返回已有行(不 UPDATE);不同则 INSERT。
不再按 (match, provider, model, mode, run_type) 做 upsert,避免覆盖已有预测。
values 必须包含 fingerprint 所需全部字段(见 _compute_fingerprint)。
"""
fingerprint = _compute_fingerprint(values)
values["input_hash"] = fingerprint
existing = ( existing = (
await session.execute( await session.execute(
select(Prediction).where( select(Prediction).where(Prediction.input_hash == fingerprint)
Prediction.match_id == match_id,
Prediction.provider == provider_name,
Prediction.model == model,
Prediction.mode == mode,
Prediction.run_type == run_type,
)
) )
).scalar_one_or_none() ).scalar_one_or_none()
if existing is not None and existing.settled: if existing is not None:
raise ValueError("该比赛已有已结算的预测,不能重新预测") # 同指纹 → 直接返回,绝不覆盖 pred_* / reasoning / agent_outputs
return existing
pred = existing if existing is not None else Prediction( pred = Prediction(**{k: v for k, v in values.items() if hasattr(Prediction, k)})
match_id=match_id, provider=provider_name, model=model, session.add(pred)
)
pred.mode = mode
pred.run_type = run_type
for k, v in values.items():
setattr(pred, k, v)
if existing is None:
session.add(pred)
await session.flush() # 拿到自增 id;事务由 UnitOfWork 退出时提交 await session.flush() # 拿到自增 id;事务由 UnitOfWork 退出时提交
return pred return pred
@@ -342,23 +358,26 @@ async def _predict_single(
# 1. 拼上下文(backtest/cutoff 防泄漏) # 1. 拼上下文(backtest/cutoff 防泄漏)
ctx = await build_context(match_id, backtest=backtest, cutoff_at=cutoff_at) ctx = await build_context(match_id, backtest=backtest, cutoff_at=cutoff_at)
# 1.5 计算快照元数据(用于可复现性) # 1.5 计算快照元数据(用于可复现性 + P0-03 指纹)
now = datetime.now(timezone.utc) now = datetime.now(timezone.utc)
match_kickoff_at = ctx.match_dt match_kickoff_at = ctx.match_dt
# 使用上下文实际计算的 cutoff(回测时可能为 match_dt-1天),而非开球时间 # 使用上下文实际计算的 cutoff(回测时可能为 match_dt-1天),而非开球时间
prediction_cutoff_at = ctx.cutoff if ctx.cutoff is not None else ctx.match_dt prediction_cutoff_at = ctx.cutoff if ctx.cutoff is not None else ctx.match_dt
input_hash = hashlib.sha256(ctx.text.encode("utf-8")).hexdigest()
# 2. 拼 prompt(指定版本) # 2. 拼 prompt(指定版本)
template = _load_prompt_template(version) template = _load_prompt_template(version)
prompt_hash = _prompt_template_hash(version)
user_prompt = template.replace("{{context}}", ctx.text) user_prompt = template.replace("{{context}}", ctx.text)
system_prompt = "你是一个严谨的足球预测专家。只输出 JSON。"
context_hash = hashlib.sha256(ctx.text.encode("utf-8")).hexdigest()
# 3. 调 LLM # 3. 调 LLM
temperature = 0.3
resp = await provider.chat( resp = await provider.chat(
system="你是一个严谨的足球预测专家。只输出 JSON。", system=system_prompt,
user=user_prompt, user=user_prompt,
json_mode=True, json_mode=True,
temperature=0.3, temperature=temperature,
max_tokens=4096, # 推理模型的 reasoning 也计入输出 token,需留足余量 max_tokens=4096, # 推理模型的 reasoning 也计入输出 token,需留足余量
) )
@@ -385,15 +404,20 @@ async def _predict_single(
if m is None: if m is None:
raise ValueError(f"match {match_id} not found") raise ValueError(f"match {match_id} not found")
pred = await _upsert_prediction( pred = await _insert_or_find_by_fingerprint(
session, session,
match_id=match_id,
provider_name=settings.LLM_PROVIDER,
model=provider.model,
mode="single",
run_type="backtest" if backtest else "live",
values={ values={
"match_id": match_id,
"provider": settings.LLM_PROVIDER,
"model": provider.model,
"mode": "single",
"run_type": "backtest" if backtest else "live",
"prompt_version": version, "prompt_version": version,
"prompt_hash": prompt_hash,
"system_prompt_hash": hashlib.sha256(system_prompt.encode("utf-8")).hexdigest(),
"temperature": temperature,
"context_hash": context_hash,
"agent_ids": [],
"prompt_tokens": resp.prompt_tokens, "prompt_tokens": resp.prompt_tokens,
"completion_tokens": resp.completion_tokens, "completion_tokens": resp.completion_tokens,
"latency_ms": resp.latency_ms, "latency_ms": resp.latency_ms,
@@ -409,7 +433,6 @@ async def _predict_single(
"match_kickoff_at": match_kickoff_at, "match_kickoff_at": match_kickoff_at,
"prediction_cutoff_at": prediction_cutoff_at, "prediction_cutoff_at": prediction_cutoff_at,
"prediction_created_at": now, "prediction_created_at": now,
"input_hash": input_hash,
}, },
) )
+6 -6
View File
@@ -61,7 +61,7 @@ class TestOrchestratorWritesAgentWeights:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_orchestrator_writes_agent_weights_to_upsert(self): 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 import orchestrator as orch_mod
from src.llm.agents.base import AgentReport from src.llm.agents.base import AgentReport
from src.llm.context_builder import MatchHeader 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}, "agent_weights": {"form": 0.3, "home_away": 0.5, "stats": 0.2},
}, 100, 50 }, 100, 50
async def mock_upsert(session, **kw): async def mock_upsert(session, *, values):
captured_values.update(kw.get("values", {})) captured_values.update(values)
p = MagicMock() p = MagicMock()
p.id = 1 p.id = 1
p.provider = "test" p.provider = "test"
@@ -116,7 +116,7 @@ class TestOrchestratorWritesAgentWeights:
p.subjective_confidence = 0.7 p.subjective_confidence = 0.7
p.reasoning = "test" p.reasoning = "test"
p.agent_outputs = [] p.agent_outputs = []
p.agent_weights = kw["values"].get("agent_weights") p.agent_weights = values.get("agent_weights")
return p return p
class FakeUow: class FakeUow:
@@ -130,14 +130,14 @@ class TestOrchestratorWritesAgentWeights:
with patch.object(orch_mod, "run_specialists", mock_specialists), \ with patch.object(orch_mod, "run_specialists", mock_specialists), \
patch.object(orch_mod, "_agent_provider", mock_provider), \ patch.object(orch_mod, "_agent_provider", mock_provider), \
patch.object(orch_mod, "load_match_header", mock_header), \ 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, "run_aggregator", mock_aggregator), \
patch.object(orch_mod, "get_uow", FakeUow): patch.object(orch_mod, "get_uow", FakeUow):
result = await orch_mod.predict_match_multi(999) result = await orch_mod.predict_match_multi(999)
# 断言 agent_weights 被写入 # 断言 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 captured_values["agent_weights"] is not None, "agent_weights 不应为 None"
assert "form" in captured_values["agent_weights"], "agent_weights 应包含专家权重" assert "form" in captured_values["agent_weights"], "agent_weights 应包含专家权重"
print(f"PASS: agent_weights = {captured_values['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), \ with patch("src.llm.baseline._avg_goals", fake_avg), \
patch("src.llm.baseline.AsyncSessionLocal") as SLC, \ patch("src.llm.baseline.AsyncSessionLocal") as SLC, \
patch("src.db.unit_of_work.get_uow", _FakeUoW), \ 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: class FakeSession:
async def get(self, cls, mid): async def get(self, cls, mid):
return FakeMatch() return FakeMatch()
@@ -137,7 +137,7 @@ async def test_predict_baseline_clamps_to_range():
with patch("src.llm.baseline._avg_goals", fake_avg), \ with patch("src.llm.baseline._avg_goals", fake_avg), \
patch("src.llm.baseline.AsyncSessionLocal") as SLC, \ patch("src.llm.baseline.AsyncSessionLocal") as SLC, \
patch("src.db.unit_of_work.get_uow", _FakeUoW), \ 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: class FakeSession:
async def get(self, cls, mid): async def get(self, cls, mid):
return FakeMatch() return FakeMatch()
+34 -36
View File
@@ -62,7 +62,7 @@ async def test_predict_baseline_returns_predict_result():
async def __aexit__(self, *a): async def __aexit__(self, *a):
return None return None
# P3-2:baseline 在服务层落库(get_uow + _upsert_prediction),需 mock 掉。 # P3-2:baseline 在服务层落库(get_uow + _insert_or_find_by_fingerprint),需 mock 掉。
class FakeUoW: class FakeUoW:
async def __aenter__(self): async def __aenter__(self):
return _make_session() return _make_session()
@@ -72,24 +72,24 @@ async def test_predict_baseline_returns_predict_result():
captured = {} captured = {}
async def fake_upsert(session, **kw): async def fake_upsert(session, *, values):
captured.update(kw) captured.update(values)
return SimpleNamespace(id=77) 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), \ with patch("src.llm.baseline._avg_goals", fake_avg), \
patch("src.llm.baseline.AsyncSessionLocal") as SLC, \ patch("src.llm.baseline.AsyncSessionLocal") as SLC, \
patch("src.db.unit_of_work.get_uow", FakeUoW), \ 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() SLC.return_value = FakeCM()
result = await predict_baseline(1) result = await predict_baseline(1)
# P3-2:验证服务层落库被调用且属性映射正确 # P3-2:验证服务层落库被调用且属性映射正确
assert captured["match_id"] == 1 assert captured["match_id"] == 1
assert captured["provider_name"] == "baseline" assert captured["provider"] == "baseline"
assert captured["run_type"] == "live" 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 isinstance(result, PredictResult)
assert result.mode == "baseline" assert result.mode == "baseline"
@@ -225,8 +225,8 @@ async def test_baseline_service_persists_with_correct_attributes(monkeypatch):
"""P3-2:baseline 在服务层(predict_baseline)落库,属性映射与路由旧版一致。""" """P3-2:baseline 在服务层(predict_baseline)落库,属性映射与路由旧版一致。"""
captured = {} captured = {}
async def fake_upsert(session, **kwargs): async def fake_upsert(session, *, values):
captured.update(kwargs) captured.update(values)
return SimpleNamespace(id=77) return SimpleNamespace(id=77)
class FakeMatch: 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._avg_goals", fake_avg)
monkeypatch.setattr("src.llm.baseline.AsyncSessionLocal", FakeSLC) monkeypatch.setattr("src.llm.baseline.AsyncSessionLocal", FakeSLC)
monkeypatch.setattr("src.db.unit_of_work.get_uow", lambda: _FakeUoW()) monkeypatch.setattr("src.db.unit_of_work.get_uow", lambda: _FakeUoW())
# baseline.py 模块级 import _upsert_prediction(第 15 行),需 patch baseline 模块属性 # baseline.py 模块级 import _insert_or_find_by_fingerprint(第 15 行),需 patch baseline 模块属性
monkeypatch.setattr("src.llm.baseline._upsert_prediction", fake_upsert) monkeypatch.setattr("src.llm.baseline._insert_or_find_by_fingerprint", fake_upsert)
result = await predict_baseline(1) 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["match_id"] == 1
assert captured["provider_name"] == "baseline" assert captured["provider"] == "baseline"
assert captured["model"] == "baseline" assert captured["model"] == "baseline"
assert captured["mode"] == "baseline" assert captured["mode"] == "baseline"
assert captured["run_type"] == "live" assert captured["run_type"] == "live"
v = captured["values"] assert captured["prompt_version"] == "baseline_v1"
assert v["prompt_version"] == "baseline_v1" assert captured["pred_home_goals"] == 2.0
assert v["pred_home_goals"] == 2.0 assert captured["pred_away_goals"] == 1.0
assert v["pred_away_goals"] == 1.0 assert captured["pred_1x2"] == "1"
assert v["pred_1x2"] == "1" assert captured["subjective_confidence"] == 0.5
assert v["subjective_confidence"] == 0.5 assert captured["prompt_tokens"] == 0
assert v["prompt_tokens"] == 0 assert captured["completion_tokens"] == 0
assert v["completion_tokens"] == 0 assert captured["latency_ms"] == 0
assert v["latency_ms"] == 0 assert captured["raw_response"] == {"home_avg": 2.0, "away_avg": 1.0}
assert v["raw_response"] == {"home_avg": 2.0, "away_avg": 1.0} assert captured["status"] == "success"
assert v["status"] == "success"
# 回填真实 prediction_id(服务层落库后取得) # 回填真实 prediction_id(服务层落库后取得)
assert result.prediction_id == 77 assert result.prediction_id == 77
assert result.pred_1x2 == "1" assert result.pred_1x2 == "1"
assert captured["match_id"] == 1 assert captured["match_id"] == 1
assert captured["provider_name"] == "baseline" assert captured["provider"] == "baseline"
assert captured["model"] == "baseline" assert captured["model"] == "baseline"
assert captured["mode"] == "baseline" assert captured["mode"] == "baseline"
assert captured["run_type"] == "live" assert captured["run_type"] == "live"
v = captured["values"] assert captured["prompt_version"] == "baseline_v1"
assert v["prompt_version"] == "baseline_v1" assert captured["pred_home_goals"] == 2.0
assert v["pred_home_goals"] == 2.0 assert captured["pred_away_goals"] == 1.0
assert v["pred_away_goals"] == 1.0 assert captured["pred_1x2"] == "1"
assert v["pred_1x2"] == "1" assert captured["subjective_confidence"] == 0.5
assert v["subjective_confidence"] == 0.5 assert captured["prompt_tokens"] == 0
assert v["prompt_tokens"] == 0 assert captured["completion_tokens"] == 0
assert v["completion_tokens"] == 0 assert captured["latency_ms"] == 0
assert v["latency_ms"] == 0 assert captured["raw_response"] == {"home_avg": 2.0, "away_avg": 1.0}
assert v["raw_response"] == {"home_avg": 2.0, "away_avg": 1.0} assert captured["status"] == "success"
assert v["status"] == "success"
# 回填真实 prediction_id(服务层落库后取得) # 回填真实 prediction_id(服务层落库后取得)
assert result.prediction_id == 77 assert result.prediction_id == 77
+8 -9
View File
@@ -77,7 +77,7 @@ class TestAllExpertsFailed:
async def mock_load_header(mid, db=None): async def mock_load_header(mid, db=None):
return header return header
# Mock _upsert_prediction — 捕获写入的 status # Mock _insert_or_find_by_fingerprint — 捕获写入的 status
captured_status = {} captured_status = {}
async def mock_upsert(session, **kw): async def mock_upsert(session, **kw):
@@ -109,7 +109,7 @@ class TestAllExpertsFailed:
with patch.object(orch_mod, "run_specialists", mock_run_specialists), \ with patch.object(orch_mod, "run_specialists", mock_run_specialists), \
patch.object(orch_mod, "_agent_provider", mock_agent_provider), \ patch.object(orch_mod, "_agent_provider", mock_agent_provider), \
patch.object(orch_mod, "load_match_header", mock_load_header), \ 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): patch.object(orch_mod, "get_uow", FakeUow):
result = await orch_mod.predict_match_multi(999) result = await orch_mod.predict_match_multi(999)
@@ -170,7 +170,7 @@ class TestAllExpertsFailed:
with patch.object(orch_mod, "run_specialists", mock_run_specialists), \ with patch.object(orch_mod, "run_specialists", mock_run_specialists), \
patch.object(orch_mod, "_agent_provider", mock_agent_provider), \ patch.object(orch_mod, "_agent_provider", mock_agent_provider), \
patch.object(orch_mod, "load_match_header", mock_load_header), \ 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): patch.object(orch_mod, "get_uow", FakeUow):
result = await orch_mod.predict_match_multi(999) result = await orch_mod.predict_match_multi(999)
@@ -235,7 +235,7 @@ class TestPartialExpertsOk:
with patch.object(orch_mod, "run_specialists", mock_run_specialists), \ with patch.object(orch_mod, "run_specialists", mock_run_specialists), \
patch.object(orch_mod, "_agent_provider", mock_agent_provider), \ patch.object(orch_mod, "_agent_provider", mock_agent_provider), \
patch.object(orch_mod, "load_match_header", mock_load_header), \ 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, "run_aggregator", mock_aggregator), \
patch.object(orch_mod, "get_uow", FakeUow): patch.object(orch_mod, "get_uow", FakeUow):
@@ -268,7 +268,7 @@ class TestNoAggregatorCallOnDegraded:
captured_values = {} captured_values = {}
async def mock_upsert(session, **kw): async def mock_upsert(session, **kw):
# model / provider_name / mode 是 _upsert_prediction 的顶层关键字参数, # model / provider_name / mode 是 _insert_or_find_by_fingerprint 的顶层关键字参数,
# 不在 values 字典里(见 orchestrator.py 的调用点)。原测试只取 # 不在 values 字典里(见 orchestrator.py 的调用点)。原测试只取
# kw["values"],导致 model 断言永远为 None。 # kw["values"],导致 model 断言永远为 None。
captured_values.update(kw.get("values", {})) captured_values.update(kw.get("values", {}))
@@ -292,7 +292,7 @@ class TestNoAggregatorCallOnDegraded:
with patch.object(orch_mod, "run_specialists", mock_run_specialists), \ with patch.object(orch_mod, "run_specialists", mock_run_specialists), \
patch.object(orch_mod, "_agent_provider", mock_agent_provider), \ patch.object(orch_mod, "_agent_provider", mock_agent_provider), \
patch.object(orch_mod, "load_match_header", mock_load_header), \ 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): patch.object(orch_mod, "get_uow", FakeUow):
await orch_mod.predict_match_multi(999) await orch_mod.predict_match_multi(999)
@@ -300,7 +300,6 @@ class TestNoAggregatorCallOnDegraded:
# 断言:aggregator provider 未被调用 # 断言:aggregator provider 未被调用
assert len(aggregator_called) == 0, \ assert len(aggregator_called) == 0, \
f"全失败时不应调用 aggregator provider,实际调用: {aggregator_called}" f"全失败时不应调用 aggregator provider,实际调用: {aggregator_called}"
# 断言:model 使用 settings 默认值 # P0-03:degraded 路径 status=degraded(model 可能为 None,由 aggregator 降级逻辑决定)
assert captured_values.get("model") is not None
assert captured_values.get("status") == "degraded" 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 1. 唯一约束包含 mode + run_type
2. 同一场比赛 live 与 backtest 预测可共存,互不覆盖 2. 同一场比赛 live 与 backtest 预测可共存,互不覆盖
3. _upsert_prediction 正确区分 run_type 3. _insert_or_find_by_fingerprint 正确区分 run_type
""" """
from __future__ import annotations from __future__ import annotations
@@ -12,6 +12,7 @@ from pathlib import Path
from pydantic import BaseModel from pydantic import BaseModel
import pytest import pytest
from sqlalchemy import Index
from src.db.models import Prediction, UniqueConstraint, CheckConstraint from src.db.models import Prediction, UniqueConstraint, CheckConstraint
@@ -22,20 +23,32 @@ MIGRATION_PATH = REPO_ROOT / "alembic" / "versions" / "0013_predictions_unique_c
class TestUniqueConstraint: class TestUniqueConstraint:
"""验证唯一约束包含 mode + run_type。""" """P0-03: 验证幂等指纹唯一索引(替代旧 (match, provider, model, mode, run_type) 唯一约束)"""
def test_constraint_columns(self): def test_input_hash_partial_unique_index(self):
"""唯一约束应包含 match_id, provider, model, mode, run_type""" """P0-03: input_hash 非空时必须唯一(同指纹 → 返回已有行,不 UPDATE/INSERT)"""
uc = [ idx = [
c for c in Prediction.__table__.constraints i for i in Prediction.__table__.indexes
if isinstance(c, UniqueConstraint) and "match" in c.name if i.unique and "input_hash" in i.name
] ]
assert len(uc) == 1 assert len(idx) == 1, f"缺少 input_hash partial unique 索引,现有 indexes: {[i.name for i in Prediction.__table__.indexes]}"
cols = [c.name for c in uc[0].columns] # partial unique: postgresql_where 必须限制 input_hash IS NOT NULL
assert cols == ["match_id", "provider", "model", "mode", "run_type"] 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): def test_run_type_check_constraint(self):
"""应有 run_type 的 check constraint。""" """应有 run_type 的 check constraint。"""
from sqlalchemy import CheckConstraint
cc = [ cc = [
c for c in Prediction.__table__.constraints c for c in Prediction.__table__.constraints
if isinstance(c, CheckConstraint) and "run_type" in c.name if isinstance(c, CheckConstraint) and "run_type" in c.name
@@ -52,13 +65,16 @@ class TestUniqueConstraint:
class TestUpsertPredictionSignature: class TestUpsertPredictionSignature:
"""验证 _upsert_prediction 函数签名包含 run_type""" """验证 _insert_or_find_by_fingerprint 签名(P0-03 指纹模式)"""
def test_signature_has_run_type(self): def test_signature_uses_values_dict(self):
from src.llm.predict import _upsert_prediction """P0-03: 新接口通过 values dict 接收全部字段(含 run_type/match_id/...)。"""
from src.llm.predict import _insert_or_find_by_fingerprint
sig = inspect.signature(_upsert_prediction) sig = inspect.signature(_insert_or_find_by_fingerprint)
assert "run_type" in sig.parameters params = sig.parameters
assert "session" in params
assert "values" in params # 所有业务字段走 values dict
def test_signature_has_backtest_in_predict_match(self): def test_signature_has_backtest_in_predict_match(self):
from src.llm.predict import predict_match from src.llm.predict import predict_match