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
114 changes: 22 additions & 92 deletions api/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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",
Expand All @@ -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:
Expand Down Expand Up @@ -279,88 +264,31 @@ 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,
*,
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"
Expand Down Expand Up @@ -401,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 []


Expand Down Expand Up @@ -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(
Expand All @@ -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:
Expand Down
20 changes: 11 additions & 9 deletions demo/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,11 +13,11 @@
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
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

Expand All @@ -29,6 +29,14 @@
}


def _make_media_store() -> MediaStore:
return MediaStore(config.MEDIA_STORAGE_PATH)


def _make_embedder() -> TextEmbedder:
return GeminiEmbedder()


def _make_bus() -> EventBus:
bus = EventBus()
ConsoleLogger().register(bus)
Expand All @@ -47,15 +55,9 @@ 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)


def _make_embedder() -> TextEmbedder:
return GeminiEmbedder()


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),
Expand Down
16 changes: 12 additions & 4 deletions retrieval/retriever.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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:
Expand Down Expand Up @@ -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(
Expand Down
Loading
Loading