Files

249 lines
9.9 KiB
Python
Raw Permalink Normal View History

2026-09-09 21:55:37 +08:00
from datetime import UTC, datetime
from uuid import uuid4
import pytest
from sqlalchemy import delete, select
from sqlalchemy.exc import IntegrityError
from app.core.contracts import AgentResult, CoreResult
from app.infrastructure.db import SessionFactory
from app.model.audit import InteractionAudit
from app.model.conversation import ConversationMessage
from app.model.platform import AgentRun, DomainEventOutbox, RequestIdempotency
from app.model.risk import RiskUser
from app.model.session import ConversationSession
2026-09-09 21:55:37 +08:00
from app.service.agent_persistence_service import AgentPersistenceService
@pytest.mark.integration
@pytest.mark.asyncio
async def test_complete_run_writes_memory_event_in_same_transaction() -> None:
now = datetime.now(UTC).replace(tzinfo=None)
session_id, trace_id, run_id = f"complete-{uuid4()}", str(uuid4()), str(uuid4())
async with SessionFactory() as session:
idem = RequestIdempotency(
user_id=1,
session_id=session_id,
agent_type="customer_service",
idempotency_key=f"complete-key-{uuid4()}",
request_hash="b" * 64,
trace_id=trace_id,
status="processing",
expire_at=now,
created_at=now,
updated_at=now,
)
user_message = ConversationMessage(
session_id=session_id,
customer_id=1,
portal="api",
role="user",
content="hello",
trace_id=trace_id,
created_at=now,
)
session.add_all([idem, user_message])
await session.flush()
run = AgentRun(
run_id=run_id,
idempotency_id=idem.id,
session_id=session_id,
user_id=1,
agent_type="customer_service",
trace_id=trace_id,
request_message_id=user_message.id,
created_at=now,
updated_at=now,
)
session.add(run)
await session.commit()
try:
result = AgentResult(run_id=run_id, result=CoreResult(text="done"))
message_id = await AgentPersistenceService(session).complete_run(run_id, result)
events = list(
(
await session.scalars(
select(DomainEventOutbox).where(DomainEventOutbox.aggregate_id == run_id)
)
).all()
)
assert message_id > 0
assert {event.event_type for event in events} == {
"agent.run_completed",
"memory.extraction_requested",
}
finally:
await session.execute(
delete(DomainEventOutbox).where(DomainEventOutbox.aggregate_id == run_id)
)
await session.execute(delete(AgentRun).where(AgentRun.run_id == run_id))
await session.execute(
delete(RequestIdempotency).where(RequestIdempotency.id == idem.id)
)
await session.execute(
delete(ConversationMessage).where(ConversationMessage.session_id == session_id)
)
await session.commit()
@pytest.mark.integration
@pytest.mark.asyncio
async def test_complete_run_rolls_back_every_write_on_outbox_conflict(monkeypatch) -> None:
now = datetime.now(UTC).replace(tzinfo=None)
session_id, trace_id, run_id = f"rollback-{uuid4()}", str(uuid4()), str(uuid4())
idem_id = 0
async with SessionFactory() as session:
idem = RequestIdempotency(
user_id=1,
session_id=session_id,
agent_type="customer_service",
idempotency_key=f"rollback-key-{uuid4()}",
request_hash="c" * 64,
trace_id=trace_id,
status="processing",
expire_at=now,
created_at=now,
updated_at=now,
)
user_message = ConversationMessage(
session_id=session_id,
customer_id=1,
portal="api",
role="user",
content="rollback",
trace_id=trace_id,
created_at=now,
)
session.add_all([idem, user_message])
await session.flush()
idem_id = idem.id
session.add(
AgentRun(
run_id=run_id,
idempotency_id=idem.id,
session_id=session_id,
user_id=1,
agent_type="customer_service",
trace_id=trace_id,
request_message_id=user_message.id,
created_at=now,
updated_at=now,
)
)
await session.commit()
fixed_event_id = uuid4()
monkeypatch.setattr("app.service.agent_persistence_service.uuid4", lambda: fixed_event_id)
try:
async with SessionFactory() as session:
result = AgentResult(run_id=run_id, result=CoreResult(text="must rollback"))
with pytest.raises(IntegrityError):
await AgentPersistenceService(session).complete_run(run_id, result)
async with SessionFactory() as session:
run = await session.scalar(select(AgentRun).where(AgentRun.run_id == run_id))
idem = await session.get(RequestIdempotency, idem_id)
messages = list(
await session.scalars(
select(ConversationMessage).where(ConversationMessage.session_id == session_id)
)
)
events = list(
await session.scalars(
select(DomainEventOutbox).where(DomainEventOutbox.aggregate_id == run_id)
)
)
audits = list(
await session.scalars(
select(InteractionAudit).where(InteractionAudit.session_id == session_id)
)
)
assert run is not None and run.status == "queued" and run.result_message_id is None
assert idem is not None and idem.status == "processing"
assert len(messages) == 1 and messages[0].role == "user"
assert events == []
assert audits == []
finally:
async with SessionFactory() as session:
await session.execute(
delete(InteractionAudit).where(InteractionAudit.session_id == session_id)
)
await session.execute(
delete(DomainEventOutbox).where(DomainEventOutbox.aggregate_id == run_id)
)
await session.execute(delete(AgentRun).where(AgentRun.run_id == run_id))
await session.execute(
delete(RequestIdempotency).where(RequestIdempotency.id == idem_id)
)
await session.execute(
delete(ConversationMessage).where(ConversationMessage.session_id == session_id)
)
await session.commit()
@pytest.mark.integration
@pytest.mark.asyncio
async def test_customer_service_clarification_advances_real_session_round() -> None:
"""澄清计数必须在客服运行成功持久化的同一 MySQL 事务内递增。"""
now = datetime.now(UTC).replace(tzinfo=None)
session_id, trace_id, run_id = f"clarify-{uuid4()}", str(uuid4()), str(uuid4())
idem_id = 0
async with SessionFactory() as session:
# 会话表有真实的 sys_user 外键;测试只复用本地已有用户,绝不为此功能伪造账号。
user_id = await session.scalar(select(RiskUser.id).limit(1))
if user_id is None:
pytest.skip("本地 MySQL 没有可用 sys_user,无法验证会话外键链路")
session.add(ConversationSession(
session_id=session_id, user_id=user_id, agent_type="customer_service", portal="api",
status="active", clarification_round=0,
))
idem = RequestIdempotency(
user_id=user_id, session_id=session_id, agent_type="customer_service",
idempotency_key=f"clarify-key-{uuid4()}", request_hash="d" * 64,
trace_id=trace_id, status="processing", expire_at=now,
created_at=now, updated_at=now,
)
user_message = ConversationMessage(
session_id=session_id, customer_id=user_id, portal="api", role="user",
content="它的费率是多少", trace_id=trace_id, created_at=now,
)
session.add_all([idem, user_message])
await session.flush()
idem_id = idem.id
session.add(AgentRun(
run_id=run_id, idempotency_id=idem.id, session_id=session_id, user_id=user_id,
agent_type="customer_service", trace_id=trace_id,
request_message_id=user_message.id, created_at=now, updated_at=now,
))
await session.commit()
try:
async with SessionFactory() as session:
await AgentPersistenceService(session).complete_run(
run_id,
AgentResult(
run_id=run_id,
result=CoreResult(text="请提供产品名称或代码。", clarification_required=True),
),
memory_extraction_requested=False,
)
row = await session.scalar(select(ConversationSession).where(
ConversationSession.session_id == session_id
))
assert row is not None and row.clarification_round == 1
finally:
async with SessionFactory() as session:
await session.execute(
delete(DomainEventOutbox).where(DomainEventOutbox.aggregate_id == run_id)
)
await session.execute(delete(AgentRun).where(AgentRun.run_id == run_id))
await session.execute(
delete(RequestIdempotency).where(RequestIdempotency.id == idem_id)
)
await session.execute(
delete(ConversationMessage).where(ConversationMessage.session_id == session_id)
)
await session.execute(
delete(ConversationSession).where(ConversationSession.session_id == session_id)
)
await session.commit()