Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 8 additions & 1 deletion catalog.json
Original file line number Diff line number Diff line change
Expand Up @@ -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).",
Expand Down Expand Up @@ -176,4 +183,4 @@
"description": "No LLM. Returns the provider's retrieved memories verbatim; the dataset scores the returned ID set."
}
}
}
}
1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
2 changes: 2 additions & 0 deletions src/memory_bench/memory/__init__.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -14,6 +15,7 @@

REGISTRY: dict[str, type[MemoryProvider]] = {
"vanilla": NoMemoryProvider,
"atmem": AtMemMemoryProvider,
"hindsight-coding": HsCodingProvider,
"bm25": BM25MemoryProvider,
"cognee": CogneeMemoryProvider,
Expand Down
122 changes: 122 additions & 0 deletions src/memory_bench/memory/atmem.py
Original file line number Diff line number Diff line change
@@ -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
54 changes: 54 additions & 0 deletions tests/test_atmem_provider.py
Original file line number Diff line number Diff line change
@@ -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."
)
24 changes: 24 additions & 0 deletions uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.