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