Spaces:
Sleeping
Sleeping
| """ | |
| Agent Memory Store | |
| SQLite-backed persistent memory for agents: episodic events (goals, actions, | |
| observations, final answers) and key/value semantic facts, scoped per agent_id | |
| and optionally per task_id. Thread-safe; safe to share across agents. | |
| """ | |
| from __future__ import annotations | |
| import json | |
| import sqlite3 | |
| import threading | |
| import time | |
| import uuid | |
| from contextlib import contextmanager | |
| from dataclasses import dataclass, field | |
| from pathlib import Path | |
| from typing import Any, Optional | |
| DEFAULT_DB_PATH = Path(__file__).resolve().parents[2] / "data" / "agent_memory.db" | |
| _SCHEMA = """ | |
| CREATE TABLE IF NOT EXISTS episodic_memory ( | |
| id TEXT PRIMARY KEY, | |
| agent_id TEXT NOT NULL, | |
| task_id TEXT, | |
| role TEXT NOT NULL, | |
| content TEXT NOT NULL, | |
| metadata TEXT, | |
| created_at REAL NOT NULL | |
| ); | |
| CREATE INDEX IF NOT EXISTS idx_episodic_agent ON episodic_memory(agent_id, created_at); | |
| CREATE INDEX IF NOT EXISTS idx_episodic_task ON episodic_memory(task_id); | |
| CREATE TABLE IF NOT EXISTS semantic_memory ( | |
| agent_id TEXT NOT NULL, | |
| key TEXT NOT NULL, | |
| value TEXT NOT NULL, | |
| updated_at REAL NOT NULL, | |
| PRIMARY KEY (agent_id, key) | |
| ); | |
| """ | |
| class MemoryEvent: | |
| id: str | |
| agent_id: str | |
| role: str | |
| content: str | |
| task_id: Optional[str] = None | |
| metadata: dict[str, Any] = field(default_factory=dict) | |
| created_at: float = field(default_factory=time.time) | |
| class AgentMemoryStore: | |
| """Shared SQLite-backed memory store used by every agent in the runtime.""" | |
| def __init__(self, db_path: Path | str = DEFAULT_DB_PATH): | |
| self.db_path = Path(db_path) | |
| self.db_path.parent.mkdir(parents=True, exist_ok=True) | |
| self._lock = threading.Lock() | |
| with self._connect() as conn: | |
| conn.executescript(_SCHEMA) | |
| def _connect(self): | |
| conn = sqlite3.connect(self.db_path, timeout=30) | |
| try: | |
| yield conn | |
| conn.commit() | |
| finally: | |
| conn.close() | |
| def record_event( | |
| self, | |
| agent_id: str, | |
| role: str, | |
| content: str, | |
| task_id: Optional[str] = None, | |
| metadata: Optional[dict[str, Any]] = None, | |
| ) -> MemoryEvent: | |
| event = MemoryEvent( | |
| id=str(uuid.uuid4()), | |
| agent_id=agent_id, | |
| role=role, | |
| content=content, | |
| task_id=task_id, | |
| metadata=metadata or {}, | |
| ) | |
| with self._lock, self._connect() as conn: | |
| conn.execute( | |
| "INSERT INTO episodic_memory (id, agent_id, task_id, role, content, metadata, created_at) " | |
| "VALUES (?, ?, ?, ?, ?, ?, ?)", | |
| ( | |
| event.id, | |
| event.agent_id, | |
| event.task_id, | |
| event.role, | |
| event.content, | |
| json.dumps(event.metadata), | |
| event.created_at, | |
| ), | |
| ) | |
| return event | |
| def recent_events( | |
| self, agent_id: str, limit: int = 20, task_id: Optional[str] = None | |
| ) -> list[MemoryEvent]: | |
| query = ( | |
| "SELECT id, agent_id, task_id, role, content, metadata, created_at " | |
| "FROM episodic_memory WHERE agent_id = ?" | |
| ) | |
| params: list[Any] = [agent_id] | |
| if task_id: | |
| query += " AND task_id = ?" | |
| params.append(task_id) | |
| query += " ORDER BY created_at DESC LIMIT ?" | |
| params.append(limit) | |
| with self._lock, self._connect() as conn: | |
| rows = conn.execute(query, params).fetchall() | |
| events = [ | |
| MemoryEvent( | |
| id=r[0], agent_id=r[1], task_id=r[2], role=r[3], content=r[4], | |
| metadata=json.loads(r[5]) if r[5] else {}, created_at=r[6], | |
| ) | |
| for r in rows | |
| ] | |
| events.reverse() | |
| return events | |
| def set_fact(self, agent_id: str, key: str, value: Any) -> None: | |
| with self._lock, self._connect() as conn: | |
| conn.execute( | |
| "INSERT INTO semantic_memory (agent_id, key, value, updated_at) VALUES (?, ?, ?, ?) " | |
| "ON CONFLICT(agent_id, key) DO UPDATE SET value = excluded.value, updated_at = excluded.updated_at", | |
| (agent_id, key, json.dumps(value), time.time()), | |
| ) | |
| def get_fact(self, agent_id: str, key: str, default: Any = None) -> Any: | |
| with self._lock, self._connect() as conn: | |
| row = conn.execute( | |
| "SELECT value FROM semantic_memory WHERE agent_id = ? AND key = ?", | |
| (agent_id, key), | |
| ).fetchone() | |
| return json.loads(row[0]) if row else default | |
| def all_facts(self, agent_id: str) -> dict[str, Any]: | |
| with self._lock, self._connect() as conn: | |
| rows = conn.execute( | |
| "SELECT key, value FROM semantic_memory WHERE agent_id = ?", (agent_id,) | |
| ).fetchall() | |
| return {k: json.loads(v) for k, v in rows} | |
| def clear_agent(self, agent_id: str) -> None: | |
| with self._lock, self._connect() as conn: | |
| conn.execute("DELETE FROM episodic_memory WHERE agent_id = ?", (agent_id,)) | |
| conn.execute("DELETE FROM semantic_memory WHERE agent_id = ?", (agent_id,)) | |