jarvis-mvp / jarvis /memory.py
Turojin's picture
Upload jarvis/memory.py
fb283e5 verified
Raw History Blame
8.18 kB
"""Memory layer: SQLite + ChromaDB for J.A.R.V.I.S."""
import sqlite3
import json
import uuid
from datetime import datetime
from pathlib import Path
from typing import List, Dict, Optional, Any
import chromadb
from chromadb.config import Settings
from sentence_transformers import SentenceTransformer
from jarvis.config import MemoryConfig
class SQLiteStore:
"""Structured data: user state, open loops, preferences, mirror log."""
def __init__(self, db_path: str):
self.db_path = db_path
Path(db_path).parent.mkdir(parents=True, exist_ok=True)
self._init_tables()
def _connect(self):
return sqlite3.connect(self.db_path)
def _init_tables(self):
with self._connect() as conn:
conn.executescript("""
CREATE TABLE IF NOT EXISTS user_state (
key TEXT PRIMARY KEY,
value TEXT,
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
);
CREATE TABLE IF NOT EXISTS open_loops (
id TEXT PRIMARY KEY,
title TEXT,
description TEXT,
status TEXT DEFAULT 'open',
priority INTEGER DEFAULT 5,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
last_nudged_at TIMESTAMP
);
CREATE TABLE IF NOT EXISTS mirror_check_log (
id INTEGER PRIMARY KEY AUTOINCREMENT,
turn_id TEXT,
contracts TEXT,
dials TEXT,
drift_detected BOOLEAN,
notes TEXT,
checked_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
);
CREATE TABLE IF NOT EXISTS exchanges (
id TEXT PRIMARY KEY,
user_message TEXT,
jarvis_response TEXT,
contract TEXT,
dial INTEGER,
archetype_mix TEXT,
score REAL,
timestamp TIMESTAMP DEFAULT CURRENT_TIMESTAMP
);
""")
# ---- User State ----
def set_state(self, key: str, value: Any):
with self._connect() as conn:
conn.execute(
"INSERT INTO user_state (key, value) VALUES (?, ?) ON CONFLICT(key) DO UPDATE SET value=excluded.value, updated_at=CURRENT_TIMESTAMP",
(key, json.dumps(value)),
)
def get_state(self, key: str, default: Any = None) -> Any:
with self._connect() as conn:
row = conn.execute("SELECT value FROM user_state WHERE key=?", (key,)).fetchone()
if row:
return json.loads(row[0])
return default
# ---- Open Loops ----
def add_loop(self, title: str, description: str = "", priority: int = 5) -> str:
lid = str(uuid.uuid4())
with self._connect() as conn:
conn.execute(
"INSERT INTO open_loops (id, title, description, priority) VALUES (?, ?, ?, ?)",
(lid, title, description, priority),
)
return lid
def list_loops(self, status: Optional[str] = None) -> List[Dict]:
with self._connect() as conn:
if status:
rows = conn.execute("SELECT * FROM open_loops WHERE status=? ORDER BY priority DESC", (status,)).fetchall()
else:
rows = conn.execute("SELECT * FROM open_loops ORDER BY priority DESC").fetchall()
cols = [d[0] for d in conn.execute("SELECT * FROM open_loops LIMIT 0").description]
return [dict(zip(cols, row)) for row in rows]
def update_loop(self, loop_id: str, **kwargs):
sets = ", ".join(f"{k}=?" for k in kwargs)
vals = list(kwargs.values()) + [loop_id]
with self._connect() as conn:
conn.execute(f"UPDATE open_loops SET {sets}, updated_at=CURRENT_TIMESTAMP WHERE id=?", vals)
# ---- Exchanges ----
def log_exchange(self, user_msg: str, response: str, contract: str, dial: int,
archetype_mix: Dict[str, float], score: Optional[float] = None) -> str:
eid = str(uuid.uuid4())
with self._connect() as conn:
conn.execute(
"INSERT INTO exchanges (id, user_message, jarvis_response, contract, dial, archetype_mix, score) VALUES (?, ?, ?, ?, ?, ?, ?)",
(eid, user_msg, response, contract, dial, json.dumps(archetype_mix), score),
)
return eid
def get_recent_exchanges(self, n: int = 5) -> List[Dict]:
with self._connect() as conn:
rows = conn.execute(
"SELECT * FROM exchanges ORDER BY timestamp DESC LIMIT ?", (n,)
).fetchall()
if not rows:
return []
cols = [d[0] for d in conn.execute("SELECT * FROM exchanges LIMIT 0").description]
return [dict(zip(cols, row)) for row in reversed(rows)]
class ChromaStore:
"""Vector memory for semantic recall of conversation history."""
def __init__(self, persist_dir: str, embedding_model: str = "all-MiniLM-L6-v2"):
Path(persist_dir).mkdir(parents=True, exist_ok=True)
self.client = chromadb.PersistentClient(
path=persist_dir,
settings=Settings(anonymized_telemetry=False),
)
self.collection = self.client.get_or_create_collection("jarvis_memory")
self.embedder = SentenceTransformer(embedding_model)
def add(self, text: str, metadata: Optional[Dict] = None, doc_id: Optional[str] = None):
doc_id = doc_id or str(uuid.uuid4())
embedding = self.embedder.encode(text).tolist()
self.collection.add(
ids=[doc_id],
embeddings=[embedding],
documents=[text],
metadatas=[metadata or {}],
)
def query(self, query_text: str, n_results: int = 5) -> List[Dict]:
embedding = self.embedder.encode(query_text).tolist()
results = self.collection.query(
query_embeddings=[embedding],
n_results=n_results,
include=["documents", "metadatas", "distances"],
)
out = []
for i in range(len(results["ids"][0])):
out.append({
"id": results["ids"][0][i],
"document": results["documents"][0][i],
"metadata": results["metadatas"][0][i],
"distance": results["distances"][0][i],
})
return out
class Memory:
"""Unified memory interface."""
def __init__(self, config: MemoryConfig):
self.sqlite = SQLiteStore(config.sqlite_path)
self.chroma = None
if config.use_chroma:
self.chroma = ChromaStore(config.chroma_path, config.embedding_model)
self.max_context_turns = config.max_context_turns
def log_turn(self, user_msg: str, response: str, contract: str, dial: int,
archetype_mix: Dict[str, float], score: Optional[float] = None):
eid = self.sqlite.log_exchange(user_msg, response, contract, dial, archetype_mix, score)
if self.chroma:
self.chroma.add(
text=f"User: {user_msg}\nJ.A.R.V.I.S.: {response}",
metadata={
"exchange_id": eid,
"contract": contract,
"dial": dial,
"timestamp": datetime.utcnow().isoformat(),
},
doc_id=eid,
)
def get_recent_context(self, n: Optional[int] = None) -> str:
n = n or self.max_context_turns
exchanges = self.sqlite.get_recent_exchanges(n)
lines = []
for ex in exchanges:
lines.append(f"User: {ex['user_message']}")
lines.append(f"J.A.R.V.I.S.: {ex['jarvis_response']}")
return "\n".join(lines)
def semantic_recall(self, query: str, n: int = 3) -> List[Dict]:
if self.chroma is None:
return []
return self.chroma.query(query, n_results=n)