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()