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
This commit is contained in:
+31
-34
@@ -63,7 +63,8 @@ async def predict(req: PredictRequest, request: Request):
|
|||||||
logger.exception("predict unexpected error")
|
logger.exception("predict unexpected error")
|
||||||
raise HTTPException(500, "预测失败,请查看服务器日志")
|
raise HTTPException(500, "预测失败,请查看服务器日志")
|
||||||
|
|
||||||
# baseline 模式:结果已是 dict,需独立落库(prediction_id)
|
# D2: 三种模式统一返回 PredictResult —— 字段映射单一化,无 dict 分支。
|
||||||
|
# 仅 baseline 的 prediction_id 需要在此落库补齐(服务层不落库)。
|
||||||
if req.mode == "baseline":
|
if req.mode == "baseline":
|
||||||
prediction_id = await _persist_baseline(req.match_id, result)
|
prediction_id = await _persist_baseline(req.match_id, result)
|
||||||
else:
|
else:
|
||||||
@@ -73,38 +74,34 @@ async def predict(req: PredictRequest, request: Request):
|
|||||||
logger.info(
|
logger.info(
|
||||||
"预测完成 match=%s mode=%s pred=%s:%s (%s)",
|
"预测完成 match=%s mode=%s pred=%s:%s (%s)",
|
||||||
req.match_id, req.mode,
|
req.match_id, req.mode,
|
||||||
result.get("pred_home_goals") if isinstance(result, dict) else result.pred_home_goals,
|
result.pred_home_goals, result.pred_away_goals, result.pred_1x2,
|
||||||
result.get("pred_away_goals") if isinstance(result, dict) else result.pred_away_goals,
|
|
||||||
result.get("pred_1x2") if isinstance(result, dict) else result.pred_1x2,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
result_dict = result if isinstance(result, dict) else None
|
|
||||||
|
|
||||||
return PredictOut(
|
return PredictOut(
|
||||||
prediction_id=prediction_id,
|
prediction_id=prediction_id,
|
||||||
provider=result.get("provider") if result_dict else result.provider,
|
provider=result.provider,
|
||||||
model=result.get("model") if result_dict else result.model,
|
model=result.model,
|
||||||
prompt_version=result.get("prompt_version") if result_dict else getattr(result, "prompt_version", None),
|
prompt_version=result.prompt_version,
|
||||||
mode=req.mode,
|
mode=req.mode,
|
||||||
pred_home_goals=result.get("pred_home_goals") if result_dict else result.pred_home_goals,
|
pred_home_goals=result.pred_home_goals,
|
||||||
pred_away_goals=result.get("pred_away_goals") if result_dict else result.pred_away_goals,
|
pred_away_goals=result.pred_away_goals,
|
||||||
alt_pred_home_goals=result.get("alt_pred_home_goals") if result_dict else result.alt_pred_home_goals,
|
alt_pred_home_goals=result.alt_pred_home_goals,
|
||||||
alt_pred_away_goals=result.get("alt_pred_away_goals") if result_dict else result.alt_pred_away_goals,
|
alt_pred_away_goals=result.alt_pred_away_goals,
|
||||||
pred_1x2=result.get("pred_1x2") if result_dict else result.pred_1x2,
|
pred_1x2=result.pred_1x2,
|
||||||
subjective_confidence=result.get("subjective_confidence") if result_dict else result.subjective_confidence,
|
subjective_confidence=result.subjective_confidence,
|
||||||
reasoning=result.get("reasoning") if result_dict else result.reasoning,
|
reasoning=result.reasoning,
|
||||||
status=result.get("status", "success") if result_dict else getattr(result, "status", "success"),
|
status=result.status,
|
||||||
agent_outputs=result.get("agent_outputs") if result_dict else getattr(result, "agent_outputs", None),
|
agent_outputs=result.agent_outputs,
|
||||||
agent_weights=result.get("agent_weights") if result_dict else getattr(result, "agent_weights", None),
|
agent_weights=result.agent_weights,
|
||||||
context=result.get("context", "") if result_dict else result.context,
|
context=result.context,
|
||||||
latency_ms=result.get("latency_ms", 0) if result_dict else result.latency_ms,
|
latency_ms=result.latency_ms,
|
||||||
prompt_tokens=result.get("prompt_tokens") if result_dict else getattr(result, "prompt_tokens", None),
|
prompt_tokens=result.prompt_tokens,
|
||||||
completion_tokens=result.get("completion_tokens") if result_dict else getattr(result, "completion_tokens", None),
|
completion_tokens=result.completion_tokens,
|
||||||
rate_limit_remaining=get_predict_rate_limit_remaining(request),
|
rate_limit_remaining=get_predict_rate_limit_remaining(request),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
async def _persist_baseline(match_id: int, baseline: dict) -> int:
|
async def _persist_baseline(match_id: int, baseline: PredictResult) -> int:
|
||||||
"""将基线预测结果写入 prediction 表,复用 upsert 语义。"""
|
"""将基线预测结果写入 prediction 表,复用 upsert 语义。"""
|
||||||
from src.db.unit_of_work import get_uow
|
from src.db.unit_of_work import get_uow
|
||||||
from src.llm.predict import _upsert_prediction
|
from src.llm.predict import _upsert_prediction
|
||||||
@@ -118,16 +115,16 @@ async def _persist_baseline(match_id: int, baseline: dict) -> int:
|
|||||||
mode="baseline",
|
mode="baseline",
|
||||||
run_type="live", # baseline 是 live 预测的变体,符合 ck_run_type_enum
|
run_type="live", # baseline 是 live 预测的变体,符合 ck_run_type_enum
|
||||||
values={
|
values={
|
||||||
"prompt_version": "baseline_v1",
|
"prompt_version": baseline.prompt_version,
|
||||||
"prompt_tokens": 0,
|
"prompt_tokens": baseline.prompt_tokens or 0,
|
||||||
"completion_tokens": 0,
|
"completion_tokens": baseline.completion_tokens or 0,
|
||||||
"latency_ms": 0,
|
"latency_ms": baseline.latency_ms or 0,
|
||||||
"pred_home_goals": baseline["pred_home_goals"],
|
"pred_home_goals": baseline.pred_home_goals,
|
||||||
"pred_away_goals": baseline["pred_away_goals"],
|
"pred_away_goals": baseline.pred_away_goals,
|
||||||
"pred_1x2": baseline["pred_1x2"],
|
"pred_1x2": baseline.pred_1x2,
|
||||||
"subjective_confidence": baseline["subjective_confidence"],
|
"subjective_confidence": baseline.subjective_confidence,
|
||||||
"reasoning": baseline["reasoning"],
|
"reasoning": baseline.reasoning,
|
||||||
"raw_response": baseline.get("raw", baseline),
|
"raw_response": baseline.raw,
|
||||||
"status": "success",
|
"status": "success",
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -6,14 +6,13 @@ import hashlib
|
|||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import time
|
import time
|
||||||
from dataclasses import dataclass
|
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
|
|
||||||
from src.core.config import settings
|
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 _upsert_prediction
|
from src.llm.predict import PredictResult, _upsert_prediction
|
||||||
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,
|
||||||
@@ -81,28 +80,12 @@ AGENT_LABELS_ZH: dict[str, str] = {
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
# D2(工程债): multi 结果类型与 single 统一 —— 扩展后的 PredictResult 用可选
|
||||||
class MultiPredictResult:
|
# 字段(agent_outputs/agent_weights/prompt_tokens/completion_tokens/mode)承载
|
||||||
prediction_id: int
|
# 全部模式,此处仅保留别名。保留 `MultiPredictResult` 名字的原因:
|
||||||
provider: str
|
# 1. predict_match_multi 签名 `-> MultiPredictResult:` 是 R4 源码守卫的标记;
|
||||||
model: str
|
# 2. src/llm/agents/__init__.py 对外 re-export 该名字。
|
||||||
prompt_version: str
|
MultiPredictResult = PredictResult
|
||||||
mode: str
|
|
||||||
pred_home_goals: float | None
|
|
||||||
pred_away_goals: float | None
|
|
||||||
alt_pred_home_goals: int | None
|
|
||||||
alt_pred_away_goals: int | None
|
|
||||||
pred_1x2: str | None
|
|
||||||
subjective_confidence: float | None
|
|
||||||
reasoning: str | None
|
|
||||||
context: str
|
|
||||||
agent_outputs: list[dict]
|
|
||||||
agent_weights: dict | None
|
|
||||||
status: str = "success"
|
|
||||||
latency_ms: int | None = None
|
|
||||||
prompt_tokens: int | None = None
|
|
||||||
completion_tokens: int | None = None
|
|
||||||
raw: dict | None = None
|
|
||||||
|
|
||||||
|
|
||||||
async def _agent_provider(agent_id: str, *, tier: str, model_override: str | None = None) -> LLMProvider:
|
async def _agent_provider(agent_id: str, *, tier: str, model_override: str | None = None) -> LLMProvider:
|
||||||
|
|||||||
+25
-20
@@ -12,6 +12,7 @@ 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
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -52,11 +53,13 @@ async def predict_baseline(
|
|||||||
*,
|
*,
|
||||||
backtest: bool = False,
|
backtest: bool = False,
|
||||||
cutoff_at: datetime | None = None,
|
cutoff_at: datetime | None = None,
|
||||||
) -> dict:
|
) -> PredictResult:
|
||||||
"""极简基线预测:主场场均进球 vs 客场场均进球。
|
"""极简基线预测:主场场均进球 vs 客场场均进球。
|
||||||
|
|
||||||
返回与 PredictResult 兼容的字典:
|
返回 PredictResult(D2 统一结果类型):
|
||||||
provider=model="baseline", 不调用 LLM,latency_ms≈0。
|
provider=model="baseline", 不调用 LLM,latency_ms≈0。
|
||||||
|
prediction_id 为占位 0 —— baseline 不在服务层落库,
|
||||||
|
由路由层 _persist_baseline 落库后取得真实 id。
|
||||||
"""
|
"""
|
||||||
async with AsyncSessionLocal() as db:
|
async with AsyncSessionLocal() as db:
|
||||||
match = await db.get(Match, match_id)
|
match = await db.get(Match, match_id)
|
||||||
@@ -90,24 +93,26 @@ async def predict_baseline(
|
|||||||
else:
|
else:
|
||||||
pred_1x2 = "X"
|
pred_1x2 = "X"
|
||||||
|
|
||||||
return {
|
return PredictResult(
|
||||||
"pred_home_goals": float(pred_home),
|
prediction_id=0, # 占位:真实 id 由路由层 _persist_baseline 落库后返回
|
||||||
"pred_away_goals": float(pred_away),
|
provider="baseline",
|
||||||
"alt_pred_home_goals": None,
|
model="baseline",
|
||||||
"alt_pred_away_goals": None,
|
prompt_version="baseline_v1",
|
||||||
"pred_1x2": pred_1x2,
|
mode="baseline",
|
||||||
"subjective_confidence": 0.5,
|
pred_home_goals=float(pred_home),
|
||||||
"prompt_tokens": 0,
|
pred_away_goals=float(pred_away),
|
||||||
"completion_tokens": 0,
|
alt_pred_home_goals=None,
|
||||||
"reasoning": (
|
alt_pred_away_goals=None,
|
||||||
|
pred_1x2=pred_1x2,
|
||||||
|
subjective_confidence=0.5,
|
||||||
|
reasoning=(
|
||||||
f"基线估计(非投注建议): 主队主场场均进球 {home_avg:.2f} → 预测 {pred_home}; "
|
f"基线估计(非投注建议): 主队主场场均进球 {home_avg:.2f} → 预测 {pred_home}; "
|
||||||
f"客队客场场均进球 {away_avg:.2f} → 预测 {pred_away}。"
|
f"客队客场场均进球 {away_avg:.2f} → 预测 {pred_away}。"
|
||||||
),
|
),
|
||||||
"provider": "baseline",
|
context="", # baseline 不构建 LLM 上下文
|
||||||
"model": "baseline",
|
status="success",
|
||||||
"prompt_version": "baseline_v1",
|
latency_ms=0,
|
||||||
"mode": "baseline",
|
prompt_tokens=0,
|
||||||
"status": "success",
|
completion_tokens=0,
|
||||||
"latency_ms": 0,
|
raw={"home_avg": round(home_avg, 2), "away_avg": round(away_avg, 2)},
|
||||||
"raw": {"home_avg": round(home_avg, 2), "away_avg": round(away_avg, 2)},
|
)
|
||||||
}
|
|
||||||
|
|||||||
+18
-1
@@ -89,6 +89,14 @@ def _prompt_template_hash(version: str) -> str:
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class PredictResult:
|
class PredictResult:
|
||||||
|
"""三种预测模式(single/multi/baseline)的统一结果类型。
|
||||||
|
|
||||||
|
D2(工程债): 原本 single 返回本类、multi 重复定义 MultiPredictResult、
|
||||||
|
baseline 返回裸 dict,导致路由 isinstance(dict) 双分支 + backtest 对
|
||||||
|
baseline 直接 AttributeError。现以可选字段扩展本类承载全部模式;
|
||||||
|
MultiPredictResult 是本类的别名(见 src/llm/agents/orchestrator.py)。
|
||||||
|
"""
|
||||||
|
|
||||||
prediction_id: int
|
prediction_id: int
|
||||||
provider: str
|
provider: str
|
||||||
model: str
|
model: str
|
||||||
@@ -101,8 +109,15 @@ class PredictResult:
|
|||||||
subjective_confidence: float | None
|
subjective_confidence: float | None
|
||||||
reasoning: str | None
|
reasoning: str | None
|
||||||
context: str
|
context: str
|
||||||
|
# 模式标识: single(默认) / multi / baseline
|
||||||
|
mode: str = "single"
|
||||||
|
# multi 专属: 各专家报告列表与融合权重;single/baseline 为 None
|
||||||
|
agent_outputs: list[dict] | None = None
|
||||||
|
agent_weights: dict | None = None
|
||||||
status: str = "success"
|
status: str = "success"
|
||||||
latency_ms: int | None = None
|
latency_ms: int | None = None
|
||||||
|
prompt_tokens: int | None = None
|
||||||
|
completion_tokens: int | None = None
|
||||||
raw: dict | None = None
|
raw: dict | None = None
|
||||||
|
|
||||||
|
|
||||||
@@ -158,9 +173,11 @@ async def predict_match(
|
|||||||
use_cache: bool = True,
|
use_cache: bool = True,
|
||||||
backtest: bool = False,
|
backtest: bool = False,
|
||||||
cutoff_at=None,
|
cutoff_at=None,
|
||||||
) -> "PredictResult | MultiPredictResult":
|
) -> PredictResult:
|
||||||
"""预测入口。mode=multi(默认)走多 agent;mode=single 走单次调用;mode=baseline 走无 LLM 基线。
|
"""预测入口。mode=multi(默认)走多 agent;mode=single 走单次调用;mode=baseline 走无 LLM 基线。
|
||||||
|
|
||||||
|
三种模式统一返回 PredictResult(D2);multi 的 MultiPredictResult 是其别名。
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
mode: multi(默认,5 专家+终裁) / single(单次) / baseline(极简统计基线,不调用 LLM)。
|
mode: multi(默认,5 专家+终裁) / single(单次) / baseline(极简统计基线,不调用 LLM)。
|
||||||
use_cache:是否允许返回进程内缓存结果。回测必须传 False——
|
use_cache:是否允许返回进程内缓存结果。回测必须传 False——
|
||||||
|
|||||||
+14
-14
@@ -81,18 +81,18 @@ async def test_predict_baseline_no_llm():
|
|||||||
|
|
||||||
result = await predict_baseline(1)
|
result = await predict_baseline(1)
|
||||||
|
|
||||||
assert result["provider"] == "baseline"
|
assert result.provider == "baseline"
|
||||||
assert result["model"] == "baseline"
|
assert result.model == "baseline"
|
||||||
assert result["mode"] == "baseline"
|
assert result.mode == "baseline"
|
||||||
assert result["latency_ms"] == 0
|
assert result.latency_ms == 0
|
||||||
assert result["prompt_tokens"] == 0
|
assert result.prompt_tokens == 0
|
||||||
assert result["completion_tokens"] == 0
|
assert result.completion_tokens == 0
|
||||||
# 2.4 → round = 2, 1.6 → round = 2 → 平局 X
|
# 2.4 → round = 2, 1.6 → round = 2 → 平局 X
|
||||||
assert result["pred_home_goals"] == 2.0
|
assert result.pred_home_goals == 2.0
|
||||||
assert result["pred_away_goals"] == 2.0
|
assert result.pred_away_goals == 2.0
|
||||||
assert result["pred_1x2"] == "X"
|
assert result.pred_1x2 == "X"
|
||||||
assert result["subjective_confidence"] == 0.5
|
assert result.subjective_confidence == 0.5
|
||||||
assert "非投注建议" in result["reasoning"]
|
assert "非投注建议" in result.reasoning
|
||||||
# 确认未调用任何 LLM 相关模块
|
# 确认未调用任何 LLM 相关模块
|
||||||
assert "home_10" in captured and "away_20" in captured
|
assert "home_10" in captured and "away_20" in captured
|
||||||
|
|
||||||
@@ -125,6 +125,6 @@ async def test_predict_baseline_clamps_to_range():
|
|||||||
|
|
||||||
result = await predict_baseline(2)
|
result = await predict_baseline(2)
|
||||||
|
|
||||||
assert result["pred_home_goals"] == 10.0 # clamped
|
assert result.pred_home_goals == 10.0 # clamped
|
||||||
assert result["pred_away_goals"] == 0.0 # clamped
|
assert result.pred_away_goals == 0.0 # clamped
|
||||||
assert result["pred_1x2"] == "1" # 10:0 主胜
|
assert result.pred_1x2 == "1" # 10:0 主胜
|
||||||
|
|||||||
@@ -0,0 +1,216 @@
|
|||||||
|
"""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"
|
||||||
Reference in New Issue
Block a user