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")
|
||||
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",
|
||||
},
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user