mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-26 09:03:54 +08:00
Merge remote-tracking branch 'upstream/py_test_Memory_powermem' into add-powermem
# Conflicts: # main/xiaozhi-server/core/providers/memory/powermem/powermem.py
This commit is contained in:
@@ -1,5 +1,6 @@
|
||||
import time
|
||||
import json
|
||||
import uuid
|
||||
import random
|
||||
import asyncio
|
||||
from core.utils.dialogue import Message
|
||||
@@ -105,6 +106,9 @@ async def checkWakeupWords(conn, text):
|
||||
# 播放唤醒词回复
|
||||
conn.client_abort = False
|
||||
|
||||
# 将唤醒词回复视为新会话,生成新的 sentence_id,确保流控器重置
|
||||
conn.sentence_id = str(uuid.uuid4().hex)
|
||||
|
||||
conn.logger.bind(tag=TAG).info(f"播放唤醒词回复: {response.get('text')}")
|
||||
await sendAudioMessage(conn, SentenceType.FIRST, opus_packets, response.get("text"))
|
||||
await sendAudioMessage(conn, SentenceType.LAST, [], None)
|
||||
|
||||
@@ -133,12 +133,27 @@ def _get_or_create_rate_controller(conn, frame_duration, is_single_packet):
|
||||
Returns:
|
||||
(rate_controller, flow_control)
|
||||
"""
|
||||
# 判断是否需要重置:单包模式且 sentence_id 变化,或者控制器不存在
|
||||
need_reset = (
|
||||
is_single_packet
|
||||
and getattr(conn, "audio_flow_control", {}).get("sentence_id")
|
||||
!= conn.sentence_id
|
||||
) or not hasattr(conn, "audio_rate_controller")
|
||||
# 检查是否需要重置控制器
|
||||
need_reset = False
|
||||
|
||||
if not hasattr(conn, "audio_rate_controller"):
|
||||
# 控制器不存在,需要创建
|
||||
need_reset = True
|
||||
else:
|
||||
rate_controller = conn.audio_rate_controller
|
||||
|
||||
# 后台发送任务已停止, 则需要重置
|
||||
if (
|
||||
not rate_controller.pending_send_task
|
||||
or rate_controller.pending_send_task.done()
|
||||
):
|
||||
need_reset = True
|
||||
# 当sentence_id 变化,需要重置
|
||||
elif (
|
||||
getattr(conn, "audio_flow_control", {}).get("sentence_id")
|
||||
!= conn.sentence_id
|
||||
):
|
||||
need_reset = True
|
||||
|
||||
if need_reset:
|
||||
# 创建或获取 rate_controller
|
||||
|
||||
@@ -46,7 +46,7 @@ class MemoryProvider(MemoryProviderBase):
|
||||
try:
|
||||
# Check if user profile mode is enabled
|
||||
self.enable_user_profile = config.get("enable_user_profile", False)
|
||||
|
||||
|
||||
# Get configuration parameters
|
||||
database_provider = config.get("database_provider", "sqlite")
|
||||
llm_provider = config.get("llm_provider", "qwen")
|
||||
@@ -134,12 +134,12 @@ class MemoryProvider(MemoryProviderBase):
|
||||
memory_mode = "AsyncMemory (普通记忆模式)"
|
||||
|
||||
self.use_powermem = True
|
||||
|
||||
|
||||
logger.bind(tag=TAG).info(
|
||||
f"PowerMem initialized successfully: mode={memory_mode}, "
|
||||
f"database={database_provider}, llm={llm_provider}, embedding={embedding_provider}"
|
||||
)
|
||||
|
||||
|
||||
except ImportError as e:
|
||||
logger.bind(tag=TAG).error(
|
||||
f"PowerMem not installed. Please install with: pip install powermem. Error: {e}"
|
||||
@@ -150,13 +150,15 @@ class MemoryProvider(MemoryProviderBase):
|
||||
logger.bind(tag=TAG).debug(f"Detailed error: {traceback.format_exc()}")
|
||||
self.use_powermem = False
|
||||
|
||||
async def save_memory(self, msgs):
|
||||
async def save_memory(self, msgs, session_id=None):
|
||||
"""
|
||||
Save conversation messages to PowerMem.
|
||||
|
||||
Args:
|
||||
msgs: List of message objects with 'role' and 'content' attributes
|
||||
|
||||
session_id: Session identifier (optional, for compatibility)
|
||||
|
||||
Returns:
|
||||
Result from PowerMem API or None if failed
|
||||
"""
|
||||
@@ -186,13 +188,13 @@ class MemoryProvider(MemoryProviderBase):
|
||||
result = await 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
|
||||
|
||||
except Exception as e:
|
||||
@@ -306,3 +308,6 @@ class MemoryProvider(MemoryProviderBase):
|
||||
|
||||
return ""
|
||||
|
||||
|
||||
# Register the memory provider instance
|
||||
powermem = MemoryProvider({})
|
||||
|
||||
@@ -185,20 +185,21 @@ class TTSProvider(TTSProviderBase):
|
||||
# 过滤Markdown
|
||||
filtered_text = MarkdownCleaner.clean_markdown(text)
|
||||
|
||||
# 发送continue-task消息
|
||||
continue_task_message = {
|
||||
"header": {
|
||||
"action": "continue-task",
|
||||
"task_id": self.conn.sentence_id,
|
||||
"streaming": "duplex",
|
||||
},
|
||||
"payload": {"input": {"text": filtered_text}},
|
||||
}
|
||||
|
||||
await self.ws.send(json.dumps(continue_task_message))
|
||||
self.last_active_time = time.time()
|
||||
logger.bind(tag=TAG).debug(f"已发送文本: {filtered_text}")
|
||||
if filtered_text:
|
||||
# 发送continue-task消息
|
||||
continue_task_message = {
|
||||
"header": {
|
||||
"action": "continue-task",
|
||||
"task_id": self.conn.sentence_id,
|
||||
"streaming": "duplex",
|
||||
},
|
||||
"payload": {"input": {"text": filtered_text}},
|
||||
}
|
||||
|
||||
await self.ws.send(json.dumps(continue_task_message))
|
||||
self.last_active_time = time.time()
|
||||
logger.bind(tag=TAG).debug(f"已发送文本: {filtered_text}")
|
||||
return
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"发送TTS文本失败: {str(e)}")
|
||||
if self.ws:
|
||||
|
||||
@@ -288,18 +288,19 @@ class TTSProvider(TTSProviderBase):
|
||||
logger.bind(tag=TAG).warning(f"WebSocket连接不存在,终止发送文本")
|
||||
return
|
||||
filtered_text = MarkdownCleaner.clean_markdown(text)
|
||||
run_request = {
|
||||
"header": {
|
||||
"message_id": uuid.uuid4().hex,
|
||||
"task_id": self.task_id,
|
||||
"namespace": "FlowingSpeechSynthesizer",
|
||||
"name": "RunSynthesis",
|
||||
"appkey": self.appkey,
|
||||
},
|
||||
"payload": {"text": filtered_text},
|
||||
}
|
||||
await self.ws.send(json.dumps(run_request))
|
||||
self.last_active_time = time.time()
|
||||
if filtered_text:
|
||||
run_request = {
|
||||
"header": {
|
||||
"message_id": uuid.uuid4().hex,
|
||||
"task_id": self.task_id,
|
||||
"namespace": "FlowingSpeechSynthesizer",
|
||||
"name": "RunSynthesis",
|
||||
"appkey": self.appkey,
|
||||
},
|
||||
"payload": {"text": filtered_text},
|
||||
}
|
||||
await self.ws.send(json.dumps(run_request))
|
||||
self.last_active_time = time.time()
|
||||
return
|
||||
|
||||
except Exception as e:
|
||||
|
||||
@@ -4,15 +4,15 @@ import json
|
||||
import queue
|
||||
import asyncio
|
||||
import traceback
|
||||
from typing import Callable, Any
|
||||
import websockets
|
||||
|
||||
from typing import Callable, Any
|
||||
from core.utils.tts import MarkdownCleaner
|
||||
from config.logger import setup_logging
|
||||
from core.utils import opus_encoder_utils
|
||||
from core.utils.util import check_model_key
|
||||
from core.providers.tts.base import TTSProviderBase
|
||||
from core.providers.tts.dto.dto import SentenceType, ContentType, InterfaceType
|
||||
from asyncio import Task
|
||||
|
||||
|
||||
TAG = __name__
|
||||
@@ -340,8 +340,9 @@ class TTSProvider(TTSProviderBase):
|
||||
# 过滤Markdown
|
||||
filtered_text = MarkdownCleaner.clean_markdown(text)
|
||||
|
||||
# 发送文本
|
||||
await self.send_text(self.voice, filtered_text, self.conn.sentence_id)
|
||||
if filtered_text:
|
||||
# 发送文本
|
||||
await self.send_text(self.voice, filtered_text, self.conn.sentence_id)
|
||||
return
|
||||
except Exception as e:
|
||||
logger.bind(tag=TAG).error(f"发送TTS文本失败: {str(e)}")
|
||||
|
||||
@@ -237,10 +237,10 @@ class TTSProvider(TTSProviderBase):
|
||||
return
|
||||
|
||||
filtered_text = MarkdownCleaner.clean_markdown(text)
|
||||
|
||||
# 发送文本合成请求
|
||||
run_request = self._build_base_request(status=1,text=filtered_text)
|
||||
await self.ws.send(json.dumps(run_request))
|
||||
if filtered_text:
|
||||
# 发送文本合成请求
|
||||
run_request = self._build_base_request(status=1,text=filtered_text)
|
||||
await self.ws.send(json.dumps(run_request))
|
||||
return
|
||||
|
||||
except Exception as e:
|
||||
|
||||
@@ -2,7 +2,7 @@ import json
|
||||
|
||||
TAG = __name__
|
||||
EMOJI_MAP = {
|
||||
"😂": "laughing",
|
||||
"😂": "funny",
|
||||
"😭": "crying",
|
||||
"😠": "angry",
|
||||
"😔": "sad",
|
||||
|
||||
@@ -1,9 +1,11 @@
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
from config.logger import setup_logging
|
||||
import importlib
|
||||
|
||||
from config.logger import setup_logging
|
||||
from core.utils.textUtils import check_emoji
|
||||
|
||||
logger = setup_logging()
|
||||
|
||||
punctuation_set = {
|
||||
@@ -135,4 +137,8 @@ class MarkdownCleaner:
|
||||
|
||||
for regex, replacement in MarkdownCleaner.REGEXES:
|
||||
text = regex.sub(replacement, text)
|
||||
|
||||
# 去除emoji表情
|
||||
text = check_emoji(text)
|
||||
|
||||
return text.strip()
|
||||
Reference in New Issue
Block a user