mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-28 01:53:53 +08:00
Merge pull request #3101 from xinnan-tech/fix-title-summary
fix:修复PowerMem和Mem0AI记忆模型不会总结会话标题的问题
This commit is contained in:
@@ -2,6 +2,7 @@ import json
|
|||||||
import traceback
|
import traceback
|
||||||
|
|
||||||
from ..base import MemoryProviderBase, logger
|
from ..base import MemoryProviderBase, logger
|
||||||
|
from config.manage_api_client import generate_and_save_chat_summary
|
||||||
from mem0 import MemoryClient
|
from mem0 import MemoryClient
|
||||||
from core.utils.util import check_model_key
|
from core.utils.util import check_model_key
|
||||||
|
|
||||||
@@ -30,38 +31,40 @@ class MemoryProvider(MemoryProviderBase):
|
|||||||
self.use_mem0 = False
|
self.use_mem0 = False
|
||||||
|
|
||||||
async def save_memory(self, msgs, session_id=None):
|
async def save_memory(self, msgs, session_id=None):
|
||||||
if not self.use_mem0:
|
|
||||||
return None
|
|
||||||
if len(msgs) < 2:
|
|
||||||
return None
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Format the content as a message list for mem0
|
if self.use_mem0 and len(msgs) >= 2:
|
||||||
messages = []
|
# Format the content as a message list for mem0
|
||||||
for message in msgs:
|
messages = []
|
||||||
if message.role == "system":
|
for message in msgs:
|
||||||
continue
|
if message.role == "system":
|
||||||
|
continue
|
||||||
|
|
||||||
content = message.content
|
content = message.content
|
||||||
|
|
||||||
# Extract content from JSON format if present (for ASR with emotion/language tags)
|
# Extract content from JSON format if present (for ASR with emotion/language tags)
|
||||||
# Same logic as in query_memory method
|
# Same logic as in query_memory method
|
||||||
try:
|
try:
|
||||||
if content and content.strip().startswith("{") and content.strip().endswith("}"):
|
if content and content.strip().startswith("{") and content.strip().endswith("}"):
|
||||||
data = json.loads(content)
|
data = json.loads(content)
|
||||||
if "content" in data:
|
if "content" in data:
|
||||||
content = data["content"]
|
content = data["content"]
|
||||||
except (json.JSONDecodeError, KeyError, TypeError):
|
except (json.JSONDecodeError, KeyError, TypeError):
|
||||||
# If parsing fails, use original content
|
# If parsing fails, use original content
|
||||||
pass
|
pass
|
||||||
|
|
||||||
messages.append({"role": message.role, "content": content})
|
messages.append({"role": message.role, "content": content})
|
||||||
|
|
||||||
result = self.client.add(messages, user_id=self.role_id)
|
result = self.client.add(messages, user_id=self.role_id)
|
||||||
logger.bind(tag=TAG).debug(f"Save memory result: {result}")
|
logger.bind(tag=TAG).debug(f"Save memory result: {result}")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"保存记忆失败: {str(e)}")
|
logger.bind(tag=TAG).error(f"保存记忆失败: {str(e)}")
|
||||||
return None
|
|
||||||
|
# Generate and save chat summary (SLM summarizes session title)
|
||||||
|
# This is independent of Mem0AI's memory saving
|
||||||
|
if session_id:
|
||||||
|
await generate_and_save_chat_summary(session_id)
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
async def query_memory(self, query: str) -> str:
|
async def query_memory(self, query: str) -> str:
|
||||||
if not self.use_mem0:
|
if not self.use_mem0:
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ import traceback
|
|||||||
from typing import Optional, Dict, Any
|
from typing import Optional, Dict, Any
|
||||||
|
|
||||||
from ..base import MemoryProviderBase, logger
|
from ..base import MemoryProviderBase, logger
|
||||||
|
from config.manage_api_client import generate_and_save_chat_summary
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
|
|
||||||
@@ -162,59 +163,60 @@ class MemoryProvider(MemoryProviderBase):
|
|||||||
Returns:
|
Returns:
|
||||||
Result from PowerMem API or None if failed
|
Result from PowerMem API or None if failed
|
||||||
"""
|
"""
|
||||||
if not self.use_powermem or self.memory_client is None:
|
|
||||||
logger.bind(tag=TAG).warning("PowerMem is not available, skipping save_memory")
|
|
||||||
return None
|
|
||||||
|
|
||||||
if len(msgs) < 2:
|
|
||||||
logger.bind(tag=TAG).debug("Not enough messages to save (need at least 2)")
|
|
||||||
return None
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Format the content as a message list for PowerMem
|
if self.use_powermem and self.memory_client is not None and len(msgs) >= 2:
|
||||||
messages = []
|
# Format the content as a message list for PowerMem
|
||||||
for message in msgs:
|
messages = []
|
||||||
if message.role == "system":
|
for message in msgs:
|
||||||
continue
|
if message.role == "system":
|
||||||
|
continue
|
||||||
|
|
||||||
content = message.content
|
content = message.content
|
||||||
|
|
||||||
# Extract content from JSON format if present (for ASR with emotion/language tags)
|
# Extract content from JSON format if present (for ASR with emotion/language tags)
|
||||||
# Same logic as in query_memory method
|
# Same logic as in query_memory method
|
||||||
try:
|
try:
|
||||||
if content and content.strip().startswith("{") and content.strip().endswith("}"):
|
if content and content.strip().startswith("{") and content.strip().endswith("}"):
|
||||||
data = json.loads(content)
|
data = json.loads(content)
|
||||||
if "content" in data:
|
if "content" in data:
|
||||||
content = data["content"]
|
content = data["content"]
|
||||||
except (json.JSONDecodeError, KeyError, TypeError):
|
except (json.JSONDecodeError, KeyError, TypeError):
|
||||||
# If parsing fails, use original content
|
# If parsing fails, use original content
|
||||||
pass
|
pass
|
||||||
|
|
||||||
messages.append({"role": message.role, "content": content})
|
messages.append({"role": message.role, "content": content})
|
||||||
|
|
||||||
# Add memory using PowerMem SDK
|
# Add memory using PowerMem SDK
|
||||||
result = self.memory_client.add(
|
result = self.memory_client.add(
|
||||||
messages=messages,
|
messages=messages,
|
||||||
user_id=self.role_id
|
user_id=self.role_id
|
||||||
)
|
)
|
||||||
# Handle both sync and async returns
|
# Handle both sync and async returns
|
||||||
if asyncio.iscoroutine(result):
|
if asyncio.iscoroutine(result):
|
||||||
result = await result
|
result = await result
|
||||||
|
|
||||||
logger.bind(tag=TAG).debug(f"Save memory result: {result}")
|
logger.bind(tag=TAG).debug(f"Save memory result: {result}")
|
||||||
|
|
||||||
# Cache user profile if UserMemory mode and profile was extracted
|
|
||||||
if self.enable_user_profile and result:
|
|
||||||
if result.get('profile_extracted'):
|
|
||||||
self.last_profile_content = result.get('profile_content', '')
|
|
||||||
logger.bind(tag=TAG).debug(f"User profile extracted: {self.last_profile_content}")
|
|
||||||
|
|
||||||
return result
|
|
||||||
|
|
||||||
|
# Cache user profile if UserMemory mode and profile was extracted
|
||||||
|
if self.enable_user_profile and result:
|
||||||
|
if result.get('profile_extracted'):
|
||||||
|
self.last_profile_content = result.get('profile_content', '')
|
||||||
|
logger.bind(tag=TAG).debug(f"User profile extracted: {self.last_profile_content}")
|
||||||
|
else:
|
||||||
|
if not self.use_powermem or self.memory_client is None:
|
||||||
|
logger.bind(tag=TAG).warning("PowerMem is not available, skipping save_memory")
|
||||||
|
elif len(msgs) < 2:
|
||||||
|
logger.bind(tag=TAG).debug("Not enough messages to save (need at least 2)")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"Error saving memory: {str(e)}")
|
logger.bind(tag=TAG).error(f"Error saving memory: {str(e)}")
|
||||||
logger.bind(tag=TAG).debug(f"Detailed error: {traceback.format_exc()}")
|
logger.bind(tag=TAG).debug(f"Detailed error: {traceback.format_exc()}")
|
||||||
return None
|
|
||||||
|
# Generate and save chat summary (SLM summarizes session title)
|
||||||
|
# This is independent of PowerMem's memory saving
|
||||||
|
if session_id:
|
||||||
|
await generate_and_save_chat_summary(session_id)
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
async def query_memory(self, query: str) -> str:
|
async def query_memory(self, query: str) -> str:
|
||||||
"""
|
"""
|
||||||
|
|||||||
Reference in New Issue
Block a user