AgentEvalTool/tests/unit/test_campaign_comparison.py

443 lines
19 KiB
Python
Raw Permalink 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.

"""Auto-baseline pairing + mechanical metric diff + narrative orchestration (v0.8)."""
import json
from datetime import timedelta
import pytest
from agenteval.evaluation.comparison import (
ComparisonError,
campaign_plan_fingerprint,
compute_metric_diff,
narrate_period_comparison,
resolve_auto_baseline,
)
from agenteval.evaluation.intelligence_jobs import execute_campaign_comparison_job
from agenteval.models import Campaign, CampaignPlanEntry, CampaignStatus, EvalRun, RunStatus, RunSummary
from agenteval.storage.db import (
CampaignAnalysisDB,
CampaignPeriodComparisonDB,
ModelConfigDB,
utc_now,
)
from agenteval.storage.model_config_repository import ModelConfigRepository
from agenteval.storage.repository import CampaignRepository, RunRepository
from sqlmodel import select
T0 = utc_now().replace(tzinfo=None) - timedelta(hours=1)
def _campaign(
campaign_id: str = "camp-1",
*,
target_id: str = "t-1",
window_seconds: int = 86400,
time_scale: float = 1.0,
plan: list[CampaignPlanEntry] | None = None,
status: CampaignStatus = CampaignStatus.COMPLETED,
completed_at=T0 + timedelta(hours=1),
) -> Campaign:
return Campaign(
id=campaign_id,
name=f"campaign-{campaign_id}",
target_id=target_id,
window_seconds=window_seconds,
time_scale=time_scale,
plan=plan or [
CampaignPlanEntry(scenario_id="s-1", offset_seconds=0, count=2),
CampaignPlanEntry(scenario_id="s-2", offset_seconds=3600, count=1),
],
status=status,
completed_at=completed_at,
)
# ── 计划指纹 ──────────────────────────────────────────────────────────────
def test_fingerprint_ignores_plan_entry_order():
a = _campaign()
b = _campaign(plan=list(reversed(a.plan)))
assert campaign_plan_fingerprint(a) == campaign_plan_fingerprint(b)
def test_fingerprint_changes_with_any_plan_dimension():
base = campaign_plan_fingerprint(_campaign())
assert campaign_plan_fingerprint(_campaign(target_id="t-2")) != base
assert campaign_plan_fingerprint(_campaign(window_seconds=3600)) != base
assert campaign_plan_fingerprint(
_campaign(plan=[CampaignPlanEntry(scenario_id="s-9", offset_seconds=0, count=2),
CampaignPlanEntry(scenario_id="s-2", offset_seconds=3600, count=1)])
) != base
assert campaign_plan_fingerprint(
_campaign(plan=[CampaignPlanEntry(scenario_id="s-1", offset_seconds=60, count=2),
CampaignPlanEntry(scenario_id="s-2", offset_seconds=3600, count=1)])
) != base
assert campaign_plan_fingerprint(
_campaign(plan=[CampaignPlanEntry(scenario_id="s-1", offset_seconds=0, count=3),
CampaignPlanEntry(scenario_id="s-2", offset_seconds=3600, count=1)])
) != base
def test_fingerprint_ignores_time_scale():
assert campaign_plan_fingerprint(_campaign()) == campaign_plan_fingerprint(_campaign(time_scale=24.0))
# ── 自动基线解析 ──────────────────────────────────────────────────────────
def _seed_baseline(session, campaign: Campaign, *, analysis_status: str | None = "completed") -> None:
CampaignRepository(session).create(campaign)
if analysis_status is not None:
session.add(CampaignAnalysisDB(campaign_id=campaign.id, status=analysis_status))
session.commit()
def test_baseline_picks_most_recent_with_completed_analysis(db_session):
old = _campaign("camp-old", completed_at=T0)
recent = _campaign("camp-recent", completed_at=T0 + timedelta(minutes=30))
current = _campaign("camp-cur", completed_at=T0 + timedelta(hours=2))
_seed_baseline(db_session, old)
_seed_baseline(db_session, recent)
resolved = resolve_auto_baseline(current, db_session)
assert resolved is not None and resolved.id == "camp-recent"
def test_baseline_skips_candidates_without_completed_analysis(db_session):
no_analysis = _campaign("camp-no", completed_at=T0)
failed_analysis = _campaign("camp-failed", completed_at=T0 + timedelta(minutes=30))
current = _campaign("camp-cur", completed_at=T0 + timedelta(hours=2))
_seed_baseline(db_session, no_analysis, analysis_status=None)
_seed_baseline(db_session, failed_analysis, analysis_status="failed")
assert resolve_auto_baseline(current, db_session) is None
def test_baseline_skips_accelerated_candidates_and_accelerated_current(db_session):
accelerated = _campaign("camp-fast", time_scale=24.0, completed_at=T0)
current = _campaign("camp-cur", completed_at=T0 + timedelta(hours=2))
_seed_baseline(db_session, accelerated)
assert resolve_auto_baseline(current, db_session) is None
# 本期自身为加速线 → 无自动基线
production = _campaign("camp-prod", completed_at=T0)
_seed_baseline(db_session, production)
assert resolve_auto_baseline(_campaign("camp-cur", time_scale=12.0), db_session) is None
def test_baseline_skips_different_fingerprint(db_session):
other_plan = _campaign("camp-other", window_seconds=3600, completed_at=T0)
current = _campaign("camp-cur", completed_at=T0 + timedelta(hours=2))
_seed_baseline(db_session, other_plan)
assert resolve_auto_baseline(current, db_session) is None
def test_baseline_skips_later_or_same_moment_completion(db_session):
later = _campaign("camp-later", completed_at=T0 + timedelta(hours=3))
current = _campaign("camp-cur", completed_at=T0 + timedelta(hours=2))
_seed_baseline(db_session, later)
assert resolve_auto_baseline(current, db_session) is None
same_moment = _campaign("camp-same", completed_at=current.completed_at)
_seed_baseline(db_session, same_moment)
assert resolve_auto_baseline(current, db_session) is None
def test_baseline_skips_uncompleted_candidates(db_session):
uncompleted = _campaign("camp-run", status=CampaignStatus.RUNNING, completed_at=None)
current = _campaign("camp-cur", completed_at=T0 + timedelta(hours=2))
_seed_baseline(db_session, uncompleted)
assert resolve_auto_baseline(current, db_session) is None
# ── 机械指标 diff ─────────────────────────────────────────────────────────
def _report(summary: dict, capability: list[dict]) -> dict:
return {"summary": summary, "capability_summary": capability}
def test_metric_diff_computes_overall_and_scenario_deltas():
baseline = _report(
{"overall_pass_rate": 0.5, "overall_availability": 1.0, "avg_latency_ms": 800.0},
[
{"scenario_id": "s-1", "scenario_name": "售前", "pass_rate": 0.4,
"availability": 1.0, "avg_latency_ms": 700.0},
{"scenario_id": "s-2", "scenario_name": "售后", "pass_rate": 0.6,
"availability": 1.0, "avg_latency_ms": 900.0},
],
)
current = _report(
{"overall_pass_rate": 0.8, "overall_availability": 0.75, "avg_latency_ms": 500.0},
[
{"scenario_id": "s-1", "scenario_name": "售前", "pass_rate": 0.9,
"availability": 0.5, "avg_latency_ms": 400.0},
{"scenario_id": "s-2", "scenario_name": "售后", "pass_rate": 0.6,
"availability": 1.0, "avg_latency_ms": 900.0},
],
)
diff = compute_metric_diff(baseline, current)
assert diff["overall"]["pass_rate"] == {"baseline": 0.5, "current": 0.8, "delta": 0.3}
assert diff["overall"]["availability"] == {"baseline": 1.0, "current": 0.75, "delta": -0.25}
assert diff["overall"]["avg_latency_ms"] == {"baseline": 800.0, "current": 500.0, "delta": -300.0}
scenarios = {s["scenario_id"]: s for s in diff["scenarios"]}
assert scenarios["s-1"]["pass_rate"] == {"baseline": 0.4, "current": 0.9, "delta": 0.5}
assert scenarios["s-1"]["availability"] == {"baseline": 1.0, "current": 0.5, "delta": -0.5}
assert scenarios["s-1"]["avg_latency_ms"] == {"baseline": 700.0, "current": 400.0, "delta": -300.0}
assert scenarios["s-2"]["pass_rate"]["delta"] == 0.0
assert [s["scenario_name"] for s in diff["scenarios"]] == ["售前", "售后"]
def test_metric_diff_union_of_scenarios_missing_side_yields_none():
baseline = _report(
{"overall_pass_rate": 0.5, "overall_availability": 1.0, "avg_latency_ms": 800.0},
[{"scenario_id": "s-1", "scenario_name": "售前", "pass_rate": 0.4,
"availability": 1.0, "avg_latency_ms": 700.0}],
)
current = _report(
{"overall_pass_rate": 0.8, "overall_availability": 1.0, "avg_latency_ms": 500.0},
[{"scenario_id": "s-2", "scenario_name": "售后", "pass_rate": 0.9,
"availability": 1.0, "avg_latency_ms": 400.0}],
)
diff = compute_metric_diff(baseline, current)
scenarios = {s["scenario_id"]: s for s in diff["scenarios"]}
assert set(scenarios) == {"s-1", "s-2"}
# s-1 仅基线有current/delta 为 None
assert scenarios["s-1"]["pass_rate"] == {"baseline": 0.4, "current": None, "delta": None}
assert scenarios["s-1"]["avg_latency_ms"]["delta"] is None
# s-2 仅本期有baseline/delta 为 None
assert scenarios["s-2"]["pass_rate"] == {"baseline": None, "current": 0.9, "delta": None}
assert scenarios["s-2"]["availability"]["baseline"] is None
def test_metric_diff_none_metrics_propagate():
baseline = _report(
{"overall_pass_rate": None, "overall_availability": None, "avg_latency_ms": None},
[{"scenario_id": "s-1", "scenario_name": "售前", "pass_rate": None,
"availability": None, "avg_latency_ms": None}],
)
current = _report(
{"overall_pass_rate": 0.8, "overall_availability": 1.0, "avg_latency_ms": 500.0},
[{"scenario_id": "s-1", "scenario_name": "售前", "pass_rate": 0.9,
"availability": 1.0, "avg_latency_ms": 400.0}],
)
diff = compute_metric_diff(baseline, current)
assert diff["overall"]["pass_rate"]["delta"] is None
assert diff["scenarios"][0]["pass_rate"]["delta"] is None
assert diff["scenarios"][0]["pass_rate"]["current"] == 0.9
def test_metric_diff_empty_capability_summaries():
diff = compute_metric_diff(
_report({"overall_pass_rate": None, "overall_availability": None, "avg_latency_ms": None}, []),
_report({"overall_pass_rate": None, "overall_availability": None, "avg_latency_ms": None}, []),
)
assert diff["scenarios"] == []
assert diff["overall"]["pass_rate"]["delta"] is None
# ── 叙述编排(单次 LLM 调用) ─────────────────────────────────────────────
class FakeChatClient:
"""Queued-response fake for the comparison LLM seam."""
def __init__(self, *responses):
self._responses = list(responses)
self.calls: list[list[dict]] = []
async def __call__(self, messages: list[dict]) -> str:
self.calls.append(messages)
if not self._responses:
raise AssertionError("unexpected extra LLM call")
item = self._responses.pop(0)
if isinstance(item, Exception):
raise item
return item
def _analysis(overall: str) -> dict:
return {
"overall": overall,
"problems": [{"severity": "high", "title": "答非所问", "scenario_ids": ["s-1"]}],
"suggestions": [{"priority": 1, "text": "补充意图语料"}],
}
def _diff() -> dict:
return compute_metric_diff(
{"summary": {"overall_pass_rate": 0.5, "overall_availability": 1.0, "avg_latency_ms": 800.0},
"capability_summary": [{"scenario_id": "s-1", "pass_rate": 0.5, "availability": 1.0, "avg_latency_ms": 800.0}]},
{"summary": {"overall_pass_rate": 0.9, "overall_availability": 1.0, "avg_latency_ms": 500.0},
"capability_summary": [{"scenario_id": "s-1", "pass_rate": 0.9, "availability": 1.0, "avg_latency_ms": 500.0}]},
)
NARRATION = json.dumps({
"trend": "improving",
"summary": "整体通过率显著提升,售前答非所问问题缓解",
"problem_evolution": [
{"status": "resolved", "title": "答非所问", "detail": "语料补充后恢复",
"scenario_ids": ["s-1", "ghost-scenario"]},
{"status": "nonsense", "title": "新问题", "detail": "...", "scenario_ids": []},
],
"suggestion_tracking": [
{"text": "补充意图语料", "status": "addressed", "note": "已落实"},
{"text": "排查上游", "status": "nonsense", "note": "状态未知"},
{"text": "本期新增建议", "status": "new", "note": ""},
],
})
async def test_narration_input_carries_both_analyses_and_diff():
client = FakeChatClient(NARRATION)
await narrate_period_comparison(
baseline_analysis=_analysis("上期整体不达标"),
current_analysis=_analysis("本期整体改善"),
metric_diff=_diff(),
valid_scenario_ids={"s-1"},
chat_client=client,
)
assert len(client.calls) == 1
prompt = str(client.calls[0])
assert "上期整体不达标" in prompt
assert "本期整体改善" in prompt
assert "pass_rate" in prompt # 机械 diff 进入输入
async def test_narration_normalizes_and_whitelists():
client = FakeChatClient(NARRATION)
result = await narrate_period_comparison(
baseline_analysis=_analysis("a"),
current_analysis=_analysis("b"),
metric_diff=_diff(),
valid_scenario_ids={"s-1"},
chat_client=client,
)
assert result["trend"] == "improving"
assert result["summary"].startswith("整体通过率显著提升")
resolved = result["problem_evolution"][0]
assert resolved["status"] == "resolved"
assert resolved["scenario_ids"] == ["s-1"] # ghost-scenario 剔除
assert result["problem_evolution"][1]["status"] == "persisting" # 非法枚举归一
statuses = [s["status"] for s in result["suggestion_tracking"]]
assert statuses == ["addressed", "unaddressed", "new"] # nonsense → unaddressed
async def test_narration_normalizes_invalid_trend_to_stable():
client = FakeChatClient(json.dumps({"trend": "wild", "summary": "结论"}))
result = await narrate_period_comparison(
baseline_analysis={}, current_analysis={}, metric_diff={},
valid_scenario_ids=set(), chat_client=client,
)
assert result["trend"] == "stable"
assert result["problem_evolution"] == []
assert result["suggestion_tracking"] == []
async def test_unparseable_narration_raises_comparison_error():
client = FakeChatClient("这不是 JSON")
with pytest.raises(ComparisonError):
await narrate_period_comparison(
baseline_analysis={}, current_analysis={}, metric_diff={},
valid_scenario_ids=set(), chat_client=client,
)
async def test_narration_missing_summary_raises():
client = FakeChatClient(json.dumps({"trend": "stable"}))
with pytest.raises(ComparisonError):
await narrate_period_comparison(
baseline_analysis={}, current_analysis={}, metric_diff={},
valid_scenario_ids=set(), chat_client=client,
)
# ── 后台执行状态机 ───────────────────────────────────────────────────────
def _seed_config(session, config_id: str) -> None:
ModelConfigRepository(session).create(ModelConfigDB(
id=config_id, name=f"cfg-{config_id}", provider="openai_compatible", capability="chat",
endpoint_url="https://models.example.com/v1/chat/completions", model_name="m",
is_analysis_default=True,
))
def _seed_campaign_with_analysis(session, campaign: Campaign, analysis_result: dict) -> None:
CampaignRepository(session).create(campaign)
row = CampaignAnalysisDB(campaign_id=campaign.id, status="completed")
row.set_result(analysis_result)
session.add(row)
session.commit()
RunRepository(session).create(EvalRun(
id=f"run-{campaign.id}", target_id="t-1", scenario_id="s-1",
campaign_id=campaign.id, status=RunStatus.COMPLETED, started_at=utc_now(),
summary=RunSummary(total_cases=2, pass_rate=0.5, avg_latency_ms=700),
))
async def test_execute_writes_completed_row_with_baseline_snapshot(db_session):
_seed_config(db_session, "mc-default")
baseline = _campaign("camp-base", completed_at=T0)
current = _campaign("camp-cur", completed_at=T0 + timedelta(hours=2))
_seed_campaign_with_analysis(db_session, baseline, _analysis("上期"))
_seed_campaign_with_analysis(db_session, current, _analysis("本期"))
await execute_campaign_comparison_job(
"camp-cur", triggered_by="manual", chat_client=FakeChatClient(NARRATION),
session_factory=lambda: db_session,
)
row = db_session.exec(
select(CampaignPeriodComparisonDB).where(CampaignPeriodComparisonDB.campaign_id == "camp-cur")
).one()
assert row.status == "completed"
assert row.baseline_campaign_id == "camp-base"
assert row.model_config_id == "mc-default"
assert row.get_result()["trend"] == "improving"
async def test_execute_records_failure_on_unparseable_output(db_session):
_seed_config(db_session, "mc-default")
baseline = _campaign("camp-base", completed_at=T0)
current = _campaign("camp-cur", completed_at=T0 + timedelta(hours=2))
_seed_campaign_with_analysis(db_session, baseline, _analysis("上期"))
_seed_campaign_with_analysis(db_session, current, _analysis("本期"))
await execute_campaign_comparison_job(
"camp-cur", triggered_by="manual", chat_client=FakeChatClient("garbage"),
session_factory=lambda: db_session,
)
row = db_session.exec(
select(CampaignPeriodComparisonDB).where(CampaignPeriodComparisonDB.campaign_id == "camp-cur")
).one()
assert row.status == "failed"
assert row.error
assert row.baseline_campaign_id == "camp-base"
async def test_execute_fails_without_baseline_analysis(db_session):
_seed_config(db_session, "mc-default")
baseline = _campaign("camp-base", completed_at=T0)
current = _campaign("camp-cur", completed_at=T0 + timedelta(hours=2))
CampaignRepository(db_session).create(baseline) # 基线无分析行
_seed_campaign_with_analysis(db_session, current, _analysis("本期"))
await execute_campaign_comparison_job(
"camp-cur", triggered_by="manual", chat_client=FakeChatClient(NARRATION),
session_factory=lambda: db_session,
)
row = db_session.exec(
select(CampaignPeriodComparisonDB).where(CampaignPeriodComparisonDB.campaign_id == "camp-cur")
).one()
assert row.status == "failed"
assert "基线" in row.error