Give EvalRun.summary a typed RunSummary value (unified RunError, lenient legacy parsing) so readers stop reaching into a schemaless dict, and route every cross-run rollup — dashboard, scenario ranking, trend, campaign report — through one aggregate_runs seam. Fixes the divergence where stats averaged pass_rate over completed-only runs while the campaign report counted faults as 0.0. Cross-run rule (ADR-0004): genuine faults count 0.0, user-cancelled runs are excluded from both denominators.
132 lines
4.7 KiB
Python
132 lines
4.7 KiB
Python
"""Integration tests for /api/stats/dashboard aggregation."""
|
||
|
||
import pytest
|
||
from fastapi.testclient import TestClient
|
||
from sqlmodel import Session, SQLModel, create_engine
|
||
|
||
from agenteval.models import (
|
||
Case, CaseType, ChannelType, EvalRun, EvalTarget, PlatformType,
|
||
RunStatus, RunTrigger, Scenario, TargetStatus,
|
||
)
|
||
from agenteval.storage.repository import RunRepository, ScenarioRepository, TargetRepository
|
||
from agenteval.web.app import app
|
||
from agenteval.web.deps import get_db
|
||
|
||
|
||
@pytest.fixture()
|
||
def client_with_db(tmp_path):
|
||
from agenteval.storage.db import ( # noqa: F401
|
||
EvalResultDB, EvalRunDB, EvalTargetDB, ModelConfigDB, ScenarioDB, TurnDB,
|
||
)
|
||
engine = create_engine(
|
||
f"sqlite:///{tmp_path / 'stats_api.db'}",
|
||
connect_args={"check_same_thread": False},
|
||
)
|
||
SQLModel.metadata.create_all(engine)
|
||
session = Session(engine)
|
||
|
||
def override_get_db():
|
||
try:
|
||
yield session
|
||
finally:
|
||
pass
|
||
|
||
app.dependency_overrides[get_db] = override_get_db
|
||
client = TestClient(app)
|
||
yield client, session
|
||
app.dependency_overrides.clear()
|
||
session.close()
|
||
|
||
|
||
def _seed(session: Session) -> None:
|
||
target = TargetRepository(session).create(EvalTarget(
|
||
name="对象A", platform=PlatformType.AI_DIGITAL_EMPLOYEE,
|
||
channel_type=ChannelType.TUTU_API, channel_config={}, status=TargetStatus.ACTIVE,
|
||
))
|
||
scenario = ScenarioRepository(session).create(Scenario(
|
||
name="场景A", cases=[Case(id="c1", type=CaseType.SINGLE, messages=["hi"])],
|
||
))
|
||
repo = RunRepository(session)
|
||
for pass_rate, trigger in [(1.0, RunTrigger.MANUAL), (0.5, RunTrigger.AI_ASSISTANT)]:
|
||
run = repo.create(EvalRun(
|
||
target_id=target.id, scenario_id=scenario.id,
|
||
status=RunStatus.COMPLETED, triggered_by=trigger,
|
||
))
|
||
run.summary = {"total_cases": 2, "passed_cases": 1, "failed_cases": 1,
|
||
"total_rules": 2, "passed_rules": 1, "pass_rate": pass_rate}
|
||
repo.update(run)
|
||
repo.create(EvalRun(
|
||
target_id=target.id, scenario_id=scenario.id,
|
||
status=RunStatus.RUNNING, triggered_by=RunTrigger.MANUAL,
|
||
))
|
||
|
||
|
||
def test_dashboard_aggregates(client_with_db):
|
||
client, session = client_with_db
|
||
_seed(session)
|
||
|
||
data = client.get("/api/stats/dashboard").json()
|
||
assert data["targets_count"] == 1
|
||
assert data["scenarios_count"] == 1
|
||
assert data["runs_count"] == 3
|
||
assert data["model_configs_count"] == 0
|
||
assert data["running_count"] == 1
|
||
assert data["today_runs"] == 3
|
||
assert data["overall_pass_rate"] == pytest.approx(0.75)
|
||
assert data["trigger_breakdown"] == {"manual": 2, "ai_assistant": 1}
|
||
|
||
assert len(data["scenario_stats"]) == 1
|
||
stat = data["scenario_stats"][0]
|
||
assert stat["scenario_name"] == "场景A"
|
||
assert stat["run_count"] == 2
|
||
assert stat["avg_pass_rate"] == pytest.approx(0.75)
|
||
assert stat["last_run_at"] is not None
|
||
|
||
assert len(data["recent_runs"]) == 3
|
||
assert data["recent_runs"][0]["scenario_name"] == "场景A"
|
||
assert data["recent_runs"][0]["target_name"] == "对象A"
|
||
assert "triggered_by" in data["recent_runs"][0]
|
||
|
||
|
||
def test_dashboard_empty_db(client_with_db):
|
||
client, _ = client_with_db
|
||
data = client.get("/api/stats/dashboard").json()
|
||
assert data["runs_count"] == 0
|
||
assert data["overall_pass_rate"] is None
|
||
assert data["scenario_stats"] == []
|
||
assert data["recent_runs"] == []
|
||
|
||
|
||
def test_trend_returns_daily_points(client_with_db):
|
||
client, session = client_with_db
|
||
_seed(session)
|
||
points = client.get("/api/stats/trend").json()
|
||
assert len(points) == 1
|
||
assert points[0]["run_count"] == 2
|
||
assert points[0]["pass_rate"] == pytest.approx(75.0)
|
||
|
||
|
||
def test_dashboard_follows_adr_0004(client_with_db):
|
||
"""故障 run 计 0.0 进分母;用户取消的 run 整体排除(ADR-0004)。"""
|
||
client, session = client_with_db
|
||
_seed(session) # two completed runs: 1.0 and 0.5
|
||
repo = RunRepository(session)
|
||
target_id = repo.list_all()[0].target_id
|
||
scenario_id = repo.list_all()[0].scenario_id
|
||
|
||
faulted = repo.create(EvalRun(
|
||
target_id=target_id, scenario_id=scenario_id, status=RunStatus.FAILED,
|
||
))
|
||
faulted.summary = {"error": "channel exploded"}
|
||
repo.update(faulted)
|
||
|
||
cancelled = repo.create(EvalRun(
|
||
target_id=target_id, scenario_id=scenario_id, status=RunStatus.FAILED,
|
||
))
|
||
cancelled.summary = {"error": {"code": "cancelled_by_user", "message": "stop"}}
|
||
repo.update(cancelled)
|
||
|
||
data = client.get("/api/stats/dashboard").json()
|
||
# (1.0 + 0.5 + 0.0[fault]) / 3 — cancelled run out of the denominator
|
||
assert data["overall_pass_rate"] == pytest.approx(0.5)
|