jarvis-cloud / backend /memory /episodic_memory.py
Jarvis2345's picture
Squash history — remove all prior commits (secret hygiene, S4)
a31f556
Raw
History Blame
5.85 kB
# backend/memory/episodic_memory.py
import logging
import asyncio
from uuid import uuid4
try:
import chromadb
from chromadb.config import Settings
except ImportError:
chromadb = None
Settings = None
try:
import google.generativeai as genai
from config import GEMINI_API_KEY
# Initialize genai explicitly
genai.configure(api_key=GEMINI_API_KEY)
except ImportError:
pass
async def embed(text: str) -> list[float]:
"""Generates an embedding for the given text using Gemini."""
try:
def _get_embedding():
result = genai.embed_content(
model="models/embedding-001",
content=text,
task_type="retrieval_document",
)
return result['embedding']
return await asyncio.to_thread(_get_embedding)
except Exception as e:
logging.error(f"Embedding failed: {e}")
return [0.0] * 768
class EpisodicMemory:
def __init__(self, persist_dir: str):
# ITEM: Cross-session persistence confirmed
# Real ChromaDB persistent client — data survives app restarts
if chromadb is None:
logging.error("chromadb is not installed. Episodic memory will not work.")
self.client = None
return
import logging
logging.getLogger("chromadb").setLevel(logging.ERROR)
self.client = chromadb.PersistentClient(
path=persist_dir,
settings=Settings(anonymized_telemetry=False)
)
def _get_collection(self, persona: str):
if not self.client:
return None
return self.client.get_or_create_collection(f"persona_{persona.lower()}")
async def add(self, text: str, metadata: dict, persona: str = "jarvis"):
# ITEM: Memory write — EXE (exact add() call with collection=f"persona_{active_persona}")
collection = self._get_collection(persona)
if not collection:
return
embedding = await embed(text)
metadata["persona"] = persona.lower()
metadata["pinned"] = metadata.get("pinned", False)
def _add():
collection.add(
documents=[text],
embeddings=[embedding],
metadatas=[metadata],
ids=[str(uuid4())]
)
await asyncio.to_thread(_add)
async def query(self, text: str, top_k=5, persona: str = "jarvis") -> list[dict]:
# ITEM: Memory read/query — EXE (exact query() call filtered to active persona collection)
collection = self._get_collection(persona)
if not collection:
return []
embedding = await embed(text)
def _query():
return collection.query(
query_embeddings=[embedding],
n_results=top_k
)
results = await asyncio.to_thread(_query)
out = []
if results and results.get("documents") and results["documents"][0]:
docs = results["documents"][0]
metas = results["metadatas"][0] if results.get("metadatas") and results["metadatas"][0] else [{}] * len(docs)
ids = results["ids"][0]
for d, m, i in zip(docs, metas, ids):
out.append({"id": i, "text": d, "metadata": m})
return out
async def pin(self, memory_id: str, persona: str = "jarvis"):
# ITEM: Memory pin — EXE (exact metadata update to mark pinned=True)
collection = self._get_collection(persona)
if not collection:
return
def _pin():
result = collection.get(ids=[memory_id])
if result and result['metadatas'] and len(result['metadatas']) > 0:
meta = result['metadatas'][0]
meta['pinned'] = True
collection.update(
ids=[memory_id],
metadatas=[meta]
)
await asyncio.to_thread(_pin)
async def delete(self, memory_id: str, persona: str = "jarvis"):
# ITEM: Memory delete — EXE (exact delete() call)
collection = self._get_collection(persona)
if not collection:
return
def _delete():
collection.delete(ids=[memory_id])
await asyncio.to_thread(_delete)
async def embed_model_description(self, description: str, model_path: str, persona: str):
collection = self._get_collection(f"{persona}_models")
if not collection:
return
def _add():
collection.add(
documents=[description],
metadatas=[{"model_path": model_path, "engine": "stable_fast_3d_local"}],
ids=[str(uuid4())]
)
try:
await asyncio.to_thread(_add)
except Exception as e:
logging.error(f"ChromaDB embed error: {e}")
async def query_similar_model(self, description: str, threshold: float = 0.85, persona: str = "JARVIS") -> dict | None:
collection = self._get_collection(f"{persona}_models")
if not collection:
return None
def _query():
return collection.query(
query_texts=[description],
n_results=1
)
try:
results = await asyncio.to_thread(_query)
if results and results.get('distances') and results['distances'][0]:
dist = results['distances'][0][0]
if dist <= (1.0 - threshold):
meta = results['metadatas'][0][0]
return {"model_path": meta.get("model_path")}
except Exception as e:
logging.error(f"ChromaDB query error: {e}")
return None