142 lines
5.5 KiB
Python
142 lines
5.5 KiB
Python
from __future__ import annotations
|
|
|
|
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
|
|
from owner_voice_pet.agent_memory import (
|
|
AgentConversationManager,
|
|
DisabledMemoryManager,
|
|
FaissIndexManifest,
|
|
FakeMemoryManager,
|
|
MemoryManagementPlan,
|
|
MemoryRecordInput,
|
|
MemoryWritePolicy,
|
|
SQLiteMemoryManager,
|
|
is_sensitive_memory_text,
|
|
)
|
|
from owner_voice_pet.conversation import ConversationContext
|
|
from owner_voice_pet.models import ErrorCode, ProviderError
|
|
|
|
|
|
class AgentMemoryTests(unittest.TestCase):
|
|
def test_fake_memory_manager_saves_and_searches_normal_memory(self) -> None:
|
|
memory = FakeMemoryManager()
|
|
saved = memory.save(MemoryRecordInput("preference", "用户喜欢 Python 项目"))
|
|
|
|
results = memory.search("Python", top_k=1)
|
|
|
|
self.assertEqual(results, [saved])
|
|
self.assertEqual(memory.search_queries, ["Python"])
|
|
self.assertEqual(saved.type, "preference")
|
|
self.assertEqual(saved.sensitivity, "normal")
|
|
self.assertTrue(saved.checksum)
|
|
|
|
def test_disabled_memory_manager_does_not_read_or_write(self) -> None:
|
|
memory = DisabledMemoryManager()
|
|
|
|
self.assertEqual(memory.search("Python"), [])
|
|
with self.assertRaises(ProviderError) as raised:
|
|
memory.save(MemoryRecordInput("fact", "不会保存"))
|
|
|
|
self.assertEqual(raised.exception.code, ErrorCode.VALIDATION_FAILED)
|
|
self.assertTrue(memory.health_check().ok)
|
|
|
|
def test_sqlite_memory_manager_persists_across_instances(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
db_path = Path(tmp) / "memory.sqlite3"
|
|
first = SQLiteMemoryManager(db_path)
|
|
saved = first.save(MemoryRecordInput("project", "Owner 项目正在做语音助手"))
|
|
|
|
second = SQLiteMemoryManager(db_path)
|
|
results = second.search("Owner 语音", top_k=3)
|
|
|
|
self.assertEqual([record.id for record in results], [saved.id])
|
|
self.assertTrue(second.health_check().ok)
|
|
|
|
def test_sqlite_memory_manager_rejects_empty_text(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
manager = SQLiteMemoryManager(Path(tmp) / "memory.sqlite3")
|
|
|
|
with self.assertRaises(ProviderError) as raised:
|
|
manager.save(MemoryRecordInput("fact", " "))
|
|
|
|
self.assertEqual(raised.exception.code, ErrorCode.VALIDATION_FAILED)
|
|
|
|
def test_faiss_manifest_detects_missing_and_mismatched_records(self) -> None:
|
|
memory = FakeMemoryManager()
|
|
first = memory.save(MemoryRecordInput("fact", "第一条"))
|
|
second = memory.save(MemoryRecordInput("fact", "第二条"))
|
|
manifest = FaissIndexManifest.from_records([first], embedding_model="fake-embedding")
|
|
broken = FaissIndexManifest(
|
|
embedding_model=manifest.embedding_model,
|
|
record_ids=(first.id, "missing-id"),
|
|
checksums={first.id: "wrong"},
|
|
)
|
|
|
|
errors = broken.consistency_errors([first, second])
|
|
|
|
self.assertTrue(any("missing-id" in error for error in errors))
|
|
self.assertTrue(any(first.id in error for error in errors))
|
|
self.assertTrue(any(second.id in error for error in errors))
|
|
|
|
def test_memory_write_policy_requires_confirmation_for_sensitive_text(self) -> None:
|
|
policy = MemoryWritePolicy()
|
|
|
|
safe = policy.evaluate(MemoryRecordInput("preference", "用户喜欢 Python"))
|
|
sensitive = policy.evaluate(MemoryRecordInput("fact", "我的 api key 是 sk-abcdefghijklmnop"))
|
|
|
|
self.assertTrue(safe.should_save)
|
|
self.assertFalse(safe.requires_confirmation)
|
|
self.assertFalse(sensitive.should_save)
|
|
self.assertTrue(sensitive.requires_confirmation)
|
|
self.assertTrue(is_sensitive_memory_text("密码是 123456"))
|
|
|
|
def test_conversation_manager_injects_memory_context(self) -> None:
|
|
context = ConversationContext()
|
|
context.append_assistant("我是小杰。")
|
|
memory = FakeMemoryManager()
|
|
memory.save(MemoryRecordInput("preference", "用户喜欢 Python"))
|
|
manager = AgentConversationManager(
|
|
context=context,
|
|
memory=memory,
|
|
memory_enabled=True,
|
|
memory_top_k=2,
|
|
)
|
|
|
|
messages = manager.build_messages_for_user("Python 项目怎么做?")
|
|
|
|
self.assertEqual(messages[0].role, "system")
|
|
self.assertEqual(messages[1].role, "system")
|
|
self.assertIn("长期记忆", messages[1].content)
|
|
self.assertIn("用户喜欢 Python", messages[1].content)
|
|
self.assertEqual(messages[-1].role, "user")
|
|
self.assertEqual(messages[-1].content, "Python 项目怎么做?")
|
|
|
|
def test_conversation_manager_skips_memory_when_disabled(self) -> None:
|
|
context = ConversationContext()
|
|
memory = FakeMemoryManager(enabled=False)
|
|
manager = AgentConversationManager(
|
|
context=context,
|
|
memory=memory,
|
|
memory_enabled=False,
|
|
)
|
|
|
|
messages = manager.build_messages_for_user("Python")
|
|
|
|
self.assertEqual([message.role for message in messages], ["system", "user"])
|
|
self.assertEqual(memory.search_queries, [])
|
|
|
|
def test_memory_management_plan_lists_user_controls(self) -> None:
|
|
plan = MemoryManagementPlan()
|
|
|
|
self.assertIn("list", plan.supported_commands)
|
|
self.assertIn("delete", plan.supported_commands)
|
|
self.assertIn("disable", plan.supported_commands)
|
|
self.assertIn("export", plan.supported_commands)
|
|
self.assertFalse(plan.default_enabled)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|