""" 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()