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