449 lines
19 KiB
Python
449 lines
19 KiB
Python
"""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,
|
||
execute_campaign_comparison,
|
||
narrate_period_comparison,
|
||
resolve_auto_baseline,
|
||
)
|
||
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, monkeypatch):
|
||
from agenteval.evaluation import comparison as comparison_module
|
||
|
||
monkeypatch.setattr(comparison_module, "get_session", lambda: 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(
|
||
"camp-cur", triggered_by="manual", chat_client=FakeChatClient(NARRATION),
|
||
)
|
||
|
||
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, monkeypatch):
|
||
from agenteval.evaluation import comparison as comparison_module
|
||
|
||
monkeypatch.setattr(comparison_module, "get_session", lambda: 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(
|
||
"camp-cur", triggered_by="manual", chat_client=FakeChatClient("garbage"),
|
||
)
|
||
|
||
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, monkeypatch):
|
||
from agenteval.evaluation import comparison as comparison_module
|
||
|
||
monkeypatch.setattr(comparison_module, "get_session", lambda: 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(
|
||
"camp-cur", triggered_by="manual", chat_client=FakeChatClient(NARRATION),
|
||
)
|
||
|
||
row = db_session.exec(
|
||
select(CampaignPeriodComparisonDB).where(CampaignPeriodComparisonDB.campaign_id == "camp-cur")
|
||
).one()
|
||
assert row.status == "failed"
|
||
assert "基线" in row.error
|