Engram-protocol/engram
0
1"""2ENGRAM Protocol — Manifold Index Tests3Tests for FAISS IndexFlatIP add/search/remove/persist (D2, D4).4"""5 6from __future__ import annotations7 8from pathlib import Path9 10import numpy as np11import pytest12import torch13 14from kvcos.core.manifold_index import IndexEntry, ManifoldIndex15 16 17def _entry(cid: str = "c1", model: str = "llama") -> IndexEntry:18 return IndexEntry(19 cache_id=cid, task_description="test",20 model_id=model, created_at="2026-01-01T00:00:00Z",21 context_len=256, l2_norm=1.0,22 )23 24 25class TestAddAndSearch:26 """Add vectors, search via MIPS."""27 28 def test_add_increments(self) -> None:29 idx = ManifoldIndex(dim=8)30 idx.add(torch.randn(8), _entry("a"))31 idx.add(torch.randn(8), _entry("b"))32 assert idx.n_entries == 233 34 def test_search_returns_correct_order(self) -> None:35 idx = ManifoldIndex(dim=4)36 v1 = torch.tensor([1.0, 0.0, 0.0, 0.0])37 v2 = torch.tensor([0.0, 1.0, 0.0, 0.0])38 idx.add(v1, _entry("close"))39 idx.add(v2, _entry("far"))40 41 query = torch.tensor([1.0, 0.0, 0.0, 0.0])42 results = idx.search(query, top_k=2)43 assert results[0]["cache_id"] == "close"44 assert results[0]["similarity"] > results[1]["similarity"]45 46 def test_search_empty_returns_empty(self) -> None:47 idx = ManifoldIndex(dim=4)48 results = idx.search(torch.randn(4), top_k=5)49 assert results == []50 51 def test_model_filter(self) -> None:52 idx = ManifoldIndex(dim=4)53 idx.add(torch.randn(4), _entry("a", model="llama"))54 idx.add(torch.randn(4), _entry("b", model="phi"))55 results = idx.search(torch.randn(4), top_k=10, model_id="phi")56 assert all(r["model_id"] == "phi" for r in results)57 58 59class TestRemoveAndRebuild:60 """Remove entries and rebuild index."""61 62 def test_remove_hides_from_search(self) -> None:63 idx = ManifoldIndex(dim=4)64 v = torch.tensor([1.0, 0.0, 0.0, 0.0])65 idx.add(v, _entry("target"))66 assert idx.remove("target")67 results = idx.search(v, top_k=1)68 assert len(results) == 069 70 def test_rebuild_compacts(self) -> None:71 idx = ManifoldIndex(dim=4)72 for i in range(5):73 idx.add(torch.randn(4), _entry(f"c{i}"))74 idx.remove("c1")75 idx.remove("c3")76 active = idx.rebuild()77 assert active == 378 79 80class TestPersistence:81 """Save/load round-trip (D2: serialize_index/deserialize_index)."""82 83 def test_save_load_round_trip(self, tmp_index_dir: Path) -> None:84 idx = ManifoldIndex(dim=4)85 v1 = torch.tensor([1.0, 0.0, 0.0, 0.0])86 idx.add(v1, _entry("persisted"))87 idx.save(tmp_index_dir / "test.faiss")88 89 idx2 = ManifoldIndex(dim=4, index_path=tmp_index_dir / "test.faiss")90 assert idx2.n_entries == 191 results = idx2.search(v1, top_k=1)92 assert results[0]["cache_id"] == "persisted"93 94 def test_dim_mismatch_raises(self) -> None:95 idx = ManifoldIndex(dim=4)96 with pytest.raises(ValueError, match="dim"):97 idx.add(torch.randn(8), _entry("wrong"))98 