fix(critical): 修复评审确认的 R1-R5 五处缺陷
R1 429 key 轮换路径调用不存在的 _km → NameError(且位于凭证脱敏日志行): 改用 src/data/key_ring._mask,全树不再有 _km 调用。 [注: 该改动已随另一工作流提交9ccda4b一并入库] R2 ingest_bzzoiro_standings 被截断(return 前无 upsert 逻辑,standings 表 永不写入、total_upserted 恒为 0):从f05dc1a移植完整实现,按 (league_id, season, team_id) upsert,保留逐联赛错误隔离,不加 db.commit()。 [注: 同上,已随9ccda4b入库] R3 src/llm/agents/orchestrator.py 完成日志把 list 喂给 %d → logging TypeError: 改为 len(ok_reports),与同文件降级日志行写法一致。 R4 mode=multi 静默丢弃调用方传入的 model:predict_match_multi 新增 model 关键字参数,经 _ACTIVE_MODEL_OVERRIDE 下传至 _agent_provider, 显式 model 优先级最高;override 生效时跳过 provider 缓存读写以免串味, 并在 predict.py 派发点透传。 R5 回测把字符串日期直接与 timestamptz 列比较:新增 _parse_date_bound 助手, 支持 YYYY-MM-DD / 完整 ISO / datetime / None,裸日期按 UTC 锚定, 结束日取当天末刻(闭区间,避免最后一天被静默排除),非法输入抛 ValueError。 新增 tests/test_review_required_fixes.py 覆盖 R1-R5(R2/R4 为行为测试), 20 项全通过。
This commit is contained in:
@@ -34,6 +34,11 @@ logger = logging.getLogger(__name__)
|
||||
_AGENT_PROVIDER_CACHE: dict[str, tuple[float, LLMProvider]] = {}
|
||||
_AGENT_PROVIDER_CACHE_TTL = 60.0
|
||||
|
||||
# 当前预测的模型覆盖(调用方显式指定)。用模块级变量而非新增形参,因为
|
||||
# run_specialists 会被既有测试 mock,加形参会破坏那些测试的调用签名。
|
||||
# 由 predict_match_multi 在进入时 set、退出时 reset。
|
||||
_ACTIVE_MODEL_OVERRIDE: str | None = None
|
||||
|
||||
|
||||
# ── 5 个专家 agent 定义 ──
|
||||
# A=近期状态 B=攻防数据 C=主客因素 D=联赛排名 E=历史交锋
|
||||
@@ -105,21 +110,24 @@ class MultiPredictResult:
|
||||
raw: dict | None = None
|
||||
|
||||
|
||||
async def _agent_provider(agent_id: str, *, tier: str) -> LLMProvider:
|
||||
async def _agent_provider(agent_id: str, *, tier: str, model_override: str | None = None) -> LLMProvider:
|
||||
"""构造某 agent 专属 provider。
|
||||
|
||||
覆盖优先级:
|
||||
模型: AGENT_MODEL_{ID}(运行时) → 层级默认(LLM_SPECIALIST/AGGREGATOR_MODEL) → 全局 LLM_MODEL
|
||||
模型: 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}"
|
||||
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
|
||||
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()
|
||||
@@ -129,6 +137,9 @@ async def _agent_provider(agent_id: str, *, tier: str) -> LLMProvider:
|
||||
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
|
||||
@@ -136,10 +147,11 @@ async def _agent_provider(agent_id: str, *, tier: str) -> LLMProvider:
|
||||
if key:
|
||||
p.api_key = key
|
||||
|
||||
_AGENT_PROVIDER_CACHE[cache_key] = (time.time(), p)
|
||||
# 简单淘汰:超过 20 条时清空(60s TTL 下不会累积太多)
|
||||
if len(_AGENT_PROVIDER_CACHE) > 20:
|
||||
_AGENT_PROVIDER_CACHE.clear()
|
||||
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
|
||||
|
||||
|
||||
@@ -152,9 +164,14 @@ async def run_specialists(
|
||||
"""并行执行 5 个专家 agent。fail-open: 单个失败不影响其他。
|
||||
|
||||
before: 数据截止时间(回测防泄漏)。None 表示不限制。
|
||||
|
||||
模型覆盖通过 _ACTIVE_MODEL_OVERRIDE 传递,而不是新增形参:既有测试
|
||||
会 mock 本函数(见 tests/test_agent_weights_persist.py),加形参会破坏
|
||||
它们的调用签名。predict_match_multi 在调用前后 set/reset 该变量。
|
||||
"""
|
||||
model_override = _ACTIVE_MODEL_OVERRIDE
|
||||
tasks = [
|
||||
_run_one(spec, header, await _agent_provider(spec.name, tier="specialist"), version=version, before=before)
|
||||
_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)
|
||||
@@ -222,11 +239,13 @@ async def predict_match_multi(
|
||||
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()
|
||||
|
||||
@@ -247,7 +266,14 @@ async def predict_match_multi(
|
||||
prediction_cutoff_at = cutoff
|
||||
|
||||
# 2. 并行专家(各自独立配置,使用统一 cutoff)
|
||||
reports = await run_specialists(header, version=version, before=cutoff)
|
||||
# 显式 model 覆盖通过模块级变量下传,避免改动 run_specialists 的签名
|
||||
global _ACTIVE_MODEL_OVERRIDE
|
||||
_prev_override = _ACTIVE_MODEL_OVERRIDE
|
||||
_ACTIVE_MODEL_OVERRIDE = model
|
||||
try:
|
||||
reports = await run_specialists(header, version=version, before=cutoff)
|
||||
finally:
|
||||
_ACTIVE_MODEL_OVERRIDE = _prev_override
|
||||
|
||||
# 2.5 统计有效专家报告数量
|
||||
ok_reports = [r for r in reports if r.status == "ok"]
|
||||
@@ -279,7 +305,7 @@ async def predict_match_multi(
|
||||
}
|
||||
agg_prompt_tokens = 0
|
||||
agg_completion_tokens = 0
|
||||
aggregator_model = settings.LLM_MODEL # 占位,无实际 LLM 调用
|
||||
aggregator_model = model or settings.LLM_MODEL # 占位,无实际 LLM 调用
|
||||
|
||||
latency_ms = int((time.perf_counter() - start) * 1000)
|
||||
|
||||
@@ -347,7 +373,7 @@ async def predict_match_multi(
|
||||
"预测完成 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, ok_reports, len(reports), pred.id,
|
||||
latency_ms, len(ok_reports), len(reports), pred.id,
|
||||
)
|
||||
|
||||
return MultiPredictResult(
|
||||
|
||||
+36
-5
@@ -10,7 +10,7 @@ from __future__ import annotations
|
||||
import asyncio
|
||||
import logging
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import selectinload
|
||||
@@ -79,6 +79,35 @@ class BacktestSummary:
|
||||
results: list[BacktestMatchResult] = field(default_factory=list)
|
||||
|
||||
|
||||
def _parse_date_bound(value, *, end_of_day: bool) -> datetime | None:
|
||||
"""把日期入参解析成可与 timestamptz 列比较的 aware datetime。
|
||||
|
||||
支持 "YYYY-MM-DD"、完整 ISO 串(可带偏移)以及 datetime 对象;None 原样返回。
|
||||
裸日期按 UTC 锚定 —— Match.match_date 是 timestamptz,naive datetime 与之
|
||||
比较会因时区不同而偏移;start 取当天 00:00,end 取当天 23:59:59.999999
|
||||
(闭区间,否则最后一天会被静默排除)。
|
||||
|
||||
解析失败抛 ValueError(不静默吞掉):fromisoformat 对非法输入统一抛 ValueError,
|
||||
这里包一层以带上原始值,便于定位是哪个参数写错了。
|
||||
"""
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, datetime):
|
||||
dt = value
|
||||
else:
|
||||
try:
|
||||
dt = datetime.fromisoformat(str(value))
|
||||
except ValueError as e:
|
||||
raise ValueError(f"无法解析日期: {value!r}(应为 YYYY-MM-DD 或 ISO 格式)") from e
|
||||
|
||||
if dt.tzinfo is None:
|
||||
dt = dt.replace(tzinfo=timezone.utc)
|
||||
# 闭区间上界:日期串解析出来是 00:00,取当天末刻才能让最后一天参与回测
|
||||
if end_of_day:
|
||||
dt = dt.replace(hour=23, minute=59, second=59, microsecond=999999)
|
||||
return dt
|
||||
|
||||
|
||||
async def _get_historical_matches(
|
||||
db,
|
||||
*,
|
||||
@@ -106,10 +135,12 @@ async def _get_historical_matches(
|
||||
)
|
||||
if league_id is not None:
|
||||
stmt = stmt.where(Match.league_id == league_id)
|
||||
if date_from:
|
||||
stmt = stmt.where(Match.match_date >= date_from)
|
||||
if date_to:
|
||||
stmt = stmt.where(Match.match_date <= date_to)
|
||||
dt_from = _parse_date_bound(date_from, end_of_day=False)
|
||||
if dt_from is not None:
|
||||
stmt = stmt.where(Match.match_date >= dt_from)
|
||||
dt_to = _parse_date_bound(date_to, end_of_day=True)
|
||||
if dt_to is not None:
|
||||
stmt = stmt.where(Match.match_date <= dt_to)
|
||||
|
||||
stmt = stmt.order_by(Match.match_date.desc()).limit(limit)
|
||||
result = await db.execute(stmt)
|
||||
|
||||
+2
-1
@@ -187,13 +187,14 @@ async def predict_match(
|
||||
)
|
||||
from src.llm.agents.orchestrator import predict_match_multi
|
||||
|
||||
# 回测参数完整传递到 multi-agent 路径
|
||||
# 回测参数 + 模型覆盖完整传递到 multi-agent 路径
|
||||
return await predict_match_multi(
|
||||
match_id,
|
||||
provider=provider,
|
||||
version=(prompt_version or "v1").removeprefix("multi_"),
|
||||
backtest=backtest,
|
||||
cutoff_at=cutoff_at,
|
||||
model=model,
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user