"""多 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, )