refactor(report): unify campaign report loading behind one read model
「活动报告取数三件套」此前在报告/markdown/分析/对比等 7 处手写重复, 唯一深化产物 build_campaign_report_dict 被锁在周期对比私有角落。 升位为 report.py 的 load_campaign_report(session, campaign) 单一出口 (探索线 summarize_campaign_exploration 同口径),并把 8 处 scenario_names 推导式收敛为 ScenarioRepository.name_map() 窄方法。 纯结构重排、零行为变更,572 项测试全绿。
This commit is contained in:
parent
a665b496b0
commit
f8d8450b1e
@ -13,7 +13,7 @@ from typing import Any, Awaitable, Callable, Optional
|
||||
|
||||
from sqlmodel import Session
|
||||
|
||||
from agenteval.evaluation.report import generate_campaign_report
|
||||
from agenteval.evaluation.report import load_campaign_report
|
||||
from agenteval.exploration.summary import summarize_campaign_exploration
|
||||
from agenteval.model_gateway import ModelGateway
|
||||
from agenteval.models import Campaign, ModelCapability, RunStatus
|
||||
@ -28,7 +28,6 @@ from agenteval.storage.repository import (
|
||||
CampaignAnalysisRepository,
|
||||
CampaignRepository,
|
||||
RunRepository,
|
||||
ScenarioRepository,
|
||||
)
|
||||
from agenteval.utils.llm import extract_reply_text, parse_json_from_llm_text
|
||||
|
||||
@ -306,8 +305,7 @@ async def execute_campaign_analysis(
|
||||
try:
|
||||
client = chat_client or gateway_chat_client(runtime)
|
||||
runs = RunRepository(session).list_by_campaign(campaign_id)
|
||||
scenario_names = {s.id: s.name for s in ScenarioRepository(session).list_all()}
|
||||
report = generate_campaign_report(campaign, runs, scenario_names=scenario_names)
|
||||
report = load_campaign_report(session, campaign)
|
||||
result = await analyze_campaign(
|
||||
campaign=campaign,
|
||||
report=report,
|
||||
|
||||
@ -15,15 +15,13 @@ from typing import Any, Optional
|
||||
from sqlmodel import Session
|
||||
|
||||
from agenteval.evaluation.analysis import ChatClient, gateway_chat_client, resolve_analysis_model
|
||||
from agenteval.evaluation.report import generate_campaign_report
|
||||
from agenteval.models import Campaign, EvalRun
|
||||
from agenteval.evaluation.report import load_campaign_report
|
||||
from agenteval.models import Campaign
|
||||
from agenteval.storage.db import get_session, iso_utc, utc_now
|
||||
from agenteval.storage.repository import (
|
||||
CampaignAnalysisRepository,
|
||||
CampaignPeriodComparisonRepository,
|
||||
CampaignRepository,
|
||||
RunRepository,
|
||||
ScenarioRepository,
|
||||
)
|
||||
from agenteval.utils.llm import parse_json_from_llm_text
|
||||
|
||||
@ -154,21 +152,14 @@ def compute_metric_diff(
|
||||
return {"overall": overall, "scenarios": scenarios}
|
||||
|
||||
|
||||
def build_campaign_report_dict(campaign: Campaign, session: Session) -> dict[str, Any]:
|
||||
"""为周期对比现算一期活动的报告 dict(读路径,不重算聚合口径)。"""
|
||||
runs: list[EvalRun] = RunRepository(session).list_by_campaign(campaign.id)
|
||||
scenario_names = {s.id: s.name for s in ScenarioRepository(session).list_all()}
|
||||
return generate_campaign_report(campaign, runs, scenario_names=scenario_names)
|
||||
|
||||
|
||||
def build_comparison_payload(campaign: Campaign, session: Session) -> dict[str, Any]:
|
||||
"""GET 返回体:自动基线信息 + 机械 diff(无基线时两者均为 null)。"""
|
||||
baseline = resolve_auto_baseline(campaign, session)
|
||||
if baseline is None:
|
||||
return {"auto_baseline": None, "metric_diff": None}
|
||||
diff = compute_metric_diff(
|
||||
build_campaign_report_dict(baseline, session),
|
||||
build_campaign_report_dict(campaign, session),
|
||||
load_campaign_report(session, baseline),
|
||||
load_campaign_report(session, campaign),
|
||||
)
|
||||
return {
|
||||
"auto_baseline": {
|
||||
@ -358,8 +349,8 @@ async def execute_campaign_comparison(
|
||||
)
|
||||
try:
|
||||
diff = compute_metric_diff(
|
||||
build_campaign_report_dict(baseline, session),
|
||||
build_campaign_report_dict(campaign, session),
|
||||
load_campaign_report(session, baseline),
|
||||
load_campaign_report(session, campaign),
|
||||
)
|
||||
valid_scenario_ids = {s["scenario_id"] for s in diff["scenarios"]}
|
||||
result = await narrate_period_comparison(
|
||||
|
||||
@ -9,6 +9,8 @@ from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
from sqlmodel import Session
|
||||
|
||||
from agenteval.evaluation.case_verdict import CaseEvidence, resolve_case_verdicts
|
||||
from agenteval.evaluation.metrics import aggregate_runs
|
||||
from agenteval.evaluation.report_render import render_html, render_json, render_markdown
|
||||
@ -359,6 +361,12 @@ def generate_campaign_report(
|
||||
}
|
||||
|
||||
|
||||
def load_campaign_report(session: Session, campaign: Campaign) -> dict[str, Any]:
|
||||
"""取数 + 聚合一步完成:报告 / 分析 / 对比 / 导出共用的活动报告 dict 取法。"""
|
||||
runs = RunRepository(session).list_by_campaign(campaign.id)
|
||||
return generate_campaign_report(campaign, runs, scenario_names=ScenarioRepository(session).name_map())
|
||||
|
||||
|
||||
def save_report(run_id: str, fmt: str = "html", output_dir: Optional[Path] = None) -> Path:
|
||||
"""Generate a run report and save it to disk in the requested format."""
|
||||
output_dir = output_dir or DATA_DIR / "reports"
|
||||
|
||||
@ -247,6 +247,10 @@ class ScenarioRepository(BaseRepository[Scenario, ScenarioDB]):
|
||||
raise
|
||||
return True
|
||||
|
||||
def name_map(self) -> dict[str, str]:
|
||||
"""scenario_id → 名称映射:报告 / 时间线 / 列表等读路径共用的场景名取法。"""
|
||||
return {sid: name for sid, name in self.session.exec(select(ScenarioDB.id, ScenarioDB.name)).all()}
|
||||
|
||||
|
||||
class RunRepository(BaseRepository[EvalRun, EvalRunDB]):
|
||||
"""Repository for evaluation runs."""
|
||||
|
||||
@ -14,7 +14,6 @@ from sqlmodel import Session
|
||||
from agenteval.evaluation.analysis import resolve_analysis_model, start_campaign_analysis
|
||||
from agenteval.evaluation.campaign_runner import campaign_progress, request_cancel, start_campaign
|
||||
from agenteval.evaluation.comparison import (
|
||||
build_campaign_report_dict,
|
||||
build_comparison_payload,
|
||||
compute_metric_diff,
|
||||
resolve_auto_baseline,
|
||||
@ -22,7 +21,7 @@ from agenteval.evaluation.comparison import (
|
||||
)
|
||||
from agenteval.evaluation.report import (
|
||||
build_campaign_timeline,
|
||||
generate_campaign_report,
|
||||
load_campaign_report,
|
||||
summarize_campaign_progress,
|
||||
)
|
||||
from agenteval.evaluation.report_render import render_campaign_markdown
|
||||
@ -131,9 +130,7 @@ async def get_campaign_report(campaign_id: str, session: Session = Depends(get_d
|
||||
campaign = CampaignRepository(session).get(campaign_id)
|
||||
if not campaign:
|
||||
raise HTTPException(status_code=404, detail="campaign not found")
|
||||
runs = RunRepository(session).list_by_campaign(campaign_id)
|
||||
scenario_names = {s.id: s.name for s in ScenarioRepository(session).list_all()}
|
||||
report = generate_campaign_report(campaign, runs, scenario_names=scenario_names)
|
||||
report = load_campaign_report(session, campaign)
|
||||
exploration = summarize_campaign_exploration(session, campaign_id)
|
||||
if exploration is not None:
|
||||
report["exploration"] = exploration
|
||||
@ -145,8 +142,8 @@ async def get_campaign_report_markdown(campaign_id: str, session: Session = Depe
|
||||
campaign = CampaignRepository(session).get(campaign_id)
|
||||
if not campaign:
|
||||
raise HTTPException(status_code=404, detail="campaign not found")
|
||||
runs = RunRepository(session).list_by_campaign(campaign_id)
|
||||
scenario_names = {s.id: s.name for s in ScenarioRepository(session).list_all()}
|
||||
report = load_campaign_report(session, campaign)
|
||||
scenario_names = ScenarioRepository(session).name_map()
|
||||
analysis_row = CampaignAnalysisRepository(session).get_by_campaign(campaign_id)
|
||||
analysis = analysis_row.get_result() if analysis_row and analysis_row.status == "completed" else None
|
||||
target = TargetRepository(session).get(campaign.target_id)
|
||||
@ -158,8 +155,8 @@ async def get_campaign_report_markdown(campaign_id: str, session: Session = Depe
|
||||
baseline = CampaignRepository(session).get(cmp_row.baseline_campaign_id)
|
||||
metric_diff = (
|
||||
compute_metric_diff(
|
||||
build_campaign_report_dict(baseline, session),
|
||||
build_campaign_report_dict(campaign, session),
|
||||
load_campaign_report(session, baseline),
|
||||
load_campaign_report(session, campaign),
|
||||
)
|
||||
if baseline is not None
|
||||
else None
|
||||
@ -178,7 +175,7 @@ async def get_campaign_report_markdown(campaign_id: str, session: Session = Depe
|
||||
}
|
||||
|
||||
md = render_campaign_markdown(
|
||||
generate_campaign_report(campaign, runs, scenario_names=scenario_names),
|
||||
report,
|
||||
analysis=analysis,
|
||||
comparison=comparison,
|
||||
exploration=summarize_campaign_exploration(session, campaign_id),
|
||||
@ -198,7 +195,7 @@ async def get_campaign_timeline(campaign_id: str, session: Session = Depends(get
|
||||
if not campaign:
|
||||
raise HTTPException(status_code=404, detail="campaign not found")
|
||||
runs = RunRepository(session).list_by_campaign(campaign_id)
|
||||
scenario_names = {s.id: s.name for s in ScenarioRepository(session).list_all()}
|
||||
scenario_names = ScenarioRepository(session).name_map()
|
||||
return {"entries": build_campaign_timeline(campaign, runs, scenario_names=scenario_names)}
|
||||
|
||||
|
||||
@ -254,8 +251,8 @@ async def get_campaign_comparison(campaign_id: str, session: Session = Depends(g
|
||||
baseline = CampaignRepository(session).get(row.baseline_campaign_id)
|
||||
metric_diff = (
|
||||
compute_metric_diff(
|
||||
build_campaign_report_dict(baseline, session),
|
||||
build_campaign_report_dict(campaign, session),
|
||||
load_campaign_report(session, baseline),
|
||||
load_campaign_report(session, campaign),
|
||||
)
|
||||
if baseline is not None
|
||||
else None
|
||||
|
||||
@ -131,7 +131,7 @@ async def patrol(session: Session = Depends(get_db)) -> dict:
|
||||
campaign_repo = CampaignRepository(session)
|
||||
run_repo = RunRepository(session)
|
||||
exploration_repo = ExplorationSessionRepository(session)
|
||||
scenario_names = {s.id: s.name for s in ScenarioRepository(session).list_all()}
|
||||
scenario_names = ScenarioRepository(session).name_map()
|
||||
target_names = {t.id: t.name for t in TargetRepository(session).list_all()}
|
||||
|
||||
patrolled_at = utc_now()
|
||||
|
||||
@ -68,7 +68,7 @@ async def _run_evaluation(
|
||||
|
||||
@router.get("")
|
||||
async def list_runs(session: Session = Depends(get_db)) -> list[dict]:
|
||||
scenario_names = {s.id: s.name for s in ScenarioRepository(session).list_all()}
|
||||
scenario_names = ScenarioRepository(session).name_map()
|
||||
target_names = {t.id: t.name for t in TargetRepository(session).list_all()}
|
||||
return [
|
||||
{
|
||||
|
||||
@ -36,11 +36,10 @@ def _settled(runs: list[EvalRun]) -> list[EvalRun]:
|
||||
@router.get("/dashboard")
|
||||
def dashboard(session: Session = Depends(get_db)) -> dict:
|
||||
targets = TargetRepository(session).list_all()
|
||||
scenarios = ScenarioRepository(session).list_all()
|
||||
scenario_names = ScenarioRepository(session).name_map()
|
||||
runs = RunRepository(session).list_all()
|
||||
model_configs = ModelConfigRepository(session).list_all()
|
||||
|
||||
scenario_names = {s.id: s.name for s in scenarios}
|
||||
target_names = {t.id: t.name for t in targets}
|
||||
|
||||
settled_runs = _settled(runs)
|
||||
@ -82,7 +81,7 @@ def dashboard(session: Session = Depends(get_db)) -> dict:
|
||||
|
||||
return {
|
||||
"targets_count": len(targets),
|
||||
"scenarios_count": len(scenarios),
|
||||
"scenarios_count": len(scenario_names),
|
||||
"runs_count": len(runs),
|
||||
"model_configs_count": len(model_configs),
|
||||
"today_runs": today_runs,
|
||||
|
||||
@ -178,6 +178,7 @@ async def test_full_lifecycle_create_message_close(seeded_db, mock_channel, clie
|
||||
async def test_dict_reply_content_is_flattened_to_text(seeded_db, monkeypatch, client):
|
||||
"""通道回复 content 为对象(如 tutu msgBody)时应提取文本而非存 str(dict)。"""
|
||||
from agenteval.channels.base import Reply
|
||||
|
||||
from tests.unit.mock_channel import MockChannel
|
||||
|
||||
class _DictReplyChannel(MockChannel):
|
||||
|
||||
96
tests/unit/test_campaign_report_loader.py
Normal file
96
tests/unit/test_campaign_report_loader.py
Normal file
@ -0,0 +1,96 @@
|
||||
"""活动报告读模型 load_campaign_report 与 ScenarioRepository.name_map 直测。
|
||||
|
||||
读模型是报告 / 分析 / 对比 / 导出共用的「取数 + 聚合」单一出口:
|
||||
返回结果必须与手写三件套(list_by_campaign + scenario_names + generate_campaign_report)完全一致。
|
||||
"""
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
import pytest
|
||||
from agenteval.evaluation.report import generate_campaign_report, load_campaign_report
|
||||
from agenteval.models import Campaign, CampaignPlanEntry, CampaignStatus, Case, CaseType, EvalRun, RunStatus, Scenario
|
||||
from agenteval.storage.repository import CampaignRepository, RunRepository, ScenarioRepository
|
||||
from sqlmodel import Session, SQLModel, create_engine
|
||||
|
||||
T0 = datetime(2026, 1, 1, 0, 0, 0, tzinfo=timezone.utc)
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def db_session(tmp_path):
|
||||
from agenteval.storage.db import ( # noqa: F401
|
||||
CampaignDB,
|
||||
EvalResultDB,
|
||||
EvalRunDB,
|
||||
EvalTargetDB,
|
||||
FileCategoryDB,
|
||||
FileRecordDB,
|
||||
ScenarioDB,
|
||||
TurnDB,
|
||||
)
|
||||
engine = create_engine(
|
||||
f"sqlite:///{tmp_path / 'loader.db'}",
|
||||
connect_args={"check_same_thread": False},
|
||||
)
|
||||
SQLModel.metadata.create_all(engine)
|
||||
session = Session(engine)
|
||||
try:
|
||||
yield session
|
||||
finally:
|
||||
session.close()
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def _seed(session) -> Campaign:
|
||||
ScenarioRepository(session).create(Scenario(
|
||||
id="s-a", name="夜间问诊",
|
||||
cases=[Case(id="c1", type=CaseType.SINGLE, messages=["hi"])],
|
||||
))
|
||||
campaign = CampaignRepository(session).create(Campaign(
|
||||
name="cycle", target_id="t-1", window_seconds=12, time_scale=1.0,
|
||||
plan=[CampaignPlanEntry(scenario_id="s-a", offset_seconds=0, count=1)],
|
||||
))
|
||||
campaign.status = CampaignStatus.RUNNING
|
||||
campaign.started_at = T0
|
||||
campaign = CampaignRepository(session).update(campaign)
|
||||
RunRepository(session).create(EvalRun(
|
||||
target_id="t-1", scenario_id="s-a", campaign_id=campaign.id,
|
||||
status=RunStatus.COMPLETED, started_at=T0 + timedelta(seconds=1),
|
||||
summary={"total_cases": 1, "passed_cases": 1, "pass_rate": 1.0, "avg_latency_ms": 100},
|
||||
))
|
||||
return campaign
|
||||
|
||||
|
||||
def test_load_campaign_report_matches_manual_trio(db_session):
|
||||
campaign = _seed(db_session)
|
||||
report = load_campaign_report(db_session, campaign)
|
||||
|
||||
runs = RunRepository(db_session).list_by_campaign(campaign.id)
|
||||
scenario_names = {s.id: s.name for s in ScenarioRepository(db_session).list_all()}
|
||||
expected = generate_campaign_report(campaign, runs, scenario_names=scenario_names)
|
||||
|
||||
assert report == expected
|
||||
cap = next(c for c in report["capability_summary"] if c["scenario_id"] == "s-a")
|
||||
assert cap["scenario_name"] == "夜间问诊"
|
||||
|
||||
|
||||
def test_load_campaign_report_empty_campaign(db_session):
|
||||
campaign = CampaignRepository(db_session).create(Campaign(
|
||||
name="empty", target_id="t-1", window_seconds=12,
|
||||
plan=[CampaignPlanEntry(scenario_id="s-x", offset_seconds=0, count=1)],
|
||||
))
|
||||
report = load_campaign_report(db_session, campaign)
|
||||
assert report["summary"]["total_runs"] == 0
|
||||
assert report["capability_summary"] == []
|
||||
|
||||
|
||||
def test_scenario_name_map(db_session):
|
||||
_seed(db_session)
|
||||
ScenarioRepository(db_session).create(Scenario(
|
||||
id="s-b", name="缴费引导",
|
||||
cases=[Case(id="c2", type=CaseType.SINGLE, messages=["缴费"])],
|
||||
))
|
||||
assert ScenarioRepository(db_session).name_map() == {"s-a": "夜间问诊", "s-b": "缴费引导"}
|
||||
|
||||
|
||||
def test_scenario_name_map_empty(db_session):
|
||||
assert ScenarioRepository(db_session).name_map() == {}
|
||||
Loading…
Reference in New Issue
Block a user