store.py 3.1 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889
  1. from langchain.memory import ConversationSummaryBufferMemory, VectorStoreRetrieverMemory
  2. from langchain_openai import ChatOpenAI, OpenAIEmbeddings
  3. from langchain_chroma import Chroma
  4. from app.config import settings
  5. from typing import Optional
  6. import logging
  7. logger = logging.getLogger(__name__)
  8. class MemoryManager:
  9. """三层记忆管理器
  10. Layer 1 - 工作记忆: 最近 20 轮 + 超出自动摘要
  11. Layer 2 - 长期事实: VectorStoreRetrieverMemory, 跨会话相似召回
  12. Layer 3 - 语义记忆: 每次对话后向量化存储, 供后续会话召回
  13. """
  14. def __init__(self):
  15. self.llm = ChatOpenAI(
  16. model=settings.llm_model,
  17. api_key=settings.llm_api_key,
  18. base_url=settings.llm_base_url,
  19. )
  20. self.embeddings = OpenAIEmbeddings(
  21. model=settings.embedding_model,
  22. api_key=settings.effective_embedding_api_key,
  23. base_url=settings.effective_embedding_base_url,
  24. )
  25. # Layer 3 向量库: 历史对话记忆
  26. self.memory_vectorstore = Chroma(
  27. collection_name="user_memory",
  28. embedding_function=self.embeddings,
  29. persist_directory=settings.chroma_db_path + "_memory",
  30. )
  31. def get_working_memory(self) -> ConversationSummaryBufferMemory:
  32. """Layer 1: 工作记忆 (当前会话)"""
  33. return ConversationSummaryBufferMemory(
  34. llm=self.llm,
  35. max_token_limit=2000,
  36. memory_key="history",
  37. return_messages=True,
  38. )
  39. def get_longterm_memory(self, user_id: int) -> VectorStoreRetrieverMemory:
  40. """Layer 2+3: 长期 + 语义记忆"""
  41. return VectorStoreRetrieverMemory(
  42. retriever=self.memory_vectorstore.as_retriever(
  43. search_kwargs={
  44. "k": 3,
  45. "filter": {"user_id": user_id},
  46. }
  47. ),
  48. memory_key="long_term_memory",
  49. input_key="input",
  50. )
  51. async def save_conversation(self, user_id: int, conversation_id: str,
  52. messages: list[dict]):
  53. """会话结束后保存到向量记忆库"""
  54. texts = []
  55. for msg in messages:
  56. role = msg.get("role", "unknown")
  57. content = msg.get("content", "")
  58. texts.append(f"[{role}] {content}")
  59. full_text = "\n".join(texts)
  60. metadata = {
  61. "user_id": user_id,
  62. "conversation_id": conversation_id,
  63. "timestamp": str(__import__("datetime").datetime.now()),
  64. }
  65. await self.memory_vectorstore.aadd_texts(
  66. texts=[full_text],
  67. metadatas=[metadata],
  68. )
  69. self.memory_vectorstore.persist()
  70. logger.info("已保存对话到向量记忆: conv=%s, user=%s", conversation_id, user_id)
  71. async def recall(self, user_id: int, query: str, k: int = 3) -> list[str]:
  72. """语义召回: 查询与 query 最相似的历史对话片段"""
  73. results = self.memory_vectorstore.similarity_search(
  74. query,
  75. k=k,
  76. filter={"user_id": user_id},
  77. )
  78. return [doc.page_content for doc in results]