From 943009407211812ad0acc994f3cae30614b21ea3 Mon Sep 17 00:00:00 2001 From: hrz <1710360675@qq.com> Date: Tue, 1 Apr 2025 00:03:05 +0800 Subject: [PATCH] =?UTF-8?q?update:=E4=BF=AE=E5=A4=8D=E9=80=80=E5=87=BA?= =?UTF-8?q?=E4=BF=9D=E5=AD=98=E8=AE=B0=E5=BF=86bug=20(#619)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- main/xiaozhi-server/core/connection.py | 6 +- main/xiaozhi-server/core/websocket_server.py | 98 +++++++++++++------- 2 files changed, 68 insertions(+), 36 deletions(-) diff --git a/main/xiaozhi-server/core/connection.py b/main/xiaozhi-server/core/connection.py index df0c2aa3..ea3206d6 100644 --- a/main/xiaozhi-server/core/connection.py +++ b/main/xiaozhi-server/core/connection.py @@ -195,7 +195,7 @@ class ConnectionHandler: self.logger.bind(tag=TAG).error(f"Connection error: {str(e)}-{stack_trace}") return finally: - self._save_and_close(ws) + await self._save_and_close(ws) async def _save_and_close(self, ws): """保存记忆并关闭连接""" @@ -253,13 +253,9 @@ class ConnectionHandler: else: # 否则使用主LLM self.intent.set_llm(self.llm) - self.logger.bind(tag=TAG).info("意图识别使用主LLM") # 记录意图识别LLM初始化耗时 intent_llm_init_time = time.time() - intent_llm_init_start - self.logger.bind(tag=TAG).info( - f"意图识别LLM初始化完成,耗时: {intent_llm_init_time:.4f}秒" - ) """加载位置信息""" self.client_ip_info = get_ip_info(self.client_ip) diff --git a/main/xiaozhi-server/core/websocket_server.py b/main/xiaozhi-server/core/websocket_server.py index 6aa9a06d..98b60133 100644 --- a/main/xiaozhi-server/core/websocket_server.py +++ b/main/xiaozhi-server/core/websocket_server.py @@ -12,50 +12,72 @@ class WebSocketServer: def __init__(self, config: dict): self.config = config self.logger = setup_logging() - self._vad, self._asr, self._llm, self._tts, self._memory, self.intent = self._create_processing_instances() + self._vad, self._asr, self._llm, self._tts, self._memory, self.intent = ( + self._create_processing_instances() + ) self.active_connections = set() # 添加全局连接记录 def _create_processing_instances(self): - memory_cls_name = self.config["selected_module"].get("Memory", "nomem") # 默认使用nomem - has_memory_cfg = self.config.get("Memory") and memory_cls_name in self.config["Memory"] + memory_cls_name = self.config["selected_module"].get( + "Memory", "nomem" + ) # 默认使用nomem + 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( self.config["selected_module"]["VAD"], - self.config["VAD"][self.config["selected_module"]["VAD"]] + self.config["VAD"][self.config["selected_module"]["VAD"]], ), asr.create_instance( - self.config["selected_module"]["ASR"] - if not 'type' in self.config["ASR"][self.config["selected_module"]["ASR"]] - else - self.config["ASR"][self.config["selected_module"]["ASR"]]["type"], + ( + self.config["selected_module"]["ASR"] + if not "type" + in self.config["ASR"][self.config["selected_module"]["ASR"]] + else self.config["ASR"][self.config["selected_module"]["ASR"]][ + "type" + ] + ), self.config["ASR"][self.config["selected_module"]["ASR"]], - self.config["delete_audio"] + self.config["delete_audio"], ), llm.create_instance( - self.config["selected_module"]["LLM"] - if not 'type' in self.config["LLM"][self.config["selected_module"]["LLM"]] - else - self.config["LLM"][self.config["selected_module"]["LLM"]]['type'], + ( + self.config["selected_module"]["LLM"] + if not "type" + in self.config["LLM"][self.config["selected_module"]["LLM"]] + else self.config["LLM"][self.config["selected_module"]["LLM"]][ + "type" + ] + ), self.config["LLM"][self.config["selected_module"]["LLM"]], ), tts.create_instance( - self.config["selected_module"]["TTS"] - if not 'type' in self.config["TTS"][self.config["selected_module"]["TTS"]] - else - self.config["TTS"][self.config["selected_module"]["TTS"]]["type"], + ( + self.config["selected_module"]["TTS"] + if not "type" + in self.config["TTS"][self.config["selected_module"]["TTS"]] + else self.config["TTS"][self.config["selected_module"]["TTS"]][ + "type" + ] + ), self.config["TTS"][self.config["selected_module"]["TTS"]], - self.config["delete_audio"] + self.config["delete_audio"], ), memory.create_instance(memory_cls_name, memory_cfg), intent.create_instance( - self.config["selected_module"]["Intent"] - if not 'type' in self.config["Intent"][self.config["selected_module"]["Intent"]] - else - self.config["Intent"][self.config["selected_module"]["Intent"]]["type"], - self.config["Intent"][self.config["selected_module"]["Intent"]] + ( + self.config["selected_module"]["Intent"] + if not "type" + in self.config["Intent"][self.config["selected_module"]["Intent"]] + else self.config["Intent"][ + self.config["selected_module"]["Intent"] + ]["type"] + ), + self.config["Intent"][self.config["selected_module"]["Intent"]], ), ) @@ -64,19 +86,33 @@ class WebSocketServer: host = server_config["ip"] port = server_config["port"] - self.logger.bind(tag=TAG).info("Server is running at ws://{}:{}", get_local_ip(), port) - self.logger.bind(tag=TAG).info("=======上面的地址是websocket协议地址,请勿用浏览器访问=======") - async with websockets.serve( - self._handle_connection, - host, - port - ): + self.logger.bind(tag=TAG).info( + "Server is running at ws://{}:{}/xiaozhi/v1/", get_local_ip(), port + ) + self.logger.bind(tag=TAG).info( + "=======上面的地址是websocket协议地址,请勿用浏览器访问=======" + ) + self.logger.bind(tag=TAG).info( + "如想测试websocket请用谷歌浏览器打开test目录下的test_page.html" + ) + self.logger.bind(tag=TAG).info( + "=============================================================\n" + ) + async with websockets.serve(self._handle_connection, host, port): await asyncio.Future() async def _handle_connection(self, websocket): """处理新连接,每次创建独立的ConnectionHandler""" # 创建ConnectionHandler时传入当前server实例 - handler = ConnectionHandler(self.config, self._vad, self._asr, self._llm, self._tts, self._memory, self.intent) + handler = ConnectionHandler( + self.config, + self._vad, + self._asr, + self._llm, + self._tts, + self._memory, + self.intent, + ) self.active_connections.add(handler) try: await handler.handle_connection(websocket)