From 1f12f9742a22f888b0f182753cce6021ce7b16c8 Mon Sep 17 00:00:00 2001 From: Atharva Date: Wed, 1 Apr 2026 16:36:14 +0530 Subject: [PATCH] Address Python SDK review feedback --- agentic_memory.py | 31 +++++++++++++++++++++++-------- tests/test_sdk.py | 37 +++++++++++++++++++++++++++---------- 2 files changed, 50 insertions(+), 18 deletions(-) diff --git a/agentic_memory.py b/agentic_memory.py index 35309c1..6b5b9ea 100644 --- a/agentic_memory.py +++ b/agentic_memory.py @@ -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. @@ -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) @@ -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( @@ -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 @@ -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) @@ -342,12 +351,12 @@ 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 = { @@ -355,13 +364,19 @@ def get(self, record_id: str, *, type: str) -> MemoryRecord | None: "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()) diff --git a/tests/test_sdk.py b/tests/test_sdk.py index 960fc3e..fc1b0ae 100644 --- a/tests/test_sdk.py +++ b/tests/test_sdk.py @@ -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" @@ -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 @@ -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 @@ -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): @@ -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"): @@ -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 @@ -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 @@ -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): @@ -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