rag_project_backend / memory_store.py
anhkhoiphan's picture
deploy: 2026-06-05 11:40:44
4d85858
Raw
History Blame Contribute Delete
6.98 kB
import json
import os
from datetime import datetime, timezone
from typing import Any
from langchain_core.documents import Document
from src.llm import non_reasoning_llm
_SUMMARIZE_PROMPT = (
"Tóm tắt ngắn gọn (3-5 câu) nội dung cuộc hội thoại sau, "
"giữ lại các thông tin quan trọng:\n\n{history}"
)
_DEFAULT_BUFFER_SIZE = 4
_DEFAULT_TTL_SECONDS = 7200
_redis_client = None
def _memory_ttl_seconds() -> int:
return int(os.environ.get("MEMORY_TTL_SECONDS", _DEFAULT_TTL_SECONDS))
def _memory_buffer_size() -> int:
return int(os.environ.get("MEMORY_BUFFER_SIZE", _DEFAULT_BUFFER_SIZE))
def _get_redis():
global _redis_client
if _redis_client is None:
from upstash_redis import Redis
_redis_client = Redis(
url=os.environ["UPSTASH_REDIS_REST_URL"],
token=os.environ["UPSTASH_REDIS_REST_TOKEN"],
)
return _redis_client
def _session_id(session_id: str | None) -> str:
return session_id or "default"
def _key(user_id: str, session_id: str | None, name: str) -> str:
return f"memory:{user_id}:{_session_id(session_id)}:{name}"
def _utc_now() -> str:
return datetime.now(timezone.utc).isoformat()
def _decode_json(value: Any, default: Any) -> Any:
if value is None:
return default
if isinstance(value, bytes):
value = value.decode("utf-8")
if isinstance(value, str):
return json.loads(value)
return value
def _load_json(key: str, default: Any) -> Any:
return _decode_json(_get_redis().get(key), default)
def _save_json(key: str, value: Any) -> None:
_get_redis().set(key, json.dumps(value, ensure_ascii=False))
_get_redis().expire(key, _memory_ttl_seconds())
def _touch_memory_keys(user_id: str, session_id: str | None) -> None:
redis = _get_redis()
ttl = _memory_ttl_seconds()
for name in ("summary", "buffer", "prev_docs"):
redis.expire(_key(user_id, session_id, name), ttl)
def _compact_doc(doc: Any) -> dict[str, Any]:
metadata = getattr(doc, "metadata", {}) or {}
page_content = getattr(doc, "page_content", None)
content = metadata.get("content") or page_content or ""
return {
"doc_id": metadata.get("doc_id") or metadata.get("id"),
"title": metadata.get("title"),
"source": metadata.get("source"),
"source_url": metadata.get("source_url"),
"chunk_id": metadata.get("chunk_id"),
"content": content,
"score": metadata.get("score"),
}
def _doc_id(doc: Any) -> str | None:
metadata = getattr(doc, "metadata", {}) or {}
return metadata.get("doc_id") or metadata.get("id") or metadata.get("chunk_id")
def _doc_from_compact(value: dict[str, Any]) -> Document:
metadata = {
key: value.get(key)
for key in ("doc_id", "title", "source", "source_url", "chunk_id", "score")
if value.get(key) is not None
}
return Document(page_content=value.get("content") or "", metadata=metadata)
class _RedisSummaryBufferMemory:
"""Session-scoped ConversationSummaryBufferMemory replacement backed by Redis."""
def __init__(self, user_id: str, session_id: str | None = None) -> None:
self.user_id = user_id
self.session_id = session_id
@property
def _summary_key(self) -> str:
return _key(self.user_id, self.session_id, "summary")
@property
def _buffer_key(self) -> str:
return _key(self.user_id, self.session_id, "buffer")
def _load_summary(self) -> str:
value = _load_json(self._summary_key, {})
return value.get("summary", "") if isinstance(value, dict) else ""
def _save_summary(self, summary: str) -> None:
_save_json(self._summary_key, {"summary": summary, "updated_at": _utc_now()})
def _ensure_summary_key(self) -> None:
if _get_redis().get(self._summary_key) is None:
self._save_summary("")
def _load_buffer(self) -> list[dict[str, Any]]:
value = _load_json(self._buffer_key, [])
return value if isinstance(value, list) else []
def _save_buffer(self, buffer: list[dict[str, Any]]) -> None:
_save_json(self._buffer_key, buffer)
def save_context(self, inputs: dict, outputs: dict) -> None:
buffer = self._load_buffer()
buffer.append({
"query": inputs["input"],
"answer": outputs["output"],
"doc_ids": inputs.get("doc_ids", []),
"created_at": _utc_now(),
})
if len(buffer) > _memory_buffer_size():
buffer = self._flush(buffer)
self._ensure_summary_key()
self._save_buffer(buffer)
_touch_memory_keys(self.user_id, self.session_id)
def _flush(self, buffer: list[dict[str, Any]]) -> list[dict[str, Any]]:
to_summarise = buffer[:-2]
remaining = buffer[-2:]
history_text = "\n".join(
f"Người dùng: {turn.get('query', '')}\nAI: {turn.get('answer', '')}"
for turn in to_summarise
)
existing_summary = self._load_summary()
existing = (
f"Tóm tắt trước: {existing_summary}\n\n"
if existing_summary
else ""
)
response = non_reasoning_llm.invoke(
_SUMMARIZE_PROMPT.format(history=existing + history_text)
)
self._save_summary(response.content)
return remaining
def load_memory_variables(self, _: dict) -> dict:
parts = []
summary = self._load_summary()
buffer = self._load_buffer()
if summary:
parts.append(f"[Tóm tắt]: {summary}")
for turn in buffer:
parts.append(
f"Người dùng: {turn.get('query', '')}\n"
f"AI: {turn.get('answer', '')}"
)
return {"history": "\n\n".join(parts)}
def get_memory(user_id: str, session_id: str | None = None) -> _RedisSummaryBufferMemory:
return _RedisSummaryBufferMemory(user_id=user_id, session_id=session_id)
def get_prev_docs(user_id: str, session_id: str | None = None) -> list[Document]:
docs = _load_json(_key(user_id, session_id, "prev_docs"), [])
if not isinstance(docs, list):
return []
return [_doc_from_compact(doc) for doc in docs if isinstance(doc, dict)]
def save_turn(
user_id: str,
query: str,
answer: str,
docs: list[Document],
session_id: str | None = None,
) -> None:
get_memory(user_id, session_id=session_id).save_context(
{"input": query, "doc_ids": [_doc_id(doc) for doc in docs if _doc_id(doc)]},
{"output": answer},
)
_save_json(_key(user_id, session_id, "prev_docs"), [_compact_doc(doc) for doc in docs])
_touch_memory_keys(user_id, session_id)
def clear_memory(user_id: str, session_id: str | None = None) -> None:
redis = _get_redis()
for name in ("summary", "buffer", "prev_docs"):
redis.delete(_key(user_id, session_id, name))