From f2e68060deec8fcbd41e01206003309e71d24aee Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8E=84=E5=87=A4=E7=A7=91=E6=8A=80?= Date: Mon, 3 Mar 2025 15:00:04 +0800 Subject: [PATCH] =?UTF-8?q?=E4=BD=BF=E7=94=A8mem0ai=20api=E5=AE=9E?= =?UTF-8?q?=E7=8E=B0=E8=AE=B0=E5=BF=86=E5=8A=9F=E8=83=BD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- config.yaml | 7 ++++ core/connection.py | 25 +++++++++++-- core/providers/memory/base.py | 23 ++++++++++++ core/providers/memory/mem0ai/mem0ai.py | 49 ++++++++++++++++++++++++++ core/utils/dialogue.py | 24 +++++++++++++ core/utils/memory.py | 17 +++++++++ core/websocket_server.py | 13 ++++--- 7 files changed, 151 insertions(+), 7 deletions(-) create mode 100644 core/providers/memory/base.py create mode 100644 core/providers/memory/mem0ai/mem0ai.py create mode 100644 core/utils/memory.py diff --git a/config.yaml b/config.yaml index ce065b17..157a23e8 100644 --- a/config.yaml +++ b/config.yaml @@ -80,7 +80,14 @@ selected_module: LLM: ChatGLMLLM # TTS将根据配置名称对应的type调用实际的TTS适配器 TTS: EdgeTTS + Memory: mem0ai +Memory: + mem0ai: + type: mem0ai + # https://app.mem0.ai/dashboard/api-keys,注册可获得试用credits + api_key: 你的mem0ai api key + ASR: FunASR: type: fun_local diff --git a/core/connection.py b/core/connection.py index 5532e3a1..28be3218 100644 --- a/core/connection.py +++ b/core/connection.py @@ -23,7 +23,7 @@ TAG = __name__ class ConnectionHandler: - def __init__(self, config: Dict[str, Any], _vad, _asr, _llm, _tts, _music): + def __init__(self, config: Dict[str, Any], _vad, _asr, _llm, _tts, _music, _memory): self.config = config self.logger = setup_logging() self.auth = AuthMiddleware(config) @@ -50,7 +50,7 @@ class ConnectionHandler: self.asr = _asr self.llm = _llm self.tts = _tts - self.dialogue = None + self.memory = _memory # vad相关变量 self.client_audio_buffer = bytes() @@ -99,6 +99,7 @@ class ConnectionHandler: await self.auth.authenticate(self.headers) device_id = self.headers.get("device-id", None) + self.memory.set_role_id(device_id) # Load private configuration if device_id is provided bUsePrivateConfig = self.config.get("use_private_config", False) @@ -161,6 +162,8 @@ class ConnectionHandler: self.logger.bind(tag=TAG).error(f"Connection error: {str(e)}-{stack_trace}") await ws.close() return + finally: + await self.memory.save_memory(self.dialogue.dialogue) async def _route_message(self, message): """消息路由""" @@ -199,6 +202,15 @@ class ConnectionHandler: return False return not self.is_device_verified + + def async_run(self, coro): + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + try: + return loop.run_until_complete(coro) + finally: + loop.close() + def chat(self, query): if self.isNeedAuth(): self.llm_finish_task = True @@ -215,7 +227,14 @@ class ConnectionHandler: processed_chars = 0 # 跟踪已处理的字符位置 try: start_time = time.time() - llm_responses = self.llm.response(self.session_id, self.dialogue.get_llm_dialogue()) + # 使用带记忆的对话 + memory_str = self.async_run(self.memory.query_memory(query)) + + self.logger.bind(tag=TAG).info(f"记忆内容: {memory_str}") + llm_responses = self.llm.response( + self.session_id, + self.dialogue.get_llm_dialogue_with_memory(memory_str) + ) except Exception as e: self.logger.bind(tag=TAG).error(f"LLM 处理出错 {query}: {e}") return None diff --git a/core/providers/memory/base.py b/core/providers/memory/base.py new file mode 100644 index 00000000..00ab6650 --- /dev/null +++ b/core/providers/memory/base.py @@ -0,0 +1,23 @@ +from abc import ABC, abstractmethod +from config.logger import setup_logging + +TAG = __name__ +logger = setup_logging() + +class MemoryProviderBase(ABC): + def __init__(self, config): + self.config = config + self.role_id = None + + @abstractmethod + async def save_memory(self, msgs): + """Save a new memory for specific role and return memory ID""" + print("this is base func", msgs) + + @abstractmethod + async def query_memory(self, query: str) -> str: + """Query memories for specific role based on similarity""" + return "please implement query method" + + def set_role_id(self, role_id: str): + self.role_id = role_id \ No newline at end of file diff --git a/core/providers/memory/mem0ai/mem0ai.py b/core/providers/memory/mem0ai/mem0ai.py new file mode 100644 index 00000000..de4eb29e --- /dev/null +++ b/core/providers/memory/mem0ai/mem0ai.py @@ -0,0 +1,49 @@ +from ..base import MemoryProviderBase, logger +from mem0 import MemoryClient + +TAG = __name__ + +class MemoryProvider(MemoryProviderBase): + def __init__(self, config): + super().__init__(config) + self.api_key = config.get("api_key", "") + if len(self.api_key) == 0 or "你" in self.api_key: + logger.bind(tag=TAG).error("你还没配置Mem0ai的密钥,请在配置文件中配置密钥,否则无法提供记忆服务") + self.use_mem0 = False + return + else: + self.use_mem0 = True + self.client = MemoryClient(api_key=self.api_key) + + async def save_memory(self, msgs): + if not self.use_mem0: + return None + if len(msgs) < 2: + return None + + try: + # Format the content as a message list for mem0 + messages = [ + {"role": message.role, "content": message.content} + for message in msgs if message.role != "system" + ] + result = self.client.add(messages, user_id=self.role_id, output_format="v1.1") + logger.bind(tag=TAG).debug(f"Save memory result: {result}") + except Exception as e: + logger.bind(tag=TAG).error(f"保存记忆失败: {str(e)}") + return None + + async def query_memory(self, query: str)-> str: + if not self.use_mem0: + return "" + try: + results = self.client.search( + query, + user_id= self.role_id, + ) + memories_str = "\n".join(f"- {entry['memory']}" for entry in results) + logger.bind(tag=TAG).debug(f"Query results: {memories_str}") + return memories_str + except Exception as e: + logger.bind(tag=TAG).error(f"查询记忆失败: {str(e)}") + return "" \ No newline at end of file diff --git a/core/utils/dialogue.py b/core/utils/dialogue.py index 703f12c0..ed5e27fe 100644 --- a/core/utils/dialogue.py +++ b/core/utils/dialogue.py @@ -24,3 +24,27 @@ class Dialogue: for m in self.dialogue: dialogue.append({"role": m.role, "content": m.content}) return dialogue + + def get_llm_dialogue_with_memory(self, memory_str: str = None) -> List[Dict[str, str]]: + # 构建带记忆的对话 + dialogue = [] + + # 添加系统提示和记忆 + system_message = next( + (msg for msg in self.dialogue if msg.role == "system"), None + ) + + + if system_message: + enhanced_system_prompt = ( + f"{system_message.content}\n\n" + f"相关记忆:\n{memory_str}" + ) + dialogue.append({"role": "system", "content": enhanced_system_prompt}) + + # 添加用户和助手的对话 + for msg in self.dialogue: + if msg.role != "system": # 跳过原始的系统消息 + dialogue.append({"role": msg.role, "content": msg.content}) + + return dialogue diff --git a/core/utils/memory.py b/core/utils/memory.py new file mode 100644 index 00000000..c750e9aa --- /dev/null +++ b/core/utils/memory.py @@ -0,0 +1,17 @@ +import os +import sys +import importlib +from config.logger import setup_logging +from core.utils.util import read_config, get_project_dir + +logger = setup_logging() + +def create_instance(class_name, *args, **kwargs): + if os.path.exists(os.path.join('core', 'providers', 'memory', class_name, f'{class_name}.py')): + lib_name = f'core.providers.memory.{class_name}.{class_name}' + if lib_name not in sys.modules: + sys.modules[lib_name] = importlib.import_module(f'{lib_name}') + return sys.modules[lib_name].MemoryProvider(*args, **kwargs) + + raise ValueError(f"不支持的记忆服务类型: {class_name}") + diff --git a/core/websocket_server.py b/core/websocket_server.py index 907ef22f..41629932 100644 --- a/core/websocket_server.py +++ b/core/websocket_server.py @@ -4,7 +4,7 @@ from config.logger import setup_logging from core.connection import ConnectionHandler from core.handle.musicHandler import MusicHandler from core.utils.util import get_local_ip -from core.utils import asr, vad, llm, tts +from core.utils import asr, vad, llm, tts, memory TAG = __name__ @@ -13,10 +13,14 @@ class WebSocketServer: def __init__(self, config: dict): self.config = config self.logger = setup_logging() - self._vad, self._asr, self._llm, self._tts, self._music = self._create_processing_instances() + self._vad, self._asr, self._llm, self._tts, self._music, self._memory = self._create_processing_instances() self.active_connections = set() # 添加全局连接记录 def _create_processing_instances(self): + memory_cls_name = self.config["selected_module"].get("Memory", "mem0ai") # 默认使用mem0ai + has_memory_cfg = self.config.get("Memory") and memory_cls_name in self.config["Memory"] + memory_cfg = self.config["Memory"][memory_cls_name] if has_memory_cfg else {} + """创建处理模块实例""" return ( vad.create_instance( @@ -46,7 +50,8 @@ class WebSocketServer: self.config["TTS"][self.config["selected_module"]["TTS"]], self.config["delete_audio"] ), - MusicHandler(self.config) + MusicHandler(self.config), + memory.create_instance(memory_cls_name, memory_cfg), ) async def start(self): @@ -66,7 +71,7 @@ class WebSocketServer: async def _handle_connection(self, websocket): """处理新连接,每次创建独立的ConnectionHandler""" # 创建ConnectionHandler时传入当前server实例 - handler = ConnectionHandler(self.config, self._vad, self._asr, self._llm, self._tts, self._music) + handler = ConnectionHandler(self.config, self._vad, self._asr, self._llm, self._tts, self._music, self._memory) self.active_connections.add(handler) try: await handler.handle_connection(websocket)