From c900498ce865d144571f7a81d0d0246e472b5713 Mon Sep 17 00:00:00 2001 From: hrz <1710360675@qq.com> Date: Wed, 21 May 2025 15:55:40 +0800 Subject: [PATCH] =?UTF-8?q?update:=E5=90=88=E5=B9=B6main=E5=88=86=E6=94=AF?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- main/xiaozhi-server/core/connection.py | 544 +++++++++++++----- .../core/handle/intentHandler.py | 75 ++- 2 files changed, 450 insertions(+), 169 deletions(-) diff --git a/main/xiaozhi-server/core/connection.py b/main/xiaozhi-server/core/connection.py index 0aa85053..6edf00bf 100644 --- a/main/xiaozhi-server/core/connection.py +++ b/main/xiaozhi-server/core/connection.py @@ -1,4 +1,8 @@ +import os +import copy import json +import subprocess +import sys import uuid import time import queue @@ -16,17 +20,22 @@ from core.handle.textHandle import handleTextMessage from core.utils.util import ( get_string_no_punctuation_or_emoji, extract_json_from_string, - get_ip_info, + initialize_modules, + check_vad_update, + check_asr_update, + filter_sensitive_info, ) from concurrent.futures import ThreadPoolExecutor, TimeoutError from core.handle.sendAudioHandle import sendAudioMessage from core.handle.receiveAudioHandle import handleAudioMessage from core.handle.functionHandler import FunctionHandler from plugins_func.register import Action, ActionResponse -from config.private_config import PrivateConfig from core.auth import AuthMiddleware, AuthenticationError -from core.utils.auth_code_gen import AuthCodeGenerator from core.mcp.manager import MCPManager +from config.config_loader import get_private_config_from_api +from config.manage_api_client import DeviceNotFoundException, DeviceBindException +from core.utils.output_counter import add_device_output +from core.handle.reportHandle import enqueue_tts_report, report TAG = __name__ @@ -39,22 +48,39 @@ class TTSException(RuntimeError): class ConnectionHandler: def __init__( - self, config: Dict[str, Any], _vad, _asr, _llm, _tts, _memory, _intent + self, + config: Dict[str, Any], + _vad, + _asr, + _llm, + _tts, + _memory, + _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.server = server # 保存server实例的引用 + self.auth = AuthMiddleware(config) + self.need_bind = False + self.bind_code = None + self.read_config_from_api = self.config.get("read_config_from_api", False) self.tts_stream = self.config.get("TTS_SET", {}).get("TTS_STREAM", False) self.websocket = None self.headers = None + self.device_id = None self.client_ip = None self.client_ip_info = {} - self.session_id = None self.prompt = None self.welcome_msg = None self.u_id = None + self.max_output_size = 0 + self.chat_history_conf = 0 # 客户端状态相关 self.client_abort = False @@ -70,9 +96,18 @@ class ConnectionHandler: self.executor = ThreadPoolExecutor(max_workers=max_workers) self.start_tts_request_flag = False + # 上报线程 + self.report_queue = queue.Queue() + self.report_thread = None + # TODO(haotian): 2025/5/12 可以通过修改此处,调节asr的上报和tts的上报 + self.report_asr_enable = self.read_config_from_api + self.report_tts_enable = self.read_config_from_api + # 依赖的组件 - 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 @@ -102,24 +137,47 @@ class ConnectionHandler: self.iot_descriptors = {} self.func_handler = None - self.cmd_exit = self.config["CMD_exit"] + self.cmd_exit = self.config["exit_commands"] self.max_cmd_length = 0 for cmd in self.cmd_exit: if len(cmd) > self.max_cmd_length: self.max_cmd_length = len(cmd) - self.private_config = None - self.auth_code_gen = AuthCodeGenerator.get_instance() - self.is_device_verified = False # 添加设备验证状态标志 - self.close_after_chat = False # 是否在聊天结束后关闭连接 - self.use_function_call_mode = False - if self.config["selected_module"]["Intent"] == "function_call": - self.use_function_call_mode = True + # 是否在聊天结束后关闭连接 + self.close_after_chat = False + self.load_function_plugin = False + self.intent_type = "nointent" + + self.timeout_task = None + self.timeout_seconds = ( + int(self.config.get("close_connection_no_voice_time", 120)) + 60 + ) # 在原来第一道关闭的基础上加60秒,进行二道关闭 + + self.audio_format = "opus" async def handle_connection(self, ws): try: # 获取并验证headers self.headers = dict(ws.request.headers) + + if self.headers.get("device-id", None) is None: + # 尝试从 URL 的查询参数中获取 device-id + from urllib.parse import parse_qs, urlparse + + # 从 WebSocket 请求中获取路径 + request_path = ws.request.path + if not request_path: + self.logger.bind(tag=TAG).error("无法获取请求路径") + return + parsed_url = urlparse(request_path) + query_params = parse_qs(parsed_url.query) + if "device-id" in query_params: + self.headers["device-id"] = query_params["device-id"][0] + self.headers["client-id"] = query_params["client-id"][0] + else: + await ws.send("端口正常,如需测试连接,请使用test_page.html") + await self.close(ws) + return # 获取客户端ip地址 self.client_ip = ws.remote_address[0] self.logger.bind(tag=TAG).info( @@ -128,49 +186,20 @@ class ConnectionHandler: # 进行认证 await self.auth.authenticate(self.headers) - device_id = self.headers.get("device-id", None) # 认证通过,继续处理 self.websocket = ws - self.session_id = str(uuid.uuid4()) + self.device_id = self.headers.get("device-id", None) + + # 启动超时检查任务 + self.timeout_task = asyncio.create_task(self._check_timeout()) self.welcome_msg = self.config["xiaozhi"] self.welcome_msg["session_id"] = self.session_id await self.websocket.send(json.dumps(self.welcome_msg)) - # Load private configuration if device_id is provided - bUsePrivateConfig = self.config.get("use_private_config", False) - if bUsePrivateConfig and device_id: - try: - self.private_config = PrivateConfig( - device_id, self.config, self.auth_code_gen - ) - await self.private_config.load_or_create() - # 判断是否已经绑定 - owner = self.private_config.get_owner() - self.is_device_verified = owner is not None - - if self.is_device_verified: - await self.private_config.update_last_chat_time() - - llm, tts = self.private_config.create_private_instances() - if all([llm, tts]): - self.llm = llm - self.tts = tts - self.logger.bind(tag=TAG).info( - f"Loaded private config and instances for device {device_id}" - ) - else: - self.logger.bind(tag=TAG).error( - f"Failed to create instances for device {device_id}" - ) - self.private_config = None - except Exception as e: - self.logger.bind(tag=TAG).error( - f"Error initializing private config: {e}" - ) - self.private_config = None - raise + # 获取差异化配置 + self._initialize_private_config() # 异步初始化 self.executor.submit(self._initialize_components) @@ -202,55 +231,254 @@ class ConnectionHandler: async def _save_and_close(self, ws): """保存记忆并关闭连接""" try: - await self.memory.save_memory(self.dialogue.dialogue) + if self.memory: + # 使用线程池异步保存记忆 + def save_memory_task(): + try: + # 创建新事件循环(避免与主循环冲突) + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + loop.run_until_complete( + self.memory.save_memory(self.dialogue.dialogue) + ) + except Exception as e: + self.logger.bind(tag=TAG).error(f"保存记忆失败: {e}") + finally: + loop.close() + + # 启动线程保存记忆,不等待完成 + threading.Thread(target=save_memory_task, daemon=True).start() except Exception as e: self.logger.bind(tag=TAG).error(f"保存记忆失败: {e}") finally: + # 立即关闭连接,不等待记忆保存完成 await self.close(ws) + async def reset_timeout(self): + """重置超时计时器""" + if self.timeout_task: + self.timeout_task.cancel() + self.timeout_task = asyncio.create_task(self._check_timeout()) + async def _route_message(self, message): """消息路由""" + # 重置超时计时器 + await self.reset_timeout() + if isinstance(message, str): await handleTextMessage(self, message) elif isinstance(message, bytes): await handleAudioMessage(self, message) - def _initialize_components(self): - """加载提示词""" - self.prompt = self.config["prompt"] - if self.private_config: - self.prompt = self.private_config.private_config.get("prompt", self.prompt) - self.dialogue.put(Message(role="system", content=self.prompt)) + async def handle_restart(self, message): + """处理服务器重启请求""" + try: + self.logger.bind(tag=TAG).info("收到服务器重启指令,准备执行...") + + # 发送确认响应 + await self.websocket.send( + json.dumps( + { + "type": "server", + "status": "success", + "message": "服务器重启中...", + "content": {"action": "restart"}, + } + ) + ) + + # 异步执行重启操作 + def restart_server(): + """实际执行重启的方法""" + time.sleep(1) + self.logger.bind(tag=TAG).info("执行服务器重启...") + subprocess.Popen( + [sys.executable, "app.py"], + stdin=sys.stdin, + stdout=sys.stdout, + stderr=sys.stderr, + start_new_session=True, + ) + os._exit(0) + + # 使用线程执行重启避免阻塞事件循环 + threading.Thread(target=restart_server, daemon=True).start() + + except Exception as e: + self.logger.bind(tag=TAG).error(f"重启失败: {str(e)}") + await self.websocket.send( + json.dumps( + { + "type": "server", + "status": "error", + "message": f"Restart failed: {str(e)}", + "content": {"action": "restart"}, + } + ) + ) + + def _initialize_components(self): + """初始化组件""" + 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() """加载意图识别""" self._initialize_intent() - """加载位置信息""" - self.client_ip_info = get_ip_info(self.client_ip) - if self.client_ip_info is not None and "city" in self.client_ip_info: - self.logger.bind(tag=TAG).info(f"Client ip info: {self.client_ip_info}") - self.prompt = self.prompt + f"\nuser location:{self.client_ip_info}" + """初始化上报线程""" + self._init_report_threads() - self.dialogue.update_system_message(self.prompt) + def _init_report_threads(self): + """初始化ASR和TTS上报线程""" + if not self.read_config_from_api or self.need_bind: + return + if self.chat_history_conf == 0: + return + if self.report_thread is None or not self.report_thread.is_alive(): + self.report_thread = threading.Thread( + target=self._report_worker, daemon=True + ) + self.report_thread.start() + self.logger.bind(tag=TAG).info("TTS上报线程已启动") + + def _initialize_private_config(self): + """如果是从配置文件获取,则进行二次实例化""" + if not self.read_config_from_api: + return + """从接口获取差异化的配置进行二次实例化,非全量重新实例化""" + try: + begin_time = time.time() + private_config = get_private_config_from_api( + self.config, + self.headers.get("device-id"), + self.headers.get("client-id", self.headers.get("device-id")), + ) + private_config["delete_audio"] = bool(self.config.get("delete_audio", True)) + self.logger.bind(tag=TAG).info( + f"{time.time() - begin_time} 秒,获取差异化配置成功: {json.dumps(filter_sensitive_info(private_config), ensure_ascii=False)}" + ) + except DeviceNotFoundException as e: + self.need_bind = True + private_config = {} + except DeviceBindException as e: + self.need_bind = True + self.bind_code = e.bind_code + private_config = {} + except Exception as e: + self.need_bind = True + self.logger.bind(tag=TAG).error(f"获取差异化配置失败: {e}") + private_config = {} + + init_llm, init_tts, init_memory, init_intent = ( + False, + False, + False, + False, + ) + + 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 + self.config["LLM"] = private_config["LLM"] + self.config["selected_module"]["LLM"] = private_config["selected_module"][ + "LLM" + ] + if private_config.get("Memory", None) is not None: + init_memory = True + self.config["Memory"] = private_config["Memory"] + self.config["selected_module"]["Memory"] = private_config[ + "selected_module" + ]["Memory"] + if private_config.get("Intent", None) is not None: + init_intent = True + self.config["Intent"] = private_config["Intent"] + 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("summaryMemory", None) is not None: + self.config["summaryMemory"] = private_config["summaryMemory"] + if private_config.get("device_max_output_size", None) is not None: + self.max_output_size = int(private_config["device_max_output_size"]) + if private_config.get("chat_history_conf", None) is not None: + self.chat_history_conf = int(private_config["chat_history_conf"]) + try: + modules = initialize_modules( + self.logger, + private_config, + init_vad, + init_asr, + init_llm, + 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: + self.asr = modules["asr"] + if modules.get("llm", None) is not None: + self.llm = modules["llm"] + if modules.get("intent", None) is not None: + self.intent = modules["intent"] + if modules.get("memory", None) is not None: + self.memory = modules["memory"] def _initialize_memory(self): """初始化记忆模块""" - device_id = self.headers.get("device-id", None) - self.memory.init_memory(device_id, self.llm) + self.memory.init_memory( + role_id=self.device_id, + llm=self.llm, + summary_memory=self.config.get("summaryMemory", None), + save_to_file=not self.read_config_from_api, + ) def _initialize_intent(self): + self.intent_type = self.config["Intent"][ + self.config["selected_module"]["Intent"] + ]["type"] + if self.intent_type == "function_call" or self.intent_type == "intent_llm": + self.load_function_plugin = True """初始化意图识别模块""" # 获取意图识别配置 intent_config = self.config["Intent"] - intent_type = self.config["selected_module"]["Intent"] + intent_type = self.config["Intent"][self.config["selected_module"]["Intent"]][ + "type" + ] # 如果使用 nointent,直接返回 if intent_type == "nointent": return # 使用 intent_llm 模式 elif intent_type == "intent_llm": - intent_llm_name = intent_config["intent_llm"]["llm"] + intent_llm_name = intent_config[self.config["selected_module"]["Intent"]][ + "llm" + ] if intent_llm_name and intent_llm_name in self.config["LLM"]: # 如果配置了专用LLM,则创建独立的LLM实例 @@ -281,47 +509,22 @@ class ConnectionHandler: def change_system_prompt(self, prompt): self.prompt = prompt - # 找到原来的role==system,替换原来的系统提示 - for m in self.dialogue.dialogue: - if m.role == "system": - m.content = prompt - - async def _check_and_broadcast_auth_code(self): - """检查设备绑定状态并广播认证码""" - if not self.private_config.get_owner(): - auth_code = self.private_config.get_auth_code() - if auth_code: - # 发送验证码语音提示 - text = f"请在后台输入验证码:{' '.join(auth_code)}" - return False - return True - - def isNeedAuth(self): - bUsePrivateConfig = self.config.get("use_private_config", False) - if not bUsePrivateConfig: - # 如果不使用私有配置,就不需要验证 - return False - return not self.is_device_verified + # 更新系统prompt至上下文 + self.dialogue.update_system_message(self.prompt) def chat(self, query): - if self.isNeedAuth(): - self.llm_finish_task = True - future = asyncio.run_coroutine_threadsafe( - self._check_and_broadcast_auth_code(), self.loop - ) - future.result() - return True self.dialogue.put(Message(role="user", content=query)) response_message = [] try: - 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).debug(f"记忆内容: {memory_str}") llm_responses = self.llm.response( @@ -376,13 +579,6 @@ class ConnectionHandler: def chat_with_function_calling(self, query, tool_call=False): self.logger.bind(tag=TAG).debug(f"Chat with function calling start: {query}") """Chat with function calling for intent detection using streaming""" - if self.isNeedAuth(): - self.llm_finish_task = True - future = asyncio.run_coroutine_threadsafe( - self._check_and_broadcast_auth_code(), self.loop - ) - future.result() - return True if not tool_call: self.dialogue.put(Message(role="user", content=query)) @@ -397,10 +593,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)}") @@ -427,14 +625,16 @@ class ConnectionHandler: self.u_id = uuid_str for response in llm_responses: content, tools_call = response + if "content" in response: content = response["content"] tools_call = None if content is not None and len(content) > 0: - if len(response_message) <= 0 and ( - content == "```" or "" in content - ): - tool_call_flag = True + content_arguments += content + + if not tool_call_flag and content_arguments.startswith(""): + # print("content_arguments", content_arguments) + tool_call_flag = True if tools_call is not None: tool_call_flag = True @@ -446,9 +646,7 @@ class ConnectionHandler: function_arguments += tools_call[0].function.arguments if content is not None and len(content) > 0: - if tool_call_flag: - content_arguments += content - else: + if not tool_call_flag: response_message.append(content) if self.client_abort: @@ -507,10 +705,9 @@ class ConnectionHandler: self.logger.bind(tag=TAG).error( f"function call error: {content_arguments}" ) - else: - function_arguments = json.loads(function_arguments) if not bHasError: - self.logger.bind(tag=TAG).info( + response_message.clear() + self.logger.bind(tag=TAG).debug( f"function_name={function_name}, function_id={function_id}, function_arguments={function_arguments}" ) function_call_data = { @@ -614,10 +811,16 @@ class ConnectionHandler: ) self.dialogue.put( - Message(role="tool", tool_call_id=function_id, content=text) + Message( + role="tool", + tool_call_id=( + str(uuid.uuid4()) if function_id is None else function_id + ), + content=text, + ) ) self.chat_with_function_calling(text, tool_call=True) - elif result.action == Action.NOTFOUND: + elif result.action == Action.NOTFOUND or result.action == Action.ERROR: text = result.result self.recode_first_last_text(text, text_index) self.tts.tts_one_sentence(self, text) @@ -646,6 +849,45 @@ class ConnectionHandler: f"audio_play_priority priority_thread: {text} {e}" ) + def _report_worker(self): + """聊天记录上报工作线程""" + while not self.stop_event.is_set(): + try: + # 从队列获取数据,设置超时以便定期检查停止事件 + item = self.report_queue.get(timeout=1) + if item is None: # 检测毒丸对象 + break + + type, text, audio_data = item + + try: + # 执行上报(传入二进制数据) + report(self, type, text, audio_data) + except Exception as e: + self.logger.bind(tag=TAG).error(f"聊天记录上报线程异常: {e}") + finally: + # 标记任务完成 + self.report_queue.task_done() + except queue.Empty: + continue + except Exception as e: + self.logger.bind(tag=TAG).error(f"聊天记录上报工作线程异常: {e}") + + self.logger.bind(tag=TAG).info("聊天记录上报线程已退出") + + def speak_and_play(self, text, text_index=0): + if text is None or len(text) <= 0: + self.logger.bind(tag=TAG).info(f"无需tts转换,query为空,{text}") + return None, text, text_index + tts_file = self.tts.to_tts(text) + if tts_file is None: + self.logger.bind(tag=TAG).error(f"tts转换失败,{text}") + return None, text, text_index + self.logger.bind(tag=TAG).debug(f"TTS 文件生成完毕: {tts_file}") + if self.max_output_size > 0: + add_device_output(self.headers.get("device-id"), len(text)) + return tts_file, text, text_index + def clearSpeakStatus(self): self.logger.bind(tag=TAG).debug(f"清除服务端讲话状态") self.asr_server_receive = True @@ -660,42 +902,56 @@ class ConnectionHandler: async def close(self, ws=None): """资源清理方法""" + + # 取消超时任务 + if self.timeout_task: + self.timeout_task.cancel() + self.timeout_task = None + # 清理MCP资源 if hasattr(self, "mcp_manager") and self.mcp_manager: await self.mcp_manager.cleanup_all() - # 触发停止事件并清理资源 + # 触发停止事件 if self.stop_event: self.stop_event.set() - # 立即关闭线程池 - if self.executor: - self.executor.shutdown(wait=False, cancel_futures=True) - self.executor = None - # 清空任务队列 - self._clear_queues() + self.clear_queues() + # 关闭WebSocket连接 if ws: await ws.close() elif self.websocket: await self.websocket.close() await self.tts.close() + + # 最后关闭线程池(避免阻塞) + if self.executor: + self.executor.shutdown(wait=False) + self.executor = None + self.logger.bind(tag=TAG).info("连接资源已释放") - def _clear_queues(self): - # 清空所有任务队列 + def clear_queues(self): + """清空所有任务队列""" + self.logger.bind(tag=TAG).debug( + f"开始清理: TTS队列大小={self.tts_queue.qsize()}, 音频队列大小={self.audio_play_queue.qsize()}" + ) + + # 使用非阻塞方式清空队列 for q in [self.tts_queue, self.audio_play_queue]: if not q: continue - while not q.empty(): + while True: try: q.get_nowait() except queue.Empty: - continue - q.queue.clear() - # 添加毒丸信号到队列,确保线程退出 - # q.queue.put(None) + break + + self.logger.bind(tag=TAG).debug( + f"清理结束: TTS队列大小={self.tts_queue.qsize()}, 音频队列大小={self.audio_play_queue.qsize()}" + ) def reset_vad_states(self): self.client_audio_buffer = bytearray() @@ -714,3 +970,15 @@ class ConnectionHandler: self.close_after_chat = True except Exception as e: self.logger.bind(tag=TAG).error(f"Chat and close error: {str(e)}") + + async def _check_timeout(self): + """检查连接超时""" + try: + while not self.stop_event.is_set(): + await asyncio.sleep(self.timeout_seconds) + if not self.stop_event.is_set(): + self.logger.bind(tag=TAG).info("连接超时,准备关闭") + await self.close(self.websocket) + break + except Exception as e: + self.logger.bind(tag=TAG).error(f"超时检查任务出错: {e}") diff --git a/main/xiaozhi-server/core/handle/intentHandler.py b/main/xiaozhi-server/core/handle/intentHandler.py index 168844b7..e8538aa1 100644 --- a/main/xiaozhi-server/core/handle/intentHandler.py +++ b/main/xiaozhi-server/core/handle/intentHandler.py @@ -4,10 +4,10 @@ import uuid from core.handle.sendAudioHandle import send_stt_message from core.utils.util import remove_punctuation_and_length from core.utils.dialogue import Message +from plugins_func.register import Action from loguru import logger TAG = __name__ -logger = setup_logging() async def handle_user_intent(conn, text): @@ -19,7 +19,7 @@ async def handle_user_intent(conn, text): # if await checkWakeupWords(conn, text): # return True - if conn.use_function_call_mode: + if conn.intent_type == "function_call": # 使用支持function calling的聊天方法,不再进行意图分析 return False # 使用LLM进行意图分析 @@ -36,7 +36,7 @@ async def check_direct_exit(conn, text): cmd_exit = conn.cmd_exit for cmd in cmd_exit: if text == cmd: - logger.bind(tag=TAG).info(f"识别到明确的退出命令: {text}") + conn.logger.bind(tag=TAG).info(f"识别到明确的退出命令: {text}") await send_stt_message(conn, text) await conn.close() return True @@ -46,7 +46,7 @@ async def check_direct_exit(conn, text): async def analyze_intent_with_llm(conn, text): """使用LLM分析用户意图""" if not hasattr(conn, "intent") or not conn.intent: - logger.bind(tag=TAG).warning("意图识别服务未初始化") + conn.logger.bind(tag=TAG).warning("意图识别服务未初始化") return None # 对话历史记录 @@ -55,7 +55,7 @@ async def analyze_intent_with_llm(conn, text): intent_result = await conn.intent.detect_intent(conn, dialogue.dialogue, text) return intent_result except Exception as e: - logger.bind(tag=TAG).error(f"意图识别失败: {str(e)}") + conn.logger.bind(tag=TAG).error(f"意图识别失败: {str(e)}") return None @@ -69,13 +69,18 @@ async def process_intent_result(conn, intent_result, original_text): # 检查是否有function_call if "function_call" in intent_data: # 直接从意图识别获取了function_call - logger.bind(tag=TAG).debug( + conn.logger.bind(tag=TAG).debug( f"检测到function_call格式的意图结果: {intent_data['function_call']['name']}" ) function_name = intent_data["function_call"]["name"] if function_name == "continue_chat": return False + if function_name == "play_music": + funcItem = conn.func_handler.get_function(function_name) + if not funcItem: + conn.func_handler.function_registry.register_function("play_music") + function_args = None if "arguments" in intent_data["function_call"]: function_args = intent_data["function_call"]["arguments"] @@ -97,37 +102,45 @@ async def process_intent_result(conn, intent_result, original_text): result = conn.func_handler.handle_llm_function_call( conn, function_call_data ) - if result and function_name != "play_music": - # 获取当前最新的文本索引 - text = result.response - if text is None: + logger.bind(tag=TAG).debug(f"检测到Action : {result.action}") + + if result: + if result.action == Action.RESPONSE: # 直接回复前端 + text = result.response + if text is not None: + speak_and_play(conn, text) + elif result.action == Action.REQLLM: # 调用函数后再请求llm生成回复 text = result.result - if text is not None: - conn.tts.tts_one_sentence(conn, text) + conn.dialogue.put(Message(role="tool", content=text)) + llm_result = conn.intent.replyResult(text, original_text) + if llm_result is None: + llm_result = text + speak_and_play(conn, llm_result) + elif ( + result.action == Action.NOTFOUND + or result.action == Action.ERROR + ): + text = result.result + if text is not None: + speak_and_play(conn, text) + elif function_name != "play_music": + # For backward compatibility with original code + # 获取当前最新的文本索引 + text = result.response + if text is None: + text = result.result + if text is not None: + speak_and_play(conn, text) + # 将函数执行放在线程池中 conn.executor.submit(process_function_call) return True return False except json.JSONDecodeError as e: - logger.bind(tag=TAG).error(f"处理意图结果时出错: {e}") + conn.logger.bind(tag=TAG).error(f"处理意图结果时出错: {e}") return False -def extract_text_in_brackets(s): - """ - 从字符串中提取中括号内的文字 - - :param s: 输入字符串 - :return: 中括号内的文字,如果不存在则返回空字符串 - """ - left_bracket_index = s.find("[") - right_bracket_index = s.find("]") - - if ( - left_bracket_index != -1 - and right_bracket_index != -1 - and left_bracket_index < right_bracket_index - ): - return s[left_bracket_index + 1 : right_bracket_index] - else: - return "" +def speak_and_play(conn, text): + conn.tts.tts_one_sentence(conn, text) + conn.dialogue.put(Message(role="assistant", content=text))