265 lines
12 KiB
Python
265 lines
12 KiB
Python
"""记忆抽取集成测试(真实 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])
|