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_candidate_never_updates_active_memory() -> None: """候选写入不读取或覆盖同键正式记忆,等待确认后再晋升。""" session = fake_session() session.scalar.side_effect = [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, status="candidate", ) assert created.status == "candidate" assert created.content == "稳健型" assert added == [created] @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