Spaces:
Paused
Paused
| """Long-term memory implementation using Qdrant with Redis/local fallback.""" | |
| from __future__ import annotations | |
| import json | |
| import logging | |
| from datetime import UTC, datetime | |
| from pathlib import Path | |
| from typing import Any | |
| from hermes.config.settings import get_settings | |
| logger = logging.getLogger(__name__) | |
| _embedding_model: Any = None | |
| def _get_embedding_model() -> Any: | |
| """Get or create singleton embedding model.""" | |
| global _embedding_model | |
| if _embedding_model is None: | |
| try: | |
| from sentence_transformers import SentenceTransformer | |
| settings = get_settings() | |
| _embedding_model = SentenceTransformer(settings.model.embedding_model) | |
| except Exception: | |
| _embedding_model = False # Sentinel: tried and failed | |
| return _embedding_model if _embedding_model is not False else None | |
| class LongTermMemory: | |
| """Long-term memory using vector database with fallback chain: Qdrant → Redis → JSON file.""" | |
| def __init__(self) -> None: | |
| self.settings = get_settings() | |
| self._client: Any = None | |
| self._redis: Any = None | |
| self._collection = self.settings.database.qdrant_collection | |
| self._local_store: list[dict[str, Any]] = [] | |
| self._json_path = Path("data/long_term_memory.json") | |
| self._load_json() | |
| def _load_json(self) -> None: | |
| """Load persisted memories from JSON file.""" | |
| try: | |
| if self._json_path.exists(): | |
| with open(self._json_path, encoding="utf-8") as f: | |
| self._local_store = json.load(f) | |
| except Exception as e: | |
| logger.warning(f"Could not load memory JSON: {e}") | |
| def _save_json(self) -> None: | |
| """Persist memories to JSON file.""" | |
| try: | |
| self._json_path.parent.mkdir(parents=True, exist_ok=True) | |
| with open(self._json_path, "w", encoding="utf-8") as f: | |
| json.dump(self._local_store, f, indent=2, default=str) | |
| except Exception as e: | |
| logger.warning(f"Could not save memory JSON: {e}") | |
| async def initialize(self) -> None: | |
| """Initialize the long-term memory with fallback chain.""" | |
| # Try Qdrant first | |
| try: | |
| from qdrant_client import QdrantClient | |
| from qdrant_client.models import Distance, VectorParams | |
| self._client = QdrantClient(url=self.settings.database.qdrant_url) | |
| collections = self._client.get_collections().collections | |
| collection_names = [c.name for c in collections] | |
| if self._collection not in collection_names: | |
| self._client.create_collection( | |
| collection_name=self._collection, | |
| vectors_config=VectorParams( | |
| size=self.settings.model.embedding_dimension, | |
| distance=Distance.COSINE, | |
| ), | |
| ) | |
| logger.info(f"Created Qdrant collection: {self._collection}") | |
| return | |
| except Exception as e: | |
| logger.warning(f"Qdrant unavailable: {e}") | |
| self._client = None | |
| # Try Redis as fallback | |
| try: | |
| import redis.asyncio as aioredis | |
| self._redis = aioredis.from_url( | |
| self.settings.database.redis_url, | |
| decode_responses=True, | |
| ) | |
| await self._redis.ping() | |
| logger.info("Using Redis for long-term memory") | |
| return | |
| except Exception as e: | |
| logger.warning(f"Redis unavailable: {e}") | |
| self._redis = None | |
| # Final fallback: JSON file | |
| logger.info("Using JSON file for long-term memory") | |
| async def store( | |
| self, | |
| key: str, | |
| value: Any, | |
| category: str = "general", | |
| metadata: dict[str, Any] | None = None, | |
| ) -> None: | |
| """Store a memory entry.""" | |
| entry = { | |
| "key": key, | |
| "value": value, | |
| "category": category, | |
| "metadata": metadata or {}, | |
| "timestamp": datetime.now(UTC).isoformat(), | |
| } | |
| self._local_store.append(entry) | |
| if self._client: | |
| try: | |
| from qdrant_client.models import PointStruct | |
| embedding = await self._generate_embedding(str(value)) | |
| point = PointStruct( | |
| id=len(self._local_store), | |
| vector=embedding, | |
| payload=entry, | |
| ) | |
| self._client.upsert(collection_name=self._collection, points=[point]) | |
| return | |
| except Exception as e: | |
| logger.error(f"Qdrant store failed: {e}") | |
| if self._redis: | |
| try: | |
| await self._redis.hset( | |
| f"memory:{category}", | |
| key, | |
| json.dumps(entry, default=str), | |
| ) | |
| return | |
| except Exception as e: | |
| logger.error(f"Redis store failed: {e}") | |
| self._save_json() | |
| async def retrieve( | |
| self, | |
| query: str | None = None, | |
| key: str | None = None, | |
| category: str | None = None, | |
| limit: int = 10, | |
| ) -> list[dict[str, Any]]: | |
| """Retrieve memory entries.""" | |
| if key: | |
| for entry in reversed(self._local_store): | |
| if entry.get("key") == key: | |
| return [entry] | |
| return [] | |
| if self._client and query: | |
| try: | |
| embedding = await self._generate_embedding(query) | |
| results = self._client.search( | |
| collection_name=self._collection, | |
| query_vector=embedding, | |
| limit=limit, | |
| ) | |
| return [r.payload for r in results if r.payload] | |
| except Exception as e: | |
| logger.error(f"Qdrant search failed: {e}") | |
| if self._redis and query: | |
| try: | |
| pattern = f"memory:{category or '*'}" | |
| keys = await self._redis.keys(pattern) | |
| results = [] | |
| for redis_key in keys: | |
| data = await self._redis.hgetall(redis_key) | |
| for _, val in data.items(): | |
| entry = json.loads(val) | |
| if query.lower() in str(entry.get("value", "")).lower(): | |
| results.append(entry) | |
| return results[-limit:] | |
| except Exception as e: | |
| logger.error(f"Redis search failed: {e}") | |
| results = self._local_store | |
| if category: | |
| results = [e for e in results if e.get("category") == category] | |
| if query: | |
| results = [e for e in results if query.lower() in str(e.get("value", "")).lower()] | |
| return results[-limit:] | |
| async def delete(self, key: str) -> bool: | |
| """Delete a memory entry.""" | |
| self._local_store = [e for e in self._local_store if e.get("key") != key] | |
| self._save_json() | |
| return True | |
| async def list_all(self, category: str | None = None) -> list[dict[str, Any]]: | |
| """List all memory entries.""" | |
| if category: | |
| return [e for e in self._local_store if e.get("category") == category] | |
| return self._local_store.copy() | |
| async def _generate_embedding(self, text: str) -> list[float]: | |
| """Generate embedding for text using singleton model.""" | |
| model = _get_embedding_model() | |
| if model is not None: | |
| try: | |
| embedding = model.encode(text) | |
| return embedding.tolist() | |
| except Exception as e: | |
| logger.warning(f"Embedding generation failed: {e}") | |
| import hashlib | |
| hash_val = hashlib.md5(text.encode()).hexdigest() | |
| return [float(int(hash_val[i : i + 2], 16)) / 255.0 for i in range(0, 32, 2)] | |
| async def get_stats(self) -> dict[str, Any]: | |
| """Get memory statistics.""" | |
| categories: dict[str, int] = {} | |
| for entry in self._local_store: | |
| cat = entry.get("category", "unknown") | |
| categories[cat] = categories.get(cat, 0) + 1 | |
| return { | |
| "total_entries": len(self._local_store), | |
| "categories": categories, | |
| "qdrant_connected": self._client is not None, | |
| "redis_connected": self._redis is not None, | |
| "backend": "qdrant" if self._client else ("redis" if self._redis else "json"), | |
| } | |