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