Files
group_fqcd_jr/tests/unit/service/test_memory_service.py
T

210 lines
7.8 KiB
Python
Raw Normal View History

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