使用mem0ai api实现记忆功能

This commit is contained in:
玄凤科技
2025-03-03 15:00:04 +08:00
parent effd79b465
commit f2e68060de
7 changed files with 151 additions and 7 deletions
+7
View File
@@ -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
+22 -3
View File
@@ -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
+23
View File
@@ -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
+49
View File
@@ -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 ""
+24
View File
@@ -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
+17
View File
@@ -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}")
+9 -4
View File
@@ -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)