from __future__ import annotations import unittest from tempfile import TemporaryDirectory import torch from dynamic_memory_lab.memory_os_v2 import ( KVBudgetManagerV2, MemoryOSV2, MemoryRouterV2, PagedMemoryBankV2, STATUS_ACTIVE, STATUS_QUARANTINED, STATUS_SUPERSEDED, ) from dynamic_memory_lab.tiered_memory_store_v2 import TieredMemoryStoreV2 class MemoryOSV2Test(unittest.TestCase): def setUp(self) -> None: torch.manual_seed(7) self.router = MemoryRouterV2(16, router_dim=8, num_heads=2, max_hops=3) self.bank = PagedMemoryBankV2( 16, router=self.router, page_capacity=2, max_pages=512, hot_pages=2, top_k_pages=2, top_k_records=4, max_hops=3, coarse_index_bits=8, ) def test_router_scores_and_compressed_address(self) -> None: query = torch.randn(4, 16) candidates = torch.randn(4, 5, 16) output = self.router(query, candidates) self.assertEqual(tuple(output["scores"].shape), (4, 5)) self.assertEqual(tuple(output["head_scores"].shape), (4, 5, 2)) self.assertEqual(tuple(self.router.encode_key(query).shape), (4, 8)) def test_write_version_and_conflict_resolution(self) -> None: first, first_action = self.bank.write( text="我住在上海", key=torch.randn(16), entity="user", attribute="city", value="上海", confidence=0.9, ) second, second_action = self.bank.write( text="我搬到了杭州", key=torch.randn(16), entity="user", attribute="city", value="杭州", confidence=0.95, ) self.assertEqual(first_action, "inserted") self.assertEqual(second_action, "updated") self.assertEqual(first.status, STATUS_SUPERSEDED) self.assertEqual(second.status, STATUS_ACTIVE) self.assertEqual(second.version, 1) self.assertEqual(self.bank.active_by_conflict["user::city"], second.record_id) def test_quarantine_and_approval(self) -> None: record, action = self.bank.write( text="未经确认的推断", key=torch.randn(16), trusted=False, confidence=0.1, ) self.assertEqual(action, "quarantined") self.assertEqual(record.status, STATUS_QUARANTINED) self.assertNotIn(record.record_id, self.bank.records) approved = self.bank.approve(record.record_id) self.assertEqual(approved.status, STATUS_ACTIVE) self.assertIn(approved.record_id, self.bank.records) def test_multi_hop_and_slot_replacement(self) -> None: second, _ = self.bank.write(text="项目的第二个节点", key=torch.randn(16), slot_index=2) third, _ = self.bank.write(text="项目的第三个节点", key=torch.randn(16), slot_index=3) first, _ = self.bank.write( text="项目的第一个节点", key=torch.randn(16), related_ids=[second.record_id, third.record_id], slot_index=1, ) replacement, action = self.bank.write( text="项目的第一个节点修正版", key=torch.randn(16), related_ids=[second.record_id], slot_index=1, ) self.assertEqual(action, "updated") self.assertEqual(first.status, STATUS_SUPERSEDED) self.assertEqual(replacement.status, STATUS_ACTIVE) records, decision = self.bank.query( query_key=replacement.key, top_k_pages=2, top_k_records=4, max_hops=3, ) ids = {record.record_id for record in records} self.assertIn(replacement.record_id, ids) self.assertGreaterEqual(decision.hop_count, 1) def test_coarse_index_bounds_candidate_pages(self) -> None: for index in range(300): key = torch.zeros(16) key[index % 16] = 1.0 key[(index * 7 + 3) % 16] += 0.05 self.bank.write(text=f"memory-{index}", key=key, importance=0.2) query_key = torch.zeros(16) query_key[3] = 1.0 self.bank.query(query_key=query_key, top_k_pages=2, top_k_records=2) stats = self.bank.stats() self.assertGreater(stats["pages"], 128) self.assertLess(stats["last_coarse_candidates"], stats["pages"]) def test_export_and_restore(self) -> None: record, _ = self.bank.write( text="可持久化事实", key=torch.randn(16), token_ids=torch.tensor([4, 5, 6]), token_mask=torch.tensor([True, True, True]), ) payload = self.bank.export_payload() restored = PagedMemoryBankV2.from_payload(payload, router=self.router) self.assertEqual(restored.stats()["active_records"], 1) self.assertTrue(torch.equal(restored.records[record.record_id].token_ids, torch.tensor([4, 5, 6]))) self.assertEqual(restored.records[record.record_id].page_id, record.page_id) def test_lazy_capacity_is_bounded(self) -> None: bank = PagedMemoryBankV2( 16, router=self.router, page_capacity=1, max_pages=2, hot_pages=0, coarse_index_bits=8, ) bank.write(text="容量一", key=torch.randn(16)) bank.write(text="容量二", key=torch.randn(16)) self.assertEqual(bank.stats()["pages"], 2) with self.assertRaises(RuntimeError): bank.write(text="容量三", key=torch.randn(16)) def test_memory_os_and_kv_budget(self) -> None: os_v2 = MemoryOSV2(16, router=self.router) record, action = os_v2.write( text="可靠事实", key=torch.randn(16), importance=0.9, confidence=0.9, ) self.assertEqual(action, "inserted") self.assertIn(record.record_id, os_v2.bank.records) budget = KVBudgetManagerV2(max_tokens=128, hard_max_tokens=512, keep_recent_tokens=32) self.assertFalse(budget.needs_compaction(100)) self.assertTrue(budget.needs_compaction(120)) self.assertEqual(budget.overflow(140), 12) def test_batch_context_records_keep_all_chunks_active(self) -> None: os_v2 = MemoryOSV2(16, router=self.router) output = os_v2.write_batch( [ { "text": "context_chunk:0:0:4", "key": torch.randn(16), "memory_type": "context_chunk", "importance": 0.55, "confidence": 0.8, "trusted": True, "force": True, }, { "text": "context_chunk:0:4:8", "key": torch.randn(16), "memory_type": "context_chunk", "importance": 0.55, "confidence": 0.8, "trusted": True, "force": True, }, ] ) self.assertEqual(len(output), 2) self.assertEqual(os_v2.stats()["active_records"], 2) def test_management_list_edit_retract_and_audit(self) -> None: os_v2 = MemoryOSV2(16, router=self.router) record, _ = os_v2.write( text="用户喜欢蓝色", key=torch.randn(16), entity="user", attribute="color", value="蓝色", confidence=0.95, importance=0.9, ) listed = os_v2.list_records(query_text="蓝色", status="active", limit=10) self.assertEqual([item.record_id for item in listed], [record.record_id]) edited = os_v2.edit_record( record.record_id, text="用户喜欢绿色", entity="user", attribute="color", value="绿色", evidence=["user_correction"], ) self.assertEqual(edited.version, 1) self.assertEqual(edited.supersedes, record.record_id) self.assertEqual(os_v2.bank.records[record.record_id].status, STATUS_SUPERSEDED) self.assertEqual(os_v2.list_records(query_text="绿色")[0].record_id, edited.record_id) os_v2.retract_record(edited.record_id) self.assertEqual(os_v2.bank.records[edited.record_id].status, "retracted") audit = os_v2.audit() self.assertTrue(audit["healthy"], audit) all_records = os_v2.list_records(status="all", limit=10) self.assertEqual(len(all_records), 2) def test_tiered_storage_restarts_and_evicts_cold_records(self) -> None: with TemporaryDirectory() as directory: path = f"{directory}/memory.sqlite" store = TieredMemoryStoreV2(path, key_dim=8, page_capacity=2) bank = PagedMemoryBankV2( 16, router=self.router, page_capacity=2, max_pages=64, hot_pages=1, top_k_pages=2, top_k_records=2, tier_store=store, max_resident_pages=1, coarse_index_bits=8, ) for index in range(8): key = torch.zeros(16) key[index % 8] = 1.0 bank.write( text=f"tiered-memory-{index}", key=key, entity="user", attribute=f"attr-{index}", value=f"value-{index}", importance=0.1 if index < 7 else 1.0, confidence=0.95, ) stats = bank.stats() self.assertEqual(stats["storage_mode"], "tiered") self.assertGreaterEqual(stats["pages"], 4) self.assertGreater(stats["cold_pages"], 0) self.assertLess(stats["resident_records"], stats["records"]) store.close() reopened_store = TieredMemoryStoreV2(path, key_dim=8, page_capacity=2) reopened = PagedMemoryBankV2( 16, router=self.router, page_capacity=2, max_pages=64, hot_pages=1, top_k_pages=2, top_k_records=2, tier_store=reopened_store, max_resident_pages=1, coarse_index_bits=8, ) records, decision = reopened.query( query_key=torch.nn.functional.one_hot(torch.tensor(3), num_classes=16).float(), query_text="tiered-memory-3", top_k_pages=2, top_k_records=2, ) self.assertTrue(records) self.assertTrue(any(item.text == "tiered-memory-3" for item in records)) self.assertGreaterEqual(decision.hop_count, 1) quarantined, action = reopened.write( text="待审批事实", key=torch.randn(16), trusted=False, confidence=0.1, ) self.assertEqual(action, "quarantined") reopened_store.close() final_store = TieredMemoryStoreV2(path, key_dim=8, page_capacity=2) final_bank = PagedMemoryBankV2( 16, router=self.router, page_capacity=2, max_pages=64, hot_pages=1, tier_store=final_store, max_resident_pages=2, coarse_index_bits=8, ) self.assertIn(quarantined.record_id, final_bank.quarantine) approved = final_bank.approve(quarantined.record_id) self.assertEqual(approved.status, STATUS_ACTIVE) final_store.close() if __name__ == "__main__": unittest.main()