Enterprise-AI-Copilot / services /memory_service.py
Punit1's picture
Initial commit
939c0c0
Raw
History Blame Contribute Delete
3.04 kB
"""
Long-Term User Memory Service
==============================
Stores user preferences, facts, and conversation context across sessions in Qdrant `user_memory` collection.
"""
from __future__ import annotations
import os
import uuid
import structlog
from typing import List, Dict, Any, Optional
from qdrant_client import models as qmodels
from services.embedding import embedding_service
from services.qdrant_service import qdrant_service, MEMORY_COLLECTION
logger = structlog.get_logger(__name__)
class MemoryService:
"""Manages long-term user memory and preference retrieval."""
async def store_user_fact(
self,
tenant_id: str,
user_id: str,
fact_text: str,
category: str = "preference",
) -> str:
"""Embed and save a user fact/preference into long-term memory."""
vector = await embedding_service.encode_query(fact_text)
client = qdrant_service.get_client(os.environ.get("QDRANT_URL"))
point_id = str(uuid.uuid4())
client.upsert(
collection_name=MEMORY_COLLECTION,
points=[
qmodels.PointStruct(
id=point_id,
vector=vector.tolist(),
payload={
"tenant_id": tenant_id,
"user_id": user_id,
"fact": fact_text,
"category": category,
},
)
],
)
logger.info("User fact stored in long-term memory", user_id=user_id, fact=fact_text)
return point_id
async def recall_user_memories(
self,
tenant_id: str,
user_id: str,
query: str,
top_k: int = 3,
) -> List[str]:
"""Recall relevant user memories for a query using vector search."""
try:
vector = await embedding_service.encode_query(query)
client = qdrant_service.get_client(os.environ.get("QDRANT_URL"))
search_filter = qmodels.Filter(
must=[
qmodels.FieldCondition(key="tenant_id", match=qmodels.MatchValue(value=tenant_id)),
qmodels.FieldCondition(key="user_id", match=qmodels.MatchValue(value=user_id)),
]
)
hits = client.search(
collection_name=MEMORY_COLLECTION,
query_vector=vector.tolist(),
query_filter=search_filter,
limit=top_k,
)
memories = [hit.payload.get("fact", "") for hit in hits if hit.payload]
logger.info("Recalled long-term memories", count=len(memories), user_id=user_id)
return memories
except Exception as e:
logger.warning("Failed to recall memories", error=str(e))
return []
# ── Singleton instance ─────────────────────────────────────────
memory_service = MemoryService()