feat: 足球 LLM 预测服务初始提交

Profeto — 给 LLM 提供数据,让 LLM 预测足球比分。

核心模块:
- FastAPI 后端 + PostgreSQL (SQLAlchemy async)
- 多 Agent LLM 预测 (5 专家 + 终裁)
- 数据采集 (bzzoiro / understat / injuries)
- React 前端 (Vite + Tailwind)

包含:
- 数据源抽象 (DataSource 协议 + 注册表)
- Alembic 数据库迁移
- Prompt 模板 (单/多 Agent)
- 核心路径单元测试
This commit is contained in:
shangfangjian
2026-09-09 02:10:47 +08:00
commit 0a27b18c27
74 changed files with 5667 additions and 0 deletions
+224
View File
@@ -0,0 +1,224 @@
"""多 agent 预测编排: 并行专家 → 终裁 → 存库。"""
from __future__ import annotations
import asyncio
import json
import logging
import time
from dataclasses import dataclass
from src.core.config import settings
from src.db.base import AsyncSessionLocal
from src.db.models import Match, Prediction
from src.llm.agents.base import AgentReport, AgentSpec, load_agent_prompt
from src.llm.context_builder import (
MatchHeader,
form_slice,
h2h_slice,
header_text,
home_away_slice,
injuries_slice,
load_match_header,
stats_slice,
)
from src.llm.provider import LLMProvider, get_default_provider
logger = logging.getLogger(__name__)
# ── 5 个专家 agent 定义 ──
# A=近期状态 B=攻防数据 C=主客因素 D=阵容完整性 E=历史交锋
SPECIALIST_SPECS: list[AgentSpec] = [
AgentSpec(
name="form",
system_prompt="你是足球近期状态分析专家。分析比分与关键事件,输出近期走势判断。只输出 JSON。",
slice_fn=form_slice,
),
AgentSpec(
name="stats",
system_prompt="你是足球攻防数据分析专家。评估进球、射门与控球,输出攻防强度。只输出 JSON。",
slice_fn=stats_slice,
),
AgentSpec(
name="home_away",
system_prompt="你是足球主客因素分析专家。对比主场与客场表现,评估地理优势影响。只输出 JSON。",
slice_fn=home_away_slice,
),
AgentSpec(
name="injuries",
system_prompt="你是足球阵容完整性分析专家。汇总伤停与停赛名单,输出战力缺失程度。只输出 JSON。",
slice_fn=injuries_slice,
),
AgentSpec(
name="h2h",
system_prompt="你是足球历史交锋分析专家。分析过去数年以及近期的交手数据,提取交手规律。只输出 JSON。",
slice_fn=h2h_slice,
),
]
AGGREGATOR_SYSTEM = "你是足球预测终裁专家。综合各领域报告输出最终预测。只输出 JSON。"
@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
pred_1x2: str | None
confidence: float | None
reasoning: str | None
agent_outputs: list[dict]
agent_weights: dict | None
context: str
latency_ms: int | None
raw: dict | None
def _get_specialist_provider() -> LLMProvider:
"""专家模型: LLM_SPECIALIST_MODEL 回落 LLM_MODEL。"""
p = get_default_provider()
if settings.LLM_SPECIALIST_MODEL:
p.model = settings.LLM_SPECIALIST_MODEL
return p
def _get_aggregator_provider() -> LLMProvider:
"""终裁模型: LLM_AGGREGATOR_MODEL 回落 LLM_MODEL。"""
p = get_default_provider()
if settings.LLM_AGGREGATOR_MODEL:
p.model = settings.LLM_AGGREGATOR_MODEL
return p
async def run_specialists(
header: MatchHeader,
*,
provider: LLMProvider,
version: str = "v1",
) -> list[AgentReport]:
"""并行执行 5 个专家 agent。fail-open: 单个失败不影响其他。"""
tasks = [
_run_one(spec, header, provider, version=version)
for spec in SPECIALIST_SPECS
]
results = await asyncio.gather(*tasks, return_exceptions=True)
reports: list[AgentReport] = []
for spec, r in zip(SPECIALIST_SPECS, results):
if isinstance(r, Exception):
logger.warning("agent %s raised: %s", spec.name, r)
reports.append(AgentReport(agent=spec.name, status="error", analysis=str(r)[:200]))
else:
reports.append(r)
return reports
async def _run_one(spec, header, provider, *, version) -> AgentReport:
from src.llm.agents.base import run_agent
return await run_agent(spec, header, provider, before=header.match_dt, version=version)
def _reports_to_json(reports: list[AgentReport]) -> str:
return json.dumps([r.to_dict() for r in reports], ensure_ascii=False, indent=1)
async def run_aggregator(
header: MatchHeader,
reports: list[AgentReport],
*,
provider: LLMProvider,
version: str = "v1",
) -> tuple[dict, int, int]:
"""终裁: 汇总报告 → 最终 JSON。返回 (解析结果, prompt_tokens, completion_tokens)。"""
template = load_agent_prompt("aggregator", version)
user_prompt = (
template
.replace("{{match_header}}", header_text(header))
.replace("{{agent_reports}}", _reports_to_json(reports))
)
resp = await provider.chat(
system=AGGREGATOR_SYSTEM,
user=user_prompt,
json_mode=True,
temperature=0.2,
max_tokens=1000,
)
if resp.error:
raise RuntimeError(f"aggregator LLM error: {resp.error}")
if not resp.parsed:
raise RuntimeError(f"aggregator 输出无法解析: {resp.content[:200]}")
return resp.parsed, resp.prompt_tokens or 0, resp.completion_tokens or 0
async def predict_match_multi(
match_id: int,
*,
provider: LLMProvider | None = None,
version: str = "v1",
) -> MultiPredictResult:
"""多 agent 端到端预测: 切片 → 并行专家 → 终裁 → 存库。"""
start = time.perf_counter()
# 1. 比赛头(各 agent 共享;不存在则 404)
header = await load_match_header(match_id)
# 2. 并行专家
specialist_provider = _get_specialist_provider()
reports = await run_specialists(header, provider=specialist_provider, version=version)
# 3. 终裁
aggregator_provider = _get_aggregator_provider()
final, agg_prompt_tokens, agg_completion_tokens = await run_aggregator(
header, reports, provider=aggregator_provider, version=version
)
latency_ms = int((time.perf_counter() - start) * 1000)
# 4. 存库
async with AsyncSessionLocal() as db:
m = await db.get(Match, match_id)
if m is None:
raise ValueError(f"match {match_id} not found")
agent_weights = final.get("agent_weights")
pred = Prediction(
match_id=match_id,
provider=settings.LLM_PROVIDER,
model=aggregator_provider.model,
prompt_version=f"multi_{version}",
mode="multi",
prompt_tokens=sum(r.prompt_tokens or 0 for r in reports) + agg_prompt_tokens,
completion_tokens=sum(r.completion_tokens or 0 for r in reports) + agg_completion_tokens,
latency_ms=latency_ms,
pred_home_goals=final.get("pred_home_goals"),
pred_away_goals=final.get("pred_away_goals"),
pred_1x2=final.get("1x2"),
confidence=final.get("confidence"),
reasoning=final.get("reasoning"),
raw_response=final,
agent_outputs=[r.to_dict() for r in reports],
)
db.add(pred)
await db.commit()
await db.refresh(pred)
return MultiPredictResult(
prediction_id=pred.id,
provider=pred.provider,
model=pred.model,
prompt_version=pred.prompt_version,
mode="multi",
pred_home_goals=pred.pred_home_goals,
pred_away_goals=pred.pred_away_goals,
pred_1x2=pred.pred_1x2,
confidence=pred.confidence,
reasoning=pred.reasoning,
agent_outputs=pred.agent_outputs,
agent_weights=agent_weights,
context=_reports_to_json(reports),
latency_ms=latency_ms,
raw=final,
)