P2-1 no_data 门控依赖文案子串(脆弱): - context_builder 新增 SliceResult(text/has_data/n_records), 5 个切片函数改为显式声明 has_data - base._slice_has_data() 优先取结构化结果,str 返回仍走文案回退 (兼容既有测试 mock 与自定义切片) - build_context 的 has_stats/has_injuries 直接取切片声明 P2-2 agent_weights 无校验即落库: - validation 新增 AgentWeightsSchema / validate_agent_weights: 未知专家名丢弃、越界值钳制、总和非 1 时归一化 - orchestrator 落库前对 agent_weights 做校验 P2-3 1x2 与比分不一致被静默修正: - 仍以比分修正,但补 logger.warning 暴露 LLM 自相矛盾 P2-5/P2-6 prompt 缓存不可刷新 + 缓存键不含模板内容: - 新增 clear_prompt_cache() 供改模板后显式失效 - 缓存键纳入模板内容 hash,模板一改缓存自动失效 P2-7 ingest/backtest/settle 接口无鉴权: - 新增 require_admin_key 依赖(X-API-Key), ADMIN_API_KEY 未设置时放行并告警(不破坏本地开发) - 挂到 3 个 ingest 接口 + backtest + eval/settle P2-8 前端请求竞态 + 未使用游标分页: - Matches.tsx 用递增 seq 丢弃过期响应,避免旧筛选结果覆盖新筛选 - 接入后端已有的 cursor 分页 + 「加载更多」按钮 附带: .env.example 补齐 LLM_TIMEOUT / 分档模型 / ADMIN_API_KEY; tests 新增 10 个用例覆盖 P2-1/2/3。
248 lines
8.4 KiB
Python
248 lines
8.4 KiB
Python
"""多 agent 预测编排: 并行专家 → 终裁 → 存库。"""
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import hashlib
|
|
import json
|
|
import logging
|
|
import time
|
|
from dataclasses import dataclass
|
|
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.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
|
|
subjective_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)
|
|
match_kickoff_at = header.match_dt
|
|
prediction_cutoff_at = header.match_dt # 默认:比赛时间作为数据截止
|
|
now = datetime.now(timezone.utc)
|
|
|
|
# 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)
|
|
|
|
# 3.5 计算输入 hash(基于终裁报告)
|
|
input_hash = hashlib.sha256(
|
|
_reports_to_json(reports).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
|
|
try:
|
|
validated = validate_prediction_output(final)
|
|
except Exception as e:
|
|
raise RuntimeError(f"终裁输出校验失败: {e}")
|
|
|
|
# agent_weights 同样必须过校验(旧实现直接取 raw 值落库,未做任何检查)
|
|
agent_weights = validate_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,
|
|
subjective_confidence=validated.subjective_confidence,
|
|
reasoning=validated.reasoning,
|
|
raw_response=final,
|
|
agent_outputs=[r.to_dict() for r in reports],
|
|
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)
|
|
|
|
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,
|
|
subjective_confidence=pred.subjective_confidence,
|
|
reasoning=pred.reasoning,
|
|
agent_outputs=pred.agent_outputs,
|
|
agent_weights=agent_weights,
|
|
context=_reports_to_json(reports),
|
|
latency_ms=latency_ms,
|
|
raw=final,
|
|
)
|