update:合并main分支

This commit is contained in:
hrz
2025-05-21 15:55:40 +08:00
parent c4c84e44e1
commit c900498ce8
2 changed files with 450 additions and 169 deletions
+406 -138
View File
@@ -1,4 +1,8 @@
import os
import copy
import json import json
import subprocess
import sys
import uuid import uuid
import time import time
import queue import queue
@@ -16,17 +20,22 @@ from core.handle.textHandle import handleTextMessage
from core.utils.util import ( from core.utils.util import (
get_string_no_punctuation_or_emoji, get_string_no_punctuation_or_emoji,
extract_json_from_string, 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 concurrent.futures import ThreadPoolExecutor, TimeoutError
from core.handle.sendAudioHandle import sendAudioMessage from core.handle.sendAudioHandle import sendAudioMessage
from core.handle.receiveAudioHandle import handleAudioMessage from core.handle.receiveAudioHandle import handleAudioMessage
from core.handle.functionHandler import FunctionHandler from core.handle.functionHandler import FunctionHandler
from plugins_func.register import Action, ActionResponse from plugins_func.register import Action, ActionResponse
from config.private_config import PrivateConfig
from core.auth import AuthMiddleware, AuthenticationError from core.auth import AuthMiddleware, AuthenticationError
from core.utils.auth_code_gen import AuthCodeGenerator
from core.mcp.manager import MCPManager 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__ TAG = __name__
@@ -39,22 +48,39 @@ class TTSException(RuntimeError):
class ConnectionHandler: class ConnectionHandler:
def __init__( 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.logger = setup_logging()
self.server = server # 保存server实例的引用
self.auth = AuthMiddleware(config) 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.tts_stream = self.config.get("TTS_SET", {}).get("TTS_STREAM", False)
self.websocket = None self.websocket = None
self.headers = None self.headers = None
self.device_id = None
self.client_ip = None self.client_ip = None
self.client_ip_info = {} self.client_ip_info = {}
self.session_id = None
self.prompt = None self.prompt = None
self.welcome_msg = None self.welcome_msg = None
self.u_id = None self.u_id = None
self.max_output_size = 0
self.chat_history_conf = 0
# 客户端状态相关 # 客户端状态相关
self.client_abort = False self.client_abort = False
@@ -70,9 +96,18 @@ class ConnectionHandler:
self.executor = ThreadPoolExecutor(max_workers=max_workers) self.executor = ThreadPoolExecutor(max_workers=max_workers)
self.start_tts_request_flag = False 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.vad = None
self.asr = _asr self.asr = None
self._asr = _asr
self._vad = _vad
self.llm = _llm self.llm = _llm
self.tts = _tts self.tts = _tts
self.memory = _memory self.memory = _memory
@@ -102,24 +137,47 @@ class ConnectionHandler:
self.iot_descriptors = {} self.iot_descriptors = {}
self.func_handler = None self.func_handler = None
self.cmd_exit = self.config["CMD_exit"] self.cmd_exit = self.config["exit_commands"]
self.max_cmd_length = 0 self.max_cmd_length = 0
for cmd in self.cmd_exit: for cmd in self.cmd_exit:
if len(cmd) > self.max_cmd_length: if len(cmd) > self.max_cmd_length:
self.max_cmd_length = len(cmd) self.max_cmd_length = len(cmd)
self.private_config = None # 是否在聊天结束后关闭连接
self.auth_code_gen = AuthCodeGenerator.get_instance() self.close_after_chat = False
self.is_device_verified = False # 添加设备验证状态标志 self.load_function_plugin = False
self.close_after_chat = False # 是否在聊天结束后关闭连接 self.intent_type = "nointent"
self.use_function_call_mode = False
if self.config["selected_module"]["Intent"] == "function_call": self.timeout_task = None
self.use_function_call_mode = True 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): async def handle_connection(self, ws):
try: try:
# 获取并验证headers # 获取并验证headers
self.headers = dict(ws.request.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地址 # 获取客户端ip地址
self.client_ip = ws.remote_address[0] self.client_ip = ws.remote_address[0]
self.logger.bind(tag=TAG).info( self.logger.bind(tag=TAG).info(
@@ -128,49 +186,20 @@ class ConnectionHandler:
# 进行认证 # 进行认证
await self.auth.authenticate(self.headers) await self.auth.authenticate(self.headers)
device_id = self.headers.get("device-id", None)
# 认证通过,继续处理 # 认证通过,继续处理
self.websocket = ws 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 = self.config["xiaozhi"]
self.welcome_msg["session_id"] = self.session_id self.welcome_msg["session_id"] = self.session_id
await self.websocket.send(json.dumps(self.welcome_msg)) 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) self.executor.submit(self._initialize_components)
@@ -202,55 +231,254 @@ class ConnectionHandler:
async def _save_and_close(self, ws): async def _save_and_close(self, ws):
"""保存记忆并关闭连接""" """保存记忆并关闭连接"""
try: 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: except Exception as e:
self.logger.bind(tag=TAG).error(f"保存记忆失败: {e}") self.logger.bind(tag=TAG).error(f"保存记忆失败: {e}")
finally: finally:
# 立即关闭连接,不等待记忆保存完成
await self.close(ws) 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): async def _route_message(self, message):
"""消息路由""" """消息路由"""
# 重置超时计时器
await self.reset_timeout()
if isinstance(message, str): if isinstance(message, str):
await handleTextMessage(self, message) await handleTextMessage(self, message)
elif isinstance(message, bytes): elif isinstance(message, bytes):
await handleAudioMessage(self, message) await handleAudioMessage(self, message)
def _initialize_components(self): async def handle_restart(self, message):
"""加载提示词""" """处理服务器重启请求"""
self.prompt = self.config["prompt"] try:
if self.private_config:
self.prompt = self.private_config.private_config.get("prompt", self.prompt)
self.dialogue.put(Message(role="system", content=self.prompt))
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_memory()
"""加载意图识别""" """加载意图识别"""
self._initialize_intent() self._initialize_intent()
"""加载位置信息""" """初始化上报线程"""
self.client_ip_info = get_ip_info(self.client_ip) self._init_report_threads()
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.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): def _initialize_memory(self):
"""初始化记忆模块""" """初始化记忆模块"""
device_id = self.headers.get("device-id", None) self.memory.init_memory(
self.memory.init_memory(device_id, self.llm) 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): 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_config = self.config["Intent"]
intent_type = self.config["selected_module"]["Intent"] intent_type = self.config["Intent"][self.config["selected_module"]["Intent"]][
"type"
]
# 如果使用 nointent,直接返回 # 如果使用 nointent,直接返回
if intent_type == "nointent": if intent_type == "nointent":
return return
# 使用 intent_llm 模式 # 使用 intent_llm 模式
elif intent_type == "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"]: if intent_llm_name and intent_llm_name in self.config["LLM"]:
# 如果配置了专用LLM,则创建独立的LLM实例 # 如果配置了专用LLM,则创建独立的LLM实例
@@ -281,47 +509,22 @@ class ConnectionHandler:
def change_system_prompt(self, prompt): def change_system_prompt(self, prompt):
self.prompt = prompt self.prompt = prompt
# 找到原来的role==system,替换原来的系统提示 # 更新系统prompt至上下文
for m in self.dialogue.dialogue: self.dialogue.update_system_message(self.prompt)
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
def chat(self, query): 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)) self.dialogue.put(Message(role="user", content=query))
response_message = [] response_message = []
try: try:
start_time = time.time()
# 使用带记忆的对话 # 使用带记忆的对话
future = asyncio.run_coroutine_threadsafe( memory_str = None
self.memory.query_memory(query), self.loop if self.memory is not None:
) future = asyncio.run_coroutine_threadsafe(
memory_str = future.result() self.memory.query_memory(query), self.loop
)
memory_str = future.result()
self.logger.bind(tag=TAG).debug(f"记忆内容: {memory_str}") self.logger.bind(tag=TAG).debug(f"记忆内容: {memory_str}")
llm_responses = self.llm.response( llm_responses = self.llm.response(
@@ -376,13 +579,6 @@ class ConnectionHandler:
def chat_with_function_calling(self, query, tool_call=False): def chat_with_function_calling(self, query, tool_call=False):
self.logger.bind(tag=TAG).debug(f"Chat with function calling start: {query}") self.logger.bind(tag=TAG).debug(f"Chat with function calling start: {query}")
"""Chat with function calling for intent detection using streaming""" """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: if not tool_call:
self.dialogue.put(Message(role="user", content=query)) self.dialogue.put(Message(role="user", content=query))
@@ -397,10 +593,12 @@ class ConnectionHandler:
start_time = time.time() start_time = time.time()
# 使用带记忆的对话 # 使用带记忆的对话
future = asyncio.run_coroutine_threadsafe( memory_str = None
self.memory.query_memory(query), self.loop if self.memory is not None:
) future = asyncio.run_coroutine_threadsafe(
memory_str = future.result() 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)}") # 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 self.u_id = uuid_str
for response in llm_responses: for response in llm_responses:
content, tools_call = response content, tools_call = response
if "content" in response: if "content" in response:
content = response["content"] content = response["content"]
tools_call = None tools_call = None
if content is not None and len(content) > 0: if content is not None and len(content) > 0:
if len(response_message) <= 0 and ( content_arguments += content
content == "```" or "<tool_call>" in content
): if not tool_call_flag and content_arguments.startswith("<tool_call>"):
tool_call_flag = True # print("content_arguments", content_arguments)
tool_call_flag = True
if tools_call is not None: if tools_call is not None:
tool_call_flag = True tool_call_flag = True
@@ -446,9 +646,7 @@ class ConnectionHandler:
function_arguments += tools_call[0].function.arguments function_arguments += tools_call[0].function.arguments
if content is not None and len(content) > 0: if content is not None and len(content) > 0:
if tool_call_flag: if not tool_call_flag:
content_arguments += content
else:
response_message.append(content) response_message.append(content)
if self.client_abort: if self.client_abort:
@@ -507,10 +705,9 @@ class ConnectionHandler:
self.logger.bind(tag=TAG).error( self.logger.bind(tag=TAG).error(
f"function call error: {content_arguments}" f"function call error: {content_arguments}"
) )
else:
function_arguments = json.loads(function_arguments)
if not bHasError: 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}" f"function_name={function_name}, function_id={function_id}, function_arguments={function_arguments}"
) )
function_call_data = { function_call_data = {
@@ -614,10 +811,16 @@ class ConnectionHandler:
) )
self.dialogue.put( 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) 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 text = result.result
self.recode_first_last_text(text, text_index) self.recode_first_last_text(text, text_index)
self.tts.tts_one_sentence(self, text) self.tts.tts_one_sentence(self, text)
@@ -646,6 +849,45 @@ class ConnectionHandler:
f"audio_play_priority priority_thread: {text} {e}" 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): def clearSpeakStatus(self):
self.logger.bind(tag=TAG).debug(f"清除服务端讲话状态") self.logger.bind(tag=TAG).debug(f"清除服务端讲话状态")
self.asr_server_receive = True self.asr_server_receive = True
@@ -660,42 +902,56 @@ class ConnectionHandler:
async def close(self, ws=None): async def close(self, ws=None):
"""资源清理方法""" """资源清理方法"""
# 取消超时任务
if self.timeout_task:
self.timeout_task.cancel()
self.timeout_task = None
# 清理MCP资源 # 清理MCP资源
if hasattr(self, "mcp_manager") and self.mcp_manager: if hasattr(self, "mcp_manager") and self.mcp_manager:
await self.mcp_manager.cleanup_all() await self.mcp_manager.cleanup_all()
# 触发停止事件并清理资源 # 触发停止事件
if self.stop_event: if self.stop_event:
self.stop_event.set() 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: if ws:
await ws.close() await ws.close()
elif self.websocket: elif self.websocket:
await self.websocket.close() await self.websocket.close()
await self.tts.close() await self.tts.close()
# 最后关闭线程池(避免阻塞)
if self.executor:
self.executor.shutdown(wait=False)
self.executor = None
self.logger.bind(tag=TAG).info("连接资源已释放") 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]: for q in [self.tts_queue, self.audio_play_queue]:
if not q: if not q:
continue continue
while not q.empty(): while True:
try: try:
q.get_nowait() q.get_nowait()
except queue.Empty: except queue.Empty:
continue break
q.queue.clear()
# 添加毒丸信号到队列,确保线程退出 self.logger.bind(tag=TAG).debug(
# q.queue.put(None) f"清理结束: TTS队列大小={self.tts_queue.qsize()}, 音频队列大小={self.audio_play_queue.qsize()}"
)
def reset_vad_states(self): def reset_vad_states(self):
self.client_audio_buffer = bytearray() self.client_audio_buffer = bytearray()
@@ -714,3 +970,15 @@ class ConnectionHandler:
self.close_after_chat = True self.close_after_chat = True
except Exception as e: except Exception as e:
self.logger.bind(tag=TAG).error(f"Chat and close error: {str(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}")
@@ -4,10 +4,10 @@ import uuid
from core.handle.sendAudioHandle import send_stt_message from core.handle.sendAudioHandle import send_stt_message
from core.utils.util import remove_punctuation_and_length from core.utils.util import remove_punctuation_and_length
from core.utils.dialogue import Message from core.utils.dialogue import Message
from plugins_func.register import Action
from loguru import logger from loguru import logger
TAG = __name__ TAG = __name__
logger = setup_logging()
async def handle_user_intent(conn, text): async def handle_user_intent(conn, text):
@@ -19,7 +19,7 @@ async def handle_user_intent(conn, text):
# if await checkWakeupWords(conn, text): # if await checkWakeupWords(conn, text):
# return True # return True
if conn.use_function_call_mode: if conn.intent_type == "function_call":
# 使用支持function calling的聊天方法,不再进行意图分析 # 使用支持function calling的聊天方法,不再进行意图分析
return False return False
# 使用LLM进行意图分析 # 使用LLM进行意图分析
@@ -36,7 +36,7 @@ async def check_direct_exit(conn, text):
cmd_exit = conn.cmd_exit cmd_exit = conn.cmd_exit
for cmd in cmd_exit: for cmd in cmd_exit:
if text == cmd: 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 send_stt_message(conn, text)
await conn.close() await conn.close()
return True return True
@@ -46,7 +46,7 @@ async def check_direct_exit(conn, text):
async def analyze_intent_with_llm(conn, text): async def analyze_intent_with_llm(conn, text):
"""使用LLM分析用户意图""" """使用LLM分析用户意图"""
if not hasattr(conn, "intent") or not conn.intent: if not hasattr(conn, "intent") or not conn.intent:
logger.bind(tag=TAG).warning("意图识别服务未初始化") conn.logger.bind(tag=TAG).warning("意图识别服务未初始化")
return None 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) intent_result = await conn.intent.detect_intent(conn, dialogue.dialogue, text)
return intent_result return intent_result
except Exception as e: except Exception as e:
logger.bind(tag=TAG).error(f"意图识别失败: {str(e)}") conn.logger.bind(tag=TAG).error(f"意图识别失败: {str(e)}")
return None return None
@@ -69,13 +69,18 @@ async def process_intent_result(conn, intent_result, original_text):
# 检查是否有function_call # 检查是否有function_call
if "function_call" in intent_data: if "function_call" in intent_data:
# 直接从意图识别获取了function_call # 直接从意图识别获取了function_call
logger.bind(tag=TAG).debug( conn.logger.bind(tag=TAG).debug(
f"检测到function_call格式的意图结果: {intent_data['function_call']['name']}" f"检测到function_call格式的意图结果: {intent_data['function_call']['name']}"
) )
function_name = intent_data["function_call"]["name"] function_name = intent_data["function_call"]["name"]
if function_name == "continue_chat": if function_name == "continue_chat":
return False 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 function_args = None
if "arguments" in intent_data["function_call"]: if "arguments" in intent_data["function_call"]:
function_args = intent_data["function_call"]["arguments"] 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( result = conn.func_handler.handle_llm_function_call(
conn, function_call_data conn, function_call_data
) )
if result and function_name != "play_music": logger.bind(tag=TAG).debug(f"检测到Action : {result.action}")
# 获取当前最新的文本索引
text = result.response if result:
if text is None: 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 text = result.result
if text is not None: conn.dialogue.put(Message(role="tool", content=text))
conn.tts.tts_one_sentence(conn, 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) conn.executor.submit(process_function_call)
return True return True
return False return False
except json.JSONDecodeError as e: except json.JSONDecodeError as e:
logger.bind(tag=TAG).error(f"处理意图结果时出错: {e}") conn.logger.bind(tag=TAG).error(f"处理意图结果时出错: {e}")
return False return False
def extract_text_in_brackets(s): def speak_and_play(conn, text):
""" conn.tts.tts_one_sentence(conn, text)
从字符串中提取中括号内的文字 conn.dialogue.put(Message(role="assistant", content=text))
: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 ""