fix(P2): no_data 结构化、权重校验、鉴权与前端竞态
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。
This commit is contained in:
@@ -0,0 +1,43 @@
|
||||
"""API 依赖:鉴权等横切关注点。
|
||||
|
||||
审查报告 P2-7:ingest / backtest / settle 这类「写入型或高成本」接口此前
|
||||
完全无鉴权 —— 任何能访问到服务的人都可触发采集、或直接烧掉 LLM 额度。
|
||||
|
||||
策略(渐进式,不破坏本地开发):
|
||||
- `ADMIN_API_KEY` 未配置 → 直接放行,并打一次 warning。
|
||||
这样本地 `docker compose up` 无需额外配置即可用。
|
||||
- 已配置 → 必须带匹配的 `X-API-Key` 请求头,否则 401。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import secrets
|
||||
|
||||
from fastapi import Header, HTTPException
|
||||
|
||||
from src.core.config import settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_warned_unset = False
|
||||
|
||||
|
||||
async def require_admin_key(x_api_key: str | None = Header(None, alias="X-API-Key")) -> None:
|
||||
"""保护「写入型 / 高成本」接口的依赖。
|
||||
|
||||
用法: `@router.post("/ingest/bzzoiro", dependencies=[Depends(require_admin_key)])`
|
||||
"""
|
||||
global _warned_unset
|
||||
|
||||
expected = settings.ADMIN_API_KEY
|
||||
if not expected:
|
||||
if not _warned_unset:
|
||||
logger.warning(
|
||||
"ADMIN_API_KEY 未设置,采集/回测接口当前【无鉴权】。"
|
||||
"生产环境请设置该环境变量。"
|
||||
)
|
||||
_warned_unset = True
|
||||
return
|
||||
|
||||
if not x_api_key or not secrets.compare_digest(x_api_key, expected):
|
||||
raise HTTPException(status_code=401, detail="无效或缺失的 X-API-Key")
|
||||
@@ -3,9 +3,10 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from src.api.deps import require_admin_key
|
||||
from src.llm.backtest import run_backtest
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -22,10 +23,13 @@ class BacktestRequest(BaseModel):
|
||||
model: str | None = Field(None, description="指定模型 (空=默认)")
|
||||
|
||||
|
||||
@router.post("/backtest")
|
||||
@router.post("/backtest", dependencies=[Depends(require_admin_key)])
|
||||
async def backtest(req: BacktestRequest):
|
||||
"""对历史比赛运行回测。
|
||||
|
||||
该接口会对每场已完赛比赛各发起一次 LLM 预测,成本高 —— 因此需要
|
||||
`X-API-Key` 鉴权(见审查报告 P2-7)。
|
||||
|
||||
对每场已完赛比赛:
|
||||
1. 用比赛之前的数据构建上下文 (防未来信息泄漏)
|
||||
2. 调 LLM 预测
|
||||
|
||||
@@ -5,6 +5,7 @@ import logging
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
from src.api.deps import require_admin_key
|
||||
from src.api.schemas import EvalSummaryOut, SettleRequest
|
||||
from src.db.base import AsyncSession, get_db, get_db_read
|
||||
from src.llm.eval import get_eval_summary, settle_prediction
|
||||
@@ -14,7 +15,7 @@ logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/api/v1", tags=["eval"])
|
||||
|
||||
|
||||
@router.post("/eval/settle")
|
||||
@router.post("/eval/settle", dependencies=[Depends(require_admin_key)])
|
||||
async def settle(req: SettleRequest, db: AsyncSession = Depends(get_db)):
|
||||
"""回填实际结果。"""
|
||||
try:
|
||||
|
||||
@@ -3,8 +3,9 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
from src.api.deps import require_admin_key
|
||||
from src.api.schemas import IngestBzzoiroRequest, IngestResponse, IngestUnderstatRequest, IngestInjuriesRequest, IngestSimpleResponse
|
||||
from src.data.sources import get_source
|
||||
from src.data.injuries import ingest_injuries
|
||||
@@ -15,7 +16,7 @@ logger = logging.getLogger(__name__)
|
||||
router = APIRouter(prefix="/api/v1", tags=["ingest"])
|
||||
|
||||
|
||||
@router.post("/ingest/bzzoiro", response_model=IngestResponse)
|
||||
@router.post("/ingest/bzzoiro", response_model=IngestResponse, dependencies=[Depends(require_admin_key)])
|
||||
async def ingest_bzzoiro_route(req: IngestBzzoiroRequest):
|
||||
"""触发 bzzoiro 采集。"""
|
||||
source = get_source("bzzoiro")
|
||||
@@ -34,7 +35,7 @@ async def ingest_bzzoiro_route(req: IngestBzzoiroRequest):
|
||||
raise HTTPException(500, "数据采集失败,请查看服务器日志")
|
||||
|
||||
|
||||
@router.post("/ingest/understat", response_model=IngestSimpleResponse)
|
||||
@router.post("/ingest/understat", response_model=IngestSimpleResponse, dependencies=[Depends(require_admin_key)])
|
||||
async def ingest_understat_route(req: IngestUnderstatRequest):
|
||||
"""触发 understat xG 回填。"""
|
||||
source = get_source("understat")
|
||||
@@ -47,7 +48,7 @@ async def ingest_understat_route(req: IngestUnderstatRequest):
|
||||
raise HTTPException(500, "xG 回填失败,请查看服务器日志")
|
||||
|
||||
|
||||
@router.post("/ingest/injuries", response_model=IngestSimpleResponse)
|
||||
@router.post("/ingest/injuries", response_model=IngestSimpleResponse, dependencies=[Depends(require_admin_key)])
|
||||
async def ingest_injuries_route(req: IngestInjuriesRequest):
|
||||
"""触发伤停采集。"""
|
||||
try:
|
||||
|
||||
@@ -33,5 +33,11 @@ class Settings(BaseSettings):
|
||||
# --- CORS ---
|
||||
CORS_ORIGINS: str = "http://localhost:5173,http://localhost:3000"
|
||||
|
||||
# --- 管理接口鉴权 ---
|
||||
# 采集 / 回测等高成本或写入型接口需要此 Key(请求头 X-API-Key)。
|
||||
# 留空表示「未启用鉴权」(本地开发默认),生产环境必须设置。
|
||||
# 见审查报告 P2-7:ingest/backtest 无鉴权可被任意调用并烧掉 LLM 额度。
|
||||
ADMIN_API_KEY: str = ""
|
||||
|
||||
|
||||
settings = Settings()
|
||||
|
||||
+23
-5
@@ -12,7 +12,7 @@ import logging
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
|
||||
from src.llm.context_builder import MatchHeader
|
||||
from src.llm.context_builder import MatchHeader, SliceResult
|
||||
from src.llm.provider import LLMProvider, LLMResponse
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -80,7 +80,12 @@ class AgentReport:
|
||||
|
||||
|
||||
def _is_no_data(slice_text: str) -> bool:
|
||||
"""切片是否全无数据(除了标题行全是无数据)。"""
|
||||
"""兜底: 判断纯字符串切片是否全无数据。
|
||||
|
||||
仅用于 slice_fn 返回 `str`(未升级为 SliceResult)的场景。
|
||||
新代码应让切片返回 SliceResult 并显式声明 has_data —— 字符串子串匹配
|
||||
依赖具体文案(「无比分数据」「无伤停数据」等变体会漏判),不可靠。
|
||||
"""
|
||||
body = [ln.strip() for ln in slice_text.splitlines() if ln.strip()]
|
||||
# 去掉标题行(── 开头)
|
||||
content = [ln for ln in body if not ln.startswith("──")]
|
||||
@@ -89,6 +94,18 @@ def _is_no_data(slice_text: str) -> bool:
|
||||
return all(any(s in ln for s in NO_DATA_SENTINELS) for ln in content)
|
||||
|
||||
|
||||
def _slice_has_data(slice_result) -> tuple[str, bool]:
|
||||
"""把切片返回值统一成 (text, has_data)。
|
||||
|
||||
优先使用 SliceResult.has_data(结构化,可信);若切片函数仍返回 str,
|
||||
则回退到文案子串匹配(向后兼容)。
|
||||
"""
|
||||
if isinstance(slice_result, SliceResult):
|
||||
return slice_result.text, slice_result.has_data
|
||||
text = str(slice_result)
|
||||
return text, not _is_no_data(text)
|
||||
|
||||
|
||||
def _stub_no_data(agent: str) -> AgentReport:
|
||||
return AgentReport(
|
||||
agent=agent,
|
||||
@@ -145,13 +162,14 @@ async def run_agent(
|
||||
"""执行单个专家 agent: 切片 → no_data 门控 → 调 LLM → 解析报告。"""
|
||||
# 1. 数据切片
|
||||
try:
|
||||
slice_text = await spec.slice_fn(header, before=before)
|
||||
slice_result = await spec.slice_fn(header, before=before)
|
||||
except Exception as e:
|
||||
logger.exception("agent %s slice failed", spec.name)
|
||||
return AgentReport(agent=spec.name, status="error", analysis=f"数据切片失败: {e}")
|
||||
|
||||
# 2. no_data 门控: 切片无数据 → 不调 LLM
|
||||
if _is_no_data(slice_text):
|
||||
# 2. no_data 门控: 切片显式声明无数据 → 不调 LLM
|
||||
slice_text, has_data = _slice_has_data(slice_result)
|
||||
if not has_data:
|
||||
logger.debug("agent %s: slice is no_data, skipping LLM", spec.name)
|
||||
return _stub_no_data(spec.name)
|
||||
|
||||
|
||||
@@ -195,13 +195,14 @@ async def predict_match_multi(
|
||||
raise ValueError(f"match {match_id} not found")
|
||||
|
||||
# 严格校验终裁输出
|
||||
from src.llm.validation import validate_prediction_output
|
||||
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 = final.get("agent_weights")
|
||||
# agent_weights 同样必须过校验(旧实现直接取 raw 值落库,未做任何检查)
|
||||
agent_weights = validate_agent_weights(final.get("agent_weights"))
|
||||
pred = Prediction(
|
||||
match_id=match_id,
|
||||
provider=settings.LLM_PROVIDER,
|
||||
|
||||
+57
-35
@@ -39,6 +39,22 @@ def _is_stats_available(stats, before) -> bool:
|
||||
return stats.available_at <= before
|
||||
|
||||
|
||||
@dataclass
|
||||
class SliceResult:
|
||||
"""数据切片的显式结果(替代「靠文案子串猜有无数据」)。
|
||||
|
||||
旧实现用 `"无数据" in slice_text` 判断,依赖具体文案 —— 一旦某个切片
|
||||
写成「无比分数据」「无伤停数据」这类变体,判断就会静默失配
|
||||
(见审查报告 P2-1)。这里让切片函数直接声明 `has_data`,不再猜。
|
||||
"""
|
||||
text: str
|
||||
has_data: bool
|
||||
n_records: int = 0
|
||||
|
||||
def __str__(self) -> str: # 让老调用点可直接当 str 用
|
||||
return self.text
|
||||
|
||||
|
||||
@dataclass
|
||||
class MatchContext:
|
||||
match_id: int
|
||||
@@ -98,16 +114,18 @@ def header_text(h: MatchHeader) -> str:
|
||||
# 切片函数: 每个领域 agent 一个
|
||||
# ============================================================
|
||||
|
||||
async def h2h_slice(header: MatchHeader, *, limit: int = 8, before=None) -> str:
|
||||
async def h2h_slice(header: MatchHeader, *, limit: int = 8, before=None) -> SliceResult:
|
||||
"""E - 历史交锋切片: 过去数年 + 近期交手数据,提取交手规律。before=match_date 用于回测。"""
|
||||
async with AsyncSessionLocal() as db:
|
||||
h2h = await _get_h2h(db, header.home_team_id, header.away_team_id, before=before, limit=limit)
|
||||
lines = [f"── 历史交锋(近 {limit} 次) ──"]
|
||||
n_with_score = 0
|
||||
if h2h:
|
||||
home_wins = draws = away_wins = 0
|
||||
for hm in h2h:
|
||||
d = hm.match_date.strftime("%Y-%m") if hm.match_date else "?"
|
||||
if hm.home_goals is not None:
|
||||
n_with_score += 1
|
||||
if hm.home_goals > hm.away_goals: home_wins += 1
|
||||
elif hm.home_goals == hm.away_goals: draws += 1
|
||||
else: away_wins += 1
|
||||
@@ -119,15 +137,17 @@ async def h2h_slice(header: MatchHeader, *, limit: int = 8, before=None) -> str:
|
||||
lines.append(f" 总计 {total} 场: 主队 {home_wins}胜 {draws}平 {away_wins}负")
|
||||
else:
|
||||
lines.append(" 无数据")
|
||||
return "\n".join(lines)
|
||||
# has_data 以「有比分的交锋」为准:仅有对阵无比分时不足以支撑分析
|
||||
return SliceResult(text="\n".join(lines), has_data=n_with_score > 0, n_records=n_with_score)
|
||||
|
||||
|
||||
async def form_slice(header: MatchHeader, *, limit: int = 5, before=None) -> str:
|
||||
async def form_slice(header: MatchHeader, *, limit: int = 5, before=None) -> SliceResult:
|
||||
"""A - 近期状态切片: 两队近 N 场赛果、关键事件、走势判断。before=match_date 用于回测。"""
|
||||
async with AsyncSessionLocal() as db:
|
||||
home_form = await _get_form(db, header.home_team_id, before=before, limit=limit)
|
||||
away_form = await _get_form(db, header.away_team_id, before=before, limit=limit)
|
||||
lines = []
|
||||
n_scored = 0
|
||||
for label, name, form, side in (
|
||||
("主队", header.home_name, home_form, "home"),
|
||||
("客队", header.away_name, away_form, "away"),
|
||||
@@ -140,6 +160,8 @@ async def form_slice(header: MatchHeader, *, limit: int = 5, before=None) -> str
|
||||
if o == "W": wins += 1
|
||||
elif o == "D": draws += 1
|
||||
else: losses += 1
|
||||
if fm.home_goals is not None:
|
||||
n_scored += 1
|
||||
score = f"{fm.home_goals}-{fm.away_goals}" if fm.home_goals is not None else "vs"
|
||||
xg = ""
|
||||
if fm.stats and _is_stats_available(fm.stats, before) and fm.stats.home_xg is not None:
|
||||
@@ -150,15 +172,16 @@ async def form_slice(header: MatchHeader, *, limit: int = 5, before=None) -> str
|
||||
lines.append(f" 近 {len(form)} 场: {wins}胜 {draws}平 {losses}负")
|
||||
else:
|
||||
lines.append(" 无数据")
|
||||
return "\n".join(lines)
|
||||
return SliceResult(text="\n".join(lines), has_data=n_scored > 0, n_records=n_scored)
|
||||
|
||||
|
||||
async def stats_slice(header: MatchHeader, *, limit: int = 10, before=None) -> str:
|
||||
async def stats_slice(header: MatchHeader, *, limit: int = 10, before=None) -> SliceResult:
|
||||
"""B - 攻防数据切片: 进球、射门、控球,评估攻防强度。before=match_date 用于回测。"""
|
||||
async with AsyncSessionLocal() as db:
|
||||
home_form = await _get_form(db, header.home_team_id, before=before, limit=limit)
|
||||
away_form = await _get_form(db, header.away_team_id, before=before, limit=limit)
|
||||
lines = [f"── 攻防数据(近 {limit} 场) ──"]
|
||||
n_total = 0
|
||||
for label, name, form, side in (
|
||||
("主队", header.home_name, home_form, "home"),
|
||||
("客队", header.away_name, away_form, "away"),
|
||||
@@ -184,6 +207,7 @@ async def stats_slice(header: MatchHeader, *, limit: int = 10, before=None) -> s
|
||||
xg += fm.stats.home_xg if side == "home" else fm.stats.away_xg
|
||||
xga += fm.stats.away_xg if side == "home" else fm.stats.home_xg
|
||||
n_xg += 1
|
||||
n_total += n
|
||||
if n > 0:
|
||||
lines.append(f" {label} {name}:")
|
||||
lines.append(f" 场均进球 {gf/n:.2f}, 场均失球 {ga/n:.2f}")
|
||||
@@ -194,15 +218,16 @@ async def stats_slice(header: MatchHeader, *, limit: int = 10, before=None) -> s
|
||||
lines.append(f" {label} {name}: 无比分数据")
|
||||
else:
|
||||
lines.append(f" {label} {name}: 无数据")
|
||||
return "\n".join(lines)
|
||||
return SliceResult(text="\n".join(lines), has_data=n_total > 0, n_records=n_total)
|
||||
|
||||
|
||||
async def home_away_slice(header: MatchHeader, *, limit: int = 10, before=None) -> str:
|
||||
async def home_away_slice(header: MatchHeader, *, limit: int = 10, before=None) -> SliceResult:
|
||||
"""C - 主客因素切片: 主场战绩 vs 客场战绩,评估地理优势影响。before=match_date 用于回测。"""
|
||||
async with AsyncSessionLocal() as db:
|
||||
home_home = await _get_home_away(db, header.home_team_id, "home", before=before, limit=limit)
|
||||
away_away = await _get_home_away(db, header.away_team_id, "away", before=before, limit=limit)
|
||||
lines = ["── 主客因素 ──"]
|
||||
n_total = 0
|
||||
for label, name, matches, side in (
|
||||
("主队主场", header.home_name, home_home, "home"),
|
||||
("客队客场", header.away_name, away_away, "away"),
|
||||
@@ -218,6 +243,7 @@ async def home_away_slice(header: MatchHeader, *, limit: int = 10, before=None)
|
||||
gf += m.home_goals if side == "home" else m.away_goals
|
||||
ga += m.away_goals if side == "home" else m.home_goals
|
||||
n = wins + draws + losses
|
||||
n_total += n
|
||||
if n > 0:
|
||||
pct = wins / n * 100
|
||||
lines.append(f" {label} {name}(近 {n} 场): {wins}胜 {draws}平 {losses}负, 胜率 {pct:.0f}%")
|
||||
@@ -226,10 +252,10 @@ async def home_away_slice(header: MatchHeader, *, limit: int = 10, before=None)
|
||||
lines.append(f" {label} {name}: 无比分数据")
|
||||
else:
|
||||
lines.append(f" {label} {name}: 无数据")
|
||||
return "\n".join(lines)
|
||||
return SliceResult(text="\n".join(lines), has_data=n_total > 0, n_records=n_total)
|
||||
|
||||
|
||||
async def injuries_slice(header: MatchHeader, *, before=None) -> str:
|
||||
async def injuries_slice(header: MatchHeader, *, before=None) -> SliceResult:
|
||||
"""D - 阵容完整性切片: 伤停与停赛名单,评估战力缺失程度。
|
||||
|
||||
before=cutoff: 只使用 cutoff 之前已采集的伤停数据,防回测泄漏。
|
||||
@@ -242,10 +268,10 @@ async def injuries_slice(header: MatchHeader, *, before=None) -> str:
|
||||
away_injuries = await get_injuries_for_match(db, header.away_team_id, cutoff, as_of=cutoff)
|
||||
|
||||
lines = ["── 阵容完整性 ──"]
|
||||
has_data = False
|
||||
n_records = 0
|
||||
for label, injuries in (("主队", home_injuries), ("客队", away_injuries)):
|
||||
if injuries:
|
||||
has_data = True
|
||||
n_records += len(injuries)
|
||||
lines.append(f" {label}伤停({len(injuries)}人):")
|
||||
for inj in injuries[:8]: # 最多显示 8 条
|
||||
reason = inj.reason or inj.injury_type or "未知"
|
||||
@@ -255,10 +281,10 @@ async def injuries_slice(header: MatchHeader, *, before=None) -> str:
|
||||
else:
|
||||
lines.append(f" {label}: 无伤停数据")
|
||||
|
||||
if not has_data:
|
||||
return "── 阵容完整性 ──\n 无数据"
|
||||
if n_records == 0:
|
||||
return SliceResult(text="── 阵容完整性 ──\n 无数据", has_data=False, n_records=0)
|
||||
|
||||
return "\n".join(lines)
|
||||
return SliceResult(text="\n".join(lines), has_data=True, n_records=n_records)
|
||||
|
||||
|
||||
# ============================================================
|
||||
@@ -266,42 +292,38 @@ async def injuries_slice(header: MatchHeader, *, before=None) -> str:
|
||||
# ============================================================
|
||||
|
||||
async def build_context(match_id: int, *, form_last: int = 5, h2h_last: int = 5) -> MatchContext:
|
||||
"""单 agent 路径的完整上下文: 拼接全部切片(before=比赛时间,防未来信息)。"""
|
||||
"""单 agent 路径的完整上下文: 拼接全部切片(before=比赛时间,防未来信息)。
|
||||
|
||||
has_stats / has_injuries 直接取切片显式声明的 has_data,
|
||||
不再靠文案子串匹配(见审查报告 P2-1)。
|
||||
"""
|
||||
header = await load_match_header(match_id)
|
||||
parts = [header_text(header), ""]
|
||||
has_stats = False
|
||||
has_injuries = False
|
||||
|
||||
form_text = await form_slice(header, limit=form_last, before=header.match_dt)
|
||||
if "无数据" not in form_text:
|
||||
has_stats = True
|
||||
parts.append(form_text)
|
||||
form_res = await form_slice(header, limit=form_last, before=header.match_dt)
|
||||
parts.append(form_res.text)
|
||||
parts.append("")
|
||||
|
||||
h2h_text = await h2h_slice(header, limit=h2h_last, before=header.match_dt)
|
||||
parts.append(h2h_text)
|
||||
h2h_res = await h2h_slice(header, limit=h2h_last, before=header.match_dt)
|
||||
parts.append(h2h_res.text)
|
||||
parts.append("")
|
||||
|
||||
stats_text = await stats_slice(header, before=header.match_dt)
|
||||
if "无数据" not in stats_text:
|
||||
has_stats = True
|
||||
parts.append(stats_text)
|
||||
stats_res = await stats_slice(header, before=header.match_dt)
|
||||
parts.append(stats_res.text)
|
||||
parts.append("")
|
||||
|
||||
home_away_text = await home_away_slice(header, before=header.match_dt)
|
||||
parts.append(home_away_text)
|
||||
home_away_res = await home_away_slice(header, before=header.match_dt)
|
||||
parts.append(home_away_res.text)
|
||||
parts.append("")
|
||||
|
||||
injuries_text = await injuries_slice(header, before=header.match_dt)
|
||||
if "无数据" not in injuries_text:
|
||||
has_injuries = True
|
||||
parts.append(injuries_text)
|
||||
injuries_res = await injuries_slice(header, before=header.match_dt)
|
||||
parts.append(injuries_res.text)
|
||||
|
||||
return MatchContext(
|
||||
match_id=match_id,
|
||||
text="\n".join(parts),
|
||||
has_stats=has_stats,
|
||||
has_injuries=has_injuries,
|
||||
has_stats=form_res.has_data or stats_res.has_data,
|
||||
has_injuries=injuries_res.has_data,
|
||||
match_dt=header.match_dt,
|
||||
)
|
||||
|
||||
|
||||
+31
-9
@@ -27,12 +27,18 @@ _cache: dict[str, tuple[float, PredictResult]] = {}
|
||||
_cache_lock = Lock()
|
||||
|
||||
|
||||
def _cache_key(match_id: int, provider: str, model: str, version: str) -> str:
|
||||
return f"{match_id}:{provider}:{model}:{version}"
|
||||
def _cache_key(match_id: int, provider: str, model: str, version: str, tpl_hash: str) -> str:
|
||||
"""缓存键:含 prompt 模板内容 hash。
|
||||
|
||||
仅用 version 做键不够 —— 编辑器里改动 `match_prediction_v1.md` 而版本号
|
||||
不变时,进程内缓存仍会返回旧模板产生的旧结果(见审查报告 P2-6)。
|
||||
把模板内容 hash 纳入键,模板一改缓存自动失效。
|
||||
"""
|
||||
return f"{match_id}:{provider}:{model}:{version}:{tpl_hash[:12]}"
|
||||
|
||||
|
||||
def _get_cached(match_id: int, provider: str, model: str, version: str) -> PredictResult | None:
|
||||
key = _cache_key(match_id, provider, model, version)
|
||||
def _get_cached(match_id: int, provider: str, model: str, version: str, tpl_hash: str) -> PredictResult | None:
|
||||
key = _cache_key(match_id, provider, model, version, tpl_hash)
|
||||
with _cache_lock:
|
||||
if key in _cache:
|
||||
ts, result = _cache[key]
|
||||
@@ -42,12 +48,22 @@ def _get_cached(match_id: int, provider: str, model: str, version: str) -> Predi
|
||||
return None
|
||||
|
||||
|
||||
def _set_cached(match_id: int, provider: str, model: str, version: str, result: PredictResult) -> None:
|
||||
key = _cache_key(match_id, provider, model, version)
|
||||
def _set_cached(match_id: int, provider: str, model: str, version: str, tpl_hash: str, result: PredictResult) -> None:
|
||||
key = _cache_key(match_id, provider, model, version, tpl_hash)
|
||||
with _cache_lock:
|
||||
_cache[key] = (time.time(), result)
|
||||
|
||||
|
||||
def clear_prompt_cache() -> None:
|
||||
"""清空 prompt 模板缓存(供开发/热更新时手动调用)。
|
||||
|
||||
lru_cache 的模板缓存是进程级的,改完 .md 需要重启进程才能生效;
|
||||
提供显式清理入口,避免"改了模板却看不到变化"的困惑(见审查报告 P2-5)。
|
||||
"""
|
||||
_load_prompt_template.cache_clear()
|
||||
logger.info("prompt 模板缓存已清空")
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=8)
|
||||
def _load_prompt_template(version: str = "v1") -> str:
|
||||
"""缓存 prompt 模板(进程生命周期内每个版本只读一次)。"""
|
||||
@@ -58,6 +74,11 @@ def _load_prompt_template(version: str = "v1") -> str:
|
||||
return f.read()
|
||||
|
||||
|
||||
def _prompt_template_hash(version: str) -> str:
|
||||
"""prompt 模板内容 hash(用于缓存键,模板变更即失效)。"""
|
||||
return hashlib.sha256(_load_prompt_template(version).encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
@dataclass
|
||||
class PredictResult:
|
||||
prediction_id: int
|
||||
@@ -117,10 +138,11 @@ async def _predict_single(
|
||||
if model:
|
||||
provider.model = model
|
||||
version = prompt_version or "v1"
|
||||
tpl_hash = _prompt_template_hash(version)
|
||||
|
||||
# 0. 查缓存(同 match+provider+model+version 5 分钟内直接返)
|
||||
# 0. 查缓存(同 match+provider+model+version+模板hash 5 分钟内直接返)
|
||||
if use_cache:
|
||||
cached = _get_cached(match_id, settings.LLM_PROVIDER, provider.model, version)
|
||||
cached = _get_cached(match_id, settings.LLM_PROVIDER, provider.model, version, tpl_hash)
|
||||
if cached is not None:
|
||||
logger.debug("predict cache hit match=%s", match_id)
|
||||
return cached
|
||||
@@ -206,5 +228,5 @@ async def _predict_single(
|
||||
|
||||
# 5. 写入缓存(仅当允许缓存时)
|
||||
if use_cache:
|
||||
_set_cached(match_id, settings.LLM_PROVIDER, provider.model, version, result)
|
||||
_set_cached(match_id, settings.LLM_PROVIDER, provider.model, version, tpl_hash, result)
|
||||
return result
|
||||
|
||||
+69
-1
@@ -10,6 +10,9 @@ from pydantic import BaseModel, Field, field_validator, model_validator
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 已知的 5 个专家 agent 名(与 orchestrator.SPECIALIST_SPECS 保持一致)
|
||||
KNOWN_AGENT_NAMES: tuple[str, ...] = ("form", "stats", "home_away", "injuries", "h2h")
|
||||
|
||||
|
||||
class AgentReportSchema(BaseModel):
|
||||
"""单个专家 Agent 输出的校验 schema。"""
|
||||
@@ -64,9 +67,17 @@ class PredictionOutputSchema(BaseModel):
|
||||
|
||||
@model_validator(mode="after")
|
||||
def check_consistency(self) -> "PredictionOutputSchema":
|
||||
"""验证比分与胜平负一致,不一致则自动修正。"""
|
||||
"""验证比分与胜平负一致。
|
||||
|
||||
不一致时以比分为准修正 pred_1x2(比分是更结构化的输出),
|
||||
但**必须告警** —— 静默修正会掩盖 LLM 的自相矛盾,让问题无法被发现。
|
||||
"""
|
||||
expected = _score_to_1x2(self.pred_home_goals, self.pred_away_goals)
|
||||
if self.pred_1x2 != expected:
|
||||
logger.warning(
|
||||
"1x2 与比分不一致: 比分 %.1f-%.1f 推出 '%s',但 LLM 给出 '%s';以比分修正",
|
||||
self.pred_home_goals, self.pred_away_goals, expected, self.pred_1x2,
|
||||
)
|
||||
self.pred_1x2 = expected
|
||||
return self
|
||||
|
||||
@@ -80,6 +91,63 @@ def _score_to_1x2(home: float, away: float) -> str:
|
||||
return "X"
|
||||
|
||||
|
||||
class AgentWeightsSchema(BaseModel):
|
||||
"""终裁给出的各专家权重校验 schema。
|
||||
|
||||
权重含义:各专家报告在最终决策中的相对影响力。约束:
|
||||
- key 必须是已知的 5 个专家名
|
||||
- value ∈ [0, 1]
|
||||
- 总和允许有 0.05 的浮点误差(LLM 常凑不到精确 1.0),
|
||||
超出则归一化到 1.0 而不是直接拒收
|
||||
"""
|
||||
weights: dict[str, float] = Field(default_factory=dict)
|
||||
|
||||
@field_validator("weights", mode="before")
|
||||
@classmethod
|
||||
def coerce_weights(cls, v):
|
||||
if v is None:
|
||||
return {}
|
||||
if not isinstance(v, dict):
|
||||
raise ValueError(f"agent_weights 必须是 dict,得到 {type(v).__name__}")
|
||||
out: dict[str, float] = {}
|
||||
for k, raw in v.items():
|
||||
key = str(k).strip().lower()
|
||||
if key not in KNOWN_AGENT_NAMES:
|
||||
logger.warning("agent_weights 含未知专家 '%s',已忽略", k)
|
||||
continue
|
||||
f = _safe_float(raw)
|
||||
if f is None:
|
||||
logger.warning("agent_weights['%s']=%r 非数值,已忽略", k, raw)
|
||||
continue
|
||||
# 负数直接钳到 0;超过 1 的钳到 1
|
||||
out[key] = min(max(f, 0.0), 1.0)
|
||||
return out
|
||||
|
||||
@model_validator(mode="after")
|
||||
def normalize_sum(self) -> "AgentWeightsSchema":
|
||||
"""权重和不为 1 时归一化(而非拒收),并在偏离较大时告警。"""
|
||||
if not self.weights:
|
||||
return self
|
||||
total = sum(self.weights.values())
|
||||
if total <= 0:
|
||||
return self
|
||||
if abs(total - 1.0) > 0.05:
|
||||
logger.warning("agent_weights 总和为 %.3f,已归一化到 1.0", total)
|
||||
self.weights = {k: v / total for k, v in self.weights.items()}
|
||||
return self
|
||||
|
||||
|
||||
def validate_agent_weights(raw) -> dict[str, float]:
|
||||
"""校验并规范化终裁给出的 agent_weights。非法输入返回空 dict。"""
|
||||
if raw is None:
|
||||
return {}
|
||||
try:
|
||||
return AgentWeightsSchema(weights=raw).weights
|
||||
except Exception as e:
|
||||
logger.warning("agent_weights 校验失败,丢弃: %s", e)
|
||||
return {}
|
||||
|
||||
|
||||
def validate_agent_output(raw: dict) -> AgentReportSchema:
|
||||
"""校验并规范化单个 Agent 输出。"""
|
||||
return AgentReportSchema(
|
||||
|
||||
Reference in New Issue
Block a user