Files
agentos/tests/test_memory.py
T
2026-07-31 17:54:01 +08:00

63 lines
2.5 KiB
Python

import tempfile
import unittest
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
from agentos.memory import MemoryStore, estimate_tokens
class MemoryStoreTests(unittest.TestCase):
def setUp(self) -> None:
self.temporary = tempfile.TemporaryDirectory()
self.store = MemoryStore(Path(self.temporary.name) / "memory.sqlite3")
self.store.initialize()
def tearDown(self) -> None:
self.temporary.cleanup()
def test_remember_recall_and_forget(self) -> None:
record = self.store.remember("preference", "用户喜欢简洁的中文输出", importance=0.9)
recalled = self.store.recall("简洁中文", 3)
self.assertEqual(recalled[0].memory_id, record.memory_id)
self.assertTrue(self.store.forget(record.memory_id))
self.assertEqual(self.store.list_memories(), [])
def test_recall_ranks_chinese_relevance_before_unrelated_importance(self) -> None:
relevant = self.store.remember("experience", "创建日志文件时先检查目标目录", importance=0.4)
self.store.remember("preference", "用户喜欢蓝色界面", importance=0.95)
recalled = self.store.recall("创建日志", 5)
self.assertEqual(recalled[0].memory_id, relevant.memory_id)
self.assertGreater(recalled[0].relevance_score, 0)
def test_recall_respects_token_budget(self) -> None:
self.store.remember("experience", "创建日志" * 200, importance=0.7)
recalled = self.store.recall("创建日志", 5, token_budget=48)
self.assertEqual(len(recalled), 1)
self.assertLessEqual(32 + estimate_tokens(recalled[0].content), 48)
self.assertTrue(recalled[0].content.endswith("…"))
def test_principal_is_unique_while_active_and_reusable_after_release(self) -> None:
first = self.store.allocate_principal("instance-1", 280_000, 3)
second = self.store.allocate_principal("instance-2", 280_000, 3)
self.assertNotEqual(first, second)
self.store.release_principal("instance-1")
third = self.store.allocate_principal("instance-3", 280_000, 3)
self.assertEqual(third, first)
def test_principal_allocation_is_atomic_across_threads(self) -> None:
def allocate(index: int) -> int:
return self.store.allocate_principal(f"parallel-{index}", 281_000, 16)
with ThreadPoolExecutor(max_workers=8) as executor:
allocated = list(executor.map(allocate, range(8)))
self.assertEqual(len(set(allocated)), 8)
if __name__ == "__main__":
unittest.main()