"""多 agent 预测编排: 并行专家 → 终裁 → 存库。""" from __future__ import annotations import asyncio import hashlib import json import logging import time from datetime import datetime, timezone 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.predict import PredictResult, _insert_or_find_by_fingerprint 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, load_match_header, standings_slice, stats_slice, ) from src.core.runtime_config import get_runtime_value from src.llm.provider import LLMProvider, get_default_provider logger = logging.getLogger(__name__) # P3-2: agent provider 配置缓存(TTL 60s),避免每次 _agent_provider 都多次查 DB _AGENT_PROVIDER_CACHE: dict[str, tuple[float, LLMProvider]] = {} _AGENT_PROVIDER_CACHE_TTL = 60.0 # ── 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="standings", system_prompt="你是足球联赛排名分析专家。分析积分榜位置、积分走势与分区,评估两队整体实力差距。只输出 JSON。", slice_fn=standings_slice, ), AgentSpec( name="h2h", system_prompt="你是足球历史交锋分析专家。分析过去数年以及近期的交手数据,提取交手规律。只输出 JSON。", slice_fn=h2h_slice, ), ] AGGREGATOR_SYSTEM = ( "你是足球预测终裁专家。综合各领域专家报告输出最终预测。" "引用专家时必须使用报告中的专家全名(如「攻防数据分析专家」),禁止使用英文代码。" "只输出 JSON。" ) # 专家代码 → 终裁/展示统一称呼 AGENT_LABELS_ZH: dict[str, str] = { "form": "近期状态分析专家", "stats": "攻防数据分析专家", "home_away": "主客因素分析专家", "standings": "联赛排名分析专家", "h2h": "历史交锋分析专家", } # D2(工程债): multi 结果类型与 single 统一 —— 扩展后的 PredictResult 用可选 # 字段(agent_outputs/agent_weights/prompt_tokens/completion_tokens/mode)承载 # 全部模式,此处仅保留别名。保留 `MultiPredictResult` 名字的原因: # 1. predict_match_multi 签名 `-> MultiPredictResult:` 是 R4 源码守卫的标记; # 2. src/llm/agents/__init__.py 对外 re-export 该名字。 MultiPredictResult = PredictResult async def _agent_provider(agent_id: str, *, tier: str, model_override: str | None = None) -> LLMProvider: """构造某 agent 专属 provider。 覆盖优先级: 模型: model_override(调用方显式指定) → AGENT_MODEL_{ID}(运行时) → 层级默认(LLM_SPECIALIST/AGGREGATOR_MODEL) → 全局 LLM_MODEL 地址/密钥: AGENT_BASE_URL_{ID} / AGENT_API_KEY_{ID}(运行时) → 全局 LLM_BASE_URL / LLM_API_KEY P3-2: 结果缓存 60 秒,避免每次预测都多次查询运行时配置 DB。 注意:model_override 生效时跳过缓存读写 —— 否则带 override 的结果会泄漏给 不带 override 的调用(反之亦然),导致跨调用的模型串味。 """ cache_key = f"{agent_id}:{tier}" if model_override is None: cached = _AGENT_PROVIDER_CACHE.get(cache_key) if cached is not None: ts, provider = cached if time.time() - ts < _AGENT_PROVIDER_CACHE_TTL: return provider pfx = f"AGENT_{agent_id.upper()}_" p = await get_default_provider() tier_model = settings.LLM_SPECIALIST_MODEL if tier == "specialist" else settings.LLM_AGGREGATOR_MODEL if tier_model: p.model = tier_model model = await get_runtime_value(f"{pfx}MODEL") if model: p.model = model # 调用方显式传入的 model 优先级最高,高于 agent 级与层级默认 if model_override: p.model = model_override base = await get_runtime_value(f"{pfx}BASE_URL") if base: p.base_url = base key = await get_runtime_value(f"{pfx}API_KEY") if key: p.api_key = key if model_override is None: _AGENT_PROVIDER_CACHE[cache_key] = (time.time(), p) # 简单淘汰:超过 20 条时清空(60s TTL 下不会累积太多) if len(_AGENT_PROVIDER_CACHE) > 20: _AGENT_PROVIDER_CACHE.clear() return p async def run_specialists( header: MatchHeader, *, version: str = "v1", before=None, model_override: str | None = None, ) -> list[AgentReport]: """并行执行 5 个专家 agent。fail-open: 单个失败不影响其他。 before: 数据截止时间(回测防泄漏)。None 表示不限制。 model_override: 调用方显式指定的模型,覆盖各 agent 的层级默认。 注意:model_override 必须作为形参下传,不能用模块级变量中转。 backtest 会 asyncio.gather 并发 8 场预测(见 backtest.py 的 Semaphore(8)), 模块级变量会被并发调用互相覆盖,导致 A 场的预测用上 B 场的模型。 """ tasks = [ _run_one(spec, header, await _agent_provider(spec.name, tier="specialist", model_override=model_override), version=version, before=before) 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( "专家调用失败 match=%s agent=%s error=%s", header.match_id, spec.name, str(r)[:120], ) 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, before=None) -> AgentReport: from src.llm.agents.base import run_agent return await run_agent(spec, header, provider, before=before, version=version) def _reports_to_json(reports: list[AgentReport]) -> str: """报告序列化: agent 字段直接用中文专家全名,引导终裁用统一称呼引用。""" out = [] for r in reports: d = r.to_dict() d["agent"] = AGENT_LABELS_ZH.get(d.get("agent", ""), d.get("agent")) out.append(d) return json.dumps(out, 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=4096, # 推理模型需要更大余量 ) 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", backtest: bool = False, cutoff_at=None, model: str | None = None, ) -> MultiPredictResult: """多 agent 端到端预测: 切片 → 并行专家 → 终裁 → 存库。 backtest: 回测模式。True 时 cutoff 自动设为 match_dt - 1 天。 cutoff_at: 显式截止时间(优先于 backtest 自动计算)。 model: 显式指定模型,优先于 agent 级/层级默认配置(single 模式语义一致)。 """ start = time.perf_counter() # 1. 比赛头(各 agent 共享;不存在则 404) header = await load_match_header(match_id) match_kickoff_at = header.match_dt now = datetime.now(timezone.utc) # 计算真正的数据截止时间(回测防泄漏) # 优先级: 显式 cutoff_at > backtest 自动计算 > 默认(比赛时间) if cutoff_at is not None: cutoff = cutoff_at elif backtest and header.match_dt: from datetime import timedelta cutoff = header.match_dt - timedelta(days=1) else: cutoff = header.match_dt prediction_cutoff_at = cutoff # 2. 并行专家(各自独立配置,使用统一 cutoff) # model 作为形参下传,而非模块级变量:backtest 并发 8 场预测时, # 模块级变量会被并发调用互相覆盖(模型串味)。 reports = await run_specialists( header, version=version, before=cutoff, model_override=model ) # 2.5 统计有效专家报告数量 ok_reports = [r for r in reports if r.status == "ok"] has_valid_data = len(ok_reports) > 0 # 3. 终裁(仅当有有效专家报告时执行) if has_valid_data: aggregator_provider = await _agent_provider("aggregator", tier="aggregator") final, agg_prompt_tokens, agg_completion_tokens = await run_aggregator( header, reports, provider=aggregator_provider, version=version ) else: # 所有专家无数据/均失败:跳过终裁,标记 degraded logger.warning( "预测降级 match=%s mode=%s status=degraded experts=%d/%d 均无有效数据", match_id, "multi", len(ok_reports), len(reports), ) # 无有效专家时不调用 aggregator provider,避免多余开销 # model 使用 settings 默认值占位(无实际 LLM 调用) final = { "pred_home_goals": None, "pred_away_goals": None, "alt_pred_home_goals": None, "alt_pred_away_goals": None, "pred_1x2": None, "subjective_confidence": None, "reasoning": f"所有 {len(reports)} 位专家均无有效数据或预测失败(状态: {','.join(r.status for r in reports)})", "agent_weights": {}, } agg_prompt_tokens = 0 agg_completion_tokens = 0 aggregator_model = model or settings.LLM_MODEL # 占位,无实际 LLM 调用 latency_ms = int((time.perf_counter() - start) * 1000) # P0-03: 指纹输入——终裁报告 hash 作 context_hash,专家列表作 agent_ids reports_json = _reports_to_json(reports) context_hash = hashlib.sha256(reports_json.encode("utf-8")).hexdigest() agent_ids = sorted([r.agent for r in reports]) if reports else [] # 终裁模板 hash(规范:复用 prompt 版本 + 终裁 system prompt) prompt_hash = hashlib.sha256(f"multi_{version}".encode("utf-8")).hexdigest() system_prompt_hash = hashlib.sha256(AGGREGATOR_SYSTEM.encode("utf-8")).hexdigest() # 4. 存库(使用 UnitOfWork) async with get_uow() as session: m = await session.get(Match, match_id) if m is None: raise ValueError(f"match {match_id} not found") # 根据是否有有效数据决定校验策略 from src.llm.validation import validate_agent_weights, validate_prediction_output if has_valid_data: # 有有效报告:严格校验终裁输出 try: validated = validate_prediction_output(final) except Exception as e: raise RuntimeError(f"终裁输出校验失败: {e}") agent_weights = validate_agent_weights(final.get("agent_weights")) pred_status = "success" model_name = aggregator_provider.model else: # 无有效报告:跳过严格校验,直接构造降级结果 validated = None # type: ignore agent_weights = {} pred_status = "degraded" model_name = aggregator_model pred = await _insert_or_find_by_fingerprint( session, values={ "match_id": match_id, "provider": settings.LLM_PROVIDER, "model": model_name, "mode": "multi", "run_type": "backtest" if backtest else "live", "prompt_version": f"multi_{version}", "prompt_hash": prompt_hash, "system_prompt_hash": system_prompt_hash, "temperature": 0.2, "context_hash": context_hash, "agent_ids": agent_ids, "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 if validated else None, "pred_away_goals": validated.pred_away_goals if validated else None, "alt_pred_home_goals": validated.alt_pred_home_goals if validated else None, "alt_pred_away_goals": validated.alt_pred_away_goals if validated else None, "pred_1x2": validated.pred_1x2 if validated else None, "subjective_confidence": validated.subjective_confidence if validated else None, "reasoning": validated.reasoning if validated else final.get("reasoning", ""), "raw_response": final, "agent_outputs": [r.to_dict() for r in reports], "agent_weights": agent_weights, "status": pred_status, "match_kickoff_at": match_kickoff_at, "prediction_cutoff_at": prediction_cutoff_at, "prediction_created_at": now, }, ) logger.info( "预测完成 match=%s mode=%s status=%s pred=%s:%s (%s) latency=%sms, experts=%d/%d, prediction_id=%s", match_id, "multi", pred_status, pred.pred_home_goals, pred.pred_away_goals, pred.pred_1x2, latency_ms, len(ok_reports), len(reports), pred.id, ) 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, alt_pred_home_goals=pred.alt_pred_home_goals, alt_pred_away_goals=pred.alt_pred_away_goals, pred_1x2=pred.pred_1x2, subjective_confidence=pred.subjective_confidence, reasoning=pred.reasoning, status=pred_status, agent_outputs=pred.agent_outputs, agent_weights=agent_weights, context=_reports_to_json(reports), latency_ms=latency_ms, prompt_tokens=pred.prompt_tokens, completion_tokens=pred.completion_tokens, raw=final, )