Spaces:
Sleeping
Sleeping
| 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 | |
| def _summary_key(self) -> str: | |
| return _key(self.user_id, self.session_id, "summary") | |
| 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)) | |