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))