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 1/4] =?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) From 343ecb3ad466159ee19cedb70e0b8f8ecba6fdac Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8E=84=E5=87=A4=E7=A7=91=E6=8A=80?= Date: Tue, 4 Mar 2025 08:54:27 +0800 Subject: [PATCH 2/4] =?UTF-8?q?=E4=BC=98=E5=8C=96chat=E4=B8=AD=E5=BC=82?= =?UTF-8?q?=E6=AD=A5=E8=B0=83=E7=94=A8=E6=96=B9=E5=BC=8F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- core/connection.py | 19 ++++--------------- 1 file changed, 4 insertions(+), 15 deletions(-) diff --git a/core/connection.py b/core/connection.py index 28be3218..2f3e1f76 100644 --- a/core/connection.py +++ b/core/connection.py @@ -203,23 +203,11 @@ class ConnectionHandler: 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 - loop = asyncio.new_event_loop() - asyncio.set_event_loop(loop) - try: - loop.run_until_complete(self._check_and_broadcast_auth_code()) - finally: - loop.close() + future = asyncio.run_coroutine_threadsafe(self._check_and_broadcast_auth_code(), self.loop) + future.result() return True self.dialogue.put(Message(role="user", content=query)) @@ -228,7 +216,8 @@ class ConnectionHandler: try: start_time = time.time() # 使用带记忆的对话 - memory_str = self.async_run(self.memory.query_memory(query)) + future = asyncio.run_coroutine_threadsafe(self.memory.query_memory(query), self.loop) + memory_str = future.result() self.logger.bind(tag=TAG).info(f"记忆内容: {memory_str}") llm_responses = self.llm.response( From 83911d9ad1500b48064341fcab16fb9d5b15e0a5 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8E=84=E5=87=A4=E7=A7=91=E6=8A=80?= Date: Tue, 4 Mar 2025 10:28:38 +0800 Subject: [PATCH 3/4] =?UTF-8?q?=E8=AE=B0=E5=BF=86=E5=A2=9E=E5=8A=A0?= =?UTF-8?q?=E6=97=B6=E9=97=B4=EF=BC=8C=E4=BB=A5=E4=BE=BF=E5=A4=A7=E6=A8=A1?= =?UTF-8?q?=E5=9E=8B=E5=88=A4=E6=96=AD=E5=85=88=E5=90=8E=E5=85=B3=E7=B3=BB?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- core/providers/memory/mem0ai/mem0ai.py | 27 +++++++++++++++++++++++--- 1 file changed, 24 insertions(+), 3 deletions(-) diff --git a/core/providers/memory/mem0ai/mem0ai.py b/core/providers/memory/mem0ai/mem0ai.py index de4eb29e..e319b542 100644 --- a/core/providers/memory/mem0ai/mem0ai.py +++ b/core/providers/memory/mem0ai/mem0ai.py @@ -7,6 +7,7 @@ class MemoryProvider(MemoryProviderBase): def __init__(self, config): super().__init__(config) self.api_key = config.get("api_key", "") + self.api_version = config.get("api_version", "v1.1") if len(self.api_key) == 0 or "你" in self.api_key: logger.bind(tag=TAG).error("你还没配置Mem0ai的密钥,请在配置文件中配置密钥,否则无法提供记忆服务") self.use_mem0 = False @@ -27,7 +28,7 @@ class MemoryProvider(MemoryProviderBase): {"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") + result = self.client.add(messages, user_id=self.role_id, output_format=self.api_version) logger.bind(tag=TAG).debug(f"Save memory result: {result}") except Exception as e: logger.bind(tag=TAG).error(f"保存记忆失败: {str(e)}") @@ -39,9 +40,29 @@ class MemoryProvider(MemoryProviderBase): try: results = self.client.search( query, - user_id= self.role_id, + user_id=self.role_id, + output_format=self.api_version ) - memories_str = "\n".join(f"- {entry['memory']}" for entry in results) + if not results or 'results' not in results: + return "" + + # Format each memory entry with its update time up to minutes + memories = [] + for entry in results['results']: + # Split timestamp and get date + time up to minutes + timestamp = entry.get('updated_at', '').split('.')[0] # Remove milliseconds + if timestamp: + try: + # Parse and reformat the timestamp + dt = timestamp.replace('T', ' ').split(':') # Split time components + formatted_time = f"{dt[0]}:{dt[1]}:{dt[2]}" # Keep only HH:MM:SS + except: + formatted_time = timestamp + memory = entry.get('memory', '') + if timestamp and memory: + memories.append(f"[{formatted_time}] {memory}") + + memories_str = "\n".join(f"- {memory}" for memory in memories) logger.bind(tag=TAG).debug(f"Query results: {memories_str}") return memories_str except Exception as e: From 93a94dea4502e3748fae55c71d4b2331d143d433 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8E=84=E5=87=A4=E7=A7=91=E6=8A=80?= Date: Tue, 4 Mar 2025 13:58:14 +0800 Subject: [PATCH 4/4] =?UTF-8?q?=E8=AE=B0=E5=BF=86=E5=AF=B9=E6=97=B6?= =?UTF-8?q?=E9=97=B4=E6=88=B3=E6=8E=92=E5=BA=8F=EF=BC=8C=E4=BE=BF=E4=BA=8E?= =?UTF-8?q?=E6=A2=B3=E7=90=86=E5=89=8D=E5=90=8E=E5=85=B3=E7=B3=BB?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- core/providers/memory/mem0ai/mem0ai.py | 18 +++++++++++------- 1 file changed, 11 insertions(+), 7 deletions(-) diff --git a/core/providers/memory/mem0ai/mem0ai.py b/core/providers/memory/mem0ai/mem0ai.py index e319b542..482c0bc2 100644 --- a/core/providers/memory/mem0ai/mem0ai.py +++ b/core/providers/memory/mem0ai/mem0ai.py @@ -49,20 +49,24 @@ class MemoryProvider(MemoryProviderBase): # Format each memory entry with its update time up to minutes memories = [] for entry in results['results']: - # Split timestamp and get date + time up to minutes - timestamp = entry.get('updated_at', '').split('.')[0] # Remove milliseconds + timestamp = entry.get('updated_at', '') if timestamp: try: # Parse and reformat the timestamp - dt = timestamp.replace('T', ' ').split(':') # Split time components - formatted_time = f"{dt[0]}:{dt[1]}:{dt[2]}" # Keep only HH:MM:SS + dt = timestamp.split('.')[0] # Remove milliseconds + formatted_time = dt.replace('T', ' ') except: formatted_time = timestamp memory = entry.get('memory', '') if timestamp and memory: - memories.append(f"[{formatted_time}] {memory}") - - memories_str = "\n".join(f"- {memory}" for memory in memories) + # Store tuple of (timestamp, formatted_string) for sorting + memories.append((timestamp, f"[{formatted_time}] {memory}")) + + # Sort by timestamp in descending order (newest first) + memories.sort(key=lambda x: x[0], reverse=True) + + # Extract only the formatted strings + memories_str = "\n".join(f"- {memory[1]}" for memory in memories) logger.bind(tag=TAG).debug(f"Query results: {memories_str}") return memories_str except Exception as e: