249 lines
9.9 KiB
Python
249 lines
9.9 KiB
Python
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
|
|
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()
|