Files
Profeto/src/llm/predict.py
T
shangfangjian 0980c2242a feat: P0-02/P0-03/P0-05 数据库时间语义与约束
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 精度
2026-09-15 02:04:25 +08:00

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