"""记忆抽取集成测试(真实 MySQL)。 P2 之后记忆链路的语义变化: - 记忆键来自**受控词表**(如 `preference:risk_level`),不再是 `conversation.{session_id}`; - 记忆内容是模型抽取出的**结构化值**,而不是用户原文整句。 抽取依赖模型端点,而本地 `model_endpoint_config` 为 0 行,端点缺失时按设计失败关闭。 因此本文件注入确定性替身抽取器,只验证与模型供应商无关的契约:正文回查、幂等边界、 客户隔离、证据落库与清理。 """ from datetime import UTC, datetime from uuid import uuid4 import pytest from sqlalchemy import delete, select from app.infrastructure.db import SessionFactory from app.model.conversation import ConversationMessage from app.model.memory import MemoryConflict, MemoryEvidence, MemoryUnit from app.model.platform import AgentRun, DomainEventOutbox, OutboxDelivery, RequestIdempotency from app.service.memory_extraction_service import ExtractedMemory from app.service.memory_service import MemoryService from app.worker.memory_extraction_worker import MemoryExtractionWorker from app.worker.outbox_worker import OutboxWorker KEY_PREFIX = "it-memory-" MEMORY_KEY = "preference:risk_level" EXTRACTED_VALUE = "稳健型" USER_FACT = f"{KEY_PREFIX}我偏好低风险稳健型基金" OTHER_FACT = f"{KEY_PREFIX}本人可承受中等风险" THIRD_FACT = f"{KEY_PREFIX}本人只做货币基金" PLAN_FACT = f"{KEY_PREFIX}我计划两年内买房" ASSISTANT_REPLY = "已为您记录该偏好。" class StubExtractor: """确定性替身抽取器:记录收到的正文,返回固定抽取结果,不发网络请求。""" def __init__(self, value: str = EXTRACTED_VALUE) -> None: self.value = value self.messages: list[str] = [] async def extract( self, *, message: str, agent_type: str = "customer_service" ) -> ExtractedMemory: del agent_type self.messages.append(message) return ExtractedMemory( memory_key=MEMORY_KEY, value=self.value, memory_type="preference", confidence=0.9, ) class Case: def __init__(self, customer_id: int, run_id: str, session_id: str) -> None: self.customer_id = customer_id self.run_id = run_id self.session_id = session_id @property def memory_key(self) -> str: """抽取结果对应的受控键;同一客户的记忆键与客户一一对应,便于隔离断言。""" return MEMORY_KEY def new_case(offset: int) -> Case: return Case( customer_id=uuid4().int % 10**14 + offset * 10**14, run_id=str(uuid4()), session_id=f"{KEY_PREFIX}{uuid4().hex[:24]}", ) async def seed(case: Case, *, user_fact: str) -> int: """建立 run 到用户消息与助手消息的真实链路,事件只携带 message_id。""" now = datetime.now(UTC).replace(tzinfo=None) async with SessionFactory() as session, session.begin(): request_message = ConversationMessage( session_id=case.session_id, customer_id=case.customer_id, portal="api", role="user", content=user_fact, trace_id=str(uuid4()), created_at=now, ) session.add(request_message) await session.flush() result_message = ConversationMessage( session_id=case.session_id, customer_id=case.customer_id, portal="agent", role="assistant", content=ASSISTANT_REPLY, trace_id=str(uuid4()), created_at=now, ) session.add(result_message) await session.flush() idempotency = RequestIdempotency( user_id=case.customer_id, session_id=case.session_id, agent_type="customer_service", idempotency_key=str(uuid4()), request_hash=uuid4().hex * 2, trace_id=str(uuid4()), status="completed", expire_at=now, created_at=now, updated_at=now, ) session.add(idempotency) await session.flush() session.add(AgentRun( run_id=case.run_id, idempotency_id=idempotency.id, session_id=case.session_id, user_id=case.customer_id, agent_type="customer_service", trace_id=str(uuid4()), request_message_id=request_message.id, result_message_id=result_message.id, status="succeeded", attempt_count=1, created_at=now, updated_at=now, )) return result_message.id async def enqueue(case: Case, message_id: int, *, event_id: str) -> None: now = datetime.now(UTC).replace(tzinfo=None) async with SessionFactory() as session, session.begin(): session.add(DomainEventOutbox( event_id=event_id, event_type="memory.extraction_requested", aggregate_type="agent_run", aggregate_id=case.run_id, trace_id=str(uuid4()), payload={"run_id": case.run_id, "message_id": message_id, "customer_id": case.customer_id}, status="pending", retry_count=0, occurred_at=now, created_at=now, updated_at=now, )) async def consume_once(case: Case, extractor: StubExtractor) -> bool: async with SessionFactory() as session: worker = OutboxWorker(session, { "memory.extraction_requested": MemoryExtractionWorker( session, extractor=extractor ).handle, }) return await worker.publish_one(aggregate_id=case.run_id) async def realtime_seen_count(extractor: StubExtractor, fact: str) -> int: return sum(1 for message in extractor.messages if message == fact) async def cleanup(cases: list[Case]) -> None: run_ids = [case.run_id for case in cases] session_ids = [case.session_id for case in cases] customer_ids = [case.customer_id for case in cases] async with SessionFactory() as session, session.begin(): message_ids = list(await session.scalars(select(ConversationMessage.id).where( ConversationMessage.session_id.in_(session_ids)))) memory_ids = list(await session.scalars(select(MemoryUnit.id).where( MemoryUnit.customer_id.in_(customer_ids)))) # 测试证据只可能挂在本次测试的会话消息或本次测试客户的记忆上。 await session.execute(delete(MemoryEvidence).where( (MemoryEvidence.source_record_id.in_([str(item) for item in message_ids])) | (MemoryEvidence.memory_id.in_(memory_ids)))) await session.execute(delete(MemoryConflict).where( (MemoryConflict.left_memory_id.in_(memory_ids)) | (MemoryConflict.right_memory_id.in_(memory_ids)))) await session.execute(delete(MemoryUnit).where( MemoryUnit.customer_id.in_(customer_ids))) event_ids = select(DomainEventOutbox.event_id).where( DomainEventOutbox.aggregate_id.in_(run_ids)) await session.execute(delete(OutboxDelivery).where(OutboxDelivery.event_id.in_(event_ids))) await session.execute(delete(DomainEventOutbox).where( DomainEventOutbox.aggregate_id.in_(run_ids))) await session.execute(delete(AgentRun).where(AgentRun.run_id.in_(run_ids))) await session.execute(delete(RequestIdempotency).where( RequestIdempotency.session_id.in_(session_ids))) await session.execute(delete(ConversationMessage).where( ConversationMessage.session_id.in_(session_ids))) @pytest.mark.integration @pytest.mark.asyncio async def test_duplicate_consumption_produces_single_memory() -> None: case = new_case(9) extractor = StubExtractor() event_id = str(uuid4()) try: message_id = await seed(case, user_fact=USER_FACT) await enqueue(case, message_id, event_id=event_id) assert await consume_once(case, extractor) # 重复投递同一事件:事件回到 pending 后再次消费。 async with SessionFactory() as session, session.begin(): event = await session.scalar(select(DomainEventOutbox).where( DomainEventOutbox.event_id == event_id)) assert event is not None event.status = "pending" event.published_at = None assert await consume_once(case, extractor) async with SessionFactory() as session: memories = list(await session.scalars(select(MemoryUnit).where( MemoryUnit.customer_id == case.customer_id, MemoryUnit.memory_key == case.memory_key))) assert len(memories) == 1 # 落库的是抽取结果,不是用户原文整句。 assert memories[0].content == EXTRACTED_VALUE assert memories[0].content != USER_FACT assert memories[0].memory_type == "preference" assert memories[0].evidence_count <= 1 evidence = list(await session.scalars(select(MemoryEvidence).where( MemoryEvidence.idempotency_key == f"memory.extraction_requested:{event_id}"))) assert len(evidence) == 1 assert evidence[0].memory_id == memories[0].id assert evidence[0].source_record_id is not None finally: await cleanup([case]) @pytest.mark.integration @pytest.mark.asyncio async def test_memory_is_not_recallable_across_customers() -> None: owner = new_case(8) other = new_case(7) # 两个客户用不同的抽取值,才能用内容区分"谁记住了什么"——受控键相同, # 单靠 memory_key 无法判定归属。 owner_extractor = StubExtractor("稳健型") other_extractor = StubExtractor("激进型") try: owner_message = await seed(owner, user_fact=OTHER_FACT) other_message = await seed(other, user_fact=THIRD_FACT) await enqueue(owner, owner_message, event_id=str(uuid4())) await enqueue(other, other_message, event_id=str(uuid4())) assert await consume_once(owner, owner_extractor) assert await consume_once(other, other_extractor) async with SessionFactory() as session: service = MemoryService(session) owner_memories = await service.recall(owner.customer_id) other_memories = await service.recall(other.customer_id) assert [item.memory_key for item in owner_memories] == [owner.memory_key] assert [item.memory_key for item in other_memories] == [other.memory_key] assert all(item.customer_id == owner.customer_id for item in owner_memories) assert all(item.customer_id == other.customer_id for item in other_memories) assert owner_memories[0].content == "稳健型" assert other_memories[0].content == "激进型" # 客户范围过滤:即使知道他人的内容,也查不到挂在他人名下的行。 assert await session.scalar(select(MemoryUnit.id).where( MemoryUnit.customer_id == owner.customer_id, MemoryUnit.content == "激进型")) is None assert await session.scalar(select(MemoryUnit.id).where( MemoryUnit.customer_id == other.customer_id, MemoryUnit.content == "稳健型")) is None finally: await cleanup([owner, other]) @pytest.mark.integration @pytest.mark.asyncio async def test_worker_reads_user_content_not_assistant_reply() -> None: case = new_case(6) extractor = StubExtractor() try: message_id = await seed(case, user_fact=PLAN_FACT) await enqueue(case, message_id, event_id=str(uuid4())) assert await consume_once(case, extractor) # 抽取器收到的必须是库中的用户消息正文,而不是助手回复。 assert extractor.messages == [PLAN_FACT] assert ASSISTANT_REPLY not in extractor.messages[0] async with SessionFactory() as session: memory = await session.scalar(select(MemoryUnit).where( MemoryUnit.customer_id == case.customer_id, MemoryUnit.memory_key == case.memory_key)) assert memory is not None assert memory.content == EXTRACTED_VALUE finally: await cleanup([case])