AgentEvalTool/tests/integration/test_campaign_analysis_api.py

218 lines
7.6 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.

"""Integration tests for /api/campaigns/{id}/analysis (v0.7 ticket 03)."""
import pytest
from agenteval.models import (
CampaignStatus,
Case,
CaseType,
ChannelType,
EvalTarget,
PlatformType,
Scenario,
TargetStatus,
)
from agenteval.storage.db import CampaignAnalysisDB, ModelConfigDB
from agenteval.storage.model_config_repository import ModelConfigRepository
from agenteval.storage.repository import CampaignRepository, ScenarioRepository, TargetRepository
from agenteval.web.app import app
from httpx import ASGITransport, AsyncClient
from sqlmodel import select
@pytest.fixture()
def seeded_db(db_session, monkeypatch):
from agenteval.storage import db as db_module
from agenteval.storage import repository as repo_module
from agenteval.web import app as app_module
from agenteval.web.routers import campaigns as campaigns_module
monkeypatch.setattr(app_module, "init_db", lambda: None)
monkeypatch.setattr(campaigns_module.campaign_runtime, "start", lambda *a, **k: None)
monkeypatch.setattr(db_module, "get_session", lambda: db_session)
monkeypatch.setattr(repo_module, "get_session", lambda: db_session)
from agenteval.web.deps import get_db
def _test_get_db():
try:
yield db_session
finally:
pass
app.dependency_overrides[get_db] = _test_get_db
TargetRepository(db_session).create(EvalTarget(
id="t-1", name="mock-target",
platform=PlatformType.AI_DIGITAL_EMPLOYEE,
channel_type=ChannelType.TUTU_API,
channel_config={"base_url": "http://mock", "token": "x"},
status=TargetStatus.ACTIVE,
))
ScenarioRepository(db_session).create(Scenario(
id="s-1", name="mock-scenario",
cases=[Case(id="c1", type=CaseType.SINGLE, messages=["hi"])],
))
yield db_session
app.dependency_overrides.clear()
@pytest.fixture()
async def client():
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as c:
yield c
def _payload() -> dict:
return {
"name": "24h-cycle",
"target_id": "t-1",
"window_seconds": 86400,
"time_scale": 1.0,
"plan": [{"scenario_id": "s-1", "offset_seconds": 0, "count": 1}],
}
async def _create_campaign(client, session, status: CampaignStatus = CampaignStatus.COMPLETED) -> str:
campaign_id = (await client.post("/api/campaigns", json=_payload())).json()["id"]
repo = CampaignRepository(session)
campaign = repo.get(campaign_id)
campaign.status = status
repo.update(campaign)
return campaign_id
def _seed_analysis_default(session) -> None:
ModelConfigRepository(session).create(ModelConfigDB(
id="mc-1", name="analysis-cfg", provider="openai_compatible", capability="chat",
endpoint_url="https://models.example.com/v1/chat/completions", model_name="m",
is_analysis_default=True,
))
def _complete_analysis_row(session, campaign_id: str) -> None:
from agenteval.storage.db import utc_now
row = CampaignAnalysisDB(
campaign_id=campaign_id, status="completed", model_config_id="mc-1",
triggered_by="manual", updated_at=utc_now(),
)
row.set_result({
"overall": "整体达标",
"problems": [],
"scenario_narratives": [{"scenario_id": "s-1", "narrative": "表现稳定"}],
"suggestions": [{"priority": 1, "text": "保持"}],
})
existing = session.exec(
select(CampaignAnalysisDB).where(CampaignAnalysisDB.campaign_id == campaign_id)
).first()
if existing:
session.delete(existing)
session.commit()
session.add(row)
session.commit()
async def test_get_analysis_empty_state(client, seeded_db):
campaign_id = await _create_campaign(client, seeded_db)
resp = await client.get(f"/api/campaigns/{campaign_id}/analysis")
assert resp.status_code == 200
assert resp.json() == {"status": "none"}
async def test_post_analysis_rejects_non_terminal_campaign(client, seeded_db):
campaign_id = await _create_campaign(client, seeded_db, status=CampaignStatus.RUNNING)
resp = await client.post(f"/api/campaigns/{campaign_id}/analysis")
assert resp.status_code == 400
async def test_post_analysis_rejects_planned_campaign(client, seeded_db):
campaign_id = await _create_campaign(client, seeded_db, status=CampaignStatus.PLANNED)
resp = await client.post(f"/api/campaigns/{campaign_id}/analysis")
assert resp.status_code == 400
async def test_post_analysis_without_model_guides_configuration(client, seeded_db):
campaign_id = await _create_campaign(client, seeded_db)
resp = await client.post(f"/api/campaigns/{campaign_id}/analysis")
assert resp.status_code == 400
assert "分析" in resp.json()["detail"]
async def test_post_then_get_completed_analysis(client, seeded_db, monkeypatch):
from agenteval.web.routers import campaigns as campaigns_module
_seed_analysis_default(seeded_db)
campaign_id = await _create_campaign(client, seeded_db)
# 假后台任务:同步写入 completed 行(真任务的单测覆盖在 test_campaign_analysis.py
monkeypatch.setattr(
campaigns_module,
"enqueue_campaign_analysis",
lambda cid, *, triggered_by: _complete_analysis_row(seeded_db, cid),
)
resp = await client.post(f"/api/campaigns/{campaign_id}/analysis")
assert resp.status_code == 200
got = (await client.get(f"/api/campaigns/{campaign_id}/analysis")).json()
assert got["status"] == "completed"
assert got["model_config_id"] == "mc-1"
assert got["result"]["overall"] == "整体达标"
assert got["result"]["scenario_narratives"][0]["narrative"] == "表现稳定"
async def test_rerun_upserts_without_new_row(client, seeded_db, monkeypatch):
from agenteval.web.routers import campaigns as campaigns_module
_seed_analysis_default(seeded_db)
campaign_id = await _create_campaign(client, seeded_db)
monkeypatch.setattr(
campaigns_module,
"enqueue_campaign_analysis",
lambda cid, *, triggered_by: _complete_analysis_row(seeded_db, cid),
)
await client.post(f"/api/campaigns/{campaign_id}/analysis")
await client.post(f"/api/campaigns/{campaign_id}/analysis")
rows = seeded_db.exec(
select(CampaignAnalysisDB).where(CampaignAnalysisDB.campaign_id == campaign_id)
).all()
assert len(rows) == 1
async def test_get_analysis_missing_campaign_404(client, seeded_db):
resp = await client.get("/api/campaigns/nope/analysis")
assert resp.status_code == 404
async def test_markdown_export_includes_completed_analysis(client, seeded_db):
campaign_id = await _create_campaign(client, seeded_db)
_complete_analysis_row(seeded_db, campaign_id)
resp = await client.get(f"/api/campaigns/{campaign_id}/report/markdown")
assert resp.status_code == 200
assert "## 智能分析" in resp.text
assert "整体达标" in resp.text
assert "表现稳定" in resp.text
async def test_markdown_export_ignores_non_completed_analysis(client, seeded_db):
campaign_id = await _create_campaign(client, seeded_db)
row = CampaignAnalysisDB(campaign_id=campaign_id, status="failed", error="模型超时")
seeded_db.add(row)
seeded_db.commit()
resp = await client.get(f"/api/campaigns/{campaign_id}/report/markdown")
assert resp.status_code == 200
assert "智能分析" not in resp.text
async def test_markdown_export_without_analysis_unchanged(client, seeded_db):
campaign_id = await _create_campaign(client, seeded_db)
resp = await client.get(f"/api/campaigns/{campaign_id}/report/markdown")
assert resp.status_code == 200
assert "智能分析" not in resp.text