Skip to content
Merged
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
31 changes: 23 additions & 8 deletions agentic_memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,12 @@
]


def _infer_modality(media_ref: str | None, media_type: str | None) -> str:
if not media_ref:
return "text"
return media_type or "text"


class Memory:
"""High-level interface to the cognitive memory system.

Expand Down Expand Up @@ -126,7 +132,7 @@ def remember(
media_ref=media_ref,
media_type=media_type,
text_description=text_description,
modality="text" if not media_ref else (media_type or "text"),
modality=_infer_modality(media_ref, media_type),
)
return self._runtime.semantic_store.store(record)

Expand All @@ -149,6 +155,9 @@ def remember_episode(
) -> str:
"""Store an episodic memory (event, interaction, experience).

When *participants* is omitted, the SDK defaults it to
``["user", "agent"]``.

Returns the record ID.
"""
kwargs: dict[str, Any] = dict(
Expand All @@ -164,7 +173,7 @@ def remember_episode(
media_ref=media_ref,
media_type=media_type,
text_description=text_description,
modality="text" if not media_ref else (media_type or "text"),
modality=_infer_modality(media_ref, media_type),
)
if session is not None:
kwargs["session_id"] = session
Expand Down Expand Up @@ -198,7 +207,7 @@ def remember_procedure(
media_ref=media_ref,
media_type=media_type,
text_description=text_description,
modality="text" if not media_ref else (media_type or "text"),
modality=_infer_modality(media_ref, media_type),
)
return self._runtime.procedural_store.store(record)

Expand Down Expand Up @@ -342,26 +351,32 @@ def forget(self, *, dry_run: bool = False) -> ForgettingReport:

# ── direct access ─────────────────────────────────────────────────────

def get(self, record_id: str, *, type: str) -> MemoryRecord | None:
def get(self, record_id: str, *, memory_type: str) -> MemoryRecord | None:
"""Fetch a single record by ID and type.

Parameters
----------
type:
memory_type:
One of ``"semantic"``, ``"episodic"``, ``"procedural"``.
"""
store_map = {
"semantic": self._runtime.semantic_store,
"episodic": self._runtime.episodic_store,
"procedural": self._runtime.procedural_store,
}
store = store_map.get(type)
store = store_map.get(memory_type)
if store is None:
raise ValueError(f"Unknown type '{type}'. Use 'semantic', 'episodic', or 'procedural'.")
raise ValueError(
f"Unknown memory_type '{memory_type}'. Use 'semantic', 'episodic', or 'procedural'."
)
return store.get_by_id(record_id)

def overview(self) -> dict[str, Any]:
"""Return a summary of the memory system state."""
"""Return a summary of the memory system state.

Note: this currently loads all records into memory to count them, so
it may be slow for large stores.
"""
semantic_count = len(self._runtime.semantic_store.get_all_records())
episodic_count = len(self._runtime.episodic_store.get_all_records())
procedural_count = len(self._runtime.procedural_store.get_all_records())
Expand Down
37 changes: 27 additions & 10 deletions tests/test_sdk.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,7 +44,7 @@ def test_remember_with_options(self, mem: Memory):
domain="security",
confidence=0.95,
)
record = mem.get(rid, type="semantic")
record = mem.get(rid, memory_type="semantic")
assert record is not None
assert record.importance == 0.9
assert record.category == "infra"
Expand All @@ -56,7 +56,7 @@ def test_remember_episode_returns_id(self, mem: Memory):

def test_remember_episode_with_session(self, mem: Memory):
rid = mem.remember_episode("built the feature", session="sess-42", turn=3)
record = mem.get(rid, type="episodic")
record = mem.get(rid, memory_type="episodic")
assert record is not None
assert record.session_id == "sess-42"
assert record.turn_number == 3
Expand All @@ -75,7 +75,7 @@ def test_remember_procedure_with_preconditions(self, mem: Memory):
preconditions=["kubectl access"],
importance=0.8,
)
record = mem.get(rid, type="procedural")
record = mem.get(rid, memory_type="procedural")
assert record is not None
assert record.preconditions == ["kubectl access"]
assert record.importance == 0.8
Expand Down Expand Up @@ -105,6 +105,14 @@ def test_recall_respects_top_k(self, mem: Memory):
results = mem.recall("testing", top_k=2)
assert len(results) <= 2

def test_recall_by_vector_finds_stored_memory(self, mem: Memory):
mem.remember("python is a programming language")
vector = mem.runtime.embedder.embed_query("programming language")
results = mem.recall_by_vector(vector)
assert len(results) > 0
assert isinstance(results[0], RankedResult)
assert results[0].record.content == "python is a programming language"


class TestRecallEpisodes:
def test_recall_recent(self, mem: Memory):
Expand Down Expand Up @@ -156,8 +164,9 @@ def test_find_contradictions(self, mem: Memory):
id1 = mem.remember("the capital of France is Paris")
id2 = mem.remember("the capital of France is Paris and it is lovely")
candidates = mem.find_contradictions(id1)
# With the hashing embedder, similar text produces similar vectors
assert isinstance(candidates, list)
assert all(isinstance(candidate, ContradictionCandidate) for candidate in candidates)
assert any(candidate.record.id == id2 for candidate in candidates)

def test_find_contradictions_missing_id_raises(self, mem: Memory):
with pytest.raises(ValueError, match="not found"):
Expand All @@ -167,7 +176,7 @@ def test_resolve_contradiction(self, mem: Memory):
id1 = mem.remember("old fact")
id2 = mem.remember("new fact replaces old")
mem.resolve_contradiction(keep_id=id2, supersede_id=id1)
old = mem.get(id1, type="semantic")
old = mem.get(id1, memory_type="semantic")
assert old is not None
assert old.superseded_by == id2
assert old.importance == 0.0
Expand All @@ -179,7 +188,7 @@ def test_record_outcome(self, mem: Memory):
mem.record_outcome(rid, success=True)
mem.record_outcome(rid, success=True)
mem.record_outcome(rid, success=False)
record = mem.get(rid, type="procedural")
record = mem.get(rid, memory_type="procedural")
assert record.success_count == 2
assert record.failure_count == 1

Expand All @@ -201,16 +210,16 @@ def test_run(self, mem: Memory):
class TestDirectAccess:
def test_get_semantic(self, mem: Memory):
rid = mem.remember("fact")
record = mem.get(rid, type="semantic")
record = mem.get(rid, memory_type="semantic")
assert record is not None
assert record.content == "fact"

def test_get_unknown_type_raises(self, mem: Memory):
with pytest.raises(ValueError, match="Unknown type"):
mem.get("some-id", type="invalid")
with pytest.raises(ValueError, match="Unknown memory_type"):
mem.get("some-id", memory_type="invalid")

def test_get_missing_returns_none(self, mem: Memory):
result = mem.get("nonexistent", type="semantic")
result = mem.get("nonexistent", memory_type="semantic")
assert result is None

def test_overview(self, mem: Memory):
Expand Down Expand Up @@ -251,4 +260,12 @@ def test_all_exports(self):
MemoryRecord,
)
assert Memory is not None
assert MemoryRuntime is not None
assert RankedResult is not None
assert ContradictionCandidate is not None
assert ProceduralMatch is not None
assert ForgettingReport is not None
assert SemanticMemory is not None
assert EpisodicMemory is not None
assert ProceduralMemory is not None
assert MemoryRecord is not None
Loading