506 lines
19 KiB
Python
506 lines
19 KiB
Python
"""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.exploration import lifecycle as exploration_lifecycle
|
||
|
||
monkeypatch.setattr(exploration_lifecycle, "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 lifecycle's ChannelFactory at a test channel."""
|
||
from agenteval.exploration import lifecycle as lifecycle_module
|
||
|
||
class _StubFactory:
|
||
@staticmethod
|
||
def create(target):
|
||
return channel
|
||
|
||
monkeypatch.setattr(lifecycle_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):
|
||
"""活动级预算覆盖生效(票据 02):max_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.exploration import lifecycle 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.exploration import lifecycle as exploration_lifecycle
|
||
|
||
started: list[str] = []
|
||
monkeypatch.setattr(exploration_lifecycle, "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 intelligent_eval_messages"))
|
||
connection.execute(text("DROP TABLE IF EXISTS intelligent_eval_sessions"))
|
||
connection.execute(text("DROP TABLE IF EXISTS intelligent_evals"))
|
||
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
|