vora-sonnet's picture
Upload folder using huggingface_hub
0d3f7cc verified
Raw
History Blame Contribute Delete
8.55 kB
"""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"),
}