Merge branch 'main' into py-test

This commit is contained in:
hrz
2025-05-09 09:52:26 +08:00
committed by GitHub
49 changed files with 798 additions and 559 deletions
+65 -73
View File
@@ -20,6 +20,8 @@ from core.utils.util import (
get_string_no_punctuation_or_emoji,
extract_json_from_string,
initialize_modules,
check_vad_update,
check_asr_update,
)
from concurrent.futures import ThreadPoolExecutor, TimeoutError
from core.handle.sendAudioHandle import sendAudioMessage
@@ -54,7 +56,9 @@ class ConnectionHandler:
_intent,
server=None,
):
self.config = config
self.common_config = config
self.config = copy.deepcopy(config)
self.session_id = str(uuid.uuid4())
self.logger = setup_logging()
self.auth = AuthMiddleware(config)
self.server = server # 保存server实例的引用
@@ -68,7 +72,6 @@ class ConnectionHandler:
self.device_id = None
self.client_ip = None
self.client_ip_info = {}
self.session_id = None
self.prompt = None
self.welcome_msg = None
self.max_output_size = 0
@@ -89,8 +92,10 @@ class ConnectionHandler:
self.tts_report_thread = None
# 依赖的组件
self.vad = _vad
self.asr = _asr
self.vad = None
self.asr = None
self._asr = _asr
self._vad = _vad
self.llm = _llm
self.tts = _tts
self.memory = _memory
@@ -133,6 +138,8 @@ class ConnectionHandler:
int(self.config.get("close_connection_no_voice_time", 120)) + 60
) # 在原来第一道关闭的基础上加60秒,进行二道关闭
self.audio_format = "opus"
async def handle_connection(self, ws):
try:
# 获取并验证headers
@@ -168,7 +175,6 @@ class ConnectionHandler:
# 认证通过,继续处理
self.websocket = ws
self.device_id = self.headers.get("device-id", None)
self.session_id = str(uuid.uuid4())
# 启动超时检查任务
self.timeout_task = asyncio.create_task(self._check_timeout())
@@ -178,9 +184,9 @@ class ConnectionHandler:
await self.websocket.send(json.dumps(self.welcome_msg))
# 获取差异化配置
private_config = self._initialize_private_config()
self._initialize_private_config()
# 异步初始化
self.executor.submit(self._initialize_components, private_config)
self.executor.submit(self._initialize_components)
# tts 消化线程
self.tts_priority_thread = threading.Thread(
target=self._tts_priority_thread, daemon=True
@@ -273,13 +279,20 @@ class ConnectionHandler:
"message": f"Restart failed: {str(e)}"
}))
def _initialize_components(self, private_config):
def _initialize_components(self):
"""初始化组件"""
if private_config is not None:
self._initialize_models(private_config)
else:
if self.config.get("prompt") is not None:
self.prompt = self.config["prompt"]
self.change_system_prompt(self.prompt)
self.logger.bind(tag=TAG).info(
f"初始化组件: prompt成功 {self.prompt[:50]}..."
)
"""初始化本地组件"""
if self.vad is None:
self.vad = self._vad
if self.asr is None:
self.asr = self._asr
"""加载记忆"""
self._initialize_memory()
"""加载意图识别"""
@@ -326,54 +339,21 @@ class ConnectionHandler:
self.logger.bind(tag=TAG).error(f"获取差异化配置失败: {e}")
private_config = {}
init_tts = False
if private_config.get("TTS", None) is not None:
init_tts = True
self.config["TTS"] = private_config["TTS"]
self.config["selected_module"]["TTS"] = private_config["selected_module"][
"TTS"
]
try:
modules = initialize_modules(
self.logger,
private_config,
False,
False,
False,
init_tts,
False,
False,
)
except Exception as e:
self.logger.bind(tag=TAG).error(f"初始化组件失败: {e}")
modules = {}
if modules.get("tts", None) is not None:
self.tts = modules["tts"]
if modules.get("prompt", None) is not None:
self.change_system_prompt(modules["prompt"])
private_config["prompt"] = None
return private_config
def _initialize_models(self, private_config):
init_vad, init_asr, init_llm, init_memory, init_intent = (
False,
init_llm, init_tts, init_memory, init_intent = (
False,
False,
False,
False,
)
if private_config.get("VAD", None) is not None:
init_vad = True
self.config["VAD"] = private_config["VAD"]
self.config["selected_module"]["VAD"] = private_config["selected_module"][
"VAD"
]
if private_config.get("ASR", None) is not None:
init_asr = True
self.config["ASR"] = private_config["ASR"]
self.config["selected_module"]["ASR"] = private_config["selected_module"][
"ASR"
init_vad = check_vad_update(self.common_config, private_config)
init_asr = check_asr_update(self.common_config, private_config)
if private_config.get("TTS", None) is not None:
init_tts = True
self.config["TTS"] = private_config["TTS"]
self.config["selected_module"]["TTS"] = private_config["selected_module"][
"TTS"
]
if private_config.get("LLM", None) is not None:
init_llm = True
@@ -393,8 +373,11 @@ class ConnectionHandler:
self.config["selected_module"]["Intent"] = private_config[
"selected_module"
]["Intent"]
if private_config.get("prompt", None) is not None:
self.config["prompt"] = private_config["prompt"]
if private_config.get("device_max_output_size", None) is not None:
self.max_output_size = int(private_config["device_max_output_size"])
try:
modules = initialize_modules(
self.logger,
@@ -402,13 +385,15 @@ class ConnectionHandler:
init_vad,
init_asr,
init_llm,
False,
init_tts,
init_memory,
init_intent,
)
except Exception as e:
self.logger.bind(tag=TAG).error(f"初始化组件失败: {e}")
modules = {}
if modules.get("tts", None) is not None:
self.tts = modules["tts"]
if modules.get("vad", None) is not None:
self.vad = modules["vad"]
if modules.get("asr", None) is not None:
@@ -486,10 +471,12 @@ class ConnectionHandler:
processed_chars = 0 # 跟踪已处理的字符位置
try:
# 使用带记忆的对话
future = asyncio.run_coroutine_threadsafe(
self.memory.query_memory(query), self.loop
)
memory_str = future.result()
memory_str = None
if self.memory is not None:
future = asyncio.run_coroutine_threadsafe(
self.memory.query_memory(query), self.loop
)
memory_str = future.result()
self.logger.bind(tag=TAG).debug(f"记忆内容: {memory_str}")
llm_responses = self.llm.response(
@@ -550,7 +537,7 @@ class ConnectionHandler:
future = self.executor.submit(
self.speak_and_play, segment_text, text_index
)
self.tts_queue.put(future)
self.tts_queue.put((future, text_index))
self.llm_finish_task = True
self.dialogue.put(Message(role="assistant", content="".join(response_message)))
@@ -577,10 +564,12 @@ class ConnectionHandler:
start_time = time.time()
# 使用带记忆的对话
future = asyncio.run_coroutine_threadsafe(
self.memory.query_memory(query), self.loop
)
memory_str = future.result()
memory_str = None
if self.memory is not None:
future = asyncio.run_coroutine_threadsafe(
self.memory.query_memory(query), self.loop
)
memory_str = future.result()
# self.logger.bind(tag=TAG).info(f"对话记录: {self.dialogue.get_llm_dialogue_with_memory(memory_str)}")
@@ -725,7 +714,7 @@ class ConnectionHandler:
future = self.executor.submit(
self.speak_and_play, segment_text, text_index
)
self.tts_queue.put(future)
self.tts_queue.put((future, text_index))
# 存储对话内容
if len(response_message) > 0:
@@ -787,7 +776,7 @@ class ConnectionHandler:
text = result.response
self.recode_first_last_text(text, text_index)
future = self.executor.submit(self.speak_and_play, text, text_index)
self.tts_queue.put(future)
self.tts_queue.put((future, text_index))
self.dialogue.put(Message(role="assistant", content=text))
elif result.action == Action.REQLLM: # 调用函数后再请求llm生成回复
text = result.result
@@ -820,7 +809,7 @@ class ConnectionHandler:
text = result.result
self.recode_first_last_text(text, text_index)
future = self.executor.submit(self.speak_and_play, text, text_index)
self.tts_queue.put(future)
self.tts_queue.put((future, text_index))
self.dialogue.put(Message(role="assistant", content=text))
else:
pass
@@ -841,7 +830,7 @@ class ConnectionHandler:
if future is None:
continue
text = None
opus_datas, tts_file = [], None
opus_datas, tts_file = [], None
try:
self.logger.bind(tag=TAG).debug("正在处理TTS任务...")
tts_timeout = int(self.config.get("tts_timeout", 10))
@@ -859,9 +848,12 @@ class ConnectionHandler:
f"TTS生成:文件路径: {tts_file}"
)
if os.path.exists(tts_file):
opus_datas, _ = self.tts.audio_to_opus_data(tts_file)
if self.audio_format == "pcm":
audio_datas, _ = self.tts.audio_to_pcm_data(tts_file)
else:
audio_datas, _ = self.tts.audio_to_opus_data(tts_file)
# 在这里上报TTS数据(使用文件路径)
enqueue_tts_report(self, 2, text, opus_datas)
enqueue_tts_report(self, 2, text, audio_datas)
else:
self.logger.bind(tag=TAG).error(
f"TTS出错:文件不存在{tts_file}"
@@ -872,7 +864,7 @@ class ConnectionHandler:
self.logger.bind(tag=TAG).error(f"TTS出错: {e}")
if not self.client_abort:
# 如果没有中途打断就发送语音
self.audio_play_queue.put((opus_datas, text, text_index))
self.audio_play_queue.put((audio_datas, text, text_index))
if (
self.tts.delete_audio_file
and tts_file is not None
@@ -903,13 +895,13 @@ class ConnectionHandler:
text = None
try:
try:
opus_datas, text, text_index = self.audio_play_queue.get(timeout=1)
audio_datas, text, text_index = self.audio_play_queue.get(timeout=1)
except queue.Empty:
if self.stop_event.is_set():
break
continue
future = asyncio.run_coroutine_threadsafe(
sendAudioMessage(self, opus_datas, text, text_index), self.loop
sendAudioMessage(self, audio_datas, text, text_index), self.loop
)
future.result()
except Exception as e: