210 lines
7.8 KiB
Python
210 lines
7.8 KiB
Python
from datetime import UTC, datetime
|
|
from typing import Any
|
|
from unittest.mock import AsyncMock, Mock
|
|
|
|
import pytest
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.model.memory import MemoryUnit
|
|
from app.service.memory_recall_service import CACHE_KEY_PREFIX, MemoryRecallService
|
|
from app.service.memory_service import MemoryService
|
|
|
|
NOW = datetime(2026, 9, 9, 0, 0, 0, tzinfo=UTC).replace(tzinfo=None)
|
|
|
|
|
|
class RecordingCache:
|
|
"""缓存替身:记录被删除的键,可注入删除故障(故障必须不阻塞写入)。"""
|
|
|
|
def __init__(self, *, fail: bool = False) -> None:
|
|
self.fail = fail
|
|
self.deleted: list[tuple[str, ...]] = []
|
|
|
|
async def delete(self, *keys: str) -> int:
|
|
self.deleted.append(tuple(keys))
|
|
if self.fail:
|
|
raise ConnectionError("redis unavailable")
|
|
return len(keys)
|
|
|
|
|
|
def memory(
|
|
memory_id: int, *, customer_id: int = 7, key: str = "conversation.session-1",
|
|
content: str = "旧内容", version: int = 1,
|
|
) -> MemoryUnit:
|
|
return MemoryUnit(
|
|
id=memory_id, memory_uuid=f"uuid-{memory_id}", customer_id=customer_id, memory_key=key,
|
|
content=content, memory_type="事实候选", source_type="用户自述", source_confidence=0.65,
|
|
confidence=0.65, evidence_count=0, conflict_count=0, recall_count=0, status="active",
|
|
valid_from=NOW, version=version, created_at=NOW, updated_at=NOW,
|
|
)
|
|
|
|
|
|
def fake_session() -> Any:
|
|
session = AsyncMock(spec=AsyncSession)
|
|
nested = Mock()
|
|
nested.__aenter__ = AsyncMock(return_value=None)
|
|
nested.__aexit__ = AsyncMock(return_value=False)
|
|
session.begin_nested = Mock(return_value=nested)
|
|
session.add = Mock()
|
|
return session
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_recall_filters_by_customer() -> None:
|
|
"""召回必须带 customer_id 过滤,不能跨客户命中。"""
|
|
session = fake_session()
|
|
session.scalars.return_value = [memory(1)]
|
|
await MemoryService(session).recall(7)
|
|
statement = str(session.scalars.await_args.args[0])
|
|
assert "memory_unit.customer_id" in statement
|
|
assert "memory_unit.status" in statement
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_upsert_conflict_does_not_reference_itself() -> None:
|
|
existing = memory(5, content="旧内容")
|
|
session = fake_session()
|
|
session.scalar.side_effect = [existing, None]
|
|
added: list[Any] = []
|
|
session.add = Mock(side_effect=added.append)
|
|
|
|
updated = await MemoryService(session).upsert(7, "conversation.session-1", "新内容")
|
|
|
|
conflict = next(item for item in added if type(item).__name__ == "MemoryConflict")
|
|
assert conflict.left_memory_id == 5
|
|
assert conflict.right_memory_id != conflict.left_memory_id
|
|
assert conflict.status == "auto_resolved"
|
|
assert conflict.severity == "low"
|
|
assert conflict.resolution is not None
|
|
assert conflict.winner_memory_id == 5
|
|
assert updated is existing
|
|
assert updated.content == "新内容"
|
|
assert updated.conflict_count == 1
|
|
assert updated.version == 2
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_upsert_leaves_primary_key_to_database() -> None:
|
|
"""新建记忆不得显式写 id=0,主键交给自增列。"""
|
|
session = fake_session()
|
|
session.scalar.side_effect = [None, None]
|
|
added: list[Any] = []
|
|
session.add = Mock(side_effect=added.append)
|
|
|
|
created = await MemoryService(session).upsert(7, "conversation.session-1", "新内容")
|
|
|
|
assert created.id is None
|
|
assert created.status == "active"
|
|
assert created.customer_id == 7
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_upsert_persists_structured_value_when_created() -> None:
|
|
"""新建记忆时结构化值随受控键一起落库,调用方不再需要落原文。"""
|
|
session = fake_session()
|
|
session.scalar.side_effect = [None, None]
|
|
added: list[Any] = []
|
|
session.add = Mock(side_effect=added.append)
|
|
|
|
created = await MemoryService(session).upsert(
|
|
7, "preference:risk_level", "稳健型", memory_type="preference", confidence=0.9,
|
|
structured_value={"memory_key": "preference:risk_level", "value": "稳健型"},
|
|
)
|
|
|
|
assert created.structured_value == {
|
|
"memory_key": "preference:risk_level", "value": "稳健型"}
|
|
assert created.memory_type == "preference"
|
|
assert created.memory_key == "preference:risk_level"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_upsert_refreshes_structured_value_on_existing_memory() -> None:
|
|
existing = memory(5, key="preference:risk_level", content="保守型")
|
|
session = fake_session()
|
|
session.scalar.side_effect = [existing, 6]
|
|
session.add = Mock()
|
|
|
|
updated = await MemoryService(session).upsert(
|
|
7, "preference:risk_level", "稳健型", memory_type="preference",
|
|
structured_value={"value": "稳健型"},
|
|
)
|
|
|
|
assert updated is existing
|
|
assert updated.content == "稳健型"
|
|
assert updated.structured_value == {"value": "稳健型"}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_record_evidence_is_idempotent_by_key() -> None:
|
|
"""同一 idempotency_key 重复写入不再新增证据、不再增加证据计数。"""
|
|
target = memory(5)
|
|
session = fake_session()
|
|
session.scalar.side_effect = [77]
|
|
added: list[Any] = []
|
|
session.add = Mock(side_effect=added.append)
|
|
|
|
recorded = await MemoryService(session).record_evidence(
|
|
target, idempotency_key="memory.extraction_requested:event-1", evidence_type="对话",
|
|
excerpt="我偏好低风险", snapshot=None, weight=0.05)
|
|
|
|
assert recorded is False
|
|
assert added == []
|
|
assert target.evidence_count == 0
|
|
|
|
|
|
# A1:写入路径此前完全不失效召回热缓存,新记忆在 TTL(300 秒)内召回不到。
|
|
@pytest.mark.asyncio
|
|
async def test_upsert_invalidates_customer_recall_cache() -> None:
|
|
"""新建记忆后必须删除该客户的召回热缓存键,且键集与召回服务完全一致。"""
|
|
session = fake_session()
|
|
session.scalar.side_effect = [None, None]
|
|
session.add = Mock()
|
|
cache = RecordingCache()
|
|
|
|
created = await MemoryService(session, cache=cache).upsert(
|
|
7, "preference:risk_level", "稳健型")
|
|
|
|
assert created.memory_key == "preference:risk_level"
|
|
assert cache.deleted == [tuple(MemoryRecallService.cache_keys(7))]
|
|
assert all(key.startswith(f"{CACHE_KEY_PREFIX}:7:") for key in cache.deleted[0])
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_update_path_invalidates_customer_recall_cache() -> None:
|
|
"""更新既有记忆同样改变召回结果,必须一并失效缓存。"""
|
|
existing = memory(5, key="preference:risk_level", content="保守型")
|
|
session = fake_session()
|
|
session.scalar.side_effect = [existing, 6]
|
|
session.add = Mock()
|
|
cache = RecordingCache()
|
|
|
|
await MemoryService(session, cache=cache).upsert(7, "preference:risk_level", "稳健型")
|
|
|
|
assert cache.deleted == [tuple(MemoryRecallService.cache_keys(7))]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_invalidate_memory_drops_recall_cache() -> None:
|
|
"""单条失效也会让已失效记忆在 TTL 内继续被召回,因此同样需要失效缓存。"""
|
|
session = fake_session()
|
|
session.scalar.side_effect = [memory(5)]
|
|
cache = RecordingCache()
|
|
|
|
assert await MemoryService(session, cache=cache).invalidate("uuid-5", 7) is True
|
|
assert cache.deleted == [tuple(MemoryRecallService.cache_keys(7))]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_cache_failure_never_blocks_memory_write() -> None:
|
|
"""缓存只是可重建的加速层:删除失败不得阻塞写入主流程。"""
|
|
session = fake_session()
|
|
session.scalar.side_effect = [None, None]
|
|
session.add = Mock()
|
|
cache = RecordingCache(fail=True)
|
|
|
|
created = await MemoryService(session, cache=cache).upsert(
|
|
7, "preference:risk_level", "稳健型")
|
|
|
|
assert created.content == "稳健型"
|
|
assert cache.deleted # 失效动作尝试过,但异常被吞
|
|
assert await MemoryService(session, cache=None).invalidate_recall_cache(7) == 0
|