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:
2026-09-21 19:23:50 +08:00
parent 3a9f3f5a0e
commit 5c7fdce0a3
6 changed files with 311 additions and 93 deletions
+31 -34
View File
@@ -63,7 +63,8 @@ async def predict(req: PredictRequest, request: Request):
logger.exception("predict unexpected error")
raise HTTPException(500, "预测失败,请查看服务器日志")
# baseline 模式:结果已是 dict,需独立落库(prediction_id)
# D2: 三种模式统一返回 PredictResult —— 字段映射单一化,无 dict 分支。
# 仅 baseline 的 prediction_id 需要在此落库补齐(服务层不落库)。
if req.mode == "baseline":
prediction_id = await _persist_baseline(req.match_id, result)
else:
@@ -73,38 +74,34 @@ async def predict(req: PredictRequest, request: Request):
logger.info(
"预测完成 match=%s mode=%s pred=%s:%s (%s)",
req.match_id, req.mode,
result.get("pred_home_goals") if isinstance(result, dict) else result.pred_home_goals,
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.pred_home_goals, result.pred_away_goals, result.pred_1x2,
)
result_dict = result if isinstance(result, dict) else None
return PredictOut(
prediction_id=prediction_id,
provider=result.get("provider") if result_dict else result.provider,
model=result.get("model") if result_dict else result.model,
prompt_version=result.get("prompt_version") if result_dict else getattr(result, "prompt_version", None),
provider=result.provider,
model=result.model,
prompt_version=result.prompt_version,
mode=req.mode,
pred_home_goals=result.get("pred_home_goals") if result_dict else result.pred_home_goals,
pred_away_goals=result.get("pred_away_goals") if result_dict else 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_away_goals=result.get("alt_pred_away_goals") if result_dict else result.alt_pred_away_goals,
pred_1x2=result.get("pred_1x2") if result_dict else result.pred_1x2,
subjective_confidence=result.get("subjective_confidence") if result_dict else result.subjective_confidence,
reasoning=result.get("reasoning") if result_dict else result.reasoning,
status=result.get("status", "success") if result_dict else getattr(result, "status", "success"),
agent_outputs=result.get("agent_outputs") if result_dict else getattr(result, "agent_outputs", None),
agent_weights=result.get("agent_weights") if result_dict else getattr(result, "agent_weights", None),
context=result.get("context", "") if result_dict else result.context,
latency_ms=result.get("latency_ms", 0) if result_dict else result.latency_ms,
prompt_tokens=result.get("prompt_tokens") if result_dict else getattr(result, "prompt_tokens", None),
completion_tokens=result.get("completion_tokens") if result_dict else getattr(result, "completion_tokens", None),
pred_home_goals=result.pred_home_goals,
pred_away_goals=result.pred_away_goals,
alt_pred_home_goals=result.alt_pred_home_goals,
alt_pred_away_goals=result.alt_pred_away_goals,
pred_1x2=result.pred_1x2,
subjective_confidence=result.subjective_confidence,
reasoning=result.reasoning,
status=result.status,
agent_outputs=result.agent_outputs,
agent_weights=result.agent_weights,
context=result.context,
latency_ms=result.latency_ms,
prompt_tokens=result.prompt_tokens,
completion_tokens=result.completion_tokens,
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 语义。"""
from src.db.unit_of_work import get_uow
from src.llm.predict import _upsert_prediction
@@ -118,16 +115,16 @@ async def _persist_baseline(match_id: int, baseline: dict) -> int:
mode="baseline",
run_type="live", # baseline 是 live 预测的变体,符合 ck_run_type_enum
values={
"prompt_version": "baseline_v1",
"prompt_tokens": 0,
"completion_tokens": 0,
"latency_ms": 0,
"pred_home_goals": baseline["pred_home_goals"],
"pred_away_goals": baseline["pred_away_goals"],
"pred_1x2": baseline["pred_1x2"],
"subjective_confidence": baseline["subjective_confidence"],
"reasoning": baseline["reasoning"],
"raw_response": baseline.get("raw", baseline),
"prompt_version": baseline.prompt_version,
"prompt_tokens": baseline.prompt_tokens or 0,
"completion_tokens": baseline.completion_tokens or 0,
"latency_ms": baseline.latency_ms or 0,
"pred_home_goals": baseline.pred_home_goals,
"pred_away_goals": baseline.pred_away_goals,
"pred_1x2": baseline.pred_1x2,
"subjective_confidence": baseline.subjective_confidence,
"reasoning": baseline.reasoning,
"raw_response": baseline.raw,
"status": "success",
},
)
+7 -24
View File
@@ -6,14 +6,13 @@ import hashlib
import json
import logging
import time
from dataclasses import dataclass
from datetime import datetime, timezone
from src.core.config import settings
from src.db.base import AsyncSessionLocal
from src.db.models import Match, Prediction
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.context_builder import (
MatchHeader,
@@ -81,28 +80,12 @@ AGENT_LABELS_ZH: dict[str, str] = {
}
@dataclass
class MultiPredictResult:
prediction_id: int
provider: str
model: str
prompt_version: str
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
# D2(工程债): multi 结果类型与 single 统一 —— 扩展后的 PredictResult 用可选
# 字段(agent_outputs/agent_weights/prompt_tokens/completion_tokens/mode)承载
# 全部模式,此处仅保留别名。保留 `MultiPredictResult` 名字的原因:
# 1. predict_match_multi 签名 `-> MultiPredictResult:` 是 R4 源码守卫的标记;
# 2. src/llm/agents/__init__.py 对外 re-export 该名字。
MultiPredictResult = PredictResult
async def _agent_provider(agent_id: str, *, tier: str, model_override: str | None = None) -> LLMProvider:
+25 -20
View File
@@ -12,6 +12,7 @@ from sqlalchemy import case, func, select
from src.db.base import AsyncSession, AsyncSessionLocal
from src.db.models import Match
from src.llm.predict import PredictResult
logger = logging.getLogger(__name__)
@@ -52,11 +53,13 @@ async def predict_baseline(
*,
backtest: bool = False,
cutoff_at: datetime | None = None,
) -> dict:
) -> PredictResult:
"""极简基线预测:主场场均进球 vs 客场场均进球。
返回 PredictResult 兼容的字典:
返回 PredictResult(D2 统一结果类型):
provider=model="baseline", 不调用 LLM,latency_ms≈0。
prediction_id 为占位 0 —— baseline 不在服务层落库,
由路由层 _persist_baseline 落库后取得真实 id。
"""
async with AsyncSessionLocal() as db:
match = await db.get(Match, match_id)
@@ -90,24 +93,26 @@ async def predict_baseline(
else:
pred_1x2 = "X"
return {
"pred_home_goals": float(pred_home),
"pred_away_goals": float(pred_away),
"alt_pred_home_goals": None,
"alt_pred_away_goals": None,
"pred_1x2": pred_1x2,
"subjective_confidence": 0.5,
"prompt_tokens": 0,
"completion_tokens": 0,
"reasoning": (
return PredictResult(
prediction_id=0, # 占位:真实 id 由路由层 _persist_baseline 落库后返回
provider="baseline",
model="baseline",
prompt_version="baseline_v1",
mode="baseline",
pred_home_goals=float(pred_home),
pred_away_goals=float(pred_away),
alt_pred_home_goals=None,
alt_pred_away_goals=None,
pred_1x2=pred_1x2,
subjective_confidence=0.5,
reasoning=(
f"基线估计(非投注建议): 主队主场场均进球 {home_avg:.2f} → 预测 {pred_home}; "
f"客队客场场均进球 {away_avg:.2f} → 预测 {pred_away}"
),
"provider": "baseline",
"model": "baseline",
"prompt_version": "baseline_v1",
"mode": "baseline",
"status": "success",
"latency_ms": 0,
"raw": {"home_avg": round(home_avg, 2), "away_avg": round(away_avg, 2)},
}
context="", # baseline 不构建 LLM 上下文
status="success",
latency_ms=0,
prompt_tokens=0,
completion_tokens=0,
raw={"home_avg": round(home_avg, 2), "away_avg": round(away_avg, 2)},
)
+18 -1
View File
@@ -89,6 +89,14 @@ def _prompt_template_hash(version: str) -> str:
@dataclass
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
provider: str
model: str
@@ -101,8 +109,15 @@ class PredictResult:
subjective_confidence: float | None
reasoning: str | None
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"
latency_ms: int | None = None
prompt_tokens: int | None = None
completion_tokens: int | None = None
raw: dict | None = None
@@ -158,9 +173,11 @@ async def predict_match(
use_cache: bool = True,
backtest: bool = False,
cutoff_at=None,
) -> "PredictResult | MultiPredictResult":
) -> PredictResult:
"""预测入口。mode=multi(默认)走多 agent;mode=single 走单次调用;mode=baseline 走无 LLM 基线。
三种模式统一返回 PredictResult(D2);multi 的 MultiPredictResult 是其别名。
Args:
mode: multi(默认,5 专家+终裁) / single(单次) / baseline(极简统计基线,不调用 LLM)。
use_cache:是否允许返回进程内缓存结果。回测必须传 False——