diff --git a/config.yaml b/config.yaml index d0217deb..4185640b 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 8b6d1285..588a3a2c 100644 --- a/core/connection.py +++ b/core/connection.py @@ -27,7 +27,7 @@ class TTSException(RuntimeError): 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) @@ -54,7 +54,7 @@ class ConnectionHandler: self.asr = _asr self.llm = _llm self.tts = _tts - self.dialogue = None + self.memory = _memory # vad相关变量 self.client_audio_buffer = bytes() @@ -101,6 +101,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) @@ -163,6 +164,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): """消息路由""" @@ -201,15 +204,12 @@ class ConnectionHandler: return False return not self.is_device_verified + 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)) @@ -217,7 +217,15 @@ class ConnectionHandler: processed_chars = 0 # 跟踪已处理的字符位置 try: start_time = time.time() - llm_responses = self.llm.response(self.session_id, self.dialogue.get_llm_dialogue()) + # 使用带记忆的对话 + 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( + 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..482c0bc2 --- /dev/null +++ b/core/providers/memory/mem0ai/mem0ai.py @@ -0,0 +1,74 @@ +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", "") + 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 + 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=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)}") + 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, + output_format=self.api_version + ) + 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']: + timestamp = entry.get('updated_at', '') + if timestamp: + try: + # Parse and reformat the timestamp + dt = timestamp.split('.')[0] # Remove milliseconds + formatted_time = dt.replace('T', ' ') + except: + formatted_time = timestamp + memory = entry.get('memory', '') + if timestamp and memory: + # 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: + 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)