From aaeecccfcfe2c58e5f267e75c7e7f62e4560c706 Mon Sep 17 00:00:00 2001 From: Atharva Date: Wed, 1 Apr 2026 02:39:26 +0530 Subject: [PATCH] Add Python SDK with Memory class for simple programmatic access --- agentic_memory.py | 377 ++++++++++++++++++++++++++++++++++++++++++++++ pyproject.toml | 2 +- tests/test_sdk.py | 254 +++++++++++++++++++++++++++++++ 3 files changed, 632 insertions(+), 1 deletion(-) create mode 100644 agentic_memory.py create mode 100644 tests/test_sdk.py diff --git a/agentic_memory.py b/agentic_memory.py new file mode 100644 index 0000000..35309c1 --- /dev/null +++ b/agentic_memory.py @@ -0,0 +1,377 @@ +"""Agentic Memory SDK — a clean Python interface over the cognitive memory runtime. + +Usage:: + + from agentic_memory import Memory + + m = Memory() + m.remember("user prefers dark mode") + m.remember_episode("debugged auth flow for 2 hours", session="abc") + results = m.recall("what does the user prefer?") +""" + +from __future__ import annotations + +from datetime import datetime +from pathlib import Path +from typing import Any + +from forgetting.contradiction import ContradictionCandidate +from forgetting.service import ForgettingReport +from models.base import MemoryRecord +from models.episodic import EpisodicMemory +from models.procedural import ProceduralMemory +from models.semantic import SemanticMemory +from retrieval.ranking import RankedResult +from runtime import MemoryRuntime, build_runtime +from stores.procedural_store import ProceduralMatch +from utils.embeddings import TextEmbedder + +__all__ = [ + "Memory", + "MemoryRuntime", + "RankedResult", + "ContradictionCandidate", + "ProceduralMatch", + "ForgettingReport", + "SemanticMemory", + "EpisodicMemory", + "ProceduralMemory", + "MemoryRecord", +] + + +class Memory: + """High-level interface to the cognitive memory system. + + Wraps :func:`runtime.build_runtime` and exposes remember / recall / + forget operations with minimal boilerplate. + + Parameters + ---------- + chroma_path: + Directory for ChromaDB persistence. Defaults to ``./chroma_db``. + media_root: + Root directory for owned media files. Defaults to ``./data/media``. + embedder: + Custom embedder instance. Defaults to :class:`GeminiEmbedder`. + embedding_dimensions: + Vector dimensionality (must match *embedder* if both are given). + max_media_bytes: + Per-file size cap for media embedding. + + Examples + -------- + >>> m = Memory() + >>> m.remember("the deploy key rotates every 90 days") + 'some-uuid' + >>> m.recall("deploy key rotation") + [RankedResult(...)] + """ + + def __init__( + self, + *, + chroma_path: str | None = None, + media_root: str | Path | None = None, + embedder: TextEmbedder | None = None, + embedding_dimensions: int | None = None, + max_media_bytes: int | None = None, + ) -> None: + self._runtime = build_runtime( + chroma_path=chroma_path, + media_root=media_root, + embedder=embedder, + embedding_dimensions=embedding_dimensions, + max_media_bytes=max_media_bytes, + ) + + @property + def runtime(self) -> MemoryRuntime: + """Escape hatch — direct access to the underlying runtime components.""" + return self._runtime + + # ── remember ────────────────────────────────────────────────────────── + + def remember( + self, + content: str, + *, + importance: float = 0.5, + category: str = "general", + domain: str | None = None, + confidence: float = 1.0, + supersedes: str | None = None, + related_ids: list[str] | None = None, + source: str | None = None, + metadata: dict[str, Any] | None = None, + media_ref: str | None = None, + media_type: str | None = None, + text_description: str | None = None, + ) -> str: + """Store a semantic memory (fact, preference, concept). + + Returns the record ID. + """ + record = SemanticMemory( + content=content, + importance=importance, + category=category, + domain=domain, + confidence=confidence, + supersedes=supersedes, + related_ids=related_ids or [], + source=source, + metadata=metadata or {}, + media_ref=media_ref, + media_type=media_type, + text_description=text_description, + modality="text" if not media_ref else (media_type or "text"), + ) + return self._runtime.semantic_store.store(record) + + def remember_episode( + self, + content: str, + *, + session: str | None = None, + turn: int | None = None, + participants: list[str] | None = None, + summary: str | None = None, + emotional_valence: float | None = None, + emotional_profile: dict[str, float] | None = None, + importance: float = 0.5, + source: str | None = None, + metadata: dict[str, Any] | None = None, + media_ref: str | None = None, + media_type: str | None = None, + text_description: str | None = None, + ) -> str: + """Store an episodic memory (event, interaction, experience). + + Returns the record ID. + """ + kwargs: dict[str, Any] = dict( + content=content, + importance=importance, + turn_number=turn, + participants=participants or ["user", "agent"], + summary=summary, + emotional_valence=emotional_valence, + emotional_profile=emotional_profile or {}, + source=source, + metadata=metadata or {}, + media_ref=media_ref, + media_type=media_type, + text_description=text_description, + modality="text" if not media_ref else (media_type or "text"), + ) + if session is not None: + kwargs["session_id"] = session + record = EpisodicMemory(**kwargs) + return self._runtime.episodic_store.store(record) + + def remember_procedure( + self, + content: str, + *, + steps: list[str], + preconditions: list[str] | None = None, + importance: float = 0.5, + source: str | None = None, + metadata: dict[str, Any] | None = None, + media_ref: str | None = None, + media_type: str | None = None, + text_description: str | None = None, + ) -> str: + """Store a procedural memory (skill, workflow, strategy). + + Returns the record ID. + """ + record = ProceduralMemory( + content=content, + steps=steps, + preconditions=preconditions or [], + importance=importance, + source=source, + metadata=metadata or {}, + media_ref=media_ref, + media_type=media_type, + text_description=text_description, + modality="text" if not media_ref else (media_type or "text"), + ) + return self._runtime.procedural_store.store(record) + + # ── recall ──────────────────────────────────────────────────────────── + + def recall( + self, + query: str, + *, + top_k: int = 5, + types: list[str] | None = None, + relevance_weight: float = 0.4, + recency_weight: float = 0.3, + importance_weight: float = 0.3, + ) -> list[RankedResult]: + """Search across memory stores with weighted ranking. + + Parameters + ---------- + query: + Natural-language search query. + top_k: + Maximum results to return. + types: + Filter to specific memory types (``"semantic"``, ``"episodic"``, + ``"procedural"``). ``None`` searches all. + relevance_weight / recency_weight / importance_weight: + Ranking weights (should sum to 1.0). + """ + return self._runtime.retriever.query( + query, + top_k=top_k, + memory_types=types, + relevance_weight=relevance_weight, + recency_weight=recency_weight, + importance_weight=importance_weight, + ) + + def recall_by_vector( + self, + vector: list[float], + *, + top_k: int = 5, + types: list[str] | None = None, + relevance_weight: float = 0.4, + recency_weight: float = 0.3, + importance_weight: float = 0.3, + ) -> list[RankedResult]: + """Search by pre-computed embedding vector.""" + return self._runtime.retriever.query_by_vector( + vector, + top_k=top_k, + memory_types=types, + relevance_weight=relevance_weight, + recency_weight=recency_weight, + importance_weight=importance_weight, + ) + + def recall_episodes( + self, + *, + mode: str = "recent", + limit: int = 5, + session: str | None = None, + start: datetime | None = None, + end: datetime | None = None, + ) -> list[MemoryRecord]: + """Retrieve episodic memories by recency, session, or time range. + + Parameters + ---------- + mode: + ``"recent"`` — most recent episodes (use *limit*). + ``"session"`` — all episodes in a session (use *session*). + ``"time_range"`` — episodes within a window (use *start* and *end*). + """ + if mode == "recent": + return self._runtime.retriever.query_recent(limit) + if mode == "session": + if session is None: + raise ValueError("session is required when mode='session'") + return self._runtime.episodic_store.get_by_session(session) + if mode == "time_range": + if start is None or end is None: + raise ValueError("start and end are required when mode='time_range'") + return self._runtime.retriever.query_time_range(start, end) + raise ValueError(f"Unknown mode '{mode}'. Use 'recent', 'session', or 'time_range'.") + + def recall_procedures( + self, + task: str, + *, + top_k: int = 3, + ) -> list[ProceduralMatch]: + """Find the best procedures for a task, ranked by similarity + Wilson score.""" + return self._runtime.procedural_store.get_best_procedure_matches(task, top_k=top_k) + + # ── contradictions ──────────────────────────────────────────────────── + + def find_contradictions( + self, + record_id: str, + *, + threshold: float = 0.85, + top_k: int = 5, + ) -> list[ContradictionCandidate]: + """Find semantic memories that may contradict a stored record.""" + record = self._runtime.semantic_store.get_by_id(record_id) + if record is None: + raise ValueError(f"Semantic record '{record_id}' not found") + return self._runtime.contradiction_detector.find_potential_contradictions( + record, threshold=threshold, top_k=top_k + ) + + def resolve_contradiction( + self, + *, + keep_id: str, + supersede_id: str, + ) -> None: + """Mark one semantic memory as superseding another.""" + self._runtime.contradiction_detector.resolve_supersession( + superseded_id=supersede_id, kept_id=keep_id + ) + + # ── procedural outcomes ─────────────────────────────────────────────── + + def record_outcome(self, record_id: str, *, success: bool) -> None: + """Record a success or failure for a procedural memory.""" + self._runtime.procedural_store.record_outcome(record_id, success) + + # ── forgetting ──────────────────────────────────────────────────────── + + def forget(self, *, dry_run: bool = False) -> ForgettingReport: + """Run the forgetting cycle. + + With ``dry_run=True``, returns what *would* happen without modifying + any records. + """ + return self._runtime.forgetting_service.run_cycle(dry_run=dry_run) + + # ── direct access ───────────────────────────────────────────────────── + + def get(self, record_id: str, *, type: str) -> MemoryRecord | None: + """Fetch a single record by ID and type. + + Parameters + ---------- + 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) + if store is None: + raise ValueError(f"Unknown type '{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.""" + 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()) + return { + "total": semantic_count + episodic_count + procedural_count, + "semantic": semantic_count, + "episodic": episodic_count, + "procedural": procedural_count, + } + + def events(self, *, limit: int = 50) -> list[dict[str, Any]]: + """Return recent lifecycle events (stored, retrieved, faded, pruned, etc.).""" + return self._runtime.events.snapshot(limit=limit) diff --git a/pyproject.toml b/pyproject.toml index 79c48ad..f62a81c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -35,7 +35,7 @@ dev = [ agentic-memory-mcp = "mcp_server.server:main" [tool.setuptools] -py-modules = ["config", "runtime"] +py-modules = ["config", "runtime", "agentic_memory"] [tool.setuptools.packages.find] include = [ diff --git a/tests/test_sdk.py b/tests/test_sdk.py new file mode 100644 index 0000000..960fc3e --- /dev/null +++ b/tests/test_sdk.py @@ -0,0 +1,254 @@ +"""Tests for the agentic_memory SDK (Memory class).""" + +from datetime import datetime, timedelta, timezone + +import pytest + +from agentic_memory import ( + Memory, + RankedResult, + ContradictionCandidate, + ProceduralMatch, + ForgettingReport, + SemanticMemory, + EpisodicMemory, + ProceduralMemory, +) +from tests.helpers import HashingEmbedder, make_temp_chroma_dir, cleanup_dir + + +@pytest.fixture() +def mem(tmp_path): + chroma_dir = make_temp_chroma_dir("sdk_test_") + m = Memory( + chroma_path=chroma_dir, + media_root=str(tmp_path / "media"), + embedder=HashingEmbedder(dimensions=64), + embedding_dimensions=64, + ) + yield m + cleanup_dir(chroma_dir) + + +class TestRemember: + def test_remember_returns_id(self, mem: Memory): + rid = mem.remember("the sky is blue") + assert isinstance(rid, str) + assert len(rid) > 0 + + def test_remember_with_options(self, mem: Memory): + rid = mem.remember( + "deploy key rotates every 90 days", + importance=0.9, + category="infra", + domain="security", + confidence=0.95, + ) + record = mem.get(rid, type="semantic") + assert record is not None + assert record.importance == 0.9 + assert record.category == "infra" + assert record.domain == "security" + + def test_remember_episode_returns_id(self, mem: Memory): + rid = mem.remember_episode("user debugged auth flow", session="s1") + assert isinstance(rid, str) + + 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") + assert record is not None + assert record.session_id == "sess-42" + assert record.turn_number == 3 + + def test_remember_procedure_returns_id(self, mem: Memory): + rid = mem.remember_procedure( + "deploy to production", + steps=["run tests", "build image", "push"], + ) + assert isinstance(rid, str) + + def test_remember_procedure_with_preconditions(self, mem: Memory): + rid = mem.remember_procedure( + "rollback deployment", + steps=["identify version", "kubectl rollout undo"], + preconditions=["kubectl access"], + importance=0.8, + ) + record = mem.get(rid, type="procedural") + assert record is not None + assert record.preconditions == ["kubectl access"] + assert record.importance == 0.8 + + +class TestRecall: + def test_recall_finds_stored_memory(self, mem: Memory): + mem.remember("python is a programming language") + results = mem.recall("programming language") + assert len(results) > 0 + assert isinstance(results[0], RankedResult) + assert results[0].record.content == "python is a programming language" + + def test_recall_with_type_filter(self, mem: Memory): + mem.remember("fact about dogs") + mem.remember_episode("walked the dog", session="s1") + results = mem.recall("dog", types=["semantic"]) + assert all(r.record.memory_type == "semantic" for r in results) + + def test_recall_empty_returns_empty(self, mem: Memory): + results = mem.recall("nonexistent topic") + assert results == [] + + def test_recall_respects_top_k(self, mem: Memory): + for i in range(5): + mem.remember(f"fact number {i} about testing") + results = mem.recall("testing", top_k=2) + assert len(results) <= 2 + + +class TestRecallEpisodes: + def test_recall_recent(self, mem: Memory): + mem.remember_episode("first event", session="s1") + mem.remember_episode("second event", session="s1") + records = mem.recall_episodes(mode="recent", limit=1) + assert len(records) == 1 + + def test_recall_by_session(self, mem: Memory): + mem.remember_episode("alpha event", session="alpha") + mem.remember_episode("beta event", session="beta") + records = mem.recall_episodes(mode="session", session="alpha") + assert len(records) == 1 + assert records[0].content == "alpha event" + + def test_recall_session_requires_session_param(self, mem: Memory): + with pytest.raises(ValueError, match="session is required"): + mem.recall_episodes(mode="session") + + def test_recall_time_range(self, mem: Memory): + now = datetime.now(timezone.utc) + mem.remember_episode("timed event", session="s1") + records = mem.recall_episodes( + mode="time_range", + start=now - timedelta(minutes=1), + end=now + timedelta(minutes=1), + ) + assert len(records) >= 1 + + def test_recall_time_range_requires_start_end(self, mem: Memory): + with pytest.raises(ValueError, match="start and end"): + mem.recall_episodes(mode="time_range", start=datetime.now(timezone.utc)) + + def test_unknown_mode_raises(self, mem: Memory): + with pytest.raises(ValueError, match="Unknown mode"): + mem.recall_episodes(mode="invalid") + + +class TestRecallProcedures: + def test_recall_procedures(self, mem: Memory): + mem.remember_procedure("deploy app", steps=["build", "push", "apply"]) + matches = mem.recall_procedures("deploy application") + assert len(matches) > 0 + assert isinstance(matches[0], ProceduralMatch) + + +class TestContradictions: + 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) + + def test_find_contradictions_missing_id_raises(self, mem: Memory): + with pytest.raises(ValueError, match="not found"): + mem.find_contradictions("nonexistent-id") + + 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") + assert old is not None + assert old.superseded_by == id2 + assert old.importance == 0.0 + + +class TestOutcomes: + def test_record_outcome(self, mem: Memory): + rid = mem.remember_procedure("test proc", steps=["step1"]) + mem.record_outcome(rid, success=True) + mem.record_outcome(rid, success=True) + mem.record_outcome(rid, success=False) + record = mem.get(rid, type="procedural") + assert record.success_count == 2 + assert record.failure_count == 1 + + +class TestForgetting: + def test_dry_run(self, mem: Memory): + mem.remember("something to maybe forget") + report = mem.forget(dry_run=True) + assert isinstance(report, ForgettingReport) + assert report.dry_run is True + + def test_run(self, mem: Memory): + mem.remember("something to maybe forget") + report = mem.forget() + assert isinstance(report, ForgettingReport) + assert report.dry_run is False + + +class TestDirectAccess: + def test_get_semantic(self, mem: Memory): + rid = mem.remember("fact") + record = mem.get(rid, 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") + + def test_get_missing_returns_none(self, mem: Memory): + result = mem.get("nonexistent", type="semantic") + assert result is None + + def test_overview(self, mem: Memory): + mem.remember("a fact") + mem.remember_episode("an event", session="s1") + overview = mem.overview() + assert overview["semantic"] == 1 + assert overview["episodic"] == 1 + assert overview["procedural"] == 0 + assert overview["total"] == 2 + + def test_events(self, mem: Memory): + mem.remember("trigger events") + evts = mem.events() + assert isinstance(evts, list) + assert len(evts) > 0 + assert "event_type" in evts[0] + + def test_runtime_escape_hatch(self, mem: Memory): + assert mem.runtime is not None + assert hasattr(mem.runtime, "semantic_store") + + +class TestExports: + """Verify that key types are importable from the SDK.""" + + def test_all_exports(self): + from agentic_memory import ( + Memory, + MemoryRuntime, + RankedResult, + ContradictionCandidate, + ProceduralMatch, + ForgettingReport, + SemanticMemory, + EpisodicMemory, + ProceduralMemory, + MemoryRecord, + ) + assert Memory is not None + assert RankedResult is not None