mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-30 05:13:59 +08:00
update:合并main分支
This commit is contained in:
@@ -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 ""
|
|
||||||
|
|||||||
Reference in New Issue
Block a user