mem0.py 4.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104
  1. from __future__ import annotations
  2. from typing import Any
  3. from ..config import settings
  4. from ..llm import LLMUnavailable
  5. from ..schemas import ChatResponse, SystemDescriptor
  6. from .base import AdapterUnavailable, MemoryAgent
  7. class Mem0Agent(MemoryAgent):
  8. id = "mem0"
  9. def __init__(self) -> None:
  10. super().__init__()
  11. self._memory: Any | None = None
  12. self._import_error: str | None = None
  13. try:
  14. from mem0 import Memory # type: ignore
  15. self._memory_class = Memory
  16. except Exception as exc:
  17. self._memory_class = None
  18. self._import_error = str(exc)
  19. embedding_ready = bool(
  20. settings.embedding_api_key
  21. and settings.embedding_base_url
  22. and settings.embedding_model
  23. )
  24. ready = bool(self._memory_class and settings.mem0_enabled and embedding_ready)
  25. self.descriptor = SystemDescriptor(
  26. id="mem0",
  27. name="Mem0",
  28. paradigm="自动提取与冲突消解",
  29. description="将对话提炼为长期记忆,并按逻辑用户标识执行检索和更新。",
  30. available=ready,
  31. mode="real-sdk" if ready else "unavailable",
  32. status="ready" if ready else "not-configured",
  33. package="mem0ai",
  34. setup_hint=None if ready else (
  35. "当前仅保留 Mem0 适配骨架,请保持 MEM0_ENABLED=false。"
  36. "完成公共记忆投影、删除、重置、冲突和隔离闭环后再启用。"
  37. ),
  38. )
  39. async def _client(self) -> Any:
  40. if not self._memory_class or not settings.mem0_enabled:
  41. raise AdapterUnavailable(self.descriptor.setup_hint or "Mem0 不可用")
  42. if self._memory is None:
  43. if not settings.embedding_api_key or not settings.embedding_model:
  44. raise AdapterUnavailable(self.descriptor.setup_hint or "Mem0 缺少 Embedding 配置")
  45. llm_provider = "deepseek" if "deepseek" in settings.llm_base_url.lower() else "openai"
  46. llm_config = {"model": settings.llm_model, "api_key": settings.llm_api_key}
  47. llm_config["deepseek_base_url" if llm_provider == "deepseek" else "openai_base_url"] = settings.llm_base_url
  48. self._memory = self._memory_class.from_config({
  49. "llm": {"provider": llm_provider, "config": llm_config},
  50. "embedder": {
  51. "provider": "openai",
  52. "config": {
  53. "model": settings.embedding_model,
  54. "api_key": settings.embedding_api_key,
  55. "openai_base_url": settings.embedding_base_url,
  56. "embedding_dims": settings.embedding_dimensions,
  57. },
  58. },
  59. "vector_store": {
  60. "provider": "pgvector",
  61. "config": {
  62. "connection_string": settings.database_url,
  63. "collection_name": "mem0_memory",
  64. "embedding_model_dims": settings.embedding_dimensions,
  65. },
  66. },
  67. })
  68. return self._memory
  69. async def chat(self, message: str) -> ChatResponse:
  70. memory = await self._client()
  71. add_result = memory.add(
  72. [{"role": "user", "content": message}],
  73. user_id=settings.workspace_id,
  74. )
  75. if hasattr(add_result, "__await__"):
  76. add_result = await add_result
  77. search_result = memory.search(message, user_id=settings.workspace_id, limit=5)
  78. if hasattr(search_result, "__await__"):
  79. search_result = await search_result
  80. context = search_result.get("results", search_result if isinstance(search_result, list) else [])
  81. context_text = "\n".join(f"- {item.get('memory', item)}" for item in context)
  82. try:
  83. answer = await self.llm.chat(
  84. "你是使用 Mem0 记忆的 AI 编程助手。只把 Mem0 返回的内容当作记忆证据。",
  85. f"Mem0 召回:\n{context_text or '(暂无)'}\n\n用户:{message}",
  86. )
  87. mode = "real-sdk+llm"
  88. except LLMUnavailable:
  89. answer = f"Mem0 已完成写入和检索。\n\n召回内容:\n{context_text or '(暂无)'}"
  90. mode = "real-sdk"
  91. await self._audit("MEM0/add+search", details={"raw_add": str(add_result)[:1000]})
  92. return ChatResponse(
  93. system="mem0", answer=answer, mode=mode,
  94. memory_context=await self.memories(), memory_events=[{"event": "add", "result": add_result}],
  95. audit_events=await self.audit(),
  96. )