agent-memory / memory_store.py
eriquesouza
Refactor query logic in MemoryStore for improved filtering
d617b55
Raw
History Blame Contribute Delete
7.78 kB
import chromadb
import uuid
import math
import json
from datetime import datetime, timezone
from typing import List, Optional
from models import Memory, MemoryType
# Phase 1: heuristic decay rates (per day)
DECAY_RATES = {
MemoryType.STATE: 0.15, # volatile β€” facts about the world
MemoryType.EPISODIC: 0.03, # moderate β€” specific past events
MemoryType.SEMANTIC: 0.005, # slow β€” abstracted patterns
MemoryType.PROCEDURAL: 0.002, # very slow β€” how-to knowledge
}
class MemoryStore:
def __init__(self, persist_dir: str = "./chroma_db"):
self.client = chromadb.PersistentClient(path=persist_dir)
self.collection = self.client.get_or_create_collection(
name="agent_memories",
metadata={"hnsw:space": "cosine"},
)
# ─── Write ────────────────────────────────────────────────────────────────
def add_memory(
self,
content: str,
memory_type: str,
source: str = "agent",
context_tags: Optional[List[str]] = None,
summary: str = "",
created_at: Optional[str] = None,
) -> Memory:
memory_id = str(uuid.uuid4())[:8]
now = created_at or datetime.now(timezone.utc).isoformat()
mem_type = MemoryType(memory_type)
decay_rate = DECAY_RATES.get(mem_type, 0.01)
auto_summary = summary or (content[:60] + "…" if len(content) > 60 else content)
metadata = {
"type": memory_type,
"created_at": now,
"source": source,
"context_tags": json.dumps(context_tags or []),
"access_count": 0,
"last_accessed": "",
"relevance_score": 1.0,
"decay_rate": decay_rate,
"active": "true",
"summary": auto_summary,
}
self.collection.add(
ids=[memory_id],
documents=[content],
metadatas=[metadata],
)
return self._build(memory_id, content, metadata)
def update_access(self, memory_id: str) -> None:
try:
result = self.collection.get(ids=[memory_id])
if not result["ids"]:
return
meta = result["metadatas"][0]
meta["access_count"] = int(meta.get("access_count", 0)) + 1
meta["last_accessed"] = datetime.now(timezone.utc).isoformat()
# Small heuristic boost on access, capped at 1.0
meta["relevance_score"] = min(1.0, float(meta.get("relevance_score", 1.0)) + 0.05)
self.collection.update(ids=[memory_id], metadatas=[meta])
except Exception as e:
print(f"[MemoryStore] update_access error for {memory_id}: {e}")
def archive_memory(self, memory_id: str) -> bool:
try:
result = self.collection.get(ids=[memory_id])
if not result["ids"]:
return False
meta = result["metadatas"][0]
meta["active"] = "false"
self.collection.update(ids=[memory_id], metadatas=[meta])
return True
except Exception:
return False
def reset(self) -> None:
self.client.delete_collection("agent_memories")
self.collection = self.client.get_or_create_collection(
name="agent_memories",
metadata={"hnsw:space": "cosine"},
)
# ─── Read ─────────────────────────────────────────────────────────────────
def count(self) -> int:
return self.collection.count()
def search(
self,
query: str,
n_results: int = 6,
type_filter: Optional[str] = None,
) -> List[Memory]:
total = self.collection.count()
if total == 0:
return []
if type_filter:
where: dict = {"$and": [{"active": {"$eq": "true"}}, {"type": {"$eq": type_filter}}]}
else:
where = {"active": {"$eq": "true"}}
try:
results = self.collection.query(
query_texts=[query],
n_results=min(n_results, total),
where=where,
)
except Exception as e:
print(f"[MemoryStore] search error: {e}")
return []
memories = []
if results["ids"] and results["ids"][0]:
for i, mem_id in enumerate(results["ids"][0]):
content = results["documents"][0][i]
meta = results["metadatas"][0][i]
memories.append(self._build(mem_id, content, meta))
return memories
def list_memories(
self,
type_filter: Optional[str] = None,
active_only: bool = True,
) -> List[Memory]:
if active_only and type_filter:
where: dict = {"$and": [{"active": {"$eq": "true"}}, {"type": {"$eq": type_filter}}]}
elif active_only:
where = {"active": {"$eq": "true"}}
elif type_filter:
where = {"type": {"$eq": type_filter}}
else:
where = None
try:
results = self.collection.get(where=where)
except Exception as e:
print(f"[MemoryStore] list error: {e}")
return []
memories = []
for i, mem_id in enumerate(results["ids"]):
content = results["documents"][i]
meta = results["metadatas"][i]
memories.append(self._build(mem_id, content, meta))
memories.sort(key=lambda m: m.created_at, reverse=True)
return memories
# ─── Internal ─────────────────────────────────────────────────────────────
def _compute_relevance(self, created_at: str, decay_rate: float, stored_score: float) -> float:
try:
created = datetime.fromisoformat(created_at)
now = datetime.now(timezone.utc)
if created.tzinfo is None:
created = created.replace(tzinfo=timezone.utc)
days_old = max(0.0, (now - created).total_seconds() / 86400)
decayed = stored_score * math.exp(-decay_rate * days_old)
return round(max(0.01, decayed), 3)
except Exception:
return stored_score
def _build(self, mem_id: str, content: str, meta: dict) -> Memory:
decay_rate = float(meta.get("decay_rate", 0.01))
stored_score = float(meta.get("relevance_score", 1.0))
current_relevance = self._compute_relevance(
meta.get("created_at", ""), decay_rate, stored_score
)
tags = meta.get("context_tags", "[]")
if isinstance(tags, str):
try:
tags = json.loads(tags)
except Exception:
tags = []
active_raw = meta.get("active", "true")
active = active_raw == "true" if isinstance(active_raw, str) else bool(active_raw)
return Memory(
id=mem_id,
content=content,
type=MemoryType(meta.get("type", "semantic")),
created_at=meta.get("created_at", ""),
source=meta.get("source", "agent"),
context_tags=tags,
access_count=int(meta.get("access_count", 0)),
last_accessed=meta.get("last_accessed") or None,
relevance_score=current_relevance,
decay_rate=decay_rate,
active=active,
summary=meta.get("summary", content[:60]),
)