AgentEvalTool/tests/integration/test_exploration_api.py
sinohqb f8d8450b1e refactor(report): unify campaign report loading behind one read model
「活动报告取数三件套」此前在报告/markdown/分析/对比等 7 处手写重复,
唯一深化产物 build_campaign_report_dict 被锁在周期对比私有角落。
升位为 report.py 的 load_campaign_report(session, campaign) 单一出口
(探索线 summarize_campaign_exploration 同口径),并把 8 处
scenario_names 推导式收敛为 ScenarioRepository.name_map() 窄方法。
纯结构重排、零行为变更,572 项测试全绿。
2026-08-04 03:30:29 +08:00

503 lines
19 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 exploration session lifecycle + platform guardrails (v0.9 票据 01).
Uses ``httpx.AsyncClient`` with ``app=`` to drive the FastAPI app in-process.
Channel I/O is stubbed with MockChannel; tests cover the full lifecycle
(create → message → close), the three budget guardrails (409), trigger-source
gating per line tier, and experience-record normalization.
"""
from datetime import timedelta
import pytest
from agenteval.models import Campaign, ChannelType, EvalTarget, PlatformType, TargetStatus
from agenteval.storage.db import ExplorationSessionDB, utc_now
from agenteval.storage.repository import CampaignRepository, ExplorationSessionRepository, TargetRepository
from agenteval.web.app import app
from httpx import ASGITransport, AsyncClient
def _make_campaign(campaign_id: str, *, time_scale: float = 1.0, status: str = "running") -> Campaign:
return Campaign(
id=campaign_id,
name=f"campaign-{campaign_id}",
target_id="t-1",
window_seconds=86400,
time_scale=time_scale,
plan=[{"scenario_id": "s-1", "offset_seconds": 0, "count": 1}],
status=status,
started_at=utc_now(),
)
@pytest.fixture()
def seeded_db(db_session, monkeypatch):
"""Patch get_session/get_db to the test session and seed target + campaign."""
from agenteval.storage import db as db_module
from agenteval.storage import repository as repo_module
from agenteval.web import app as app_module
monkeypatch.setattr(app_module, "init_db", lambda: None)
def _test_get_session():
return db_session
monkeypatch.setattr(db_module, "get_session", _test_get_session)
monkeypatch.setattr(repo_module, "get_session", _test_get_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
# 隔离后台 judge 复核任务API 测试不真正派发后台任务
from agenteval.web.routers import exploration as exploration_router
monkeypatch.setattr(exploration_router, "start_judge_review", lambda session_id: None)
target = 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,
)
TargetRepository(db_session).create(target)
CampaignRepository(db_session).create(_make_campaign("c-1"))
yield db_session
app.dependency_overrides.clear()
def _stub_channel_factory(monkeypatch, channel) -> None:
"""Point the exploration router's ChannelFactory at a test channel."""
from agenteval.web.routers import exploration as exploration_module
class _StubFactory:
@staticmethod
def create(target):
return channel
monkeypatch.setattr(exploration_module, "ChannelFactory", _StubFactory)
@pytest.fixture()
def mock_channel(monkeypatch):
"""Stub ChannelFactory in the exploration router with a MockChannel."""
from tests.unit.mock_channel import MockChannel
channel = MockChannel(reply_text="您好,请问有什么可以帮您?")
_stub_channel_factory(monkeypatch, channel)
return channel
@pytest.fixture()
async def client():
transport = ASGITransport(app=app)
async with AsyncClient(transport=transport, base_url="http://test") as c:
yield c
def _session_payload(**overrides) -> dict:
payload = {
"campaign_id": "c-1",
"persona": {"name": "急性子用户", "traits": ["急躁", "目标导向"]},
"goal": "查询本月账单并完成缴费",
"triggered_by": "auto",
}
payload.update(overrides)
return payload
def _rewind_latest_session(db_session, minutes: int = 31) -> None:
"""Move the newest session's created_at back so the interval guardrail passes."""
repo = ExplorationSessionRepository(db_session)
sessions = repo.list_by_campaign("c-1")
latest = max(sessions, key=lambda s: s.created_at)
row = db_session.get(ExplorationSessionDB, latest.id)
row.created_at = utc_now() - timedelta(minutes=minutes)
db_session.add(row)
db_session.commit()
async def _create_session(client, **overrides):
return await client.post("/api/exploration/sessions", json=_session_payload(**overrides))
# ---------------------------------------------------------------- lifecycle
async def test_full_lifecycle_create_message_close(seeded_db, mock_channel, client):
resp = await _create_session(client)
assert resp.status_code == 200, resp.text
session = resp.json()
assert session["status"] == "running"
assert session["campaign_id"] == "c-1"
assert session["target_id"] == "t-1"
assert session["persona"]["name"] == "急性子用户"
session_id = session["id"]
resp = await client.post(
f"/api/exploration/sessions/{session_id}/messages",
json={"content": "我要查这个月的账单"},
)
assert resp.status_code == 200, resp.text
body = resp.json()
assert body["reply"] == "您好,请问有什么可以帮您?"
assert isinstance(body["latency_ms"], int)
assert body["turn_count"] == 1
repo = ExplorationSessionRepository(seeded_db)
assert repo.get(session_id).turn_count == 1
resp = await client.post(
f"/api/exploration/sessions/{session_id}/close",
json={
"experience": {
"goal_achieved": True,
"blockers": [],
"misled": [],
"emotion": "positive",
"notes": "顺利完成",
}
},
)
assert resp.status_code == 200, resp.text
closed = resp.json()
assert closed["status"] == "completed"
assert closed["experience"]["goal_achieved"] is True
assert closed["experience"]["emotion"] == "positive"
assert closed["closed_at"] is not None
async def test_dict_reply_content_is_flattened_to_text(seeded_db, monkeypatch, client):
"""通道回复 content 为对象(如 tutu msgBody时应提取文本而非存 str(dict)。"""
from agenteval.channels.base import Reply
from tests.unit.mock_channel import MockChannel
class _DictReplyChannel(MockChannel):
async def poll_reply(self, question_msg_id, timeout=30.0, poll_interval=1.0):
return Reply(
question_msg_id=question_msg_id,
content={"content": "您好,我是客服"},
raw_message={},
)
_stub_channel_factory(monkeypatch, _DictReplyChannel(reply_text="unused"))
resp = await _create_session(client)
assert resp.status_code == 200, resp.text
session_id = resp.json()["id"]
resp = await client.post(
f"/api/exploration/sessions/{session_id}/messages",
json={"content": "你好"},
)
assert resp.status_code == 200, resp.text
assert resp.json()["reply"] == "您好,我是客服"
messages_resp = await client.get(f"/api/exploration/sessions/{session_id}/messages")
assistant = [m for m in messages_resp.json()["messages"] if m["role"] == "assistant"]
assert assistant[0]["content"] == "您好,我是客服"
async def test_create_requires_existing_running_campaign(seeded_db, client):
resp = await _create_session(client, campaign_id="nope")
assert resp.status_code == 404
CampaignRepository(seeded_db).create(_make_campaign("c-done", status="completed"))
resp = await _create_session(client, campaign_id="c-done")
assert resp.status_code == 409
assert "进行中" in resp.json()["detail"]
async def test_accelerated_line_accepts_manual_only(seeded_db, mock_channel, client):
CampaignRepository(seeded_db).create(_make_campaign("c-fast", time_scale=24.0))
resp = await _create_session(client, campaign_id="c-fast", triggered_by="auto")
assert resp.status_code == 409
assert "手动" in resp.json()["detail"]
resp = await _create_session(client, campaign_id="c-fast", triggered_by="manual")
assert resp.status_code == 200
async def test_session_budget_guardrail(seeded_db, mock_channel, client):
for _ in range(8):
resp = await _create_session(client)
assert resp.status_code == 200, resp.text
_rewind_latest_session(seeded_db)
resp = await _create_session(client)
assert resp.status_code == 409
assert "预算" in resp.json()["detail"]
async def test_session_interval_guardrail(seeded_db, mock_channel, client):
resp = await _create_session(client)
assert resp.status_code == 200
resp = await _create_session(client)
assert resp.status_code == 409
assert "间隔" in resp.json()["detail"]
_rewind_latest_session(seeded_db)
resp = await _create_session(client)
assert resp.status_code == 200
async def test_campaign_budget_override_limits_sessions(seeded_db, mock_channel, client):
"""活动级预算覆盖生效(票据 02max_sessions=1 时第二个会话被拒。"""
campaign = CampaignRepository(seeded_db).get("c-1")
campaign.exploration_budget = {"max_sessions": 1}
CampaignRepository(seeded_db).update(campaign)
assert (await _create_session(client)).status_code == 200
_rewind_latest_session(seeded_db)
resp = await _create_session(client)
assert resp.status_code == 409
assert "预算" in resp.json()["detail"]
async def test_turn_budget_guardrail(seeded_db, mock_channel, client):
session_id = (await _create_session(client)).json()["id"]
for i in range(12):
resp = await client.post(
f"/api/exploration/sessions/{session_id}/messages",
json={"content": f"{i + 1} 轮问题"},
)
assert resp.status_code == 200, resp.text
resp = await client.post(
f"/api/exploration/sessions/{session_id}/messages",
json={"content": "第 13 轮问题"},
)
assert resp.status_code == 409
assert "轮数" in resp.json()["detail"]
async def test_message_rejected_when_session_not_running(seeded_db, mock_channel, client):
session_id = (await _create_session(client)).json()["id"]
resp = await client.post(
f"/api/exploration/sessions/{session_id}/close",
json={"experience": {"goal_achieved": False}},
)
assert resp.status_code == 200
resp = await client.post(
f"/api/exploration/sessions/{session_id}/messages",
json={"content": "还在吗?"},
)
assert resp.status_code == 409
assert "进行中" in resp.json()["detail"]
async def test_message_unknown_session_returns_404(seeded_db, mock_channel, client):
resp = await client.post("/api/exploration/sessions/nope/messages", json={"content": "hi"})
assert resp.status_code == 404
async def test_channel_failure_returns_502_without_consuming_turn(seeded_db, monkeypatch, client):
from tests.unit.mock_channel import MockChannel
channel = MockChannel(send_ok=False)
_stub_channel_factory(monkeypatch, channel)
session_id = (await _create_session(client)).json()["id"]
resp = await client.post(
f"/api/exploration/sessions/{session_id}/messages",
json={"content": "你好"},
)
assert resp.status_code == 502
assert ExplorationSessionRepository(seeded_db).get(session_id).turn_count == 0
async def test_poll_timeout_consumes_turn_budget(seeded_db, monkeypatch, client):
"""消息已送达但等不到回复:账本仍计一轮(超时不可绕过轮数预算)。"""
from types import SimpleNamespace
from agenteval.web.routers import exploration as exploration_module
from tests.unit.mock_channel import MockChannel
channel = MockChannel(missing_reply=True)
_stub_channel_factory(monkeypatch, channel)
monkeypatch.setattr(exploration_module, "get_settings", lambda: SimpleNamespace(poll_reply_timeout=0.05))
session_id = (await _create_session(client)).json()["id"]
resp = await client.post(
f"/api/exploration/sessions/{session_id}/messages",
json={"content": "有人在吗"},
)
assert resp.status_code == 502
assert ExplorationSessionRepository(seeded_db).get(session_id).turn_count == 1
async def test_close_normalizes_experience(seeded_db, mock_channel, client):
session_id = (await _create_session(client)).json()["id"]
resp = await client.post(
f"/api/exploration/sessions/{session_id}/close",
json={
"experience": {
"blockers": [42, {"not": "a string"}],
"misled": "不是列表",
"emotion": "暴怒!!!",
}
},
)
assert resp.status_code == 200, resp.text
experience = resp.json()["experience"]
assert experience["goal_achieved"] is False
assert experience["blockers"] == ["42"]
assert experience["misled"] == []
assert experience["emotion"] == "neutral"
async def test_close_twice_rejected(seeded_db, mock_channel, client):
session_id = (await _create_session(client)).json()["id"]
payload = {"experience": {"goal_achieved": True}}
assert (await client.post(f"/api/exploration/sessions/{session_id}/close", json=payload)).status_code == 200
resp = await client.post(f"/api/exploration/sessions/{session_id}/close", json=payload)
assert resp.status_code == 409
async def test_close_triggers_judge_review(seeded_db, mock_channel, client, monkeypatch):
from agenteval.web.routers import exploration as exploration_router
started: list[str] = []
monkeypatch.setattr(exploration_router, "start_judge_review", started.append)
session_id = (await _create_session(client)).json()["id"]
resp = await client.post(
f"/api/exploration/sessions/{session_id}/close", json={"experience": {"goal_achieved": True}}
)
assert resp.status_code == 200, resp.text
assert started == [session_id]
# ---------------------------------------------------------------- read endpoints (ticket 06)
async def test_list_campaign_sessions_returns_lifecycle_fields(seeded_db, mock_channel, client):
resp = await _create_session(client)
assert resp.status_code == 200, resp.text
session_id = resp.json()["id"]
await client.post(f"/api/exploration/sessions/{session_id}/messages", json={"content": "查账单"})
resp = await client.post(
f"/api/exploration/sessions/{session_id}/close",
json={"experience": {"goal_achieved": True, "blockers": [], "misled": ["跳转误导"], "notes": "绕了三圈"}},
)
assert resp.status_code == 200, resp.text
resp = await client.get("/api/exploration/campaigns/c-1/sessions")
assert resp.status_code == 200, resp.text
sessions = resp.json()["sessions"]
assert len(sessions) == 1
entry = sessions[0]
assert entry["id"] == session_id
assert entry["status"] == "completed"
assert entry["turn_count"] == 1
assert entry["goal"] == "查询本月账单并完成缴费"
assert entry["persona"]["name"] == "急性子用户"
assert entry["experience"]["goal_achieved"] is True
assert entry["experience"]["misled"] == ["跳转误导"]
assert entry["closed_at"] is not None
async def test_list_sessions_empty_campaign(seeded_db, client):
resp = await client.get("/api/exploration/campaigns/c-1/sessions")
assert resp.status_code == 200, resp.text
assert resp.json() == {"sessions": []}
async def test_list_session_messages_returns_conversation(seeded_db, mock_channel, client):
session_id = (await _create_session(client)).json()["id"]
await client.post(f"/api/exploration/sessions/{session_id}/messages", json={"content": "我要查账单"})
resp = await client.get(f"/api/exploration/sessions/{session_id}/messages")
assert resp.status_code == 200, resp.text
messages = resp.json()["messages"]
assert len(messages) == 2
user_msg, agent_msg = messages
assert user_msg["role"] == "user"
assert user_msg["content"] == "我要查账单"
assert user_msg["round_index"] == 1
assert agent_msg["role"] == "assistant"
assert agent_msg["content"] == "您好,请问有什么可以帮您?"
assert agent_msg["round_index"] == 1
assert isinstance(agent_msg["latency_ms"], int)
assert user_msg["created_at"] is not None
async def test_list_messages_unknown_session_returns_404(seeded_db, client):
resp = await client.get("/api/exploration/sessions/nope/messages")
assert resp.status_code == 404
# ---------------------------------------------------------------- migration
def test_exploration_migration_on_existing_db(tmp_path, monkeypatch):
"""Alembic migration applies on an existing DB at the previous head.
Brand-new DBs take the create_all path (exercised by every test above via
the db_session fixture, which creates the new tables from metadata).
"""
from pathlib import Path
from agenteval.storage import db as db_module
from alembic import command
from alembic.config import Config
from sqlalchemy import create_engine, inspect
database_url = f"sqlite:///{tmp_path / 'exploration.db'}"
monkeypatch.setattr(db_module, "DATABASE_URL", database_url)
config = Config(str(Path(__file__).resolve().parents[2] / "alembic.ini"))
# 既有库先例:基线迁移假设 create_all 建好的表已存在,先 stamp 基线前状态
from sqlmodel import SQLModel
engine = create_engine(database_url)
SQLModel.metadata.create_all(engine)
engine.dispose()
with create_engine(database_url).begin() as connection:
from sqlalchemy import text
connection.execute(text("DROP TABLE IF EXISTS exploration_sessions"))
connection.execute(text("DROP TABLE IF EXISTS exploration_messages"))
connection.execute(text("ALTER TABLE campaigns DROP COLUMN exploration_seeds"))
connection.execute(text("ALTER TABLE campaigns DROP COLUMN exploration_budget"))
connection.execute(text("ALTER TABLE campaigns DROP COLUMN last_patrolled_at"))
connection.execute(text("DROP TABLE IF EXISTS alembic_version"))
command.stamp(config, "f2a9b7c34d18")
command.upgrade(config, "head")
inspector = inspect(create_engine(database_url))
tables = set(inspector.get_table_names())
assert "exploration_sessions" in tables
assert "exploration_messages" in tables
session_cols = {c["name"] for c in inspector.get_columns("exploration_sessions")}
assert {
"campaign_id",
"target_id",
"persona",
"goal",
"seed_ref",
"status",
"triggered_by",
"experience",
"judge_review",
"turn_count",
"closed_at",
} <= session_cols
message_cols = {c["name"] for c in inspector.get_columns("exploration_messages")}
assert {"session_id", "round_index", "role", "content", "latency_ms"} <= message_cols