diff --git a/src/api/routes/predict.py b/src/api/routes/predict.py index bcf99c1..eff6e84 100644 --- a/src/api/routes/predict.py +++ b/src/api/routes/predict.py @@ -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", }, ) diff --git a/src/llm/agents/orchestrator.py b/src/llm/agents/orchestrator.py index f54621b..d1b214d 100644 --- a/src/llm/agents/orchestrator.py +++ b/src/llm/agents/orchestrator.py @@ -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: diff --git a/src/llm/baseline.py b/src/llm/baseline.py index baad4c7..d4aa690 100644 --- a/src/llm/baseline.py +++ b/src/llm/baseline.py @@ -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)}, + ) diff --git a/src/llm/predict.py b/src/llm/predict.py index 383b79f..c302e95 100644 --- a/src/llm/predict.py +++ b/src/llm/predict.py @@ -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—— diff --git a/tests/test_baseline.py b/tests/test_baseline.py index 6d76338..2eb7564 100644 --- a/tests/test_baseline.py +++ b/tests/test_baseline.py @@ -81,18 +81,18 @@ async def test_predict_baseline_no_llm(): result = await predict_baseline(1) - assert result["provider"] == "baseline" - assert result["model"] == "baseline" - assert result["mode"] == "baseline" - assert result["latency_ms"] == 0 - assert result["prompt_tokens"] == 0 - assert result["completion_tokens"] == 0 + assert result.provider == "baseline" + assert result.model == "baseline" + assert result.mode == "baseline" + assert result.latency_ms == 0 + assert result.prompt_tokens == 0 + assert result.completion_tokens == 0 # 2.4 → round = 2, 1.6 → round = 2 → 平局 X - 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 "非投注建议" in result["reasoning"] + 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 "非投注建议" in result.reasoning # 确认未调用任何 LLM 相关模块 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) - assert result["pred_home_goals"] == 10.0 # clamped - assert result["pred_away_goals"] == 0.0 # clamped - assert result["pred_1x2"] == "1" # 10:0 主胜 + assert result.pred_home_goals == 10.0 # clamped + assert result.pred_away_goals == 0.0 # clamped + assert result.pred_1x2 == "1" # 10:0 主胜 diff --git a/tests/test_d2_predict_unification.py b/tests/test_d2_predict_unification.py new file mode 100644 index 0000000..aa9a386 --- /dev/null +++ b/tests/test_d2_predict_unification.py @@ -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"