From bfae28d48236da259dff1fb4dd610aa2db4492c9 Mon Sep 17 00:00:00 2001 From: Atharva-Kanherkar Date: Mon, 30 Mar 2026 00:40:47 +0530 Subject: [PATCH 1/5] Extract shared memory runtime factory --- api/app.py | 112 ++++++----------------------- demo/cli.py | 58 ++++----------- retrieval/retriever.py | 16 +++-- runtime.py | 141 +++++++++++++++++++++++++++++++++++++ stores/episodic_store.py | 3 +- stores/procedural_store.py | 3 +- stores/semantic_store.py | 3 +- tests/test_runtime.py | 65 +++++++++++++++++ utils/embeddings.py | 14 ++-- 9 files changed, 267 insertions(+), 148 deletions(-) create mode 100644 runtime.py create mode 100644 tests/test_runtime.py diff --git a/api/app.py b/api/app.py index b0678df..2be82f2 100644 --- a/api/app.py +++ b/api/app.py @@ -4,7 +4,6 @@ import mimetypes import os import tempfile -from collections import deque from collections.abc import Mapping from datetime import datetime from pathlib import Path @@ -14,20 +13,16 @@ from fastapi.middleware.cors import CORSMiddleware import config -from events.bus import MemoryEvent -from events import EventBus from models.base import MemoryRecord, normalize_modality from models.episodic import EpisodicMemory from models.procedural import ProceduralMemory from models.semantic import SemanticMemory -from retrieval.retriever import UnifiedRetriever -from stores.episodic_store import EpisodicStore, EpisodicStoreError, MediaTooLargeError -from stores.media_store import MediaStore +from runtime import build_runtime +from stores.episodic_store import EpisodicStoreError, MediaTooLargeError from stores.procedural_store import ProceduralMatch, ProceduralStore -from stores.semantic_store import SemanticStore from forgetting.contradiction import ContradictionCandidate, ContradictionDetector from forgetting.service import ForgettingReport, ForgettingService -from utils.embeddings import EmbeddingProviderError, GeminiEmbedder, TextEmbedder +from utils.embeddings import EmbeddingProviderError, TextEmbedder DEFAULT_ALLOWED_ORIGINS = [ "http://localhost:3000", @@ -48,16 +43,6 @@ def _normalise_origins(origins: list[str] | None = None) -> list[str]: return [origin.strip() for origin in env_value.split(",") if origin.strip()] -def _jsonable(value: Any) -> Any: - if isinstance(value, Mapping): - return {key: _jsonable(val) for key, val in value.items()} - if isinstance(value, (list, tuple, set, frozenset)): - return [_jsonable(item) for item in value] - if isinstance(value, datetime): - return value.isoformat() - return value - - def _infer_media_contract(*, mime_type: str | None, filename: str | None) -> tuple[str, str] | None: guessed_mime = mime_type or mimetypes.guess_type(filename or "")[0] if guessed_mime: @@ -279,36 +264,6 @@ def _serialise_procedural_match(match: ProceduralMatch) -> dict[str, Any]: } -class EventRecorder: - def __init__(self, bus: EventBus, *, max_events: int = 200): - self._events: deque[dict[str, Any]] = deque(maxlen=max_events) - for event_type in ( - "memory.stored", - "memory.retrieved", - "memory.ranked", - "memory.accessed", - "memory.contradiction_flagged", - "memory.supersession_resolved", - "memory.faded", - "memory.forgotten", - "forgetting.cycle_completed", - "forgetting.cycle_dry_run", - ): - bus.subscribe(event_type, self._record) - - def _record(self, event: MemoryEvent) -> None: - self._events.appendleft( - { - "event_type": event.event_type, - "timestamp": event.timestamp.isoformat(), - "data": _jsonable(dict(event.data)), - } - ) - - def snapshot(self, limit: int = 50) -> list[dict[str, Any]]: - return list(self._events)[:limit] - - class MemoryAPIService: def __init__( self, @@ -316,51 +271,24 @@ def __init__( chroma_path: str | None = None, media_root: Path | None = None, embedder: TextEmbedder | None = None, + embedding_dimensions: int | None = None, ): - self.media_store = MediaStore(media_root or DEFAULT_MEDIA_DIR) - self.bus = EventBus() - self.embedder = embedder or GeminiEmbedder() - original_chroma_path = config.CHROMA_DB_PATH - try: - if chroma_path is not None: - config.CHROMA_DB_PATH = chroma_path - self.semantic_store = SemanticStore( - event_bus=self.bus, - embedder=self.embedder, - media_store=self.media_store, - ) - self.episodic_store = EpisodicStore( - event_bus=self.bus, - embedder=self.embedder, - media_store=self.media_store, - ) - self.procedural_store = ProceduralStore( - event_bus=self.bus, - embedder=self.embedder, - media_store=self.media_store, - ) - finally: - config.CHROMA_DB_PATH = original_chroma_path - self.retriever = UnifiedRetriever( - stores={ - "semantic": self.semantic_store, - "episodic": self.episodic_store, - "procedural": self.procedural_store, - }, - event_bus=self.bus, - ) - self.contradiction_detector = ContradictionDetector( - self.semantic_store, event_bus=self.bus, - ) - self.forgetting_service = ForgettingService( - semantic_store=self.semantic_store, - episodic_store=self.episodic_store, - procedural_store=self.procedural_store, - media_store=self.media_store, - event_bus=self.bus, - contradiction_detector=self.contradiction_detector, + runtime = build_runtime( + chroma_path=chroma_path, + media_root=media_root or DEFAULT_MEDIA_DIR, + embedder=embedder, + embedding_dimensions=embedding_dimensions, ) - self.events = EventRecorder(self.bus) + self.media_store = runtime.media_store + self.bus = runtime.bus + self.embedder = runtime.embedder + self.semantic_store = runtime.semantic_store + self.episodic_store = runtime.episodic_store + self.procedural_store = runtime.procedural_store + self.retriever = runtime.retriever + self.contradiction_detector = runtime.contradiction_detector + self.forgetting_service = runtime.forgetting_service + self.events = runtime.events def save_upload(self, upload: UploadFile, memory_id: str) -> tuple[str, str]: guessed_mime = upload.content_type or mimetypes.guess_type(upload.filename or "")[0] or "application/octet-stream" @@ -449,6 +377,7 @@ def create_app( media_root: str | None = None, allowed_origins: list[str] | None = None, embedder: TextEmbedder | None = None, + embedding_dimensions: int | None = None, ) -> FastAPI: app = FastAPI(title="Agentic Memory API", version="0.1.0") app.add_middleware( @@ -463,6 +392,7 @@ def create_app( "chroma_path": chroma_path, "media_root": Path(media_root) if media_root else None, "embedder": embedder, + "embedding_dimensions": embedding_dimensions, } def service() -> MemoryAPIService: diff --git a/demo/cli.py b/demo/cli.py index 6f1969b..7fa06eb 100644 --- a/demo/cli.py +++ b/demo/cli.py @@ -9,15 +9,12 @@ sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) -from events import ConsoleLogger, EventBus +from events import ConsoleLogger from models.episodic import EpisodicMemory from models.procedural import ProceduralMemory from models.semantic import SemanticMemory -from stores.episodic_store import EpisodicStore +from runtime import build_runtime from stores.media_store import MediaStore -from stores.procedural_store import ProceduralStore -from stores.semantic_store import SemanticStore -from retrieval.retriever import UnifiedRetriever from utils.embeddings import GeminiEmbedder, TextEmbedder import config @@ -29,24 +26,6 @@ } -def _make_bus() -> EventBus: - bus = EventBus() - ConsoleLogger().register(bus) - return bus - - -def _make_semantic_store(event_bus: EventBus | None = None) -> SemanticStore: - return SemanticStore(event_bus=event_bus, media_store=_make_media_store()) - - -def _make_episodic_store(event_bus: EventBus | None = None) -> EpisodicStore: - return EpisodicStore(event_bus=event_bus) - - -def _make_procedural_store(event_bus: EventBus | None = None) -> ProceduralStore: - return ProceduralStore(event_bus=event_bus, media_store=_make_media_store()) - - def _make_media_store() -> MediaStore: return MediaStore(config.MEDIA_STORAGE_PATH) @@ -55,15 +34,10 @@ def _make_embedder() -> TextEmbedder: return GeminiEmbedder() -def _make_retriever(event_bus: EventBus | None = None) -> UnifiedRetriever: - return UnifiedRetriever( - stores={ - "semantic": _make_semantic_store(event_bus=event_bus), - "episodic": _make_episodic_store(event_bus=event_bus), - "procedural": _make_procedural_store(event_bus=event_bus), - }, - event_bus=event_bus, - ) +def _make_runtime(): + runtime = build_runtime() + ConsoleLogger().register(runtime.bus) + return runtime def _guess_mime_type(path: str, modality: str) -> str: @@ -127,9 +101,8 @@ def _print_best_procedures(results) -> None: def _query_by_media(args, *, modality: str) -> None: - bus = _make_bus() + runtime = _make_runtime() embedder = _make_embedder() - retriever = _make_retriever(event_bus=bus) source_path = os.path.abspath(args.path) mime_type = _guess_mime_type(source_path, modality) if modality == "image": @@ -139,7 +112,7 @@ def _query_by_media(args, *, modality: str) -> None: else: raise ValueError(f"Unsupported query modality: {modality}") - results = retriever.query_by_vector( + results = runtime.retriever.query_by_vector( vector, top_k=args.top_k, memory_types=args.memory_types, @@ -152,8 +125,7 @@ def _query_by_media(args, *, modality: str) -> None: def cmd_store(args): - bus = _make_bus() - store = _make_semantic_store(event_bus=bus) + runtime = _make_runtime() media_path = None modality = "text" if args.image: @@ -170,15 +142,14 @@ def cmd_store(args): media_type=_infer_media_type(media_path) if media_path else None, ) try: - record_id = store.store(record) + record_id = runtime.semantic_store.store(record) except (FileNotFoundError, ValueError) as exc: _exit_with_error(str(exc)) print(f"Stored [{record_id[:8]}]: {args.content}") def cmd_store_episode(args): - bus = _make_bus() - store = _make_episodic_store(event_bus=bus) + runtime = _make_runtime() media_store = _make_media_store() if args.text is None else None if args.text is not None: @@ -202,7 +173,7 @@ def cmd_store_episode(args): record.media_ref = media_ref try: - record_id = store.store(record) + record_id = runtime.episodic_store.store(record) except Exception: if media_store is not None and record.media_ref: media_store.delete(record.media_ref) @@ -211,9 +182,8 @@ def cmd_store_episode(args): def cmd_query(args): - bus = _make_bus() - retriever = _make_retriever(event_bus=bus) - results = retriever.query(args.query, top_k=args.top_k) + runtime = _make_runtime() + results = runtime.retriever.query(args.query, top_k=args.top_k) if not results: print("No results found.") return diff --git a/retrieval/retriever.py b/retrieval/retriever.py index 2545dd2..ec75a4d 100644 --- a/retrieval/retriever.py +++ b/retrieval/retriever.py @@ -1,7 +1,7 @@ from datetime import datetime, timezone from typing import Any -from config import EMBEDDING_DIMENSIONS +import config from events.bus import EventBus from models.base import MemoryRecord from stores.base import BaseStore @@ -18,9 +18,17 @@ class UnifiedRetriever: once they are wired into that mapping. """ - def __init__(self, stores: dict[str, BaseStore], event_bus: EventBus | None = None): + def __init__( + self, + stores: dict[str, BaseStore], + event_bus: EventBus | None = None, + embedding_dimensions: int | None = None, + ): self._stores = stores self._event_bus = event_bus + self._embedding_dimensions = ( + config.EMBEDDING_DIMENSIONS if embedding_dimensions is None else embedding_dimensions + ) def _emit_event(self, event_type: str, data: dict[str, Any]) -> None: if self._event_bus is not None: @@ -121,9 +129,9 @@ def _rank_and_touch( return final def _validate_vector_dimensions(self, vector: list[float]) -> None: - if len(vector) != EMBEDDING_DIMENSIONS: + if len(vector) != self._embedding_dimensions: raise ValueError( - f"Expected query vector dimension {EMBEDDING_DIMENSIONS}, got {len(vector)}" + f"Expected query vector dimension {self._embedding_dimensions}, got {len(vector)}" ) def query( diff --git a/runtime.py b/runtime.py new file mode 100644 index 0000000..f30cbfb --- /dev/null +++ b/runtime.py @@ -0,0 +1,141 @@ +from __future__ import annotations + +from collections import deque +from collections.abc import Mapping +from dataclasses import dataclass +from datetime import datetime +from pathlib import Path +from typing import Any + +import config +from events import EventBus +from events.bus import MemoryEvent +from forgetting.contradiction import ContradictionDetector +from forgetting.service import ForgettingService +from retrieval.retriever import UnifiedRetriever +from stores.episodic_store import EpisodicStore +from stores.media_store import MediaStore +from stores.procedural_store import ProceduralStore +from stores.semantic_store import SemanticStore +from utils.embeddings import GeminiEmbedder, TextEmbedder + + +def _jsonable(value: Any) -> Any: + if isinstance(value, Mapping): + return {key: _jsonable(val) for key, val in value.items()} + if isinstance(value, (list, tuple, set, frozenset)): + return [_jsonable(item) for item in value] + if isinstance(value, datetime): + return value.isoformat() + return value + + +class EventRecorder: + def __init__(self, bus: EventBus, *, max_events: int = 200): + self._events: deque[dict[str, Any]] = deque(maxlen=max_events) + for event_type in ( + "memory.stored", + "memory.retrieved", + "memory.ranked", + "memory.accessed", + "memory.contradiction_flagged", + "memory.supersession_resolved", + "memory.faded", + "memory.forgotten", + "forgetting.cycle_completed", + "forgetting.cycle_dry_run", + ): + bus.subscribe(event_type, self._record) + + def _record(self, event: MemoryEvent) -> None: + self._events.appendleft( + { + "event_type": event.event_type, + "timestamp": event.timestamp.isoformat(), + "data": _jsonable(dict(event.data)), + } + ) + + def snapshot(self, limit: int = 50) -> list[dict[str, Any]]: + return list(self._events)[:limit] + + +@dataclass(slots=True) +class MemoryRuntime: + media_store: MediaStore + bus: EventBus + embedder: TextEmbedder + semantic_store: SemanticStore + episodic_store: EpisodicStore + procedural_store: ProceduralStore + retriever: UnifiedRetriever + contradiction_detector: ContradictionDetector + forgetting_service: ForgettingService + events: EventRecorder + + +def build_runtime( + *, + chroma_path: str | None = None, + media_root: str | Path | None = None, + embedder: TextEmbedder | None = None, + embedding_dimensions: int | None = None, + max_events: int = 200, +) -> MemoryRuntime: + resolved_media_root = Path(media_root) if media_root is not None else Path(config.MEDIA_STORAGE_PATH) + resolved_dimensions = ( + config.EMBEDDING_DIMENSIONS if embedding_dimensions is None else embedding_dimensions + ) + + bus = EventBus() + active_embedder = embedder or GeminiEmbedder(dimensions=resolved_dimensions) + media_store = MediaStore(resolved_media_root) + semantic_store = SemanticStore( + chroma_path=chroma_path, + event_bus=bus, + embedder=active_embedder, + media_store=media_store, + ) + episodic_store = EpisodicStore( + chroma_path=chroma_path, + event_bus=bus, + embedder=active_embedder, + media_store=media_store, + ) + procedural_store = ProceduralStore( + chroma_path=chroma_path, + event_bus=bus, + embedder=active_embedder, + media_store=media_store, + ) + retriever = UnifiedRetriever( + stores={ + "semantic": semantic_store, + "episodic": episodic_store, + "procedural": procedural_store, + }, + event_bus=bus, + embedding_dimensions=resolved_dimensions, + ) + contradiction_detector = ContradictionDetector(semantic_store, event_bus=bus) + forgetting_service = ForgettingService( + semantic_store=semantic_store, + episodic_store=episodic_store, + procedural_store=procedural_store, + media_store=media_store, + event_bus=bus, + contradiction_detector=contradiction_detector, + ) + events = EventRecorder(bus, max_events=max_events) + return MemoryRuntime( + media_store=media_store, + bus=bus, + embedder=active_embedder, + semantic_store=semantic_store, + episodic_store=episodic_store, + procedural_store=procedural_store, + retriever=retriever, + contradiction_detector=contradiction_detector, + forgetting_service=forgetting_service, + events=events, + ) diff --git a/stores/episodic_store.py b/stores/episodic_store.py index 2d778b7..e58358d 100644 --- a/stores/episodic_store.py +++ b/stores/episodic_store.py @@ -38,13 +38,14 @@ class EpisodicStore(BaseStore): def __init__( self, + chroma_path: str | None = None, event_bus: EventBus | None = None, embedder: TextEmbedder | None = None, media_store: MediaStore | None = None, max_media_bytes: int | None = None, ): super().__init__(event_bus=event_bus) - client = chromadb.PersistentClient(path=config.CHROMA_DB_PATH) + client = chromadb.PersistentClient(path=chroma_path or config.CHROMA_DB_PATH) self._collection = client.get_or_create_collection( name="episodic_memories", metadata={"hnsw:space": "cosine"}, diff --git a/stores/procedural_store.py b/stores/procedural_store.py index 71b4a83..d22f79e 100644 --- a/stores/procedural_store.py +++ b/stores/procedural_store.py @@ -31,12 +31,13 @@ class ProceduralStore(BaseStore): def __init__( self, + chroma_path: str | None = None, event_bus: EventBus | None = None, embedder: TextEmbedder | None = None, media_store: MediaStore | None = None, ): super().__init__(event_bus=event_bus) - client = chromadb.PersistentClient(path=config.CHROMA_DB_PATH) + client = chromadb.PersistentClient(path=chroma_path or config.CHROMA_DB_PATH) self._collection = client.get_or_create_collection( name="procedural_memories", metadata={"hnsw:space": "cosine"}, diff --git a/stores/semantic_store.py b/stores/semantic_store.py index b40df92..c994224 100644 --- a/stores/semantic_store.py +++ b/stores/semantic_store.py @@ -17,12 +17,13 @@ class SemanticStore(BaseStore): def __init__( self, + chroma_path: str | None = None, event_bus: EventBus | None = None, embedder: TextEmbedder | None = None, media_store: MediaStore | None = None, ): super().__init__(event_bus=event_bus) - client = chromadb.PersistentClient(path=config.CHROMA_DB_PATH) + client = chromadb.PersistentClient(path=chroma_path or config.CHROMA_DB_PATH) self._collection = client.get_or_create_collection( name="semantic_memories", metadata={"hnsw:space": "cosine"}, diff --git a/tests/test_runtime.py b/tests/test_runtime.py new file mode 100644 index 0000000..8dacd66 --- /dev/null +++ b/tests/test_runtime.py @@ -0,0 +1,65 @@ +"""Verify the shared runtime factory isolates Chroma paths and config state.""" + +import os +import shutil +import sys +import tempfile + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +import config +from models.semantic import SemanticMemory +from runtime import build_runtime +from tests.helpers import HashingEmbedder + + +def _make_chroma_dir(prefix: str) -> str: + path = tempfile.mkdtemp(prefix=prefix) + shutil.rmtree(path, ignore_errors=True) + return path + + +def test_build_runtime_does_not_mutate_global_chroma_path(): + original_path = config.CHROMA_DB_PATH + runtime_path = _make_chroma_dir("runtime_factory_primary_") + + try: + runtime = build_runtime( + chroma_path=runtime_path, + embedder=HashingEmbedder(dimensions=config.EMBEDDING_DIMENSIONS), + ) + runtime.semantic_store.store(SemanticMemory(content="Runtime-scoped fact")) + + assert config.CHROMA_DB_PATH == original_path + assert runtime.semantic_store.get_by_id(runtime.semantic_store.get_all_records()[0].id) is not None + finally: + shutil.rmtree(runtime_path, ignore_errors=True) + + +def test_build_runtime_keeps_chroma_paths_isolated(): + original_path = config.CHROMA_DB_PATH + first_path = _make_chroma_dir("runtime_factory_first_") + second_path = _make_chroma_dir("runtime_factory_second_") + + try: + first_runtime = build_runtime( + chroma_path=first_path, + embedder=HashingEmbedder(dimensions=config.EMBEDDING_DIMENSIONS), + ) + second_runtime = build_runtime( + chroma_path=second_path, + embedder=HashingEmbedder(dimensions=config.EMBEDDING_DIMENSIONS), + ) + + first_runtime.semantic_store.store(SemanticMemory(content="Only in first runtime")) + second_runtime.semantic_store.store(SemanticMemory(content="Only in second runtime")) + + first_records = {record.content for record in first_runtime.semantic_store.get_all_records()} + second_records = {record.content for record in second_runtime.semantic_store.get_all_records()} + + assert first_records == {"Only in first runtime"} + assert second_records == {"Only in second runtime"} + assert config.CHROMA_DB_PATH == original_path + finally: + shutil.rmtree(first_path, ignore_errors=True) + shutil.rmtree(second_path, ignore_errors=True) diff --git a/utils/embeddings.py b/utils/embeddings.py index f03ab9a..546d6ff 100644 --- a/utils/embeddings.py +++ b/utils/embeddings.py @@ -9,7 +9,8 @@ from pathlib import Path from typing import Any, Protocol -from config import EMBEDDING_DIMENSIONS, EMBEDDING_MODEL, GEMINI_API_KEY +import config +from config import EMBEDDING_MODEL, GEMINI_API_KEY from utils.retry import retry_with_exponential_backoff _IMAGE_MIME_TYPES = { @@ -79,13 +80,14 @@ class EmbeddingProviderError(RuntimeError): class GeminiEmbedder: """Converts text and local media into Gemini embedding vectors.""" - def __init__(self): + def __init__(self, *, dimensions: int | None = None): self._genai = None self._client = None self._types = None self._errors = None self._doc_config = None self._query_config = None + self._dimensions = config.EMBEDDING_DIMENSIONS if dimensions is None else dimensions def embed_text(self, text: str) -> list[float]: return self._embed([text], self._document_config())[0] @@ -294,9 +296,9 @@ def _should_retry(exc: Exception) -> bool: raise def _normalize_vector(self, vector: list[float]) -> list[float]: - if len(vector) != EMBEDDING_DIMENSIONS: + if len(vector) != self._dimensions: raise ValueError( - f"Expected embedding dimension {EMBEDDING_DIMENSIONS}, got {len(vector)}" + f"Expected embedding dimension {self._dimensions}, got {len(vector)}" ) norm = math.sqrt(sum(value * value for value in vector)) if norm == 0: @@ -463,7 +465,7 @@ def _document_config(self): if self._doc_config is None: _, types, _ = self._load_sdk() self._doc_config = types.EmbedContentConfig( - output_dimensionality=EMBEDDING_DIMENSIONS, + output_dimensionality=self._dimensions, task_type="RETRIEVAL_DOCUMENT", ) return self._doc_config @@ -472,7 +474,7 @@ def _query_config_obj(self): if self._query_config is None: _, types, _ = self._load_sdk() self._query_config = types.EmbedContentConfig( - output_dimensionality=EMBEDDING_DIMENSIONS, + output_dimensionality=self._dimensions, task_type="RETRIEVAL_QUERY", ) return self._query_config From 588b3e7fb7b3efab26db765442834a8bd533e887 Mon Sep 17 00:00:00 2001 From: Atharva-Kanherkar Date: Mon, 30 Mar 2026 00:48:26 +0530 Subject: [PATCH 2/5] Restore CLI test seams after runtime refactor --- demo/cli.py | 58 ++++++++++++++++++++++++++++++++++++++++------------- 1 file changed, 44 insertions(+), 14 deletions(-) diff --git a/demo/cli.py b/demo/cli.py index 7fa06eb..e2e24b6 100644 --- a/demo/cli.py +++ b/demo/cli.py @@ -9,12 +9,15 @@ sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) -from events import ConsoleLogger +from events import ConsoleLogger, EventBus from models.episodic import EpisodicMemory from models.procedural import ProceduralMemory from models.semantic import SemanticMemory -from runtime import build_runtime +from retrieval.retriever import UnifiedRetriever +from stores.episodic_store import EpisodicStore from stores.media_store import MediaStore +from stores.procedural_store import ProceduralStore +from stores.semantic_store import SemanticStore from utils.embeddings import GeminiEmbedder, TextEmbedder import config @@ -34,10 +37,33 @@ def _make_embedder() -> TextEmbedder: return GeminiEmbedder() -def _make_runtime(): - runtime = build_runtime() - ConsoleLogger().register(runtime.bus) - return runtime +def _make_bus() -> EventBus: + bus = EventBus() + ConsoleLogger().register(bus) + return bus + + +def _make_semantic_store(event_bus: EventBus | None = None) -> SemanticStore: + return SemanticStore(event_bus=event_bus, media_store=_make_media_store()) + + +def _make_episodic_store(event_bus: EventBus | None = None) -> EpisodicStore: + return EpisodicStore(event_bus=event_bus) + + +def _make_procedural_store(event_bus: EventBus | None = None) -> ProceduralStore: + return ProceduralStore(event_bus=event_bus, media_store=_make_media_store()) + + +def _make_retriever(event_bus: EventBus | None = None) -> UnifiedRetriever: + return UnifiedRetriever( + stores={ + "semantic": _make_semantic_store(event_bus=event_bus), + "episodic": _make_episodic_store(event_bus=event_bus), + "procedural": _make_procedural_store(event_bus=event_bus), + }, + event_bus=event_bus, + ) def _guess_mime_type(path: str, modality: str) -> str: @@ -101,8 +127,9 @@ def _print_best_procedures(results) -> None: def _query_by_media(args, *, modality: str) -> None: - runtime = _make_runtime() + bus = _make_bus() embedder = _make_embedder() + retriever = _make_retriever(event_bus=bus) source_path = os.path.abspath(args.path) mime_type = _guess_mime_type(source_path, modality) if modality == "image": @@ -112,7 +139,7 @@ def _query_by_media(args, *, modality: str) -> None: else: raise ValueError(f"Unsupported query modality: {modality}") - results = runtime.retriever.query_by_vector( + results = retriever.query_by_vector( vector, top_k=args.top_k, memory_types=args.memory_types, @@ -125,7 +152,8 @@ def _query_by_media(args, *, modality: str) -> None: def cmd_store(args): - runtime = _make_runtime() + bus = _make_bus() + store = _make_semantic_store(event_bus=bus) media_path = None modality = "text" if args.image: @@ -142,14 +170,15 @@ def cmd_store(args): media_type=_infer_media_type(media_path) if media_path else None, ) try: - record_id = runtime.semantic_store.store(record) + record_id = store.store(record) except (FileNotFoundError, ValueError) as exc: _exit_with_error(str(exc)) print(f"Stored [{record_id[:8]}]: {args.content}") def cmd_store_episode(args): - runtime = _make_runtime() + bus = _make_bus() + store = _make_episodic_store(event_bus=bus) media_store = _make_media_store() if args.text is None else None if args.text is not None: @@ -173,7 +202,7 @@ def cmd_store_episode(args): record.media_ref = media_ref try: - record_id = runtime.episodic_store.store(record) + record_id = store.store(record) except Exception: if media_store is not None and record.media_ref: media_store.delete(record.media_ref) @@ -182,8 +211,9 @@ def cmd_store_episode(args): def cmd_query(args): - runtime = _make_runtime() - results = runtime.retriever.query(args.query, top_k=args.top_k) + bus = _make_bus() + retriever = _make_retriever(event_bus=bus) + results = retriever.query(args.query, top_k=args.top_k) if not results: print("No results found.") return From 4485a324b7f71221553a65c41ad3590c2409358c Mon Sep 17 00:00:00 2001 From: Atharva-Kanherkar Date: Mon, 30 Mar 2026 01:12:03 +0530 Subject: [PATCH 3/5] Harden runtime dimension injection --- runtime.py | 31 ++++++++++++++++++++++++++-- stores/episodic_store.py | 1 + stores/procedural_store.py | 1 + stores/semantic_store.py | 1 + tests/test_runtime.py | 42 +++++++++++++++++++++++++++++++++++--- 5 files changed, 71 insertions(+), 5 deletions(-) diff --git a/runtime.py b/runtime.py index f30cbfb..63b5bda 100644 --- a/runtime.py +++ b/runtime.py @@ -74,6 +74,14 @@ class MemoryRuntime: events: EventRecorder +def _embedder_dimensions(embedder: TextEmbedder) -> int | None: + for attr in ("_dimensions", "dimensions"): + value = getattr(embedder, attr, None) + if isinstance(value, int): + return value + return None + + def build_runtime( *, chroma_path: str | None = None, @@ -82,9 +90,28 @@ def build_runtime( embedding_dimensions: int | None = None, max_events: int = 200, ) -> MemoryRuntime: - resolved_media_root = Path(media_root) if media_root is not None else Path(config.MEDIA_STORAGE_PATH) + resolved_media_root = ( + Path(media_root) if media_root is not None else Path(config.MEDIA_STORAGE_PATH) + ) + embedder_dimensions = _embedder_dimensions(embedder) if embedder is not None else None + if ( + embedding_dimensions is not None + and embedder_dimensions is not None + and embedder_dimensions != embedding_dimensions + ): + raise ValueError( + "Provided embedder dimensions do not match embedding_dimensions: " + f"embedder={embedder_dimensions} embedding_dimensions={embedding_dimensions}" + ) + resolved_dimensions = ( - config.EMBEDDING_DIMENSIONS if embedding_dimensions is None else embedding_dimensions + embedding_dimensions + if embedding_dimensions is not None + else ( + embedder_dimensions + if embedder_dimensions is not None + else config.EMBEDDING_DIMENSIONS + ) ) bus = EventBus() diff --git a/stores/episodic_store.py b/stores/episodic_store.py index e58358d..e9e38b0 100644 --- a/stores/episodic_store.py +++ b/stores/episodic_store.py @@ -38,6 +38,7 @@ class EpisodicStore(BaseStore): def __init__( self, + *, chroma_path: str | None = None, event_bus: EventBus | None = None, embedder: TextEmbedder | None = None, diff --git a/stores/procedural_store.py b/stores/procedural_store.py index d22f79e..657a36a 100644 --- a/stores/procedural_store.py +++ b/stores/procedural_store.py @@ -31,6 +31,7 @@ class ProceduralStore(BaseStore): def __init__( self, + *, chroma_path: str | None = None, event_bus: EventBus | None = None, embedder: TextEmbedder | None = None, diff --git a/stores/semantic_store.py b/stores/semantic_store.py index c994224..f925f35 100644 --- a/stores/semantic_store.py +++ b/stores/semantic_store.py @@ -17,6 +17,7 @@ class SemanticStore(BaseStore): def __init__( self, + *, chroma_path: str | None = None, event_bus: EventBus | None = None, embedder: TextEmbedder | None = None, diff --git a/tests/test_runtime.py b/tests/test_runtime.py index 8dacd66..bd8db61 100644 --- a/tests/test_runtime.py +++ b/tests/test_runtime.py @@ -14,9 +14,7 @@ def _make_chroma_dir(prefix: str) -> str: - path = tempfile.mkdtemp(prefix=prefix) - shutil.rmtree(path, ignore_errors=True) - return path + return tempfile.mkdtemp(prefix=prefix) def test_build_runtime_does_not_mutate_global_chroma_path(): @@ -63,3 +61,41 @@ def test_build_runtime_keeps_chroma_paths_isolated(): finally: shutil.rmtree(first_path, ignore_errors=True) shutil.rmtree(second_path, ignore_errors=True) + + +def test_build_runtime_uses_custom_embedding_dimensions_for_vector_validation(): + runtime_path = _make_chroma_dir("runtime_factory_dimensions_") + + try: + runtime = build_runtime( + chroma_path=runtime_path, + embedding_dimensions=8, + embedder=HashingEmbedder(dimensions=8), + ) + + runtime.retriever.query_by_vector([0.0] * 8, top_k=1) + + try: + runtime.retriever.query_by_vector([0.0] * 7, top_k=1) + raise AssertionError("Expected dimension mismatch") + except ValueError as exc: + assert "Expected query vector dimension 8" in str(exc) + finally: + shutil.rmtree(runtime_path, ignore_errors=True) + + +def test_build_runtime_rejects_mismatched_embedder_and_dimension_override(): + runtime_path = _make_chroma_dir("runtime_factory_dimension_mismatch_") + + try: + try: + build_runtime( + chroma_path=runtime_path, + embedding_dimensions=16, + embedder=HashingEmbedder(dimensions=8), + ) + raise AssertionError("Expected mismatched dimensions to raise") + except ValueError as exc: + assert "Provided embedder dimensions do not match embedding_dimensions" in str(exc) + finally: + shutil.rmtree(runtime_path, ignore_errors=True) From 3bc2f1db3e56a84228c2dc5ec3b1759ed2b4500c Mon Sep 17 00:00:00 2001 From: Atharva-Kanherkar Date: Mon, 30 Mar 2026 01:17:31 +0530 Subject: [PATCH 4/5] Address PR review follow-ups --- api/app.py | 2 +- demo/cli.py | 2 ++ runtime.py | 6 +----- tests/helpers.py | 4 ++++ tests/test_runtime.py | 4 +++- utils/embeddings.py | 7 +++++++ 6 files changed, 18 insertions(+), 7 deletions(-) diff --git a/api/app.py b/api/app.py index 2be82f2..2bc82bb 100644 --- a/api/app.py +++ b/api/app.py @@ -329,7 +329,7 @@ def _safe_contradiction_lookup( ) -> list[ContradictionCandidate]: try: return active_service.contradiction_detector.find_potential_contradictions(record) - except (ValueError, Exception): + except ValueError: return [] diff --git a/demo/cli.py b/demo/cli.py index e2e24b6..5c765af 100644 --- a/demo/cli.py +++ b/demo/cli.py @@ -56,6 +56,8 @@ def _make_procedural_store(event_bus: EventBus | None = None) -> ProceduralStore def _make_retriever(event_bus: EventBus | None = None) -> UnifiedRetriever: + # The demo CLI keeps direct wiring so tests can patch the individual factory + # helpers without constructing the full shared runtime container. return UnifiedRetriever( stores={ "semantic": _make_semantic_store(event_bus=event_bus), diff --git a/runtime.py b/runtime.py index 63b5bda..a78001c 100644 --- a/runtime.py +++ b/runtime.py @@ -75,11 +75,7 @@ class MemoryRuntime: def _embedder_dimensions(embedder: TextEmbedder) -> int | None: - for attr in ("_dimensions", "dimensions"): - value = getattr(embedder, attr, None) - if isinstance(value, int): - return value - return None + return embedder.dimensions def build_runtime( diff --git a/tests/helpers.py b/tests/helpers.py index a7e884b..4161641 100644 --- a/tests/helpers.py +++ b/tests/helpers.py @@ -12,6 +12,10 @@ class HashingEmbedder: def __init__(self, dimensions: int = 64): self._dimensions = dimensions + @property + def dimensions(self) -> int | None: + return self._dimensions + def embed_text(self, text: str) -> list[float]: return self._embed(text) diff --git a/tests/test_runtime.py b/tests/test_runtime.py index bd8db61..a90d195 100644 --- a/tests/test_runtime.py +++ b/tests/test_runtime.py @@ -27,9 +27,11 @@ def test_build_runtime_does_not_mutate_global_chroma_path(): embedder=HashingEmbedder(dimensions=config.EMBEDDING_DIMENSIONS), ) runtime.semantic_store.store(SemanticMemory(content="Runtime-scoped fact")) + records = runtime.semantic_store.get_all_records() assert config.CHROMA_DB_PATH == original_path - assert runtime.semantic_store.get_by_id(runtime.semantic_store.get_all_records()[0].id) is not None + assert len(records) == 1 + assert runtime.semantic_store.get_by_id(records[0].id) is not None finally: shutil.rmtree(runtime_path, ignore_errors=True) diff --git a/utils/embeddings.py b/utils/embeddings.py index 546d6ff..97dbc25 100644 --- a/utils/embeddings.py +++ b/utils/embeddings.py @@ -68,6 +68,9 @@ class TextEmbedder(Protocol): + @property + def dimensions(self) -> int | None: ... + def embed_text(self, text: str) -> list[float]: ... def embed_query(self, text: str) -> list[float]: ... @@ -92,6 +95,10 @@ def __init__(self, *, dimensions: int | None = None): def embed_text(self, text: str) -> list[float]: return self._embed([text], self._document_config())[0] + @property + def dimensions(self) -> int | None: + return self._dimensions + def embed_query(self, text: str) -> list[float]: return self._embed([text], self._query_config_obj())[0] From 019f7598dd2a53ebd7d2dd2a8281422748345798 Mon Sep 17 00:00:00 2001 From: Atharva-Kanherkar Date: Mon, 30 Mar 2026 01:23:20 +0530 Subject: [PATCH 5/5] Expand runtime edge-case coverage --- tests/test_api_runtime_edges.py | 75 +++++++++++++++++++++++++ tests/test_runtime.py | 98 +++++++++++++++++++++++++++++++++ 2 files changed, 173 insertions(+) create mode 100644 tests/test_api_runtime_edges.py diff --git a/tests/test_api_runtime_edges.py b/tests/test_api_runtime_edges.py new file mode 100644 index 0000000..5e4e08e --- /dev/null +++ b/tests/test_api_runtime_edges.py @@ -0,0 +1,75 @@ +"""Edge-case API tests around the shared runtime refactor.""" + +import os +import sys +import tempfile + +import httpx +import pytest + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +from api.app import create_app +from tests.helpers import HashingEmbedder + + +def make_client(*, embedding_dimensions: int = 8) -> httpx.AsyncClient: + chroma_dir = tempfile.mkdtemp(prefix="memory_api_runtime_edges_") + app = create_app( + chroma_path=chroma_dir, + allowed_origins=["http://localhost:3000"], + embedder=HashingEmbedder(dimensions=embedding_dimensions), + embedding_dimensions=embedding_dimensions, + ) + transport = httpx.ASGITransport(app=app, raise_app_exceptions=False) + client = httpx.AsyncClient(transport=transport, base_url="http://testserver") + client.app = app + return client + + +@pytest.mark.anyio +async def test_semantic_write_swallows_valueerror_from_contradiction_lookup(): + async with make_client() as client: + service = client.app.state.service + assert service is None + + await client.get("/api/overview") + client.app.state.service.contradiction_detector.find_potential_contradictions = ( + lambda record: (_ for _ in ()).throw(ValueError("synthetic contradiction miss")) + ) + + response = await client.post( + "/api/memories/semantic", + json={"content": "ValueError contradiction lookup should be tolerated"}, + ) + + assert response.status_code == 200 + assert response.json()["potential_contradictions"] == [] + + +@pytest.mark.anyio +async def test_semantic_write_surfaces_programming_error_from_contradiction_lookup(): + async with make_client() as client: + await client.get("/api/overview") + client.app.state.service.contradiction_detector.find_potential_contradictions = ( + lambda record: (_ for _ in ()).throw(TypeError("synthetic programming error")) + ) + + response = await client.post( + "/api/memories/semantic", + json={"content": "TypeError contradiction lookup should not be suppressed"}, + ) + + assert response.status_code == 500 + + +@pytest.mark.anyio +async def test_create_app_uses_runtime_embedding_dimensions_for_service_retriever(): + async with make_client(embedding_dimensions=6) as client: + await client.get("/api/overview") + + try: + client.app.state.service.retriever.query_by_vector([0.0] * 5, top_k=1) + raise AssertionError("Expected retriever dimension mismatch") + except ValueError as exc: + assert "Expected query vector dimension 6" in str(exc) diff --git a/tests/test_runtime.py b/tests/test_runtime.py index a90d195..c384fef 100644 --- a/tests/test_runtime.py +++ b/tests/test_runtime.py @@ -13,6 +13,12 @@ from tests.helpers import HashingEmbedder +class DimensionlessEmbedder(HashingEmbedder): + @property + def dimensions(self) -> int | None: + return None + + def _make_chroma_dir(prefix: str) -> str: return tempfile.mkdtemp(prefix=prefix) @@ -101,3 +107,95 @@ def test_build_runtime_rejects_mismatched_embedder_and_dimension_override(): assert "Provided embedder dimensions do not match embedding_dimensions" in str(exc) finally: shutil.rmtree(runtime_path, ignore_errors=True) + + +def test_build_runtime_uses_embedder_dimensions_when_override_is_omitted(): + runtime_path = _make_chroma_dir("runtime_factory_embedder_dimensions_") + + try: + runtime = build_runtime( + chroma_path=runtime_path, + embedder=HashingEmbedder(dimensions=12), + ) + + runtime.retriever.query_by_vector([0.0] * 12, top_k=1) + + try: + runtime.retriever.query_by_vector([0.0] * 11, top_k=1) + raise AssertionError("Expected embedder-derived dimension mismatch") + except ValueError as exc: + assert "Expected query vector dimension 12" in str(exc) + finally: + shutil.rmtree(runtime_path, ignore_errors=True) + + +def test_build_runtime_allows_dimensionless_embedder_with_explicit_override(): + runtime_path = _make_chroma_dir("runtime_factory_dimensionless_embedder_") + + try: + runtime = build_runtime( + chroma_path=runtime_path, + embedding_dimensions=10, + embedder=DimensionlessEmbedder(dimensions=10), + ) + + runtime.retriever.query_by_vector([0.0] * 10, top_k=1) + finally: + shutil.rmtree(runtime_path, ignore_errors=True) + + +def test_build_runtime_uses_custom_media_root_for_owned_files(): + runtime_path = _make_chroma_dir("runtime_factory_media_root_db_") + media_root = tempfile.mkdtemp(prefix="runtime_factory_media_root_") + source_fd, source_path = tempfile.mkstemp(suffix=".png", prefix="runtime_factory_media_source_") + os.close(source_fd) + + try: + with open(source_path, "wb") as handle: + handle.write(b"runtime-media") + + runtime = build_runtime( + chroma_path=runtime_path, + media_root=media_root, + embedder=HashingEmbedder(dimensions=config.EMBEDDING_DIMENSIONS), + ) + record = SemanticMemory( + content="Runtime media root test", + modality="image", + media_ref=source_path, + media_type="image", + ) + + runtime.semantic_store.store(record) + + assert record.media_ref is not None + assert record.media_ref.startswith(media_root) + assert os.path.exists(record.media_ref) + finally: + shutil.rmtree(runtime_path, ignore_errors=True) + shutil.rmtree(media_root, ignore_errors=True) + try: + os.remove(source_path) + except FileNotFoundError: + pass + + +def test_build_runtime_event_recorder_honors_max_events_and_latest_first(): + runtime_path = _make_chroma_dir("runtime_factory_event_limit_") + + try: + runtime = build_runtime( + chroma_path=runtime_path, + max_events=2, + embedder=HashingEmbedder(dimensions=config.EMBEDDING_DIMENSIONS), + ) + + runtime.bus.emit("memory.stored", {"record_id": "first", "memory_type": "semantic"}) + runtime.bus.emit("memory.stored", {"record_id": "second", "memory_type": "semantic"}) + runtime.bus.emit("memory.stored", {"record_id": "third", "memory_type": "semantic"}) + + snapshot = runtime.events.snapshot(10) + + assert [event["data"]["record_id"] for event in snapshot] == ["third", "second"] + finally: + shutil.rmtree(runtime_path, ignore_errors=True)