AgentEvalTool/tests/unit/test_comparison_read_model.py
sinohqb 2b6cab6cb2 feat(comparison): unify read model and validation for period comparison
周期对比读模型升位为单一出口(load_comparison_view),GET/POST/markdown
三处调用点统一走同一 view 投影,消除「取数三件套」重复。校验逻辑收敛到
validate_comparison_request,router 捕获映射 400,执行器捕获落 failed 行,
校验顺序权威不再漂移。

- 新增 load_comparison_view:无行返回 status=none + auto_baseline,有行
  返回完整 comparison dict(含 model_name 标签)+ metric_diff
- 新增 validate_comparison_request:活动终态 → 模型 → 基线 → 分析,违
  规抛 ComparisonError
- execute_campaign_comparison 内联校验替换为 validate_comparison_request
  调用,catch ComparisonError 落 failed 行
- router 三处迁移:GET /comparison、POST /comparison、markdown 导出
- 删除 build_comparison_payload(已吸收进 load_comparison_view)
- 8 个新测试覆盖读模型三态 + 校验五错
2026-08-04 10:49:14 +08:00

179 lines
5.9 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""周期对比读模型与校验合一直测(架构保养第二轮候选 3
load_comparison_view 单一出口返回 GET /comparison 完整响应形状;
validate_comparison_request 共享校验入口抛 ComparisonErrorrouter
映射 400、执行器落 failed 行。
"""
from datetime import datetime, timezone
from uuid import uuid4
import pytest
from agenteval.evaluation.comparison import (
ComparisonError,
load_comparison_view,
validate_comparison_request,
)
from agenteval.models import (
Campaign,
CampaignPlanEntry,
CampaignStatus,
EvalTarget,
Scenario,
Case,
CaseType,
)
from agenteval.storage.repository import (
CampaignAnalysisRepository,
CampaignPeriodComparisonRepository,
CampaignRepository,
ScenarioRepository,
TargetRepository,
)
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
CampaignAnalysisDB,
CampaignDB,
CampaignPeriodComparisonDB,
EvalRunDB,
EvalTargetDB,
ScenarioDB,
)
engine = create_engine(
f"sqlite:///{tmp_path / 'comparison.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_campaign(db_session, *, status=CampaignStatus.COMPLETED):
target = TargetRepository(db_session).create(
EvalTarget(id=f"t-{uuid4().hex[:8]}", name="数字员工")
)
scenario = ScenarioRepository(db_session).create(
Scenario(id=f"s-{uuid4().hex[:8]}", name="查账单", target_id=target.id, cases=[Case(id="c-1", type=CaseType.SINGLE, messages=["你好"])])
)
campaign = CampaignRepository(db_session).create(
Campaign(
name="cycle",
target_id=target.id,
window_seconds=3600,
time_scale=1.0,
plan=[CampaignPlanEntry(scenario_id=scenario.id, offset_seconds=0, count=1)],
status=status,
started_at=T0,
)
)
return campaign
def test_load_comparison_view_no_row_returns_none_status(db_session):
campaign = _seed_campaign(db_session)
view = load_comparison_view(db_session, campaign)
assert view["status"] == "none"
assert view["comparison"] is None
assert "auto_baseline" in view
assert "metric_diff" in view
def test_load_comparison_view_generating_row(db_session):
campaign = _seed_campaign(db_session)
CampaignPeriodComparisonRepository(db_session).upsert(
campaign.id,
status="generating",
baseline_campaign_id="baseline-1",
triggered_by="manual",
)
view = load_comparison_view(db_session, campaign)
assert view["status"] == "generating"
assert view["comparison"] is not None
assert view["comparison"]["baseline_campaign_id"] == "baseline-1"
def test_load_comparison_view_completed_row_includes_model_name(db_session):
campaign = _seed_campaign(db_session)
CampaignPeriodComparisonRepository(db_session).upsert(
campaign.id,
status="completed",
baseline_campaign_id="baseline-1",
result={"trend": "improving"},
model_config_id="mc-1",
triggered_by="manual",
)
view = load_comparison_view(db_session, campaign)
assert view["status"] == "completed"
assert view["comparison"]["result"] == {"trend": "improving"}
assert "model_name" in view["comparison"]
def test_validate_rejects_non_completed_campaign(db_session):
campaign = _seed_campaign(db_session, status=CampaignStatus.RUNNING)
with pytest.raises(ComparisonError, match="活动完成后"):
validate_comparison_request(db_session, campaign)
def test_validate_rejects_missing_analysis_model(db_session, monkeypatch):
campaign = _seed_campaign(db_session)
monkeypatch.setattr(
"agenteval.evaluation.comparison.resolve_analysis_model",
lambda *a, **kw: None,
)
with pytest.raises(ComparisonError, match="未配置分析模型"):
validate_comparison_request(db_session, campaign)
def test_validate_rejects_incomplete_current_analysis(db_session, monkeypatch):
campaign = _seed_campaign(db_session)
monkeypatch.setattr(
"agenteval.evaluation.comparison.resolve_analysis_model",
lambda *a, **kw: object(),
)
CampaignAnalysisRepository(db_session).upsert(
campaign.id, status="generating", triggered_by="auto"
)
with pytest.raises(ComparisonError, match="请先生成本期活动的智能分析"):
validate_comparison_request(db_session, campaign)
def test_validate_rejects_missing_explicit_baseline(db_session, monkeypatch):
campaign = _seed_campaign(db_session)
monkeypatch.setattr(
"agenteval.evaluation.comparison.resolve_analysis_model",
lambda *a, **kw: object(),
)
CampaignAnalysisRepository(db_session).upsert(
campaign.id, status="completed", result={}, triggered_by="auto"
)
with pytest.raises(ComparisonError, match="基线活动不存在"):
validate_comparison_request(db_session, campaign, explicit_baseline_id="missing")
def test_validate_rejects_incomplete_baseline_analysis(db_session, monkeypatch):
campaign = _seed_campaign(db_session)
monkeypatch.setattr(
"agenteval.evaluation.comparison.resolve_analysis_model",
lambda *a, **kw: object(),
)
baseline = _seed_campaign(db_session)
CampaignAnalysisRepository(db_session).upsert(
campaign.id, status="completed", result={}, triggered_by="auto"
)
CampaignAnalysisRepository(db_session).upsert(
baseline.id, status="generating", triggered_by="auto"
)
with pytest.raises(ComparisonError, match="基线活动没有已完成的智能分析"):
validate_comparison_request(db_session, campaign, explicit_baseline_id=baseline.id)