Spaces:
Sleeping
Sleeping
| """ | |
| 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() | |