diff --git a/catalog.json b/catalog.json index d3e338b..0c59bd1 100644 --- a/catalog.json +++ b/catalog.json @@ -56,6 +56,13 @@ } }, "providers": { + "atmem": { + "key": "atmem", + "description": "Local auditable memory with deterministic extraction, lifecycle governance, SQLite persistence, lexical/graph retrieval, and calibrated direct-support selection.", + "kind": "local", + "link": "https://github.com/aetna000/atmem", + "logo": "https://www.google.com/s2/favicons?sz=32&domain=github.com" + }, "vanilla": { "key": "none", "description": "No memory system (vanilla baseline).", @@ -176,4 +183,4 @@ "description": "No LLM. Returns the provider's retrieved memories verbatim; the dataset scores the returned ID set." } } -} \ No newline at end of file +} diff --git a/pyproject.toml b/pyproject.toml index 9a85d22..23f736a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,6 +4,7 @@ version = "0.1.0" description = "Open Memory Benchmark" requires-python = ">=3.11" dependencies = [ + "atmem==2.2.7b1", "datasets>=2.0", "typer>=0.12", "rich>=13", diff --git a/src/memory_bench/memory/__init__.py b/src/memory_bench/memory/__init__.py index 2b7e5e0..3577dc9 100644 --- a/src/memory_bench/memory/__init__.py +++ b/src/memory_bench/memory/__init__.py @@ -1,4 +1,5 @@ from .base import MemoryProvider +from .atmem import AtMemMemoryProvider from .bm25 import BM25MemoryProvider from .cognee import CogneeMemoryProvider from .hindsight import HindsightCloudMemoryProvider, HindsightHTTPMemoryProvider, HindsightMemoryProvider @@ -14,6 +15,7 @@ REGISTRY: dict[str, type[MemoryProvider]] = { "vanilla": NoMemoryProvider, + "atmem": AtMemMemoryProvider, "hindsight-coding": HsCodingProvider, "bm25": BM25MemoryProvider, "cognee": CogneeMemoryProvider, diff --git a/src/memory_bench/memory/atmem.py b/src/memory_bench/memory/atmem.py new file mode 100644 index 0000000..99ea686 --- /dev/null +++ b/src/memory_bench/memory/atmem.py @@ -0,0 +1,122 @@ +import uuid +from pathlib import Path + +from atmem import Memory +from atmem.retrieve import decide_retrieval + +from ..models import Document +from .base import MemoryProvider + + +class AtMemMemoryProvider(MemoryProvider): + name = "atmem" + description = ( + "Local auditable memory with deterministic extraction, lifecycle governance, " + "SQLite persistence, lexical/graph retrieval, and calibrated direct-support selection." + ) + kind = "local" + link = "https://github.com/aetna000/atmem" + logo = "https://www.google.com/s2/favicons?sz=32&domain=github.com" + concurrency = 1 + + def __init__(self): + self._memory: Memory | None = None + self._default_user_id = f"bench_{uuid.uuid4().hex[:8]}" + + def prepare( + self, + store_dir: Path, + unit_ids: set[str] | None = None, + reset: bool = True, + ) -> None: + self.cleanup() + database = store_dir / "atmem.db" + if reset: + database.unlink(missing_ok=True) + self._memory = Memory(database, graph_recall=True, auto_vectors=False) + + def cleanup(self) -> None: + if self._memory is not None: + self._memory.close() + self._memory = None + + def _ensure_memory(self) -> Memory: + if self._memory is None: + self._memory = Memory(":memory:", graph_recall=True, auto_vectors=False) + return self._memory + + @staticmethod + def _format_content(doc: Document) -> str: + if not doc.messages: + return doc.content + + lines = [] + if doc.timestamp: + lines.append(f"Date: {doc.timestamp}") + for message in doc.messages: + role = str(message.get("role") or "unknown").capitalize() + content = str(message.get("content") or "").strip() + if content: + lines.append(f"{role}: {content}") + return "\n".join(lines) or doc.content + + def ingest(self, documents: list[Document]) -> None: + memory = self._ensure_memory() + for doc in documents: + subject_id = doc.user_id or self._default_user_id + memory.remember( + subject_id, + self._format_content(doc), + force=True, + session_id=doc.id, + source_type="user_message", + raw={"amb_document_id": doc.id, "source_timestamp": doc.timestamp}, + ) + + def retrieve( + self, + query: str, + k: int = 10, + user_id: str | None = None, + query_timestamp: str | None = None, + ) -> tuple[list[Document], dict | None]: + memory = self._ensure_memory() + subject_id = user_id or self._default_user_id + session_id = f"amb-retrieval-{uuid.uuid4().hex}" + records = memory.recall( + subject_id, + query, + session_id=session_id, + limit=k, + use_graph=True, + include_scores=True, + ) + decision = decide_retrieval(query, records) + records_by_id = {str(record["id"]): record for record in records} + selected_records = [ + records_by_id[record_id] + for record_id in decision.ranked_record_ids + if record_id in records_by_id + ][:k] + [retrieval] = memory.get_retrieval_log(subject_id, session_id=session_id) + + documents = [] + for record in selected_records: + source_id = record.get("source_session_id") + documents.append( + Document( + id=str(record["id"]), + content=str(record["content"]), + user_id=subject_id, + source_ids=[str(source_id)] if source_id else None, + ) + ) + + raw = { + "retrieval_id": retrieval["id"], + "candidate_ids": retrieval["returned_ids"], + "returned_ids": [str(record["id"]) for record in selected_records], + "candidates": retrieval["candidates"], + "decision": decision.to_dict(), + } + return documents, raw diff --git a/tests/test_atmem_provider.py b/tests/test_atmem_provider.py new file mode 100644 index 0000000..7834f2d --- /dev/null +++ b/tests/test_atmem_provider.py @@ -0,0 +1,54 @@ +from memory_bench.memory.atmem import AtMemMemoryProvider +from memory_bench.models import Document + + +def test_atmem_provider_retrieves_scoped_memory_with_source_evidence(tmp_path): + provider = AtMemMemoryProvider() + provider.prepare(tmp_path) + try: + provider.ingest( + [ + Document( + id="alice-color", + content="My favorite color is teal.", + user_id="alice", + ), + Document( + id="bob-color", + content="My favorite color is orange.", + user_id="bob", + ), + ] + ) + + documents, raw = provider.retrieve( + "What is my favorite color?", k=5, user_id="alice" + ) + finally: + provider.cleanup() + + assert len(documents) == 1 + assert "teal" in documents[0].content + assert "orange" not in documents[0].content + assert documents[0].source_ids == ["alice-color"] + assert raw["returned_ids"] == [documents[0].id] + assert raw["retrieval_id"].startswith("ret_") + assert raw["decision"]["support_class"] == "direct_support" + + +def test_atmem_provider_formats_structured_messages(): + document = Document( + id="session-1", + content="fallback", + timestamp="2026-09-10T00:00:00Z", + messages=[ + {"role": "user", "content": "I prefer window seats."}, + {"role": "assistant", "content": "Understood."}, + ], + ) + + assert AtMemMemoryProvider._format_content(document) == ( + "Date: 2026-09-10T00:00:00Z\n" + "User: I prefer window seats.\n" + "Assistant: Understood." + ) diff --git a/uv.lock b/uv.lock index cbd392e..f6c45f8 100644 --- a/uv.lock +++ b/uv.lock @@ -210,6 +210,7 @@ name = "amb" version = "0.1.0" source = { editable = "." } dependencies = [ + { name = "atmem" }, { name = "cognee" }, { name = "datasets" }, { name = "fastapi", extra = ["standard"] }, @@ -233,6 +234,7 @@ dependencies = [ [package.metadata] requires-dist = [ + { name = "atmem", specifier = "==2.2.7b1" }, { name = "cognee", specifier = ">=0.5.4" }, { name = "datasets", specifier = ">=2.0" }, { name = "fastapi", extras = ["standard"], specifier = ">=0.135.1" }, @@ -422,6 +424,28 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/3c/d7/8fb3044eaef08a310acfe23dae9a8e2e07d305edc29a53497e52bc76eca7/asyncpg-0.31.0-cp314-cp314t-win_amd64.whl", hash = "sha256:bd4107bb7cdd0e9e65fae66a62afd3a249663b844fa34d479f6d5b3bef9c04c3", size = 706062, upload-time = "2025-11-24T23:26:44.086Z" }, ] +[[package]] +name = "atmem" +version = "2.2.7b1" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "atmem-atbot" }, + { name = "cryptography" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/70/81/6176d404c1eaaac21eea156a8419221bed69838f8abbc0aa49d5870f15c8/atmem-2.2.7b1.tar.gz", hash = "sha256:e218726d096f26ad7f28bd0ea9f74207b9f88c085d5c37fd9dbfddf75e16549d", size = 1016070, upload-time = "2026-09-09T09:59:01.166Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/91/8a/293be51574e76043a4f3dedd86b8513d3691dfbea40804ac8451503621ab/atmem-2.2.7b1-py3-none-any.whl", hash = "sha256:c08fe85c2b02ec327b5c62dff92c12d0464633d1fcb1afa6e31d67a508a605c0", size = 616704, upload-time = "2026-09-09T09:58:58.904Z" }, +] + +[[package]] +name = "atmem-atbot" +version = "0.1.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/6c/d6/9321de4e8d374e014bb37f227fdbb988e281f02423600279add7a588f75c/atmem_atbot-0.1.0.tar.gz", hash = "sha256:3324c1da545d6cd0daf0d7e32e1f50d8ca05a5494346ebe08056f09d2f8cee79", size = 26433, upload-time = "2026-09-09T09:53:35.981Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/be/69/548d9d91df6a865f9faf5022f0c8c4209f1f15c2711bdd389a118872d902/atmem_atbot-0.1.0-py3-none-any.whl", hash = "sha256:994201212ec3974a4d20adfb1231961182bcf818184fc7e57bcef5b944859c0d", size = 28961, upload-time = "2026-09-09T09:53:34.584Z" }, +] + [[package]] name = "attrs" version = "25.4.0"