P0-02: MatchStats 增加 source/source_record_id/retrieved_at/available_at
context_builder 过滤统计数据时检查 available_at <= cutoff
P0-03: Prediction 增加 match_kickoff_at/prediction_created_at/prediction_cutoff_at
明确区分比赛时间/预测创建时间/数据截止时间
P0-05: Prediction 增加 status 字段(success/failed/degraded)
数据库 CHECK 约束
Alembic: 0005_prediction_status_and_stats_provenance
P1-02: injuries as_of 不再截断为 date,保持 datetime 精度
197 lines
6.3 KiB
Python
197 lines
6.3 KiB
Python
"""预测服务:拼上下文 → 调 LLM → 存预测。"""
|
|
from __future__ import annotations
|
|
|
|
import functools
|
|
import hashlib
|
|
import logging
|
|
import time
|
|
from dataclasses import dataclass
|
|
from datetime import datetime, timezone
|
|
from pathlib import Path
|
|
from threading import Lock
|
|
|
|
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.context_builder import build_context
|
|
from src.llm.provider import LLMProvider, get_default_provider
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
_PROMPT_DIR = Path(__file__).resolve().parent / "prompts"
|
|
|
|
# ── LLM 响应缓存(match+provider+model+version → 结果) ──
|
|
_CACHE_TTL_SEC = 300 # 5 分钟
|
|
_cache: dict[str, tuple[float, PredictResult]] = {}
|
|
_cache_lock = Lock()
|
|
|
|
|
|
def _cache_key(match_id: int, provider: str, model: str, version: str) -> str:
|
|
return f"{match_id}:{provider}:{model}:{version}"
|
|
|
|
|
|
def _get_cached(match_id: int, provider: str, model: str, version: str) -> PredictResult | None:
|
|
key = _cache_key(match_id, provider, model, version)
|
|
with _cache_lock:
|
|
if key in _cache:
|
|
ts, result = _cache[key]
|
|
if time.time() - ts < _CACHE_TTL_SEC:
|
|
return result
|
|
del _cache[key]
|
|
return None
|
|
|
|
|
|
def _set_cached(match_id: int, provider: str, model: str, version: str, result: PredictResult) -> None:
|
|
key = _cache_key(match_id, provider, model, version)
|
|
with _cache_lock:
|
|
_cache[key] = (time.time(), result)
|
|
|
|
|
|
@functools.lru_cache(maxsize=8)
|
|
def _load_prompt_template(version: str = "v1") -> str:
|
|
"""缓存 prompt 模板(进程生命周期内每个版本只读一次)。"""
|
|
path = _PROMPT_DIR / f"match_prediction_{version}.md"
|
|
if not path.exists():
|
|
raise FileNotFoundError(f"prompt 模板不存在: {path}")
|
|
with open(path, encoding="utf-8") as f:
|
|
return f.read()
|
|
|
|
|
|
@dataclass
|
|
class PredictResult:
|
|
prediction_id: int
|
|
provider: str
|
|
model: str
|
|
prompt_version: str
|
|
pred_home_goals: float | None
|
|
pred_away_goals: float | None
|
|
pred_1x2: str | None
|
|
subjective_confidence: float | None
|
|
reasoning: str | None
|
|
context: str
|
|
latency_ms: int | None
|
|
raw: dict | None
|
|
|
|
|
|
async def predict_match(
|
|
match_id: int,
|
|
*,
|
|
provider: LLMProvider | None = None,
|
|
model: str | None = None,
|
|
prompt_version: str | None = None,
|
|
mode: str = "multi",
|
|
) -> "PredictResult | MultiPredictResult":
|
|
"""预测入口。mode=multi(默认)走多 agent;mode=single 走单次调用。"""
|
|
if mode == "single":
|
|
return await _predict_single(
|
|
match_id, provider=provider, model=model, prompt_version=prompt_version
|
|
)
|
|
from src.llm.agents.orchestrator import predict_match_multi
|
|
|
|
return await predict_match_multi(match_id, provider=provider, version=(prompt_version or "v1").removeprefix("multi_"))
|
|
|
|
|
|
async def _predict_single(
|
|
match_id: int,
|
|
*,
|
|
provider: LLMProvider | None = None,
|
|
model: str | None = None,
|
|
prompt_version: str | None = None,
|
|
) -> PredictResult:
|
|
"""单次调用路径(原有实现)。"""
|
|
if provider is None:
|
|
provider = get_default_provider()
|
|
if model:
|
|
provider.model = model
|
|
version = prompt_version or "v1"
|
|
|
|
# 0. 查缓存(同 match+provider+model+version 5 分钟内直接返)
|
|
cached = _get_cached(match_id, settings.LLM_PROVIDER, provider.model, version)
|
|
if cached is not None:
|
|
logger.debug("predict cache hit match=%s", match_id)
|
|
return cached
|
|
|
|
# 1. 拼上下文
|
|
ctx = await build_context(match_id)
|
|
|
|
# 1.5 计算快照元数据(用于可复现性)
|
|
now = datetime.now(timezone.utc)
|
|
match_kickoff_at = ctx.match_dt
|
|
prediction_cutoff_at = ctx.match_dt # 默认:比赛时间作为数据截止
|
|
input_hash = hashlib.sha256(ctx.text.encode("utf-8")).hexdigest()
|
|
|
|
# 2. 拼 prompt(指定版本)
|
|
template = _load_prompt_template(version)
|
|
user_prompt = template.replace("{{context}}", ctx.text)
|
|
|
|
# 3. 调 LLM
|
|
resp = await provider.chat(
|
|
system="你是一个严谨的足球预测专家。只输出 JSON。",
|
|
user=user_prompt,
|
|
json_mode=True,
|
|
temperature=0.3,
|
|
max_tokens=800,
|
|
)
|
|
|
|
if resp.error:
|
|
raise RuntimeError(f"LLM error: {resp.error}")
|
|
|
|
parsed = resp.parsed or {}
|
|
|
|
# 3.5 严格校验 LLM 输出
|
|
from src.llm.validation import validate_prediction_output
|
|
try:
|
|
validated = validate_prediction_output(parsed)
|
|
except Exception as e:
|
|
raise RuntimeError(f"LLM 输出校验失败: {e}")
|
|
|
|
# 4. 存预测(使用 UnitOfWork 统一事务)
|
|
async with get_uow() as session:
|
|
# 验证 match 存在
|
|
m = await session.get(Match, match_id)
|
|
if m is None:
|
|
raise ValueError(f"match {match_id} not found")
|
|
|
|
pred = Prediction(
|
|
match_id=match_id,
|
|
provider=settings.LLM_PROVIDER,
|
|
model=provider.model,
|
|
prompt_version=version,
|
|
prompt_tokens=resp.prompt_tokens,
|
|
completion_tokens=resp.completion_tokens,
|
|
latency_ms=resp.latency_ms,
|
|
pred_home_goals=validated.pred_home_goals,
|
|
pred_away_goals=validated.pred_away_goals,
|
|
pred_1x2=validated.pred_1x2,
|
|
subjective_confidence=validated.subjective_confidence,
|
|
reasoning=validated.reasoning,
|
|
raw_response=resp.raw,
|
|
status="success",
|
|
match_kickoff_at=match_kickoff_at,
|
|
prediction_cutoff_at=prediction_cutoff_at,
|
|
prediction_created_at=now,
|
|
input_hash=input_hash,
|
|
)
|
|
session.add(pred)
|
|
await session.refresh(pred)
|
|
|
|
result = PredictResult(
|
|
prediction_id=pred.id,
|
|
provider=pred.provider,
|
|
model=pred.model,
|
|
prompt_version=version,
|
|
pred_home_goals=pred.pred_home_goals,
|
|
pred_away_goals=pred.pred_away_goals,
|
|
pred_1x2=pred.pred_1x2,
|
|
subjective_confidence=pred.subjective_confidence,
|
|
reasoning=pred.reasoning,
|
|
context=ctx.text,
|
|
latency_ms=resp.latency_ms,
|
|
raw=resp.raw,
|
|
)
|
|
|
|
# 5. 写入缓存
|
|
_set_cached(match_id, settings.LLM_PROVIDER, provider.model, version, result)
|
|
return result
|