from langchain.memory import ConversationSummaryBufferMemory, VectorStoreRetrieverMemory from langchain_openai import ChatOpenAI, OpenAIEmbeddings from langchain_chroma import Chroma from app.config import settings from typing import Optional import logging logger = logging.getLogger(__name__) class MemoryManager: """三层记忆管理器 Layer 1 - 工作记忆: 最近 20 轮 + 超出自动摘要 Layer 2 - 长期事实: VectorStoreRetrieverMemory, 跨会话相似召回 Layer 3 - 语义记忆: 每次对话后向量化存储, 供后续会话召回 """ def __init__(self): self.llm = ChatOpenAI( model=settings.llm_model, api_key=settings.llm_api_key, base_url=settings.llm_base_url, ) self.embeddings = OpenAIEmbeddings( model=settings.embedding_model, api_key=settings.effective_embedding_api_key, base_url=settings.effective_embedding_base_url, ) # Layer 3 向量库: 历史对话记忆 self.memory_vectorstore = Chroma( collection_name="user_memory", embedding_function=self.embeddings, persist_directory=settings.chroma_db_path + "_memory", ) def get_working_memory(self) -> ConversationSummaryBufferMemory: """Layer 1: 工作记忆 (当前会话)""" return ConversationSummaryBufferMemory( llm=self.llm, max_token_limit=2000, memory_key="history", return_messages=True, ) def get_longterm_memory(self, user_id: int) -> VectorStoreRetrieverMemory: """Layer 2+3: 长期 + 语义记忆""" return VectorStoreRetrieverMemory( retriever=self.memory_vectorstore.as_retriever( search_kwargs={ "k": 3, "filter": {"user_id": user_id}, } ), memory_key="long_term_memory", input_key="input", ) async def save_conversation(self, user_id: int, conversation_id: str, messages: list[dict]): """会话结束后保存到向量记忆库""" texts = [] for msg in messages: role = msg.get("role", "unknown") content = msg.get("content", "") texts.append(f"[{role}] {content}") full_text = "\n".join(texts) metadata = { "user_id": user_id, "conversation_id": conversation_id, "timestamp": str(__import__("datetime").datetime.now()), } await self.memory_vectorstore.aadd_texts( texts=[full_text], metadatas=[metadata], ) self.memory_vectorstore.persist() logger.info("已保存对话到向量记忆: conv=%s, user=%s", conversation_id, user_id) async def recall(self, user_id: int, query: str, k: int = 3) -> list[str]: """语义召回: 查询与 query 最相似的历史对话片段""" results = self.memory_vectorstore.similarity_search( query, k=k, filter={"user_id": user_id}, ) return [doc.page_content for doc in results]