Multi-Agent-System / backend /memory /memory_store.py
jatin gyass
initial commit
2eef9ea
Raw
History Blame Contribute Delete
9.93 kB
"""
backend/memory/memory_store.py
Two-tier memory architecture:
Short-term (Redis):
- Current session context
- Recent tool results
- Working memory for active task
- TTL: configurable (default 24h)
Long-term (SQLite/PostgreSQL):
- Episodic memory: what happened in past tasks
- Semantic memory: learned facts and patterns
- Procedural memory: successful task strategies
- Persists across restarts
Memory retrieval uses simple keyword matching (production: use embeddings + vector DB).
"""
from __future__ import annotations
import hashlib
import json
import time
from datetime import datetime, timezone
from typing import Any
from sqlalchemy import Column, String, Float, Text, Integer, DateTime, create_engine, select
from sqlalchemy.orm import DeclarativeBase, Session
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker
from ..core.config import get_settings
from ..core.logger import get_logger
log = get_logger(__name__)
# ── SQLAlchemy models ─────────────────────────────────────────────────────────
class Base(DeclarativeBase):
pass
class MemoryRecord(Base):
__tablename__ = "memories"
id = Column(String, primary_key=True)
task_id = Column(String, index=True)
content = Column(Text, nullable=False)
memory_type = Column(String, default="episodic") # episodic|semantic|procedural
importance = Column(Float, default=0.5)
tags = Column(Text, default="[]") # JSON list
access_count = Column(Integer, default=0)
created_at = Column(DateTime, default=datetime.utcnow)
last_accessed = Column(DateTime, default=datetime.utcnow)
class TaskRecord(Base):
__tablename__ = "tasks"
task_id = Column(String, primary_key=True)
task = Column(Text, nullable=False)
status = Column(String, default="pending")
final_output = Column(Text)
quality_score = Column(Float)
total_tokens = Column(Integer, default=0)
created_at = Column(DateTime, default=datetime.utcnow)
completed_at = Column(DateTime)
state_json = Column(Text) # Full state snapshot
# ── Database setup ─────────────────────────────────────────────────────────────
_engine = None
_session_factory = None
async def init_db():
global _engine, _session_factory
settings = get_settings()
_engine = create_async_engine(settings.database_url, echo=False)
_session_factory = async_sessionmaker(_engine, expire_on_commit=False)
async with _engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
log.info("Database initialized", url=settings.database_url)
async def get_session() -> AsyncSession:
return _session_factory()
# ── Redis short-term memory ───────────────────────────────────────────────────
_redis_client = None
_redis_checked = False # prevents re-attempting after a failed connect
def get_redis():
global _redis_client, _redis_checked
if _redis_checked:
return _redis_client
_redis_checked = True
import redis
settings = get_settings()
try:
client = redis.from_url(
settings.redis_url, decode_responses=True,
socket_connect_timeout=2, socket_timeout=2,
)
client.ping()
_redis_client = client
except Exception as e:
log.warning("Redis unavailable β€” short-term memory disabled", error=str(e))
_redis_client = None
return _redis_client
class ShortTermMemory:
"""Redis-backed working memory for active tasks."""
PREFIX = "agent:stm:"
def set(self, task_id: str, key: str, value: Any, ttl: int | None = None) -> None:
r = get_redis()
if r is None:
return
full_key = f"{self.PREFIX}{task_id}:{key}"
r.set(full_key, json.dumps(value), ex=ttl or get_settings().redis_ttl)
def get(self, task_id: str, key: str) -> Any | None:
r = get_redis()
if r is None:
return None
raw = r.get(f"{self.PREFIX}{task_id}:{key}")
return json.loads(raw) if raw else None
def get_all(self, task_id: str) -> dict[str, Any]:
r = get_redis()
if r is None:
return {}
pattern = f"{self.PREFIX}{task_id}:*"
keys = r.keys(pattern)
result = {}
for k in keys:
sub_key = k.replace(f"{self.PREFIX}{task_id}:", "")
raw = r.get(k)
if raw:
result[sub_key] = json.loads(raw)
return result
def store_state(self, task_id: str, state: dict) -> None:
"""Cache full workflow state for resumption."""
self.set(task_id, "state", state, ttl=3600)
def get_state(self, task_id: str) -> dict | None:
return self.get(task_id, "state")
def clear(self, task_id: str) -> None:
r = get_redis()
if r is None:
return
for k in r.keys(f"{self.PREFIX}{task_id}:*"):
r.delete(k)
# ── Long-term memory ──────────────────────────────────────────────────────────
class LongTermMemory:
"""SQLite/PostgreSQL-backed episodic + semantic memory."""
async def store(self, memory: dict, task_id: str = "") -> str:
mem_id = hashlib.sha256(
(memory["content"] + str(time.time())).encode()
).hexdigest()[:12]
async with await get_session() as session:
record = MemoryRecord(
id=mem_id,
task_id=task_id,
content=memory["content"],
memory_type=memory.get("memory_type", "episodic"),
importance=memory.get("importance", 0.5),
tags=json.dumps(memory.get("tags", [])),
)
session.add(record)
await session.commit()
log.debug("Memory stored", id=mem_id, type=memory.get("memory_type"))
return mem_id
async def retrieve(
self,
query: str,
memory_type: str | None = None,
limit: int = 5,
min_importance: float = 0.3,
) -> list[dict]:
"""
Retrieve relevant memories using keyword matching.
Production upgrade: embed query + cosine similarity with pgvector.
"""
async with await get_session() as session:
result = await session.execute(
select(MemoryRecord)
.where(MemoryRecord.importance >= min_importance)
.order_by(MemoryRecord.importance.desc())
.limit(50)
)
all_memories = result.scalars().all()
# Score by keyword overlap
query_words = set(query.lower().split())
scored = []
for m in all_memories:
if memory_type and m.memory_type != memory_type:
continue
content_words = set(m.content.lower().split())
overlap = len(query_words & content_words)
if overlap > 0:
scored.append((overlap, m))
scored.sort(key=lambda x: x[0], reverse=True)
results = []
for _, m in scored[:limit]:
results.append({
"memory_id": m.id,
"content": m.content,
"memory_type": m.memory_type,
"importance": m.importance,
"tags": json.loads(m.tags),
"created_at": m.created_at.isoformat() if m.created_at else "",
})
return results
async def store_task(self, task_id: str, task: str, state: dict) -> None:
"""Persist task record to DB."""
async with await get_session() as session:
existing = await session.get(TaskRecord, task_id)
if existing:
existing.status = state.get("status", "unknown")
existing.final_output = state.get("final_output")
existing.quality_score = state.get("quality_score")
existing.total_tokens = state.get("total_tokens", 0)
existing.state_json = json.dumps(state, default=str)
if state.get("status") in ("completed", "failed"):
existing.completed_at = datetime.utcnow()
else:
record = TaskRecord(
task_id=task_id,
task=task,
status=state.get("status", "pending"),
state_json=json.dumps(state, default=str),
)
session.add(record)
await session.commit()
async def get_recent_tasks(self, limit: int = 10) -> list[dict]:
async with await get_session() as session:
result = await session.execute(
select(TaskRecord).order_by(TaskRecord.created_at.desc()).limit(limit)
)
tasks = result.scalars().all()
return [
{
"task_id": t.task_id,
"task": t.task[:100],
"status": t.status,
"quality_score": t.quality_score,
"total_tokens": t.total_tokens,
"created_at": t.created_at.isoformat() if t.created_at else "",
}
for t in tasks
]
# ── Unified memory interface ──────────────────────────────────────────────────
short_term = ShortTermMemory()
long_term = LongTermMemory()