Files
Profeto/src/llm/agents/orchestrator.py
T
shangfangjian f3160e3062 feat: Sprint 1 - 数据正确性整改
P2-01: 移除 lifespan create_all,改为仅验证连接
       新增 /health/ready 就绪检查
P0-04: LLM 输出严格 Pydantic 校验
       - Agent 输出越界/非法 → parse_error
       - 预测输出自动修正 1X2 与比分一致性
P0-01: injuries cutoff 修复
       - get_injuries_for_match 增加 as_of 参数
       - injuries_slice 使用 as_of 过滤 retrieved_at
       - 防止回测时未来采集数据泄漏
P1-12: 批量入库优化
       - 预加载 teams 到内存 dict
       - 预加载 existing matches 到内存 set
       - 消灭 N+1 查询
2026-09-14 23:36:35 +08:00

232 lines
7.6 KiB
Python

"""多 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")
# 严格校验终裁输出
from src.llm.validation import validate_prediction_output
try:
validated = validate_prediction_output(final)
except Exception as e:
raise RuntimeError(f"终裁输出校验失败: {e}")
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=validated.pred_home_goals,
pred_away_goals=validated.pred_away_goals,
pred_1x2=validated.pred_1x2,
confidence=validated.confidence,
reasoning=validated.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,
)