Files
group_fqcd_jr/tests/integration/test_complete_run.py
T

180 lines
6.7 KiB
Python
Raw 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.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()