- {{ isAllSelected ? '取消全选' : '全选' }}
+ {{ isCurrentPageAllSelected ? '取消全选' : '全选' }}
验证码绑定
@@ -128,8 +128,6 @@ export default {
return {
addDeviceDialogVisible: false,
manualAddDeviceDialogVisible: false,
- selectedDevices: [],
- isAllSelected: false,
searchKeyword: "",
activeSearchKeyword: "",
currentAgentId: this.$route.query.agentId || '',
@@ -160,6 +158,11 @@ export default {
pageCount() {
return Math.ceil(this.filteredDeviceList.length / this.pageSize);
},
+ // 计算当前页是否全选
+ isCurrentPageAllSelected() {
+ return this.paginatedDeviceList.length > 0 &&
+ this.paginatedDeviceList.every(device => device.selected);
+ },
visiblePages() {
const pages = [];
const maxVisible = 3;
@@ -205,16 +208,15 @@ export default {
},
handleSelectAll() {
- this.isAllSelected = !this.isAllSelected;
+ const shouldSelectAll = !this.isCurrentPageAllSelected;
this.paginatedDeviceList.forEach(row => {
- row.selected = this.isAllSelected;
+ row.selected = shouldSelectAll;
});
- this.selectedDevices = this.paginatedDeviceList.filter(device => device.selected);
},
deleteSelected() {
- this.selectedDevices = this.paginatedDeviceList.filter(device => device.selected);
- if (this.selectedDevices.length === 0) {
+ const selectedDevices = this.paginatedDeviceList.filter(device => device.selected);
+ if (selectedDevices.length === 0) {
this.$message.warning({
message: '请至少选择一条记录',
showClose: true
@@ -222,12 +224,12 @@ export default {
return;
}
- this.$confirm(`确认要解绑选中的 ${this.selectedDevices.length} 台设备吗?`, '警告', {
+ this.$confirm(`确认要解绑选中的 ${selectedDevices.length} 台设备吗?`, '警告', {
confirmButtonText: '确定',
cancelButtonText: '取消',
type: 'warning'
}).then(() => {
- const deviceIds = this.selectedDevices.map(device => device.device_id);
+ const deviceIds = selectedDevices.map(device => device.device_id);
this.batchUnbindDevices(deviceIds);
});
},
@@ -250,8 +252,6 @@ export default {
showClose: true
});
this.fetchBindDevices(this.currentAgentId);
- this.selectedDevices = [];
- this.isAllSelected = false;
})
.catch(error => {
this.$message.error({
@@ -355,7 +355,8 @@ export default {
isEdit: false,
_submitting: false,
otaSwitch: device.autoUpdate === 1,
- rawBindTime: new Date(device.createDate).getTime()
+ rawBindTime: new Date(device.createDate).getTime(),
+ selected: false
};
})
.sort((a, b) => a.rawBindTime - b.rawBindTime);
diff --git a/main/manager-web/src/views/ModelConfig.vue b/main/manager-web/src/views/ModelConfig.vue
index bdc92dc0..e945e015 100644
--- a/main/manager-web/src/views/ModelConfig.vue
+++ b/main/manager-web/src/views/ModelConfig.vue
@@ -79,11 +79,14 @@
-
+
修改
+
+ 创建副本
+
删除
@@ -277,6 +280,11 @@ export default {
this.editModelData = JSON.parse(JSON.stringify(model));
this.editDialogVisible = true;
},
+ duplicateModel(model) {
+ this.editModelData = JSON.parse(JSON.stringify(model));
+ this.editModelData.duplicateMode = true;
+ this.editDialogVisible = true;
+ },
// 删除单个模型
deleteModel(model) {
this.$confirm('确定要删除该模型吗?', '提示', {
@@ -313,19 +321,34 @@ export default {
const modelType = this.activeTab;
const id = formData.id;
- Api.model.updateModel(
- { modelType, provideCode, id, formData },
+ if (this.editModelData.duplicateMode) {
+ Api.model.addModel({modelType, provideCode, formData},
({ data }) => {
if (data.code === 0) {
- this.$message.success('保存成功');
+ this.$message.success('创建副本成功');
this.loadData();
this.editDialogVisible = false;
} else {
- this.$message.error(data.msg || '保存失败');
+ this.$message.error(data.msg || '创建副本失败');
}
done && done(); // 调用done回调关闭加载状态
- }
- );
+ })
+ }
+ else {
+ Api.model.updateModel(
+ { modelType, provideCode, id, formData },
+ ({ data }) => {
+ if (data.code === 0) {
+ this.$message.success('保存成功');
+ this.loadData();
+ this.editDialogVisible = false;
+ } else {
+ this.$message.error(data.msg || '保存失败');
+ }
+ done && done(); // 调用done回调关闭加载状态
+ }
+ );
+ }
},
selectAll() {
if (this.isAllSelected) {
diff --git a/main/xiaozhi-server/config.yaml b/main/xiaozhi-server/config.yaml
index f6275306..314573c4 100644
--- a/main/xiaozhi-server/config.yaml
+++ b/main/xiaozhi-server/config.yaml
@@ -61,6 +61,10 @@ delete_audio: true
close_connection_no_voice_time: 120
# TTS请求超时时间(秒)
tts_timeout: 10
+# 开启唤醒词加速
+enable_wakeup_words_response_cache: true
+# 开场是否回复唤醒词
+enable_greeting: true
# 说完话是否开启提示音
enable_stop_tts_notify: false
# 说完话是否开启提示音,音效地址
@@ -909,7 +913,7 @@ TTS:
sample_rate: 24000 # 采样率 [websocket默认24000,http默认0 自动选择]
speed: 1.0 # 语速,1.0 表示正常语速,>1 表示加快,<1 表示减慢
volume: 1.0 # 音量,1.0 表示正常音量,>1 表示增大,<1 表示减小
- save_path: ./streaming_tts.wav # 服务器生成的语音文件保存路径
+ save_path: # 保存路径
IndexStreamTTS:
# 基于Index-TTS-vLLM项目的TTS接口服务
# 参照教程:https://github.com/Ksuriuri/index-tts-vllm/blob/master/README.md
diff --git a/main/xiaozhi-server/config/logger.py b/main/xiaozhi-server/config/logger.py
index cd4d9965..5c916ba5 100644
--- a/main/xiaozhi-server/config/logger.py
+++ b/main/xiaozhi-server/config/logger.py
@@ -5,7 +5,7 @@ from config.config_loader import load_config
from config.settings import check_config_file
from datetime import datetime
-SERVER_VERSION = "0.7.5"
+SERVER_VERSION = "0.7.7"
_logger_initialized = False
diff --git a/main/xiaozhi-server/core/connection.py b/main/xiaozhi-server/core/connection.py
index 2ca2602a..4b3b9d1c 100644
--- a/main/xiaozhi-server/core/connection.py
+++ b/main/xiaozhi-server/core/connection.py
@@ -140,10 +140,6 @@ class ConnectionHandler:
self.func_handler = None
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.close_after_chat = False
diff --git a/main/xiaozhi-server/core/handle/helloHandle.py b/main/xiaozhi-server/core/handle/helloHandle.py
index de9587b5..75b9fcb2 100644
--- a/main/xiaozhi-server/core/handle/helloHandle.py
+++ b/main/xiaozhi-server/core/handle/helloHandle.py
@@ -1,5 +1,13 @@
+import time
import json
+import random
import asyncio
+from core.utils.dialogue import Message
+from core.utils.util import audio_to_data
+from core.providers.tts.dto.dto import SentenceType
+from core.utils.wakeup_word import WakeupWordsConfig
+from core.handle.sendAudioHandle import sendAudioMessage, send_stt_message
+from core.utils.util import remove_punctuation_and_length, opus_datas_to_wav_bytes
from core.providers.tools.device_mcp import (
MCPClient,
send_mcp_initialize_message,
@@ -8,6 +16,17 @@ from core.providers.tools.device_mcp import (
TAG = __name__
+WAKEUP_CONFIG = {
+ "refresh_time": 5,
+ "words": ["你好", "你好啊", "嘿,你好", "嗨"],
+}
+
+# 创建全局的唤醒词配置管理器
+wakeup_words_config = WakeupWordsConfig()
+
+# 用于防止并发调用wakeupWordsResponse的锁
+_wakeup_response_lock = asyncio.Lock()
+
async def handleHelloMessage(conn, msg_json):
"""处理hello消息"""
@@ -30,3 +49,103 @@ async def handleHelloMessage(conn, msg_json):
asyncio.create_task(send_mcp_tools_list_request(conn))
await conn.websocket.send(json.dumps(conn.welcome_msg))
+
+
+async def checkWakeupWords(conn, text):
+ enable_wakeup_words_response_cache = conn.config[
+ "enable_wakeup_words_response_cache"
+ ]
+
+ # 等待tts初始化,最多等待3秒
+ start_time = time.time()
+ while time.time() - start_time < 3:
+ if conn.tts:
+ break
+ await asyncio.sleep(0.1)
+ else:
+ return False
+
+ if not enable_wakeup_words_response_cache:
+ return False
+
+ _, filtered_text = remove_punctuation_and_length(text)
+ if filtered_text not in conn.config.get("wakeup_words"):
+ return False
+
+ conn.just_woken_up = True
+ await send_stt_message(conn, text)
+
+ # 获取当前音色
+ voice = getattr(conn.tts, "voice", "default")
+ if not voice:
+ voice = "default"
+
+ # 获取唤醒词回复配置
+ response = wakeup_words_config.get_wakeup_response(voice)
+ if not response or not response.get("file_path"):
+ response = {
+ "voice": "default",
+ "file_path": "config/assets/wakeup_words.wav",
+ "time": 0,
+ "text": "哈啰啊,我是小智啦,声音好听的台湾女孩一枚,超开心认识你耶,最近在忙啥,别忘了给我来点有趣的料哦,我超爱听八卦的啦",
+ }
+
+ # 获取音频数据
+ opus_packets = audio_to_data(response.get("file_path"))
+ # 播放唤醒词回复
+ conn.client_abort = False
+
+ conn.logger.bind(tag=TAG).info(f"播放唤醒词回复: {response.get('text')}")
+ await sendAudioMessage(conn, SentenceType.FIRST, opus_packets, response.get("text"))
+ await sendAudioMessage(conn, SentenceType.LAST, [], None)
+
+ # 补充对话
+ conn.dialogue.put(Message(role="assistant", content=response.get("text")))
+
+ # 检查是否需要更新唤醒词回复
+ if time.time() - response.get("time", 0) > WAKEUP_CONFIG["refresh_time"]:
+ if not _wakeup_response_lock.locked():
+ asyncio.create_task(wakeupWordsResponse(conn))
+ return True
+
+
+async def wakeupWordsResponse(conn):
+ if not conn.tts or not conn.llm or not conn.llm.response_no_stream:
+ return
+
+ try:
+ # 尝试获取锁,如果获取不到就返回
+ if not await _wakeup_response_lock.acquire():
+ return
+
+ # 生成唤醒词回复
+ wakeup_word = random.choice(WAKEUP_CONFIG["words"])
+ question = (
+ "此刻用户正在和你说```"
+ + wakeup_word
+ + "```。\n请你根据以上用户的内容进行20-30字回复。要符合系统设置的角色情感和态度,不要像机器人一样说话。\n"
+ + "请勿对这条内容本身进行任何解释和回应,请勿返回表情符号,仅返回对用户的内容的回复。"
+ )
+
+ result = conn.llm.response_no_stream(conn.config["prompt"], question)
+ if not result or len(result) == 0:
+ return
+
+ # 生成TTS音频
+ tts_result = await asyncio.to_thread(conn.tts.to_tts, result)
+ if not tts_result:
+ return
+
+ # 获取当前音色
+ voice = getattr(conn.tts, "voice", "default")
+
+ wav_bytes = opus_datas_to_wav_bytes(tts_result, sample_rate=16000)
+ file_path = wakeup_words_config.generate_file_path(voice)
+ with open(file_path, "wb") as f:
+ f.write(wav_bytes)
+ # 更新配置
+ wakeup_words_config.update_wakeup_response(voice, file_path, result)
+ finally:
+ # 确保在任何情况下都释放锁
+ if _wakeup_response_lock.locked():
+ _wakeup_response_lock.release()
\ No newline at end of file
diff --git a/main/xiaozhi-server/core/handle/intentHandler.py b/main/xiaozhi-server/core/handle/intentHandler.py
index 58b2cee7..72424968 100644
--- a/main/xiaozhi-server/core/handle/intentHandler.py
+++ b/main/xiaozhi-server/core/handle/intentHandler.py
@@ -1,11 +1,12 @@
import json
-import asyncio
import uuid
+import asyncio
+from core.utils.dialogue import Message
+from core.providers.tts.dto.dto import ContentType
+from core.handle.helloHandle import checkWakeupWords
+from plugins_func.register import Action, ActionResponse
from core.handle.sendAudioHandle import send_stt_message
from core.utils.util import remove_punctuation_and_length
-from core.providers.tts.dto.dto import ContentType
-from core.utils.dialogue import Message
-from plugins_func.register import Action, ActionResponse
from core.providers.tts.dto.dto import TTSMessageDTO, SentenceType
TAG = __name__
@@ -23,10 +24,14 @@ async def handle_user_intent(conn, text):
pass
# 检查是否有明确的退出命令
- filtered_text = remove_punctuation_and_length(text)[1]
+ _, filtered_text = remove_punctuation_and_length(text)
if await check_direct_exit(conn, filtered_text):
return True
+ # 检查是否是唤醒词
+ if await checkWakeupWords(conn, filtered_text):
+ return True
+
if conn.intent_type == "function_call":
# 使用支持function calling的聊天方法,不再进行意图分析
return False
diff --git a/main/xiaozhi-server/core/handle/receiveAudioHandle.py b/main/xiaozhi-server/core/handle/receiveAudioHandle.py
index 7d46685b..8db50633 100644
--- a/main/xiaozhi-server/core/handle/receiveAudioHandle.py
+++ b/main/xiaozhi-server/core/handle/receiveAudioHandle.py
@@ -1,12 +1,11 @@
import time
-import asyncio
import json
-from core.handle.sendAudioHandle import send_stt_message
+import asyncio
+from core.utils.util import audio_to_data
+from core.handle.abortHandle import handleAbortMessage
from core.handle.intentHandler import handle_user_intent
from core.utils.output_counter import check_device_output_limit
-from core.handle.abortHandle import handleAbortMessage
-from core.handle.sendAudioHandle import SentenceType
-from core.utils.util import audio_to_data_stream
+from core.handle.sendAudioHandle import send_stt_message, SentenceType
TAG = __name__
@@ -22,7 +21,6 @@ async def handleAudioMessage(conn, audio):
if not hasattr(conn, "vad_resume_task") or conn.vad_resume_task.done():
conn.vad_resume_task = asyncio.create_task(resume_vad_detection(conn))
return
-
if have_voice:
if conn.client_is_speaking:
await handleAbortMessage(conn)
@@ -31,13 +29,11 @@ async def handleAudioMessage(conn, audio):
# 接收音频
await conn.asr.receive_audio(conn, audio, have_voice)
-
async def resume_vad_detection(conn):
# 等待2秒后恢复VAD检测
await asyncio.sleep(1)
conn.just_woken_up = False
-
async def startToChat(conn, text):
# 检查输入是否是JSON格式(包含说话人信息)
speaker_name = None
@@ -118,12 +114,13 @@ async def no_voice_close_connect(conn, have_voice):
async def max_out_size(conn):
+ # 播放超出最大输出字数的提示
+ conn.client_abort = False
text = "不好意思,我现在有点事情要忙,明天这个时候我们再聊,约好了哦!明天不见不散,拜拜!"
await send_stt_message(conn, text)
file_path = "config/assets/max_output_size.wav"
- conn.tts.tts_audio_queue.put((SentenceType.FIRST, [], text))
- play_audio_frames(conn, file_path)
- conn.tts.tts_audio_queue.put((SentenceType.LAST, [], None))
+ opus_packets = audio_to_data(file_path)
+ conn.tts.tts_audio_queue.put((SentenceType.LAST, opus_packets, text))
conn.close_after_chat = True
@@ -141,35 +138,25 @@ async def check_bind_device(conn):
# 播放提示音
music_path = "config/assets/bind_code.wav"
- conn.tts.tts_audio_queue.put((SentenceType.FIRST, [], text))
- play_audio_frames(conn, music_path)
+ opus_packets = audio_to_data(music_path)
+ conn.tts.tts_audio_queue.put((SentenceType.FIRST, opus_packets, text))
# 逐个播放数字
for i in range(6): # 确保只播放6位数字
try:
digit = conn.bind_code[i]
num_path = f"config/assets/bind_code/{digit}.wav"
- play_audio_frames(conn, num_path)
+ num_packets = audio_to_data(num_path)
+ conn.tts.tts_audio_queue.put((SentenceType.MIDDLE, num_packets, None))
except Exception as e:
conn.logger.bind(tag=TAG).error(f"播放数字音频失败: {e}")
continue
conn.tts.tts_audio_queue.put((SentenceType.LAST, [], None))
else:
+ # 播放未绑定提示
+ conn.client_abort = False
text = f"没有找到该设备的版本信息,请正确配置 OTA地址,然后重新编译固件。"
await send_stt_message(conn, text)
music_path = "config/assets/bind_not_found.wav"
- conn.tts.tts_audio_queue.put((SentenceType.FIRST, [], text))
- play_audio_frames(conn, music_path)
- conn.tts.tts_audio_queue.put((SentenceType.LAST, [], None))
-
-
-def play_audio_frames(conn, file_path):
- """播放音频文件并处理发送帧数据"""
- def handle_audio_frame(frame_data):
- conn.tts.tts_audio_queue.put((SentenceType.MIDDLE, frame_data, None))
-
- audio_to_data_stream(
- file_path,
- is_opus=True,
- callback=handle_audio_frame
- )
+ opus_packets = audio_to_data(music_path)
+ conn.tts.tts_audio_queue.put((SentenceType.LAST, opus_packets, text))
diff --git a/main/xiaozhi-server/core/handle/sendAudioHandle.py b/main/xiaozhi-server/core/handle/sendAudioHandle.py
index c3c50fbc..661c4e21 100644
--- a/main/xiaozhi-server/core/handle/sendAudioHandle.py
+++ b/main/xiaozhi-server/core/handle/sendAudioHandle.py
@@ -1,6 +1,9 @@
import json
-from core.providers.tts.dto.dto import SentenceType
+import time
+import asyncio
from core.utils import textUtils
+from core.utils.util import audio_to_data
+from core.providers.tts.dto.dto import SentenceType
TAG = __name__
@@ -28,13 +31,78 @@ async def sendAudioMessage(conn, sentenceType, audios, text):
# 播放音频
-async def sendAudio(conn, audios):
- if audios is None:
+async def sendAudio(conn, audios, frame_duration=60):
+ """
+ 发送单个opus包,支持流控
+ Args:
+ conn: 连接对象
+ opus_packet: 单个opus数据包
+ pre_buffer: 快速发送音频
+ frame_duration: 帧时长(毫秒),匹配 Opus 编码
+ """
+ if audios is None or len(audios) == 0:
return
- # 如果audios不是opus数组,则不需要进行遍历,可以直接发送;这里需要进行流控管理,防止发送过快引发客户端溢出
+
if isinstance(audios, bytes):
+ if conn.client_abort:
+ return
+
+ conn.last_activity_time = time.time() * 1000
+
+ # 获取或初始化流控状态
+ if not hasattr(conn, "audio_flow_control"):
+ conn.audio_flow_control = {
+ "last_send_time": 0,
+ "packet_count": 0,
+ "start_time": time.perf_counter(),
+ }
+
+ flow_control = conn.audio_flow_control
+ current_time = time.perf_counter()
+ # 计算预期发送时间
+ expected_time = flow_control["start_time"] + (
+ flow_control["packet_count"] * frame_duration / 1000
+ )
+ delay = expected_time - current_time
+ if delay > 0:
+ await asyncio.sleep(delay)
+
+ # 发送数据包
await conn.websocket.send(audios)
+ # 更新流控状态
+ flow_control["packet_count"] += 1
+ flow_control["last_send_time"] = time.perf_counter()
+ else:
+ # 文件型音频走普通播放
+ start_time = time.perf_counter()
+ play_position = 0
+
+ # 执行预缓冲
+ pre_buffer_frames = min(3, len(audios))
+ for i in range(pre_buffer_frames):
+ await conn.websocket.send(audios[i])
+ remaining_audios = audios[pre_buffer_frames:]
+
+ # 播放剩余音频帧
+ for opus_packet in remaining_audios:
+ if conn.client_abort:
+ break
+
+ # 重置没有声音的状态
+ conn.last_activity_time = time.time() * 1000
+
+ # 计算预期发送时间
+ expected_time = start_time + (play_position / 1000)
+ current_time = time.perf_counter()
+ delay = expected_time - current_time
+ if delay > 0:
+ await asyncio.sleep(delay)
+
+ await conn.websocket.send(opus_packet)
+
+ play_position += frame_duration
+
async def send_tts_message(conn, state, text=None):
"""发送 TTS 状态消息"""
@@ -52,7 +120,7 @@ async def send_tts_message(conn, state, text=None):
stop_tts_notify_voice = conn.config.get(
"stop_tts_notify_voice", "config/assets/tts_notify.mp3"
)
- audios, _ = conn.tts.audio_to_opus_data(stop_tts_notify_voice)
+ audios = audio_to_data(stop_tts_notify_voice, is_opus=True)
await sendAudio(conn, audios)
# 清除服务端讲话状态
conn.clearSpeakStatus()
@@ -72,7 +140,7 @@ async def send_stt_message(conn, text):
display_text = text
try:
# 尝试解析JSON格式
- if text.strip().startswith('{') and text.strip().endswith('}'):
+ if text.strip().startswith("{") and text.strip().endswith("}"):
parsed_data = json.loads(text)
if isinstance(parsed_data, dict) and "content" in parsed_data:
# 如果是包含说话人信息的JSON格式,只显示content部分
diff --git a/main/xiaozhi-server/core/handle/textHandle.py b/main/xiaozhi-server/core/handle/textHandle.py
index 86363217..b5e87783 100644
--- a/main/xiaozhi-server/core/handle/textHandle.py
+++ b/main/xiaozhi-server/core/handle/textHandle.py
@@ -1,154 +1,14 @@
-import json
-import time
-from core.handle.abortHandle import handleAbortMessage
-from core.handle.helloHandle import handleHelloMessage
-from core.providers.tools.device_mcp import handle_mcp_message
-from core.utils.util import remove_punctuation_and_length, filter_sensitive_info
-from core.handle.receiveAudioHandle import startToChat, handleAudioMessage
-from core.providers.tools.device_iot import handleIotDescriptors, handleIotStatus
-from core.handle.reportHandle import enqueue_asr_report
-import asyncio
+from core.handle.textMessageHandlerRegistry import TextMessageHandlerRegistry
+from core.handle.textMessageProcessor import TextMessageProcessor
TAG = __name__
+# 全局处理器注册表
+message_registry = TextMessageHandlerRegistry()
+
+# 创建全局消息处理器实例
+message_processor = TextMessageProcessor(message_registry)
async def handleTextMessage(conn, message):
"""处理文本消息"""
- try:
- msg_json = json.loads(message)
- if isinstance(msg_json, int):
- conn.logger.bind(tag=TAG).info(f"收到文本消息:{message}")
- await conn.websocket.send(message)
- return
- if msg_json["type"] == "hello":
- conn.logger.bind(tag=TAG).info(f"收到hello消息:{message}")
- await handleHelloMessage(conn, msg_json)
- elif msg_json["type"] == "abort":
- conn.logger.bind(tag=TAG).info(f"收到abort消息:{message}")
- await handleAbortMessage(conn)
- elif msg_json["type"] == "listen":
- conn.logger.bind(tag=TAG).info(f"收到listen消息:{message}")
- if "mode" in msg_json:
- conn.client_listen_mode = msg_json["mode"]
- conn.logger.bind(tag=TAG).debug(
- f"客户端拾音模式:{conn.client_listen_mode}"
- )
- if msg_json["state"] == "start":
- conn.client_have_voice = True
- conn.client_voice_stop = False
- elif msg_json["state"] == "stop":
- conn.client_have_voice = True
- conn.client_voice_stop = True
- if len(conn.asr_audio) > 0:
- await handleAudioMessage(conn, b"")
- elif msg_json["state"] == "detect":
- conn.client_have_voice = False
- conn.asr_audio.clear()
- if "text" in msg_json:
- conn.last_activity_time = time.time() * 1000
- original_text = msg_json["text"] # 保留原始文本
- filtered_len, filtered_text = remove_punctuation_and_length(
- original_text
- )
- # 识别是否是唤醒词
- is_wakeup_words = filtered_text in conn.config.get("wakeup_words")
- if not is_wakeup_words:
- # 上报纯文字数据(复用ASR上报功能,但不提供音频数据)
- enqueue_asr_report(conn, original_text, [])
- # 否则需要LLM对文字内容进行答复
- await startToChat(conn, original_text)
- elif msg_json["type"] == "iot":
- conn.logger.bind(tag=TAG).info(f"收到iot消息:{message}")
- if "descriptors" in msg_json:
- asyncio.create_task(handleIotDescriptors(conn, msg_json["descriptors"]))
- if "states" in msg_json:
- asyncio.create_task(handleIotStatus(conn, msg_json["states"]))
- elif msg_json["type"] == "mcp":
- conn.logger.bind(tag=TAG).info(f"收到mcp消息:{message[:100]}")
- if "payload" in msg_json:
- asyncio.create_task(
- handle_mcp_message(conn, conn.mcp_client, msg_json["payload"])
- )
- elif msg_json["type"] == "server":
- # 记录日志时过滤敏感信息
- conn.logger.bind(tag=TAG).info(
- f"收到服务器消息:{filter_sensitive_info(msg_json)}"
- )
- # 如果配置是从API读取的,则需要验证secret
- if not conn.read_config_from_api:
- return
- # 获取post请求的secret
- post_secret = msg_json.get("content", {}).get("secret", "")
- secret = conn.config["manager-api"].get("secret", "")
- # 如果secret不匹配,则返回
- if post_secret != secret:
- await conn.websocket.send(
- json.dumps(
- {
- "type": "server",
- "status": "error",
- "message": "服务器密钥验证失败",
- }
- )
- )
- return
- # 动态更新配置
- if msg_json["action"] == "update_config":
- try:
- # 更新WebSocketServer的配置
- if not conn.server:
- await conn.websocket.send(
- json.dumps(
- {
- "type": "server",
- "status": "error",
- "message": "无法获取服务器实例",
- "content": {"action": "update_config"},
- }
- )
- )
- return
-
- if not await conn.server.update_config():
- await conn.websocket.send(
- json.dumps(
- {
- "type": "server",
- "status": "error",
- "message": "更新服务器配置失败",
- "content": {"action": "update_config"},
- }
- )
- )
- return
-
- # 发送成功响应
- await conn.websocket.send(
- json.dumps(
- {
- "type": "server",
- "status": "success",
- "message": "配置更新成功",
- "content": {"action": "update_config"},
- }
- )
- )
- except Exception as e:
- conn.logger.bind(tag=TAG).error(f"更新配置失败: {str(e)}")
- await conn.websocket.send(
- json.dumps(
- {
- "type": "server",
- "status": "error",
- "message": f"更新配置失败: {str(e)}",
- "content": {"action": "update_config"},
- }
- )
- )
- # 重启服务器
- elif msg_json["action"] == "restart":
- await conn.handle_restart(msg_json)
- else:
- conn.logger.bind(tag=TAG).error(f"收到未知类型消息:{message}")
- except json.JSONDecodeError:
- await conn.websocket.send(message)
+ await message_processor.process_message(conn, message)
diff --git a/main/xiaozhi-server/core/handle/textHandler/abortMessageHandler.py b/main/xiaozhi-server/core/handle/textHandler/abortMessageHandler.py
new file mode 100644
index 00000000..dc540d24
--- /dev/null
+++ b/main/xiaozhi-server/core/handle/textHandler/abortMessageHandler.py
@@ -0,0 +1,16 @@
+from typing import Dict, Any
+
+from core.handle.abortHandle import handleAbortMessage
+from core.handle.textMessageHandler import TextMessageHandler
+from core.handle.textMessageType import TextMessageType
+
+
+class AbortTextMessageHandler(TextMessageHandler):
+ """Abort消息处理器"""
+
+ @property
+ def message_type(self) -> TextMessageType:
+ return TextMessageType.ABORT
+
+ async def handle(self, conn, msg_json: Dict[str, Any]) -> None:
+ await handleAbortMessage(conn)
diff --git a/main/xiaozhi-server/core/handle/textHandler/helloMessageHandler.py b/main/xiaozhi-server/core/handle/textHandler/helloMessageHandler.py
new file mode 100644
index 00000000..1839814e
--- /dev/null
+++ b/main/xiaozhi-server/core/handle/textHandler/helloMessageHandler.py
@@ -0,0 +1,16 @@
+from typing import Dict, Any
+
+from core.handle.helloHandle import handleHelloMessage
+from core.handle.textMessageHandler import TextMessageHandler
+from core.handle.textMessageType import TextMessageType
+
+
+class HelloTextMessageHandler(TextMessageHandler):
+ """Hello消息处理器"""
+
+ @property
+ def message_type(self) -> TextMessageType:
+ return TextMessageType.HELLO
+
+ async def handle(self, conn, msg_json: Dict[str, Any]) -> None:
+ await handleHelloMessage(conn, msg_json)
\ No newline at end of file
diff --git a/main/xiaozhi-server/core/handle/textHandler/iotMessageHandler.py b/main/xiaozhi-server/core/handle/textHandler/iotMessageHandler.py
new file mode 100644
index 00000000..335d08b0
--- /dev/null
+++ b/main/xiaozhi-server/core/handle/textHandler/iotMessageHandler.py
@@ -0,0 +1,20 @@
+import asyncio
+from typing import Dict, Any
+
+from core.handle.textMessageHandler import TextMessageHandler
+from core.handle.textMessageType import TextMessageType
+from core.providers.tools.device_iot import handleIotStatus, handleIotDescriptors
+
+
+class IotTextMessageHandler(TextMessageHandler):
+ """IOT消息处理器"""
+
+ @property
+ def message_type(self) -> TextMessageType:
+ return TextMessageType.IOT
+
+ async def handle(self, conn, msg_json: Dict[str, Any]) -> None:
+ if "descriptors" in msg_json:
+ asyncio.create_task(handleIotDescriptors(conn, msg_json["descriptors"]))
+ if "states" in msg_json:
+ asyncio.create_task(handleIotStatus(conn, msg_json["states"]))
\ No newline at end of file
diff --git a/main/xiaozhi-server/core/handle/textHandler/listenMessageHandler.py b/main/xiaozhi-server/core/handle/textHandler/listenMessageHandler.py
new file mode 100644
index 00000000..97286dfe
--- /dev/null
+++ b/main/xiaozhi-server/core/handle/textHandler/listenMessageHandler.py
@@ -0,0 +1,63 @@
+import time
+from typing import Dict, Any
+
+from core.handle.receiveAudioHandle import handleAudioMessage, startToChat
+from core.handle.reportHandle import enqueue_asr_report
+from core.handle.sendAudioHandle import send_stt_message, send_tts_message
+from core.handle.textMessageHandler import TextMessageHandler
+from core.handle.textMessageType import TextMessageType
+from core.utils.util import remove_punctuation_and_length
+
+TAG = __name__
+
+class ListenTextMessageHandler(TextMessageHandler):
+ """Listen消息处理器"""
+
+ @property
+ def message_type(self) -> TextMessageType:
+ return TextMessageType.LISTEN
+
+ async def handle(self, conn, msg_json: Dict[str, Any]) -> None:
+ if "mode" in msg_json:
+ conn.client_listen_mode = msg_json["mode"]
+ conn.logger.bind(tag=TAG).debug(
+ f"客户端拾音模式:{conn.client_listen_mode}"
+ )
+ if msg_json["state"] == "start":
+ conn.client_have_voice = True
+ conn.client_voice_stop = False
+ elif msg_json["state"] == "stop":
+ conn.client_have_voice = True
+ conn.client_voice_stop = True
+ if len(conn.asr_audio) > 0:
+ await handleAudioMessage(conn, b"")
+ elif msg_json["state"] == "detect":
+ conn.client_have_voice = False
+ conn.asr_audio.clear()
+ if "text" in msg_json:
+ conn.last_activity_time = time.time() * 1000
+ original_text = msg_json["text"] # 保留原始文本
+ filtered_len, filtered_text = remove_punctuation_and_length(
+ original_text
+ )
+
+ # 识别是否是唤醒词
+ is_wakeup_words = filtered_text in conn.config.get("wakeup_words")
+ # 是否开启唤醒词回复
+ enable_greeting = conn.config.get("enable_greeting", True)
+
+ if is_wakeup_words and not enable_greeting:
+ # 如果是唤醒词,且关闭了唤醒词回复,就不用回答
+ await send_stt_message(conn, original_text)
+ await send_tts_message(conn, "stop", None)
+ conn.client_is_speaking = False
+ elif is_wakeup_words:
+ conn.just_woken_up = True
+ # 上报纯文字数据(复用ASR上报功能,但不提供音频数据)
+ enqueue_asr_report(conn, "嘿,你好呀", [])
+ await startToChat(conn, "嘿,你好呀")
+ else:
+ # 上报纯文字数据(复用ASR上报功能,但不提供音频数据)
+ enqueue_asr_report(conn, original_text, [])
+ # 否则需要LLM对文字内容进行答复
+ await startToChat(conn, original_text)
\ No newline at end of file
diff --git a/main/xiaozhi-server/core/handle/textHandler/mcpMessageHandler.py b/main/xiaozhi-server/core/handle/textHandler/mcpMessageHandler.py
new file mode 100644
index 00000000..65876f24
--- /dev/null
+++ b/main/xiaozhi-server/core/handle/textHandler/mcpMessageHandler.py
@@ -0,0 +1,20 @@
+import asyncio
+from typing import Dict, Any
+
+from core.handle.textMessageHandler import TextMessageHandler
+from core.handle.textMessageType import TextMessageType
+from core.providers.tools.device_mcp import handle_mcp_message
+
+
+class McpTextMessageHandler(TextMessageHandler):
+ """MCP消息处理器"""
+
+ @property
+ def message_type(self) -> TextMessageType:
+ return TextMessageType.MCP
+
+ async def handle(self, conn, msg_json: Dict[str, Any]) -> None:
+ if "payload" in msg_json:
+ asyncio.create_task(
+ handle_mcp_message(conn, conn.mcp_client, msg_json["payload"])
+ )
\ No newline at end of file
diff --git a/main/xiaozhi-server/core/handle/textHandler/serverMessageHandler.py b/main/xiaozhi-server/core/handle/textHandler/serverMessageHandler.py
new file mode 100644
index 00000000..b9a23588
--- /dev/null
+++ b/main/xiaozhi-server/core/handle/textHandler/serverMessageHandler.py
@@ -0,0 +1,92 @@
+import asyncio
+import json
+from typing import Dict, Any
+
+from core.handle.textMessageHandler import TextMessageHandler
+from core.handle.textMessageType import TextMessageType
+from core.providers.tools.device_mcp import handle_mcp_message
+
+TAG = __name__
+
+class ServerTextMessageHandler(TextMessageHandler):
+ """MCP消息处理器"""
+
+ @property
+ def message_type(self) -> TextMessageType:
+ return TextMessageType.SERVER
+
+ async def handle(self, conn, msg_json: Dict[str, Any]) -> None:
+ # 如果配置是从API读取的,则需要验证secret
+ if not conn.read_config_from_api:
+ return
+ # 获取post请求的secret
+ post_secret = msg_json.get("content", {}).get("secret", "")
+ secret = conn.config["manager-api"].get("secret", "")
+ # 如果secret不匹配,则返回
+ if post_secret != secret:
+ await conn.websocket.send(
+ json.dumps(
+ {
+ "type": "server",
+ "status": "error",
+ "message": "服务器密钥验证失败",
+ }
+ )
+ )
+ return
+ # 动态更新配置
+ if msg_json["action"] == "update_config":
+ try:
+ # 更新WebSocketServer的配置
+ if not conn.server:
+ await conn.websocket.send(
+ json.dumps(
+ {
+ "type": "server",
+ "status": "error",
+ "message": "无法获取服务器实例",
+ "content": {"action": "update_config"},
+ }
+ )
+ )
+ return
+
+ if not await conn.server.update_config():
+ await conn.websocket.send(
+ json.dumps(
+ {
+ "type": "server",
+ "status": "error",
+ "message": "更新服务器配置失败",
+ "content": {"action": "update_config"},
+ }
+ )
+ )
+ return
+
+ # 发送成功响应
+ await conn.websocket.send(
+ json.dumps(
+ {
+ "type": "server",
+ "status": "success",
+ "message": "配置更新成功",
+ "content": {"action": "update_config"},
+ }
+ )
+ )
+ except Exception as e:
+ conn.logger.bind(tag=TAG).error(f"更新配置失败: {str(e)}")
+ await conn.websocket.send(
+ json.dumps(
+ {
+ "type": "server",
+ "status": "error",
+ "message": f"更新配置失败: {str(e)}",
+ "content": {"action": "update_config"},
+ }
+ )
+ )
+ # 重启服务器
+ elif msg_json["action"] == "restart":
+ await conn.handle_restart(msg_json)
\ No newline at end of file
diff --git a/main/xiaozhi-server/core/handle/textMessageHandler.py b/main/xiaozhi-server/core/handle/textMessageHandler.py
new file mode 100644
index 00000000..f94a0bac
--- /dev/null
+++ b/main/xiaozhi-server/core/handle/textMessageHandler.py
@@ -0,0 +1,21 @@
+from abc import abstractmethod, ABC
+from typing import Dict, Any
+
+from core.handle.textMessageType import TextMessageType
+
+TAG = __name__
+
+
+class TextMessageHandler(ABC):
+ """消息处理器抽象基类"""
+
+ @abstractmethod
+ async def handle(self, conn, msg_json: Dict[str, Any]) -> None:
+ """处理消息的抽象方法"""
+ pass
+
+ @property
+ @abstractmethod
+ def message_type(self) -> TextMessageType:
+ """返回处理的消息类型"""
+ pass
diff --git a/main/xiaozhi-server/core/handle/textMessageHandlerRegistry.py b/main/xiaozhi-server/core/handle/textMessageHandlerRegistry.py
new file mode 100644
index 00000000..e90d7231
--- /dev/null
+++ b/main/xiaozhi-server/core/handle/textMessageHandlerRegistry.py
@@ -0,0 +1,45 @@
+from typing import Dict, Optional
+
+from core.handle.textHandler.abortMessageHandler import AbortTextMessageHandler
+from core.handle.textHandler.helloMessageHandler import HelloTextMessageHandler
+from core.handle.textHandler.iotMessageHandler import IotTextMessageHandler
+from core.handle.textHandler.listenMessageHandler import ListenTextMessageHandler
+from core.handle.textHandler.mcpMessageHandler import McpTextMessageHandler
+from core.handle.textMessageHandler import TextMessageHandler
+from core.handle.textHandler.serverMessageHandler import ServerTextMessageHandler
+
+TAG = __name__
+
+
+class TextMessageHandlerRegistry:
+ """消息处理器注册表"""
+
+ def __init__(self):
+ self._handlers: Dict[str, TextMessageHandler] = {}
+ self._register_default_handlers()
+
+ def _register_default_handlers(self) -> None:
+ """注册默认的消息处理器"""
+ handlers = [
+ HelloTextMessageHandler(),
+ AbortTextMessageHandler(),
+ ListenTextMessageHandler(),
+ IotTextMessageHandler(),
+ McpTextMessageHandler(),
+ ServerTextMessageHandler(),
+ ]
+
+ for handler in handlers:
+ self.register_handler(handler)
+
+ def register_handler(self, handler: TextMessageHandler) -> None:
+ """注册消息处理器"""
+ self._handlers[handler.message_type.value] = handler
+
+ def get_handler(self, message_type: str) -> Optional[TextMessageHandler]:
+ """获取消息处理器"""
+ return self._handlers.get(message_type)
+
+ def get_supported_types(self) -> list:
+ """获取支持的消息类型"""
+ return list(self._handlers.keys())
diff --git a/main/xiaozhi-server/core/handle/textMessageProcessor.py b/main/xiaozhi-server/core/handle/textMessageProcessor.py
new file mode 100644
index 00000000..0cae5e09
--- /dev/null
+++ b/main/xiaozhi-server/core/handle/textMessageProcessor.py
@@ -0,0 +1,41 @@
+import json
+
+from core.handle.textMessageHandlerRegistry import TextMessageHandlerRegistry
+
+TAG = __name__
+
+
+class TextMessageProcessor:
+ """消息处理器主类"""
+
+ def __init__(self, registry: TextMessageHandlerRegistry):
+ self.registry = registry
+
+ async def process_message(self, conn, message: str) -> None:
+ """处理消息的主入口"""
+ try:
+ # 解析JSON消息
+ msg_json = json.loads(message)
+
+ # 处理JSON消息
+ if isinstance(msg_json, dict):
+ message_type = msg_json.get("type")
+
+ # 记录日志
+ conn.logger.bind(tag=TAG).info(f"收到{message_type}消息:{message}")
+
+ # 获取并执行处理器
+ handler = self.registry.get_handler(message_type)
+ if handler:
+ await handler.handle(conn, msg_json)
+ else:
+ conn.logger.bind(tag=TAG).error(f"收到未知类型消息:{message}")
+ # 处理纯数字消息
+ elif isinstance(msg_json, int):
+ conn.logger.bind(tag=TAG).info(f"收到数字消息:{message}")
+ await conn.websocket.send(message)
+
+ except json.JSONDecodeError:
+ # 非JSON消息直接转发
+ conn.logger.bind(tag=TAG).error(f"解析到错误的消息:{message}")
+ await conn.websocket.send(message)
diff --git a/main/xiaozhi-server/core/handle/textMessageType.py b/main/xiaozhi-server/core/handle/textMessageType.py
new file mode 100644
index 00000000..53e71b71
--- /dev/null
+++ b/main/xiaozhi-server/core/handle/textMessageType.py
@@ -0,0 +1,11 @@
+from enum import Enum
+
+
+class TextMessageType(Enum):
+ """消息类型枚举"""
+ HELLO = "hello"
+ ABORT = "abort"
+ LISTEN = "listen"
+ IOT = "iot"
+ MCP = "mcp"
+ SERVER = "server"
diff --git a/main/xiaozhi-server/core/providers/asr/base.py b/main/xiaozhi-server/core/providers/asr/base.py
index 972818d4..250a25f2 100644
--- a/main/xiaozhi-server/core/providers/asr/base.py
+++ b/main/xiaozhi-server/core/providers/asr/base.py
@@ -1,18 +1,18 @@
import os
+import io
import wave
import uuid
+import json
+import time
import queue
import asyncio
import traceback
import threading
import opuslib_next
-import json
-import io
-import time
import concurrent.futures
from abc import ABC, abstractmethod
from config.logger import setup_logging
-from typing import Optional, Tuple, List, Dict, Any
+from typing import Optional, Tuple, List
from core.handle.receiveAudioHandle import startToChat
from core.handle.reportHandle import enqueue_asr_report
from core.utils.util import remove_punctuation_and_length
@@ -87,11 +87,9 @@ class ASRProviderBase(ABC):
# 预先准备WAV数据
wav_data = None
- # 使用连接的声纹识别提供者
if conn.voiceprint_provider and combined_pcm_data:
wav_data = self._pcm_to_wav(combined_pcm_data)
-
# 定义ASR任务
def run_asr():
start_time = time.monotonic()
@@ -132,8 +130,6 @@ class ASRProviderBase(ABC):
return None
# 使用线程池执行器并行运行
- parallel_start_time = time.monotonic()
-
with concurrent.futures.ThreadPoolExecutor(max_workers=2) as thread_executor:
asr_future = thread_executor.submit(run_asr)
@@ -151,7 +147,7 @@ class ASRProviderBase(ABC):
# 处理结果
- raw_text, file_path = results.get("asr", ("", None))
+ raw_text, _ = results.get("asr", ("", None))
speaker_name = results.get("voiceprint", None)
# 记录识别结果
diff --git a/main/xiaozhi-server/core/providers/llm/AliBL/AliBL.py b/main/xiaozhi-server/core/providers/llm/AliBL/AliBL.py
index 012035ee..a68c4348 100644
--- a/main/xiaozhi-server/core/providers/llm/AliBL/AliBL.py
+++ b/main/xiaozhi-server/core/providers/llm/AliBL/AliBL.py
@@ -1,8 +1,10 @@
from config.logger import setup_logging
from http import HTTPStatus
+import dashscope
from dashscope import Application
from core.providers.llm.base import LLMProviderBase
from core.utils.util import check_model_key
+import time
TAG = __name__
logger = setup_logging()
@@ -15,6 +17,7 @@ class LLMProvider(LLMProviderBase):
self.base_url = config.get("base_url")
self.is_No_prompt = config.get("is_no_prompt")
self.memory_id = config.get("ali_memory_id")
+ self.streaming_chunk_size = config.get("streaming_chunk_size", 3) # 每次流式返回的字符数
check_model_key("AliBLLLM", self.api_key)
def response(self, session_id, dialogue):
@@ -32,6 +35,8 @@ class LLMProvider(LLMProviderBase):
"app_id": self.app_id,
"session_id": session_id,
"messages": dialogue,
+ # 开启SDK原生流式
+ "stream": True,
}
if self.memory_id != False:
# 百练memory需要prompt参数
@@ -42,25 +47,63 @@ class LLMProvider(LLMProviderBase):
f"【阿里百练API服务】处理后的prompt: {prompt}"
)
+ # 可选地设置自定义API基地址(若配置为兼容模式URL则忽略)
+ if self.base_url and ("/api/" in self.base_url):
+ dashscope.base_http_api_url = self.base_url
+
responses = Application.call(**call_params)
- if responses.status_code != HTTPStatus.OK:
- logger.bind(tag=TAG).error(
- f"code={responses.status_code}, "
- f"message={responses.message}, "
- f"请参考文档:https://help.aliyun.com/zh/model-studio/developer-reference/error-code"
- )
- yield "【阿里百练API服务响应异常】"
- else:
- logger.bind(tag=TAG).debug(
- f"【阿里百练API服务】构造参数: {call_params}"
- )
- yield responses.output.text
+
+ # 流式处理(SDK在stream=True时返回可迭代对象;否则返回单次响应对象)
+ logger.bind(tag=TAG).debug(
+ f"【阿里百练API服务】构造参数: {dict(call_params, api_key='***')}"
+ )
+
+ last_text = ""
+ try:
+ for resp in responses:
+ if resp.status_code != HTTPStatus.OK:
+ logger.bind(tag=TAG).error(
+ f"code={resp.status_code}, message={resp.message}, 请参考文档:https://help.aliyun.com/zh/model-studio/developer-reference/error-code"
+ )
+ continue
+ current_text = getattr(getattr(resp, "output", None), "text", None)
+ if current_text is None:
+ continue
+ # SDK流式为增量覆盖,计算差量输出
+ if len(current_text) >= len(last_text):
+ delta = current_text[len(last_text):]
+ else:
+ # 避免偶发回退
+ delta = current_text
+ if delta:
+ yield delta
+ last_text = current_text
+ except TypeError:
+ # 非流式回落(一次性返回)
+ if responses.status_code != HTTPStatus.OK:
+ logger.bind(tag=TAG).error(
+ f"code={responses.status_code}, message={responses.message}, 请参考文档:https://help.aliyun.com/zh/model-studio/developer-reference/error-code"
+ )
+ yield "【阿里百练API服务响应异常】"
+ else:
+ full_text = getattr(getattr(responses, "output", None), "text", "")
+ logger.bind(tag=TAG).info(
+ f"【阿里百练API服务】完整响应长度: {len(full_text)}"
+ )
+ for i in range(0, len(full_text), self.streaming_chunk_size):
+ chunk = full_text[i:i + self.streaming_chunk_size]
+ if chunk:
+ yield chunk
except Exception as e:
logger.bind(tag=TAG).error(f"【阿里百练API服务】响应异常: {e}")
yield "【LLM服务响应异常】"
def response_with_functions(self, session_id, dialogue, functions=None):
- logger.bind(tag=TAG).error(
- f"阿里百练暂未实现完整的工具调用(function call),建议使用其他意图识别"
+ # 阿里百练当前未支持原生的 function call。为保持兼容,这里回退到普通文本流式输出。
+ # 上层会按 (content, tool_calls) 的形式消费,这里始终返回 (token, None)
+ logger.bind(tag=TAG).warning(
+ "阿里百练未实现原生 function call,已回退为纯文本流式输出"
)
+ for token in self.response(session_id, dialogue):
+ yield token, None
diff --git a/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_executor.py b/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_executor.py
index 9ae15d4b..ee0bcb2a 100644
--- a/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_executor.py
+++ b/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_executor.py
@@ -18,8 +18,8 @@ class ServerMCPExecutor(ToolExecutor):
"""初始化MCP管理器"""
if not self._initialized:
self.mcp_manager = ServerMCPManager(self.conn)
- await self.mcp_manager.initialize_servers()
self._initialized = True
+ await self.mcp_manager.initialize_servers()
async def execute(
self, conn, tool_name: str, arguments: Dict[str, Any]
diff --git a/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_manager.py b/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_manager.py
index 52ab2b79..cf3deb21 100644
--- a/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_manager.py
+++ b/main/xiaozhi-server/core/providers/tools/server_mcp/mcp_manager.py
@@ -68,6 +68,9 @@ class ServerMCPManager:
# 输出当前支持的服务端MCP工具列表
if hasattr(self.conn, "func_handler") and self.conn.func_handler:
+ # 刷新工具缓存以确保服务端MCP工具被正确加载
+ if hasattr(self.conn.func_handler, "tool_manager"):
+ self.conn.func_handler.tool_manager.refresh_tools()
self.conn.func_handler.current_support_functions()
def get_all_tools(self) -> List[Dict[str, Any]]:
diff --git a/main/xiaozhi-server/core/providers/tts/aliyun_stream.py b/main/xiaozhi-server/core/providers/tts/aliyun_stream.py
index 32e4d50f..38e70927 100644
--- a/main/xiaozhi-server/core/providers/tts/aliyun_stream.py
+++ b/main/xiaozhi-server/core/providers/tts/aliyun_stream.py
@@ -216,7 +216,6 @@ class TTSProvider(TTSProviderBase):
if message.sentence_type == SentenceType.FIRST:
self.conn.client_abort = False
- self.reset_flow_controller()
if self.conn.client_abort:
logger.bind(tag=TAG).info("收到打断信息,终止TTS文本处理线程")
@@ -479,3 +478,142 @@ class TTSProvider(TTSProviderBase):
finally:
self._monitor_task = None
+ def to_tts(self, text: str) -> list:
+ """非流式TTS处理,用于测试及保存音频文件的场景"""
+ try:
+ # 创建新的事件循环
+ loop = asyncio.new_event_loop()
+ asyncio.set_event_loop(loop)
+
+ # 生成会话ID
+ session_id = uuid.uuid4().hex
+ # 存储音频数据
+ audio_data = []
+
+ async def _generate_audio():
+ # 刷新Token(如果需要)
+ if self._is_token_expired():
+ self._refresh_token()
+
+ # 建立WebSocket连接
+ ws = await websockets.connect(
+ self.ws_url,
+ additional_headers={"X-NLS-Token": self.token},
+ ping_interval=30,
+ ping_timeout=10,
+ close_timeout=10,
+ )
+ try:
+ # 发送StartSynthesis请求
+ start_message_id = str(uuid.uuid4().hex)
+ start_request = {
+ "header": {
+ "message_id": start_message_id,
+ "task_id": session_id,
+ "namespace": "FlowingSpeechSynthesizer",
+ "name": "StartSynthesis",
+ "appkey": self.appkey,
+ },
+ "payload": {
+ "voice": self.voice,
+ "format": self.format,
+ "sample_rate": self.sample_rate,
+ "volume": self.volume,
+ "speech_rate": self.speech_rate,
+ "pitch_rate": self.pitch_rate,
+ "enable_subtitle": True,
+ },
+ }
+ await ws.send(json.dumps(start_request))
+
+ # 等待SynthesisStarted响应
+ synthesis_started = False
+ while not synthesis_started:
+ msg = await ws.recv()
+ if isinstance(msg, str):
+ data = json.loads(msg)
+ header = data.get("header", {})
+ if header.get("name") == "SynthesisStarted":
+ synthesis_started = True
+ logger.bind(tag=TAG).debug("TTS合成已启动")
+ elif header.get("name") == "TaskFailed":
+ error_info = data.get("payload", {}).get(
+ "error_info", {}
+ )
+ error_code = error_info.get("error_code")
+ error_message = error_info.get(
+ "error_message", "未知错误"
+ )
+ raise Exception(
+ f"启动合成失败: {error_code} - {error_message}"
+ )
+
+ # 发送文本合成请求
+ filtered_text = MarkdownCleaner.clean_markdown(text)
+ run_message_id = str(uuid.uuid4().hex)
+ run_request = {
+ "header": {
+ "message_id": run_message_id,
+ "task_id": session_id,
+ "namespace": "FlowingSpeechSynthesizer",
+ "name": "RunSynthesis",
+ "appkey": self.appkey,
+ },
+ "payload": {"text": filtered_text},
+ }
+ await ws.send(json.dumps(run_request))
+
+ # 发送停止合成请求
+ stop_message_id = str(uuid.uuid4().hex)
+ stop_request = {
+ "header": {
+ "message_id": stop_message_id,
+ "task_id": session_id,
+ "namespace": "FlowingSpeechSynthesizer",
+ "name": "StopSynthesis",
+ "appkey": self.appkey,
+ }
+ }
+ await ws.send(json.dumps(stop_request))
+
+ # 接收音频数据
+ synthesis_completed = False
+ while not synthesis_completed:
+ msg = await ws.recv()
+ if isinstance(msg, (bytes, bytearray)):
+ self.opus_encoder.encode_pcm_to_opus_stream(
+ msg,
+ end_of_stream=False,
+ callback=lambda opus: audio_data.append(opus)
+ )
+ elif isinstance(msg, str):
+ data = json.loads(msg)
+ header = data.get("header", {})
+ event_name = header.get("name")
+ if event_name == "SynthesisCompleted":
+ synthesis_completed = True
+ logger.bind(tag=TAG).debug("TTS合成完成")
+ elif event_name == "TaskFailed":
+ error_info = data.get("payload", {}).get(
+ "error_info", {}
+ )
+ error_code = error_info.get("error_code")
+ error_message = error_info.get(
+ "error_message", "未知错误"
+ )
+ raise Exception(
+ f"合成失败: {error_code} - {error_message}"
+ )
+ finally:
+ try:
+ await ws.close()
+ except:
+ pass
+
+ loop.run_until_complete(_generate_audio())
+ loop.close()
+
+ return audio_data
+ except Exception as e:
+ logger.bind(tag=TAG).error(f"生成音频数据失败: {str(e)}")
+ return []
\ No newline at end of file
diff --git a/main/xiaozhi-server/core/providers/tts/base.py b/main/xiaozhi-server/core/providers/tts/base.py
index 0b942862..04a7fa36 100644
--- a/main/xiaozhi-server/core/providers/tts/base.py
+++ b/main/xiaozhi-server/core/providers/tts/base.py
@@ -1,22 +1,22 @@
import os
import re
-import queue
+import time
import uuid
+import queue
import asyncio
import threading
-from typing import Callable, Any
+import traceback
from core.utils import p3
-import time
from datetime import datetime
from core.utils import textUtils
+from typing import Callable, Any
from abc import ABC, abstractmethod
from config.logger import setup_logging
-from core.utils.audio_flow_control import FlowControlConfig
-from core.utils.util import audio_bytes_to_data_stream, audio_to_data_stream
from core.utils.tts import MarkdownCleaner
from core.utils.output_counter import add_device_output
from core.handle.reportHandle import enqueue_tts_report
from core.handle.sendAudioHandle import sendAudioMessage
+from core.utils.util import audio_bytes_to_data_stream, audio_to_data_stream
from core.providers.tts.dto.dto import (
TTSMessageDTO,
SentenceType,
@@ -24,8 +24,6 @@ from core.providers.tts.dto.dto import (
InterfaceType,
)
-import traceback
-
TAG = __name__
logger = setup_logging()
@@ -34,7 +32,6 @@ class TTSProviderBase(ABC):
def __init__(self, config, delete_audio_file):
self.interface_type = InterfaceType.NON_STREAM
self.conn = None
- self.tts_timeout = 10
self.delete_audio_file = delete_audio_file
self.audio_file_type = "wav"
self.output_file = config.get("output_dir", "tmp/")
@@ -71,7 +68,6 @@ class TTSProviderBase(ABC):
self.tts_stop_request = False
self.processed_chars = 0
self.is_first_sentence = True
- self.flow_controller = FlowControlConfig.create_flow_controller()
def generate_filename(self, extension=".wav"):
return os.path.join(
@@ -147,6 +143,68 @@ class TTSProviderBase(ABC):
except Exception as e:
logger.bind(tag=TAG).error(f"Failed to generate TTS file: {e}")
return None
+
+ def to_tts(self, text):
+ text = MarkdownCleaner.clean_markdown(text)
+ max_repeat_time = 5
+ if self.delete_audio_file:
+ # 需要删除文件的直接转为音频数据
+ while max_repeat_time > 0:
+ try:
+ audio_bytes = asyncio.run(self.text_to_speak(text, None))
+ if audio_bytes:
+ audio_datas = []
+ audio_bytes_to_data_stream(
+ audio_bytes,
+ file_type=self.audio_file_type,
+ is_opus=True,
+ callback=lambda data: audio_datas.append(data)
+ )
+ return audio_datas
+ else:
+ max_repeat_time -= 1
+ except Exception as e:
+ logger.bind(tag=TAG).warning(
+ f"语音生成失败{5 - max_repeat_time + 1}次: {text},错误: {e}"
+ )
+ max_repeat_time -= 1
+ if max_repeat_time > 0:
+ logger.bind(tag=TAG).info(
+ f"语音生成成功: {text},重试{5 - max_repeat_time}次"
+ )
+ else:
+ logger.bind(tag=TAG).error(
+ f"语音生成失败: {text},请检查网络或服务是否正常"
+ )
+ return None
+ else:
+ tmp_file = self.generate_filename()
+ try:
+ while not os.path.exists(tmp_file) and max_repeat_time > 0:
+ try:
+ asyncio.run(self.text_to_speak(text, tmp_file))
+ except Exception as e:
+ logger.bind(tag=TAG).warning(
+ f"语音生成失败{5 - max_repeat_time + 1}次: {text},错误: {e}"
+ )
+ # 未执行成功,删除文件
+ if os.path.exists(tmp_file):
+ os.remove(tmp_file)
+ max_repeat_time -= 1
+
+ if max_repeat_time > 0:
+ logger.bind(tag=TAG).info(
+ f"语音生成成功: {text}:{tmp_file},重试{5 - max_repeat_time}次"
+ )
+ else:
+ logger.bind(tag=TAG).error(
+ f"语音生成失败: {text},请检查网络或服务是否正常"
+ )
+
+ return tmp_file
+ except Exception as e:
+ logger.bind(tag=TAG).error(f"Failed to generate TTS file: {e}")
+ return None
@abstractmethod
async def text_to_speak(self, text, output_file):
@@ -194,7 +252,6 @@ class TTSProviderBase(ABC):
async def open_audio_channels(self, conn):
self.conn = conn
- self.tts_timeout = conn.config.get("tts_timeout", 10)
# tts 消化线程
self.tts_priority_thread = threading.Thread(
target=self.tts_text_priority_thread, daemon=True
@@ -225,7 +282,6 @@ class TTSProviderBase(ABC):
self.tts_text_buff = []
self.is_first_sentence = True
self.tts_audio_first_sentence = True
- self.reset_flow_controller()
elif ContentType.TEXT == message.content_type:
self.tts_text_buff.append(message.content_detail)
segment_text = self._get_segment_text()
@@ -270,103 +326,43 @@ class TTSProviderBase(ABC):
if self.conn.client_abort:
logger.bind(tag=TAG).debug("收到打断信号,跳过当前音频数据")
- # 打断时丢弃未上报的音频数据
enqueue_text, enqueue_audio = None, []
continue
# 收到下一个文本开始或会话结束时进行上报
if sentence_type is not SentenceType.MIDDLE:
+ # 重置音频流控状态(新句子开始或者结束)
+ if hasattr(self.conn, 'audio_flow_control'):
+ self.conn.audio_flow_control = {
+ 'last_send_time': 0,
+ 'packet_count': 0,
+ 'start_time': time.perf_counter()
+ }
+
# 上报TTS数据
if enqueue_text is not None and enqueue_audio is not None:
enqueue_tts_report(self.conn, enqueue_text, enqueue_audio)
enqueue_audio = []
enqueue_text = text
- # 计算音频数据的帧数
- if isinstance(audio_datas, bytes):
- frame_count = 1 # 单个字节流作为一帧
+ # 收集上报音频数据
+ if isinstance(audio_datas, bytes) and enqueue_audio is not None:
enqueue_audio.append(audio_datas)
- else:
- frame_count = 0
+
+ # 发送音频
+ future = asyncio.run_coroutine_threadsafe(
+ sendAudioMessage(self.conn, sentence_type, audio_datas, text),
+ self.conn.loop,
+ )
+ future.result()
# 记录输出和报告
if self.conn.max_output_size > 0 and text:
add_device_output(self.conn.headers.get("device-id"), len(text))
- # 流控检查
- if frame_count > 0:
- max_wait_time = FlowControlConfig.DEFAULT_MAX_WAIT_TIME
- wait_start_time = time.time()
- retry_interval = FlowControlConfig.DEFAULT_RETRY_INTERVAL
-
- while not self.flow_controller.can_send_frames(frame_count):
- # 检查是否超时或需要停止
- if (
- time.time() - wait_start_time > max_wait_time
- or self.conn.stop_event.is_set()
- or self.conn.client_abort
- ):
- logger.bind(tag=TAG).debug(
- "流控等待超时或收到停止信号,跳过音频发送"
- )
- break
- # 短暂等待后重试
- time.sleep(retry_interval)
- else:
- # 可以发送,记录发送的帧数
- self.flow_controller.record_sent_frames(frame_count)
-
- # 发送音频
- future = asyncio.run_coroutine_threadsafe(
- self._send_audio_with_flow_control(
- sentence_type, audio_datas, text
- ),
- self.conn.loop,
- )
- future.result()
-
- # 输出流控状态(调试用)
- # status = self.flow_controller.get_status()
- # logger.bind(tag=TAG).debug(
- # f"流控状态: 缓冲区使用率={status['buffer_usage_percent']:.1f}%, "
- # f"可用令牌={status['available_tokens']}..."
- # f"发送帧数={status['sent_frames']}..."
- # f"消费帧数={status['consumed_frames']}..."
- # f"代播放帧数={status['sent_frames'] - status['consumed_frames']}..."
- # )
- else:
- # 没有音频数据,直接发送
- future = asyncio.run_coroutine_threadsafe(
- self._send_audio_with_flow_control(
- sentence_type, audio_datas, text
- ),
- self.conn.loop,
- )
- future.result()
-
except Exception as e:
logger.bind(tag=TAG).error(f"audio_play_priority_thread: {text} {e}")
- async def _send_audio_with_flow_control(self, sentence_type, audio_datas, text):
- """
- 带流控的音频发送方法 模拟设备消费音频帧的过程
- 实际应用中应该根据设备反馈来更新消费情况
- """
- await sendAudioMessage(self.conn, sentence_type, audio_datas, text)
-
- # 模拟设备消费(实际应用中应该从设备获取反馈)防止音字不同步
- if isinstance(audio_datas, bytes):
- # 模拟设备播放延迟(60ms per frame), 实际情况可以低一点(50ms),增加使用体验
- await asyncio.sleep(0.06)
- self.flow_controller.update_device_consumption(1)
-
- # 在类中添加流控制器重置方法
- def reset_flow_controller(self):
- """重置流控制器状态,通常在新会话开始时调用"""
- if hasattr(self, "flow_controller"):
- self.flow_controller.reset()
- logger.bind(tag=TAG).info("流控制器状态已重置")
-
async def start_session(self, session_id):
pass
diff --git a/main/xiaozhi-server/core/providers/tts/fishspeech.py b/main/xiaozhi-server/core/providers/tts/fishspeech.py
index bbf19164..83b3bc3f 100644
--- a/main/xiaozhi-server/core/providers/tts/fishspeech.py
+++ b/main/xiaozhi-server/core/providers/tts/fishspeech.py
@@ -143,8 +143,8 @@ class TTSProvider(TTSProviderBase):
data = {
"text": text,
"references": [
- ServeReferenceAudio(audio=audio if audio else b"", text=text)
- for text, audio in zip(ref_texts, byte_audios)
+ ServeReferenceAudio(audio=audio if audio else b"", text=ref_text)
+ for ref_text, audio in zip(ref_texts, byte_audios)
],
"reference_id": self.reference_id,
"normalize": self.normalize,
diff --git a/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py b/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py
index ac6301f9..ff295581 100644
--- a/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py
+++ b/main/xiaozhi-server/core/providers/tts/huoshan_double_stream.py
@@ -213,7 +213,6 @@ class TTSProvider(TTSProviderBase):
if message.sentence_type == SentenceType.FIRST:
self.conn.client_abort = False
- self.reset_flow_controller()
if self.conn.client_abort:
try:
@@ -629,3 +628,104 @@ class TTSProvider(TTSProviderBase):
def wav_to_opus_data_audio_raw_stream(self, raw_data_var, is_end=False, callback: Callable[[Any], Any]=None):
return self.opus_encoder.encode_pcm_to_opus_stream(raw_data_var, is_end, callback=callback)
+
+ def to_tts(self, text: str) -> list:
+ """非流式生成音频数据,用于生成音频及测试场景
+ Args:
+ text: 要转换的文本
+ Returns:
+ list: 音频数据列表
+ """
+ try:
+ # 创建事件循环
+ loop = asyncio.new_event_loop()
+ asyncio.set_event_loop(loop)
+
+ # 生成会话ID
+ session_id = uuid.uuid4().__str__().replace("-", "")
+
+ # 存储音频数据
+ audio_data = []
+
+ async def _generate_audio():
+ # 创建新的WebSocket连接
+ ws_header = {
+ "X-Api-App-Key": self.appId,
+ "X-Api-Access-Key": self.access_token,
+ "X-Api-Resource-Id": self.resource_id,
+ "X-Api-Connect-Id": uuid.uuid4(),
+ }
+ ws = await websockets.connect(
+ self.ws_url, additional_headers=ws_header, max_size=1000000000
+ )
+
+ try:
+ # 启动会话
+ header = Header(
+ message_type=FULL_CLIENT_REQUEST,
+ message_type_specific_flags=MsgTypeFlagWithEvent,
+ serial_method=JSON,
+ ).as_bytes()
+ optional = Optional(
+ event=EVENT_StartSession, sessionId=session_id
+ ).as_bytes()
+ payload = self.get_payload_bytes(
+ event=EVENT_StartSession, speaker=self.voice
+ )
+ await self.send_event(ws, header, optional, payload)
+
+ # 发送文本
+ header = Header(
+ message_type=FULL_CLIENT_REQUEST,
+ message_type_specific_flags=MsgTypeFlagWithEvent,
+ serial_method=JSON,
+ ).as_bytes()
+ optional = Optional(
+ event=EVENT_TaskRequest, sessionId=session_id
+ ).as_bytes()
+ payload = self.get_payload_bytes(
+ event=EVENT_TaskRequest, text=text, speaker=self.voice
+ )
+ await self.send_event(ws, header, optional, payload)
+
+ # 发送结束会话请求
+ header = Header(
+ message_type=FULL_CLIENT_REQUEST,
+ message_type_specific_flags=MsgTypeFlagWithEvent,
+ serial_method=JSON,
+ ).as_bytes()
+ optional = Optional(
+ event=EVENT_FinishSession, sessionId=session_id
+ ).as_bytes()
+ payload = str.encode("{}")
+ await self.send_event(ws, header, optional, payload)
+
+ # 接收音频数据
+ while True:
+ msg = await ws.recv()
+ res = self.parser_response(msg)
+
+ if (
+ res.optional.event == EVENT_TTSResponse
+ and res.header.message_type == AUDIO_ONLY_RESPONSE
+ ):
+ self.wav_to_opus_data_audio_raw_stream(res.payload, callback=lambda opus_frame: audio_data.append(opus_frame))
+ elif res.optional.event == EVENT_SessionFinished:
+ break
+
+ finally:
+ # 清理资源
+ try:
+ await ws.close()
+ except:
+ pass
+
+ # 运行异步任务
+ loop.run_until_complete(_generate_audio())
+ loop.close()
+
+ return audio_data
+
+ except Exception as e:
+ logger.bind(tag=TAG).error(f"生成音频数据失败: {str(e)}")
+ return []
diff --git a/main/xiaozhi-server/core/providers/tts/index_stream.py b/main/xiaozhi-server/core/providers/tts/index_stream.py
index d562e802..8abb06f2 100644
--- a/main/xiaozhi-server/core/providers/tts/index_stream.py
+++ b/main/xiaozhi-server/core/providers/tts/index_stream.py
@@ -1,8 +1,10 @@
import os
+import time
import queue
-import asyncio
-import traceback
import aiohttp
+import asyncio
+import requests
+import traceback
from config.logger import setup_logging
from core.utils.tts import MarkdownCleaner
from core.providers.tts.base import TTSProviderBase
@@ -45,7 +47,6 @@ class TTSProvider(TTSProviderBase):
self.processed_chars = 0
self.tts_text_buff = []
self.before_stop_play_files.clear()
- self.reset_flow_controller()
elif ContentType.TEXT == message.content_type:
self.tts_text_buff.append(message.content_detail)
segment_text = self._get_segment_text()
@@ -178,3 +179,57 @@ class TTSProvider(TTSProviderBase):
await super().close()
if hasattr(self, "opus_encoder"):
self.opus_encoder.close()
+
+ def to_tts(self, text: str) -> list:
+ """非流式TTS处理,用于测试及保存音频文件的场景
+ Args:
+ text: 要转换的文本
+ Returns:
+ list: 返回opus编码后的音频数据列表
+ """
+ start_time = time.time()
+ text = MarkdownCleaner.clean_markdown(text)
+
+ payload = {"text": text, "character": self.character}
+
+ try:
+ with requests.post(self.api_url, json=payload, timeout=5) as response:
+ if response.status_code != 200:
+ logger.bind(tag=TAG).error(
+ f"TTS请求失败: {response.status_code}, {response.text}"
+ )
+ return []
+
+ logger.info(f"TTS请求成功: {text}, 耗时: {time.time() - start_time}秒")
+
+ # 使用opus编码器处理PCM数据
+ opus_datas = []
+ pcm_data = response.content
+
+ # 计算每帧的字节数
+ frame_bytes = int(
+ self.opus_encoder.sample_rate
+ * self.opus_encoder.channels
+ * self.opus_encoder.frame_size_ms
+ / 1000
+ * 2
+ )
+
+ # 分帧处理PCM数据
+ for i in range(0, len(pcm_data), frame_bytes):
+ frame = pcm_data[i : i + frame_bytes]
+ if len(frame) < frame_bytes:
+ # 最后一帧可能不足,用0填充
+ frame = frame + b"\x00" * (frame_bytes - len(frame))
+
+ self.opus_encoder.encode_pcm_to_opus_stream(
+ frame,
+ end_of_stream=(i + frame_bytes >= len(pcm_data)),
+ callback=lambda opus: opus_datas.append(opus)
+ )
+
+ return opus_datas
+
+ except Exception as e:
+ logger.bind(tag=TAG).error(f"TTS请求异常: {e}")
+ return []
\ No newline at end of file
diff --git a/main/xiaozhi-server/core/providers/tts/linkerai.py b/main/xiaozhi-server/core/providers/tts/linkerai.py
index a855e8ed..9b5e946f 100644
--- a/main/xiaozhi-server/core/providers/tts/linkerai.py
+++ b/main/xiaozhi-server/core/providers/tts/linkerai.py
@@ -1,8 +1,10 @@
import os
+import time
import queue
-import asyncio
-import traceback
import aiohttp
+import asyncio
+import requests
+import traceback
from config.logger import setup_logging
from core.utils.tts import MarkdownCleaner
from core.providers.tts.base import TTSProviderBase
@@ -42,7 +44,6 @@ class TTSProvider(TTSProviderBase):
self.processed_chars = 0
self.tts_text_buff = []
self.before_stop_play_files.clear()
- self.reset_flow_controller()
elif ContentType.TEXT == message.content_type:
self.tts_text_buff.append(message.content_detail)
segment_text = self._get_segment_text()
@@ -110,10 +111,6 @@ class TTSProvider(TTSProviderBase):
finally:
return None
- ###################################################################################
- # linkerai单流式TTS重写父类的方法--结束
- ###################################################################################
-
async def text_to_speak(self, text, is_last):
"""流式处理TTS音频,每句只推送一次音频列表"""
await self._tts_request(text, is_last)
@@ -200,3 +197,71 @@ class TTSProvider(TTSProviderBase):
except Exception as e:
logger.bind(tag=TAG).error(f"TTS请求异常: {e}")
self.tts_audio_queue.put((SentenceType.LAST, [], None))
+
+ def to_tts(self, text: str) -> list:
+ """非流式TTS处理,用于测试及保存音频文件的场景
+ Args:
+ text: 要转换的文本
+ Returns:
+ list: 返回opus编码后的音频数据列表
+ """
+ start_time = time.time()
+ text = MarkdownCleaner.clean_markdown(text)
+
+ params = {
+ "tts_text": text,
+ "spk_id": self.voice,
+ "frame_duration": 60,
+ "stream": False,
+ "target_sr": 16000,
+ "audio_format": self.audio_format,
+ "instruct_text": "请生成一段自然流畅的语音",
+ }
+ headers = {
+ "Authorization": f"Bearer {self.access_token}",
+ "Content-Type": "application/json",
+ }
+
+ try:
+ with requests.get(
+ self.api_url, params=params, headers=headers, timeout=5
+ ) as response:
+ if response.status_code != 200:
+ logger.bind(tag=TAG).error(
+ f"TTS请求失败: {response.status_code}, {response.text}"
+ )
+ return []
+
+ logger.info(f"TTS请求成功: {text}, 耗时: {time.time() - start_time}秒")
+
+ # 使用opus编码器处理PCM数据
+ opus_datas = []
+ pcm_data = response.content
+
+ # 计算每帧的字节数
+ frame_bytes = int(
+ self.opus_encoder.sample_rate
+ * self.opus_encoder.channels
+ * self.opus_encoder.frame_size_ms
+ / 1000
+ * 2
+ )
+
+ # 分帧处理PCM数据
+ for i in range(0, len(pcm_data), frame_bytes):
+ frame = pcm_data[i : i + frame_bytes]
+ if len(frame) < frame_bytes:
+ # 最后一帧可能不足,用0填充
+ frame = frame + b"\x00" * (frame_bytes - len(frame))
+
+ self.opus_encoder.encode_pcm_to_opus_stream(
+ frame,
+ end_of_stream=(i + frame_bytes >= len(pcm_data)),
+ callback=lambda opus: opus_datas.append(opus)
+ )
+
+ return opus_datas
+
+ except Exception as e:
+ logger.bind(tag=TAG).error(f"TTS请求异常: {e}")
+ return []
\ No newline at end of file
diff --git a/main/xiaozhi-server/core/providers/tts/paddle_speech.py b/main/xiaozhi-server/core/providers/tts/paddle_speech.py
index 7f1c6a23..864cf505 100644
--- a/main/xiaozhi-server/core/providers/tts/paddle_speech.py
+++ b/main/xiaozhi-server/core/providers/tts/paddle_speech.py
@@ -1,13 +1,15 @@
-import asyncio
-import json
-import base64
-import aiohttp
-import numpy as np
import io
import wave
+import json
+import base64
+import asyncio
import websockets
-from core.providers.tts.base import TTSProviderBase
+import numpy as np
+from datetime import datetime
from config.logger import setup_logging
+from core.providers.tts.base import TTSProviderBase
+
+
TAG = __name__
logger = setup_logging()
@@ -18,11 +20,12 @@ class TTSProvider(TTSProviderBase):
super().__init__(config, delete_audio_file)
self.url = config.get("url", "ws://192.168.1.10:8092/paddlespeech/tts/streaming")
self.protocol = config.get("protocol", "websocket")
+
if config.get("private_voice"):
self.spk_id = int(config.get("private_voice"))
else:
- self.spk_id = int(config.get("spk_id", "0"))
-
+ self.spk_id = int(config.get("spk_id", "0"))
+
sample_rate = config.get("sample_rate", 24000)
self.sample_rate = float(sample_rate) if sample_rate else 24000
@@ -32,7 +35,21 @@ class TTSProvider(TTSProviderBase):
volume = config.get("volume", 1.0)
self.volume = float(volume) if volume else 1.0
- self.save_path = config.get("save_path", "./streaming_tts.wav")
+ self.delete_audio_file = config.get("delete_audio", True)
+ if not self.delete_audio_file:
+ timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
+ save_path = config.get("save_path")
+ if save_path:
+ if not save_path.endswith('.wav'):
+ save_path = f"{save_path}_{timestamp}.wav"
+ else:
+ other_path = save_path[:-4]
+ save_path = f"{other_path}_{timestamp}.wav"
+ self.save_path = save_path
+ else:
+ self.save_path = f"./streaming_tts_{timestamp}.wav"
+ else:
+ self.save_path = None
async def pcm_to_wav(self, pcm_data: bytes, sample_rate: int = 24000, num_channels: int = 1,
bits_per_sample: int = 16) -> bytes:
@@ -58,43 +75,9 @@ class TTSProvider(TTSProviderBase):
async def text_to_speak(self, text, output_file):
if self.protocol == "websocket":
return await self.text_streaming(text, output_file)
- elif self.protocol == "http":
- return await self.text(text, output_file)
else:
raise ValueError("Unsupported protocol. Please use 'websocket' or 'http'.")
- async def text(self, text, output_file):
- request_json = {
- "text": text,
- "spk_id": self.spk_id,
- "speed": self.speed,
- "volume": self.volume,
- "sample_rate": self.sample_rate,
- "save_path": self.save_path
- }
-
- try:
- async with aiohttp.ClientSession() as session:
- async with session.post(self.url, json=request_json) as resp:
- if resp.status == 200:
- resp_json = await resp.json()
- if resp_json.get("success"):
- data = resp_json["result"]
- audio_bytes = base64.b64decode(data["audio"])
- if output_file:
- with open(output_file, "wb") as file_to_save:
- file_to_save.write(audio_bytes)
- else:
- return audio_bytes
- else:
- raise Exception(
- f"Error: {resp_json.get('message', 'Unknown error')} while processing text: {text}")
- else:
- raise Exception(
- f"HTTP Error: {resp.status} - {await resp.text()} while processing text: {text}")
- except Exception as e:
- raise Exception(f"Error during TTS HTTP request: {e} while processing text: {text}")
-
async def text_streaming(self, text, output_file):
try:
# 使用 websockets 异步连接到 WebSocket 服务器
@@ -151,6 +134,12 @@ class TTSProvider(TTSProviderBase):
# 接收结束响应避免服务抛出异常
await ws.recv()
+ # 根据配置决定是否保存文件
+ if not self.delete_audio_file and self.save_path:
+ with open(self.save_path, "wb") as f:
+ f.write(wav_data)
+ logger.bind(tag=TAG).info(f"音频文件已保存到: {self.save_path}")
+
# 返回或保存音频数据
if output_file:
with open(output_file, "wb") as file_to_save:
@@ -159,4 +148,4 @@ class TTSProvider(TTSProviderBase):
return wav_data
except Exception as e:
- raise Exception(f"Error during TTS WebSocket request: {e} while processing text: {text}")
+ raise Exception(f"Error during TTS WebSocket request: {e} while processing text: {text}")
\ No newline at end of file
diff --git a/main/xiaozhi-server/core/utils/audio_flow_control.py b/main/xiaozhi-server/core/utils/audio_flow_control.py
deleted file mode 100644
index 98c29676..00000000
--- a/main/xiaozhi-server/core/utils/audio_flow_control.py
+++ /dev/null
@@ -1,186 +0,0 @@
-"""
-音频流控模块
-包含令牌桶算法和音频流控制器的实现
-"""
-
-import asyncio
-import time
-import threading
-from collections import deque
-from typing import Optional, Dict, Any
-
-
-class TokenBucket:
- """令牌桶实现,用于限流控制"""
-
- def __init__(self, capacity: int, refill_rate: float, initial_tokens: Optional[int] = None):
- """
- 初始化令牌桶
-
- Args:
- capacity: 桶容量(最大令牌数)
- refill_rate: 令牌补充速率(每秒补充的令牌数)
- initial_tokens: 初始令牌数,默认为桶容量
- """
- self.capacity = capacity
- self.refill_rate = refill_rate
- self.tokens = initial_tokens if initial_tokens is not None else capacity
- self.last_refill_time = time.time()
- self.lock = threading.Lock()
-
- def get_tokens(self, requested_tokens: int = 1) -> bool:
- """
- 获取指定数量的令牌
-
- Args:
- requested_tokens: 请求的令牌数量
-
- Returns:
- bool: 是否成功获取到令牌
- """
- with self.lock:
- self._refill_tokens()
-
- if self.tokens >= requested_tokens:
- self.tokens -= requested_tokens
- return True
- else:
- return False
-
- def get_available_tokens(self) -> int:
- """获取当前可用令牌数"""
- with self.lock:
- self._refill_tokens()
- return int(self.tokens)
-
- def _refill_tokens(self):
- """内部方法:补充令牌"""
- current_time = time.time()
- time_passed = current_time - self.last_refill_time
- tokens_to_add = time_passed * self.refill_rate
-
- self.tokens = min(self.capacity, self.tokens + tokens_to_add)
- self.last_refill_time = current_time
-
-
-class AudioFlowController:
- """音频流控制器,基于令牌桶算法控制音频数据发送"""
-
- def __init__(self, max_device_buffer: int = 3000, refill_rate: float = 20):
- """
- 初始化音频流控制器
-
- Args:
- max_device_buffer: 设备端最大缓冲区大小(Opus帧数)
- refill_rate: 令牌补充速率(每秒允许发送的帧数)
- """
- self.max_device_buffer = max_device_buffer
- self.token_bucket = TokenBucket(
- capacity=max_device_buffer,
- refill_rate=refill_rate,
- initial_tokens=max_device_buffer // 2 # 初始令牌为容量的一半
- )
- self.sent_frames_count = 0 # 已发送帧数计数
- self.device_consumed_frames = 0 # 设备端已消费帧数
- self.pending_queue = deque() # 等待发送的数据队列
- self._lock = threading.Lock()
-
- def can_send_frames(self, frame_count: int) -> bool:
- """
- 检查是否可以发送指定数量的帧
-
- Args:
- frame_count: 要发送的帧数
-
- Returns:
- bool: 是否可以发送
- """
- with self._lock:
- # 检查设备端缓冲区是否会溢出
- estimated_device_buffer = self.sent_frames_count - self.device_consumed_frames
- if estimated_device_buffer + frame_count > self.max_device_buffer:
- return False
-
- # 检查令牌桶是否有足够令牌
- return self.token_bucket.get_tokens(frame_count)
-
- def update_device_consumption(self, consumed_frames: int):
- """
- 更新设备端消费的帧数
-
- Args:
- consumed_frames: 设备端消费的帧数
- """
- with self._lock:
- self.device_consumed_frames += consumed_frames
-
- def record_sent_frames(self, frame_count: int):
- """
- 记录已发送的帧数
-
- Args:
- frame_count: 发送的帧数
- """
- with self._lock:
- self.sent_frames_count += frame_count
-
- def get_status(self) -> Dict[str, Any]:
- """获取流控状态信息"""
- with self._lock:
- estimated_buffer = self.sent_frames_count - self.device_consumed_frames
- return {
- "sent_frames": self.sent_frames_count,
- "consumed_frames": self.device_consumed_frames,
- "estimated_device_buffer": estimated_buffer,
- "available_tokens": self.token_bucket.get_available_tokens(),
- "pending_queue_size": len(self.pending_queue),
- "buffer_usage_percent": (estimated_buffer / self.max_device_buffer) * 100
- }
-
- def reset(self):
- """重置流控状态"""
- with self._lock:
- self.sent_frames_count = 0
- self.device_consumed_frames = 0
- self.pending_queue.clear()
- # 重新初始化令牌桶
- self.token_bucket = TokenBucket(
- capacity=self.max_device_buffer,
- refill_rate=self.token_bucket.refill_rate,
- initial_tokens=self.max_device_buffer // 2
- )
-
-
-# 流控配置常量
-class FlowControlConfig:
- """流控配置常量"""
- # Opus 编码参数
- OPUS_FRAME_DURATION_MS = 60 # Opus帧时长(毫秒)
- OPUS_FRAMES_PER_SECOND = 1000 / OPUS_FRAME_DURATION_MS # 每秒帧数
-
- # 默认流控参数
- DEFAULT_MAX_DEVICE_BUFFER = 40 # 设备端最大缓冲帧数
- DEFAULT_REFILL_RATE = OPUS_FRAMES_PER_SECOND # 默认令牌补充速率(帧/秒)
- DEFAULT_MAX_WAIT_TIME = 5.0 # 流控最大等待时间(秒)
- DEFAULT_RETRY_INTERVAL = 0.06 # 流控重试间隔(秒)
-
- # 预缓冲参数
- PRE_BUFFER_FRAMES = 3 # 预缓冲帧数
-
- @classmethod
- def create_flow_controller(cls, max_buffer: Optional[int] = None,
- refill_rate: Optional[float] = None) -> AudioFlowController:
- """
- 创建流控制器的工厂方法
-
- Args:
- max_buffer: 最大缓冲区大小,使用默认值如果为None
- refill_rate: 令牌补充速率,使用默认值如果为None
-
- Returns:
- AudioFlowController: 配置好的流控制器实例
- """
- return AudioFlowController(
- max_device_buffer=max_buffer or cls.DEFAULT_MAX_DEVICE_BUFFER,
- refill_rate=refill_rate or cls.DEFAULT_REFILL_RATE
- )
\ No newline at end of file
diff --git a/main/xiaozhi-server/core/utils/opus_encoder_utils.py b/main/xiaozhi-server/core/utils/opus_encoder_utils.py
index 8ca406d3..ae7066ce 100644
--- a/main/xiaozhi-server/core/utils/opus_encoder_utils.py
+++ b/main/xiaozhi-server/core/utils/opus_encoder_utils.py
@@ -6,10 +6,9 @@ Opus编码工具类
import logging
import traceback
import numpy as np
-from typing import Optional, Callable, Any
from opuslib_next import Encoder
from opuslib_next import constants
-
+from typing import Optional, Callable, Any
class OpusEncoderUtils:
"""PCM到Opus的编码器"""
@@ -130,4 +129,4 @@ class OpusEncoderUtils:
def close(self):
"""关闭编码器并释放资源"""
# opuslib没有明确的关闭方法,Python的垃圾回收会处理
- pass
+ pass
\ No newline at end of file
diff --git a/main/xiaozhi-server/core/utils/p3.py b/main/xiaozhi-server/core/utils/p3.py
index 415e5366..c75b968e 100644
--- a/main/xiaozhi-server/core/utils/p3.py
+++ b/main/xiaozhi-server/core/utils/p3.py
@@ -1,12 +1,15 @@
-import io
import struct
-from typing import Callable, Any
+def decode_opus_from_file(input_file):
+ """
+ 从p3文件中解码 Opus 数据,并返回一个 Opus 数据包的列表以及总时长。
+ """
+ opus_datas = []
+ total_frames = 0
+ sample_rate = 16000 # 文件采样率
+ frame_duration_ms = 60 # 帧时长
+ frame_size = int(sample_rate * frame_duration_ms / 1000)
-def decode_opus_from_file_stream(input_file, callback: Callable[[Any], Any]):
- """
- 从p3文件中解码 Opus 数据,由 callback 处理 Opus 数据包。
- """
with open(input_file, 'rb') as f:
while True:
# 读取头部(4字节):[1字节类型,1字节保留,2字节长度]
@@ -22,13 +25,23 @@ def decode_opus_from_file_stream(input_file, callback: Callable[[Any], Any]):
if len(opus_data) != data_len:
raise ValueError(f"Data length({len(opus_data)}) mismatch({data_len}) in the file.")
- callback(opus_data)
+ opus_datas.append(opus_data)
+ total_frames += 1
+ # 计算总时长
+ total_duration = (total_frames * frame_duration_ms) / 1000.0
+ return opus_datas, total_duration
-def decode_opus_from_bytes_stream(input_bytes, callback: Callable[[Any], Any]):
+def decode_opus_from_bytes(input_bytes):
"""
- 从p3二进制数据中解码 Opus 数据,由 callback 处理 Opus 数据包。
+ 从p3二进制数据中解码 Opus 数据,并返回一个 Opus 数据包的列表以及总时长。
"""
+ import io
+ opus_datas = []
+ total_frames = 0
+ sample_rate = 16000 # 文件采样率
+ frame_duration_ms = 60 # 帧时长
+ frame_size = int(sample_rate * frame_duration_ms / 1000)
f = io.BytesIO(input_bytes)
while True:
@@ -39,4 +52,8 @@ def decode_opus_from_bytes_stream(input_bytes, callback: Callable[[Any], Any]):
opus_data = f.read(data_len)
if len(opus_data) != data_len:
raise ValueError(f"Data length({len(opus_data)}) mismatch({data_len}) in the bytes.")
- callback(opus_data)
+ opus_datas.append(opus_data)
+ total_frames += 1
+
+ total_duration = (total_frames * frame_duration_ms) / 1000.0
+ return opus_datas, total_duration
\ No newline at end of file
diff --git a/main/xiaozhi-server/core/utils/util.py b/main/xiaozhi-server/core/utils/util.py
index 1eec7327..27471e58 100644
--- a/main/xiaozhi-server/core/utils/util.py
+++ b/main/xiaozhi-server/core/utils/util.py
@@ -1,16 +1,17 @@
-import json
-import socket
-import subprocess
import re
import os
-from io import BytesIO
-from typing import Callable, Any
-from core.utils import p3
-import numpy as np
-import requests
-import opuslib_next
-from pydub import AudioSegment
+import json
import copy
+import wave
+import socket
+import requests
+import subprocess
+import numpy as np
+import opuslib_next
+from io import BytesIO
+from core.utils import p3
+from pydub import AudioSegment
+from typing import Callable, Any
TAG = __name__
emoji_map = {
@@ -228,6 +229,56 @@ def audio_to_data_stream(audio_file_path, is_opus=True, callback: Callable[[Any]
raw_data = audio.raw_data
pcm_to_data_stream(raw_data, is_opus, callback)
+def audio_to_data(audio_file_path: str, is_opus: bool = True) -> list[bytes]:
+ """
+ 将音频文件转换为Opus/PCM编码的帧列表
+ Args:
+ audio_file_path: 音频文件路径
+ is_opus: 是否进行Opus编码
+ """
+ # 获取文件后缀名
+ file_type = os.path.splitext(audio_file_path)[1]
+ if file_type:
+ file_type = file_type.lstrip(".")
+ # 读取音频文件,-nostdin 参数:不要从标准输入读取数据,否则FFmpeg会阻塞
+ audio = AudioSegment.from_file(
+ audio_file_path, format=file_type, parameters=["-nostdin"]
+ )
+
+ # 转换为单声道/16kHz采样率/16位小端编码(确保与编码器匹配)
+ audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2)
+
+ # 获取原始PCM数据(16位小端)
+ raw_data = audio.raw_data
+
+ # 初始化Opus编码器
+ encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO)
+
+ # 编码参数
+ frame_duration = 60 # 60ms per frame
+ frame_size = int(16000 * frame_duration / 1000) # 960 samples/frame
+
+ datas = []
+ # 按帧处理所有音频数据(包括最后一帧可能补零)
+ for i in range(0, len(raw_data), frame_size * 2): # 16bit=2bytes/sample
+ # 获取当前帧的二进制数据
+ chunk = raw_data[i : i + frame_size * 2]
+
+ # 如果最后一帧不足,补零
+ if len(chunk) < frame_size * 2:
+ chunk += b"\x00" * (frame_size * 2 - len(chunk))
+
+ if is_opus:
+ # 转换为numpy数组处理
+ np_frame = np.frombuffer(chunk, dtype=np.int16)
+ # 编码Opus数据
+ frame_data = encoder.encode(np_frame.tobytes(), frame_size)
+ else:
+ frame_data = chunk if isinstance(chunk, bytes) else bytes(chunk)
+
+ datas.append(frame_data)
+
+ return datas
def audio_bytes_to_data_stream(audio_bytes, file_type, is_opus, callback: Callable[[Any], Any]) -> None:
"""
@@ -273,6 +324,31 @@ def pcm_to_data_stream(raw_data, is_opus=True, callback: Callable[[Any], Any] =
frame_data = chunk if isinstance(chunk, bytes) else bytes(chunk)
callback(frame_data)
+def opus_datas_to_wav_bytes(opus_datas, sample_rate=16000, channels=1):
+ """
+ 将opus帧列表解码为wav字节流
+ """
+ decoder = opuslib_next.Decoder(sample_rate, channels)
+ pcm_datas = []
+
+ frame_duration = 60 # ms
+ frame_size = int(sample_rate * frame_duration / 1000) # 960
+
+ for opus_frame in opus_datas:
+ # 解码为PCM(返回bytes,2字节/采样点)
+ pcm = decoder.decode(opus_frame, frame_size)
+ pcm_datas.append(pcm)
+
+ pcm_bytes = b"".join(pcm_datas)
+
+ # 写入wav字节流
+ wav_buffer = BytesIO()
+ with wave.open(wav_buffer, "wb") as wf:
+ wf.setnchannels(channels)
+ wf.setsampwidth(2) # 16bit
+ wf.setframerate(sample_rate)
+ wf.writeframes(pcm_bytes)
+ return wav_buffer.getvalue()
def check_vad_update(before_config, new_config):
if (
diff --git a/main/xiaozhi-server/core/utils/wakeup_word.py b/main/xiaozhi-server/core/utils/wakeup_word.py
new file mode 100644
index 00000000..d2f4fb32
--- /dev/null
+++ b/main/xiaozhi-server/core/utils/wakeup_word.py
@@ -0,0 +1,140 @@
+import os
+import re
+import yaml
+import time
+import hashlib
+import portalocker
+from typing import Dict
+
+
+class FileLock:
+ def __init__(self, file, timeout=5):
+ self.file = file
+ self.timeout = timeout
+ self.start_time = None
+
+ def __enter__(self):
+ self.start_time = time.time()
+ while True:
+ try:
+ portalocker.lock(self.file, portalocker.LOCK_EX | portalocker.LOCK_NB)
+ return self.file
+ except portalocker.LockException:
+ if time.time() - self.start_time > self.timeout:
+ raise TimeoutError("获取文件锁超时")
+ time.sleep(0.1)
+
+ def __exit__(self, exc_type, exc_val, exc_tb):
+ portalocker.unlock(self.file)
+
+
+class WakeupWordsConfig:
+ def __init__(self):
+ self.config_file = "data/.wakeup_words.yaml"
+ self.assets_dir = "config/assets/wakeup_words"
+ self._ensure_directories()
+ self._config_cache = None
+ self._last_load_time = 0
+ self._cache_ttl = 1 # 缓存有效期(秒)
+ self._lock_timeout = 5 # 文件锁超时时间(秒)
+
+ def _ensure_directories(self):
+ """确保必要的目录存在"""
+ os.makedirs(os.path.dirname(self.config_file), exist_ok=True)
+ os.makedirs(self.assets_dir, exist_ok=True)
+
+ def _load_config(self) -> Dict:
+ """加载配置文件,使用缓存机制"""
+ current_time = time.time()
+
+ # 如果缓存有效,直接返回缓存
+ if (
+ self._config_cache is not None
+ and current_time - self._last_load_time < self._cache_ttl
+ ):
+ return self._config_cache
+
+ try:
+ with open(self.config_file, "a+") as f:
+ with FileLock(f, timeout=self._lock_timeout):
+ f.seek(0)
+ content = f.read()
+ config = yaml.safe_load(content) if content else {}
+ self._config_cache = config
+ self._last_load_time = current_time
+ return config
+ except (TimeoutError, IOError) as e:
+ print(f"加载配置文件失败: {e}")
+ return {}
+ except Exception as e:
+ print(f"加载配置文件时发生未知错误: {e}")
+ return {}
+
+ def _save_config(self, config: Dict):
+ """保存配置到文件,使用文件锁保护"""
+ try:
+ with open(self.config_file, "w") as f:
+ with FileLock(f, timeout=self._lock_timeout):
+ yaml.dump(config, f, allow_unicode=True)
+ self._config_cache = config
+ self._last_load_time = time.time()
+ except (TimeoutError, IOError) as e:
+ print(f"保存配置文件失败: {e}")
+ raise
+ except Exception as e:
+ print(f"保存配置文件时发生未知错误: {e}")
+ raise
+
+ def get_wakeup_response(self, voice: str) -> Dict:
+ voice = hashlib.md5(voice.encode()).hexdigest()
+ """获取唤醒词回复配置"""
+ config = self._load_config()
+
+ if not config or voice not in config:
+ return None
+
+ # 检查文件大小
+ file_path = config[voice]["file_path"]
+ if not os.path.exists(file_path) or os.stat(file_path).st_size < (15 * 1024):
+ return None
+
+ return config[voice]
+
+ def update_wakeup_response(self, voice: str, file_path: str, text: str):
+ """更新唤醒词回复配置"""
+ try:
+ # 过滤表情符号
+ filtered_text = re.sub(r'[\U0001F600-\U0001F64F\U0001F900-\U0001F9FF]', '', text)
+
+ config = self._load_config()
+ voice_hash = hashlib.md5(voice.encode()).hexdigest()
+ config[voice_hash] = {
+ "voice": voice,
+ "file_path": file_path,
+ "time": time.time(),
+ "text": filtered_text,
+ }
+ self._save_config(config)
+ except Exception as e:
+ print(f"更新唤醒词回复配置失败: {e}")
+ raise
+
+ def generate_file_path(self, voice: str) -> str:
+ """生成音频文件路径,使用voice的哈希值作为文件名"""
+ try:
+ # 生成voice的哈希值
+ voice_hash = hashlib.md5(voice.encode()).hexdigest()
+ file_path = os.path.join(self.assets_dir, f"{voice_hash}.wav")
+
+ # 如果文件已存在,先删除
+ if os.path.exists(file_path):
+ try:
+ os.remove(file_path)
+ except Exception as e:
+ print(f"删除已存在的音频文件失败: {e}")
+ raise
+
+ return file_path
+ except Exception as e:
+ print(f"生成音频文件路径失败: {e}")
+ raise
\ No newline at end of file
diff --git a/main/xiaozhi-server/performance_tester/performance_tester_asr.py b/main/xiaozhi-server/performance_tester/performance_tester_asr.py
index e4d8c6d8..db359e28 100644
--- a/main/xiaozhi-server/performance_tester/performance_tester_asr.py
+++ b/main/xiaozhi-server/performance_tester/performance_tester_asr.py
@@ -2,35 +2,56 @@ import asyncio
import logging
import os
import time
-from typing import Dict
-
+import concurrent.futures
+from typing import Dict, Optional
import aiohttp
from tabulate import tabulate
from core.utils.asr import create_instance as create_stt_instance
-from config.settings import load_config
# 设置全局日志级别为WARNING,抑制INFO级别日志
logging.basicConfig(level=logging.WARNING)
description = "语音识别模型性能测试"
-
class ASRPerformanceTester:
def __init__(self):
- self.config = load_config()
+ self.config = self._load_config_from_data_dir()
self.test_wav_list = self._load_test_wav_files()
self.results = {"stt": {}}
-
+
# 调试日志
print(f"[DEBUG] 加载的ASR配置: {self.config.get('ASR', {})}")
print(f"[DEBUG] 音频文件数量: {len(self.test_wav_list)}")
+ def _load_config_from_data_dir(self) -> Dict:
+ """从 data 目录加载所有 .config.yaml 文件的配置"""
+ config = {"ASR": {}}
+ data_dir = os.path.join(os.getcwd(), "data")
+ print(f"[DEBUG] 扫描配置文件目录: {data_dir}")
+
+ for root, _, files in os.walk(data_dir):
+ for file in files:
+ if file.endswith(".config.yaml"):
+ file_path = os.path.join(root, file)
+ try:
+ with open(file_path, "r", encoding="utf-8") as f:
+ import yaml
+ file_config = yaml.safe_load(f)
+ # 兼容大小写的 ASR/asr 配置
+ asr_config = file_config.get("ASR") or file_config.get("asr")
+ if asr_config:
+ config["ASR"].update(asr_config)
+ print(f"[DEBUG] 从 {file_path} 加载 ASR 配置成功")
+ except Exception as e:
+ print(f" 加载配置文件 {file_path} 失败: {str(e)}")
+ return config
+
def _load_test_wav_files(self) -> list:
"""加载测试用的音频文件(添加路径调试)"""
wav_root = os.path.join(os.getcwd(), "config", "assets")
print(f"[DEBUG] 音频文件目录: {wav_root}")
test_wav_list = []
-
+
if os.path.exists(wav_root):
file_list = os.listdir(wav_root)
print(f"[DEBUG] 找到音频文件: {file_list}")
@@ -43,18 +64,46 @@ class ASRPerformanceTester:
print(f" 目录不存在: {wav_root}")
return test_wav_list
- async def _test_stt(self, stt_name: str, config: Dict) -> Dict:
- """异步测试单个STT性能(跳过无效配置)"""
+ async def _test_single_audio(self, stt_name: str, stt, audio_data: bytes) -> Optional[float]:
+ """测试单个音频文件的性能"""
try:
+ start_time = time.time()
+ text, _ = await stt.speech_to_text([audio_data], "1", stt.audio_format)
+ if text is None:
+ return None
+
+ duration = time.time() - start_time
+
+ # 检测0.000s的异常时间
+ if abs(duration) < 0.001: # 小于1毫秒视为异常
+ print(f"{stt_name} 检测到异常时间: {duration:.6f}s (视为错误)")
+ return None
+
+ return duration
+ except Exception as e:
+ error_msg = str(e).lower()
+ if "502" in error_msg or "bad gateway" in error_msg:
+ print(f"{stt_name} 遇到502错误")
+ return None
+ return None
+
+ async def _test_stt_with_timeout(self, stt_name: str, config: Dict) -> Dict:
+ """异步测试单个STT性能,带超时控制"""
+ try:
+ # 检查配置有效性
token_fields = ["access_token", "api_key", "token"]
- # 忽略值为 "none" 的情况(需根据实际需求调整)
if any(
field in config
- and str(config[field]).lower() in ["你的", "placeholder"]
+ and str(config[field]).lower() in ["你的", "placeholder", "none", "null", ""]
for field in token_fields
):
- print(f" STT {stt_name} 未配置access_token/api_key,已跳过")
- return {"name": stt_name, "type": "stt", "errors": 1}
+ print(f" STT {stt_name} 未配置有效access_token/api_key,已跳过")
+ return {
+ "name": stt_name,
+ "type": "stt",
+ "errors": 1,
+ "error_type": "配置错误"
+ }
module_type = config.get("type", stt_name)
stt = create_stt_instance(module_type, config, delete_audio_file=True)
@@ -62,56 +111,203 @@ class ASRPerformanceTester:
print(f" 测试 STT: {stt_name}")
- # 测试第一个音频文件
- text, _ = await stt.speech_to_text(
- [self.test_wav_list[0]], "1", stt.audio_format
- )
- if text is None:
- print(f" {stt_name} 连接失败")
- return {"name": stt_name, "type": "stt", "errors": 1}
+ # 使用线程池和超时控制
+ loop = asyncio.get_event_loop()
+
+ # 测试第一个音频文件作为连通性检查
+ try:
+ with concurrent.futures.ThreadPoolExecutor() as executor:
+ future = executor.submit(
+ lambda: asyncio.run(self._test_single_audio(stt_name, stt, self.test_wav_list[0]))
+ )
+ first_result = await asyncio.wait_for(
+ asyncio.wrap_future(future), timeout=10.0
+ )
+
+ if first_result is None:
+ print(f" {stt_name} 连接失败")
+ return {
+ "name": stt_name,
+ "type": "stt",
+ "errors": 1,
+ "error_type": "网络错误"
+ }
+ except asyncio.TimeoutError:
+ print(f" {stt_name} 连接超时(10秒),跳过")
+ return {
+ "name": stt_name,
+ "type": "stt",
+ "errors": 1,
+ "error_type": "超时连接"
+ }
+ except Exception as e:
+ error_msg = str(e).lower()
+ if "502" in error_msg or "bad gateway" in error_msg:
+ print(f" {stt_name} 遇到502错误,跳过")
+ return {
+ "name": stt_name,
+ "type": "stt",
+ "errors": 1,
+ "error_type": "502网络错误"
+ }
+ print(f" {stt_name} 连接异常: {str(e)}")
+ return {
+ "name": stt_name,
+ "type": "stt",
+ "errors": 1,
+ "error_type": "网络错误"
+ }
- # 全量测试
+ # 全量测试,带超时控制
total_time = 0
+ valid_tests = 0
test_count = len(self.test_wav_list)
- for i, sentence in enumerate(self.test_wav_list, 1):
- start = time.time()
- text, _ = await stt.speech_to_text([sentence], "1", stt.audio_format)
- duration = time.time() - start
- total_time += duration
- print(f" {stt_name} [{i}/{test_count}] 耗时: {duration:.2f}s")
+
+ for i, audio_data in enumerate(self.test_wav_list, 1):
+ try:
+ with concurrent.futures.ThreadPoolExecutor() as executor:
+ future = executor.submit(
+ lambda: asyncio.run(self._test_single_audio(stt_name, stt, audio_data))
+ )
+ duration = await asyncio.wait_for(
+ asyncio.wrap_future(future), timeout=10.0
+ )
+
+ if duration is not None and duration > 0.001:
+ total_time += duration
+ valid_tests += 1
+ print(f" {stt_name} [{i}/{test_count}] 耗时: {duration:.2f}s")
+ else:
+ print(f" {stt_name} [{i}/{test_count}] 测试失败(含0.000s异常)")
+
+ except asyncio.TimeoutError:
+ print(f" {stt_name} [{i}/{test_count}] 超时(10秒),跳过")
+ continue
+ except Exception as e:
+ error_msg = str(e).lower()
+ if "502" in error_msg or "bad gateway" in error_msg:
+ print(f" {stt_name} [{i}/{test_count}] 502错误,跳过")
+ return {
+ "name": stt_name,
+ "type": "stt",
+ "errors": 1,
+ "error_type": "502网络错误"
+ }
+ print(f" {stt_name} [{i}/{test_count}] 异常: {str(e)}")
+ continue
+ # 检查有效测试数量
+ if valid_tests < test_count * 0.3: # 至少30%成功率
+ print(f" {stt_name} 成功测试过少({valid_tests}/{test_count}),可能网络不稳定")
+ return {
+ "name": stt_name,
+ "type": "stt",
+ "errors": 1,
+ "error_type": "网络错误"
+ }
+ if valid_tests == 0:
+ return {
+ "name": stt_name,
+ "type": "stt",
+ "errors": 1,
+ "error_type": "网络错误"
+ }
+
+ avg_time = total_time / valid_tests
return {
"name": stt_name,
"type": "stt",
- "avg_time": total_time / test_count,
+ "avg_time": avg_time,
+ "success_rate": f"{valid_tests}/{test_count}",
"errors": 0,
}
+
except Exception as e:
+ error_msg = str(e).lower()
+ if "502" in error_msg or "bad gateway" in error_msg:
+ error_type = "502网络错误"
+ elif "timeout" in error_msg:
+ error_type = "超时连接"
+ else:
+ error_type = "网络错误"
print(f"⚠️ {stt_name} 测试失败: {str(e)}")
- return {"name": stt_name, "type": "stt", "errors": 1}
+ return {
+ "name": stt_name,
+ "type": "stt",
+ "errors": 1,
+ "error_type": error_type
+ }
def _print_results(self):
- """打印测试结果"""
- stt_table = []
+ """打印测试结果,按响应时间排序"""
+ print("\n" + "=" * 50)
+ print("ASR 性能测试结果")
+ print("=" * 50)
+
+ if not self.results.get("stt"):
+ print("没有可用的测试结果")
+ return
+
+ headers = ["模型名称", "平均耗时(s)", "成功率", "状态"]
+ table_data = []
+
+ # 收集所有数据并分类
+ valid_results = []
+ error_results = []
+
for name, data in self.results["stt"].items():
if data["errors"] == 0:
- stt_table.append([name, f"{data['avg_time']:.3f}秒"])
+ # 正常结果
+ avg_time = f"{data['avg_time']:.3f}"
+ success_rate = data.get("success_rate", "N/A")
+ status = "✅ 正常"
+
+ # 保存用于排序的值
+ sort_key = data["avg_time"]
+
+ valid_results.append({
+ "name": name,
+ "avg_time": avg_time,
+ "success_rate": success_rate,
+ "status": status,
+ "sort_key": sort_key,
+ })
+ else:
+ # 错误结果
+ avg_time = "-"
+ success_rate = "0/N"
+
+ # 获取具体错误类型
+ error_type = data.get("error_type", "网络错误")
+ status = f"❌ {error_type}"
+
+ error_results.append([name, avg_time, success_rate, status])
- if stt_table:
- print("\nASR 性能排行:\n")
- print(
- tabulate(
- stt_table,
- headers=["模型名称", "平均耗时"],
- tablefmt="github",
- colalign=("left", "right"),
- )
- )
- else:
- print("\n 没有可用的ASR模块进行测试。")
+ # 按响应时间升序排序(从快到慢)
+ valid_results.sort(key=lambda x: x["sort_key"])
+
+ # 将排序后的有效结果转换为表格数据
+ for result in valid_results:
+ table_data.append([
+ result["name"],
+ result["avg_time"],
+ result["success_rate"],
+ result["status"],
+ ])
+
+ # 将错误结果添加到表格数据末尾
+ table_data.extend(error_results)
+
+ print(tabulate(table_data, headers=headers, tablefmt="grid"))
+ print("\n测试说明:")
+ print("- 超时控制:单个音频最大等待时间为10秒")
+ print("- 错误处理:自动跳过502错误、超时和网络异常的模型")
+ print("- 成功率:成功识别的音频数量/总测试音频数量")
+ print("- 排序规则:按平均耗时从快到慢排序,错误模型排最后")
+ print("\n测试完成!")
async def run(self):
- """执行全量异步测试"""
+ """执行全量异步测试"""
print("开始筛选可用ASR模块...")
if not self.config.get("ASR"):
print("配置中未找到 ASR 模块")
@@ -119,24 +315,33 @@ class ASRPerformanceTester:
all_tasks = []
for stt_name, config in self.config["ASR"].items():
- print(f"[DEBUG] 检查 ASR 模块: {stt_name}, 配置: {config}")
- all_tasks.append(self._test_stt(stt_name, config))
+ # 检查配置有效性
+ token_fields = ["access_token", "api_key", "token"]
+ if any(
+ field in config
+ and str(config[field]).lower() in ["你的", "placeholder", "none", "null", ""]
+ for field in token_fields
+ ):
+ print(f"ASR {stt_name} 未配置有效access_token/api_key,已跳过")
+ continue
+
+ print(f"添加 ASR 测试任务: {stt_name}")
+ all_tasks.append(self._test_stt_with_timeout(stt_name, config))
if not all_tasks:
print("没有可用的ASR模块进行测试。")
return
+ print(f"\n找到 {len(all_tasks)} 个可用ASR模块")
print("\n开始并发测试所有ASR模块...")
all_results = await asyncio.gather(*all_tasks, return_exceptions=True)
# 处理结果
for result in all_results:
if isinstance(result, dict) and result.get("type") == "stt":
- if result["errors"] == 0:
- self.results["stt"][result["name"]] = result
+ self.results["stt"][result["name"]] = result
# 打印结果
- print("\n测试完成")
self._print_results()
@@ -146,4 +351,4 @@ async def main():
if __name__ == "__main__":
- asyncio.run(main())
+ asyncio.run(main())
\ No newline at end of file
diff --git a/main/xiaozhi-server/performance_tester/performance_tester_stream_asr.py b/main/xiaozhi-server/performance_tester/performance_tester_stream_asr.py
new file mode 100644
index 00000000..e0309685
--- /dev/null
+++ b/main/xiaozhi-server/performance_tester/performance_tester_stream_asr.py
@@ -0,0 +1,404 @@
+import asyncio
+import time
+import json
+import uuid
+import os
+import websockets
+import gzip
+import hmac
+import base64
+import hashlib
+import random
+from urllib import parse
+from tabulate import tabulate
+from config.settings import load_config
+description = "流式ASR首词耗时测试"
+
+class AccessToken:
+ @staticmethod
+ def _encode_text(text):
+ encoded_text = parse.quote_plus(text)
+ return encoded_text.replace("+", "%20").replace("*", "%2A").replace("%7E", "~")
+
+ @staticmethod
+ def _encode_dict(dic):
+ keys = dic.keys()
+ dic_sorted = [(key, dic[key]) for key in sorted(keys)]
+ encoded_text = parse.urlencode(dic_sorted)
+ return encoded_text.replace("+", "%20").replace("*", "%2A").replace("%7E", "~")
+
+ @staticmethod
+ def create_token(access_key_id, access_key_secret):
+ parameters = {
+ "AccessKeyId": access_key_id,
+ "Action": "CreateToken",
+ "Format": "JSON",
+ "RegionId": "cn-shanghai",
+ "SignatureMethod": "HMAC-SHA1",
+ "SignatureNonce": str(uuid.uuid1()),
+ "SignatureVersion": "1.0",
+ "Timestamp": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
+ "Version": "2019-02-28",
+ }
+ query_string = AccessToken._encode_dict(parameters)
+ string_to_sign = (
+ "GET" + "&" + AccessToken._encode_text("/") + "&" + AccessToken._encode_text(query_string)
+ )
+ secreted_string = hmac.new(
+ bytes(access_key_secret + "&", encoding="utf-8"),
+ bytes(string_to_sign, encoding="utf-8"),
+ hashlib.sha1,
+ ).digest()
+ signature = base64.b64encode(secreted_string)
+ signature = AccessToken._encode_text(signature)
+ full_url = "http://nls-meta.cn-shanghai.aliyuncs.com/?Signature=%s&%s" % (signature, query_string)
+ response = requests.get(full_url)
+ if response.ok:
+ root_obj = response.json()
+ if "Token" in root_obj:
+ return root_obj["Token"]["Id"], root_obj["Token"]["ExpireTime"]
+ return None, None
+
+
+class DoubaoStreamASRPerformanceTester:
+ def __init__(self):
+ self.config = load_config()
+ self.test_audio_files = self._load_test_audio_files()
+ self.results = []
+
+ def _load_test_audio_files(self):
+ """加载测试用的音频文件"""
+ audio_root = os.path.join(os.getcwd(), "config", "assets")
+ test_files = []
+
+ if os.path.exists(audio_root):
+ for file_name in os.listdir(audio_root):
+ if file_name.endswith('.wav') or file_name.endswith('.pcm'):
+ with open(os.path.join(audio_root, file_name), 'rb') as f:
+ test_files.append(f.read())
+ return test_files
+
+ async def test_doubao_stream_asr(self, test_count=5):
+ """测试豆包流式ASR首词响应时间"""
+ if not self.test_audio_files:
+ print("没有找到测试音频文件")
+ return
+
+ asr_config = self.config["ASR"]["DoubaoStreamASR"]
+ latencies = []
+
+ for i in range(test_count):
+ try:
+ ws_url = "wss://openspeech.bytedance.com/api/v3/sauc/bigmodel"
+ appid = asr_config["appid"]
+ access_token = asr_config["access_token"]
+ uid = asr_config.get("uid", "streaming_asr_service")
+
+ start_time = time.time()
+
+ headers = {
+ "X-Api-App-Key": appid,
+ "X-Api-Access-Key": access_token,
+ "X-Api-Resource-Id": "volc.bigasr.sauc.duration",
+ "X-Api-Connect-Id": str(uuid.uuid4())
+ }
+
+ async with websockets.connect(
+ ws_url,
+ additional_headers=headers,
+ max_size=1000000000,
+ ping_interval=None,
+ ping_timeout=None,
+ close_timeout=10
+ ) as ws:
+ # 发送初始化请求
+ request_params = {
+ "app": {
+ "appid": appid,
+ "token": access_token
+ },
+ "user": {"uid": uid},
+ "request": {
+ "reqid": str(uuid.uuid4()),
+ "workflow": "audio_in,resample,partition,vad,fe,decode,itn,nlu_punctuate",
+ "show_utterances": True,
+ "result_type": "single",
+ "sequence": 1
+ },
+ "audio": {
+ "format": "pcm",
+ "codec": "pcm",
+ "rate": 16000,
+ "language": "zh-CN",
+ "bits": 16,
+ "channel": 1,
+ "sample_rate": 16000
+ }
+ }
+
+ payload_bytes = str.encode(json.dumps(request_params))
+ payload_bytes = gzip.compress(payload_bytes)
+ full_client_request = self._generate_header()
+ full_client_request.extend((len(payload_bytes)).to_bytes(4, "big"))
+ full_client_request.extend(payload_bytes)
+
+ await ws.send(full_client_request)
+
+ init_res = await ws.recv()
+ result = self._parse_response(init_res)
+
+ if "code" in result and result["code"] != 1000:
+ raise Exception(f"ASR服务初始化失败: {result.get('payload_msg', {}).get('error', '未知错误')}")
+
+ # 发送音频数据
+ audio_data = self.test_audio_files[0]
+ if audio_data.startswith(b'RIFF'):
+ audio_data = audio_data[44:]
+
+ # 直接发送原始音频数据,不进行opus解码
+ payload = gzip.compress(audio_data)
+ audio_request = bytearray(self._generate_audio_default_header())
+ audio_request.extend(len(payload).to_bytes(4, "big"))
+ audio_request.extend(payload)
+ await ws.send(audio_request)
+
+ # 等待第一个数据块
+ first_chunk = await ws.recv()
+ latency = time.time() - start_time
+ latencies.append(latency)
+ await ws.close()
+
+ except Exception as e:
+ print(f"第{i+1}次测试: {str(e)}")
+ latencies.append(0)
+
+ return self._calculate_result("豆包流式ASR", latencies, test_count)
+
+ async def test_aliyun_stream_asr(self, test_count=5):
+ """测试阿里云流式ASR首词响应时间"""
+ if not self.test_audio_files:
+ print("没有找到测试音频文件")
+ return
+
+ asr_config = self.config["ASR"]["AliyunStreamASR"]
+ latencies = []
+
+ for i in range(test_count):
+ try:
+ access_key_id = asr_config["access_key_id"]
+ access_key_secret = asr_config["access_key_secret"]
+ appkey = asr_config["appkey"]
+ host = asr_config.get("host", "nls-gateway-cn-shanghai.aliyuncs.com")
+
+ # 获取Token
+ token, _ = AccessToken.create_token(access_key_id, access_key_secret)
+ if not token:
+ raise Exception("无法获取阿里云ASR Token")
+
+ # 确定WebSocket URL
+ if "-internal." in host:
+ ws_url = f"ws://{host}/ws/v1"
+ else:
+ ws_url = f"wss://{host}/ws/v1"
+
+ start_time = time.time()
+ async with websockets.connect(
+ ws_url,
+ additional_headers={"X-NLS-Token": token},
+ max_size=1000000000,
+ ping_interval=None,
+ ping_timeout=None,
+ close_timeout=10
+ ) as ws:
+ # 发送开始请求
+ start_request = {
+ "header": {
+ "namespace": "SpeechTranscriber",
+ "name": "StartTranscription",
+ "status": 20000000,
+ "message_id": ''.join(random.choices('0123456789abcdef', k=32)),
+ "task_id": ''.join(random.choices('0123456789abcdef', k=32)),
+ "status_text": "Gateway:SUCCESS:Success.",
+ "appkey": appkey
+ },
+ "payload": {
+ "format": "pcm",
+ "sample_rate": 16000,
+ "enable_intermediate_result": True,
+ "enable_punctuation_prediction": True,
+ "enable_inverse_text_normalization": True,
+ "max_sentence_silence": asr_config.get("max_sentence_silence", 8000),
+ "enable_voice_detection": False,
+ }
+ }
+ await ws.send(json.dumps(start_request, ensure_ascii=False))
+
+ # 等待服务器准备
+ start_response = await ws.recv()
+ response_data = json.loads(start_response)
+ if response_data["header"]["name"] != "TranscriptionStarted":
+ raise Exception("阿里云ASR服务初始化失败")
+
+ # 发送音频数据
+ audio_data = self.test_audio_files[0]
+ if audio_data.startswith(b'RIFF'):
+ audio_data = audio_data[44:] # 去掉WAV头
+
+ await ws.send(audio_data)
+
+ # 等待第一个结果
+ while True:
+ response = await ws.recv()
+ if isinstance(response, str):
+ result = json.loads(response)
+ if result["header"]["name"] == "TranscriptionResultChanged":
+ latency = time.time() - start_time
+ latencies.append(latency)
+ break
+ elif result["header"]["name"] == "TaskFailed":
+ raise Exception(f"阿里云ASR识别失败: {result.get('payload', {}).get('error_info', '未知错误')}")
+
+ # 发送停止请求
+ stop_msg = {
+ "header": {
+ "namespace": "SpeechTranscriber",
+ "name": "StopTranscription",
+ "status": 20000000,
+ "message_id": ''.join(random.choices('0123456789abcdef', k=32)),
+ "status_text": "Client:Stop",
+ "appkey": appkey
+ }
+ }
+ await ws.send(json.dumps(stop_msg, ensure_ascii=False))
+ await ws.close()
+
+ except Exception as e:
+ print(f"第{i+1}次测试: {str(e)}")
+ latencies.append(0)
+
+ return self._calculate_result("阿里云流式ASR", latencies, test_count)
+
+ def _generate_header(self):
+ """生成请求头"""
+ header = bytearray()
+ header.append((0x01 << 4) | 0x01)
+ header.append((0x01 << 4) | 0x00)
+ header.append((0x01 << 4) | 0x01)
+ header.append(0x00)
+ return header
+
+ def _generate_audio_default_header(self):
+ """生成音频请求头"""
+ return self._generate_header()
+
+ def _parse_response(self, res: bytes) -> dict:
+ """解析响应"""
+ try:
+ if len(res) < 4:
+ return {"error": "响应数据长度不足"}
+
+ header = res[:4]
+ message_type = header[1] >> 4
+
+ if message_type == 0x0F:
+ code = int.from_bytes(res[4:8], "big", signed=False)
+ msg_length = int.from_bytes(res[8:12], "big", signed=False)
+ error_msg = json.loads(res[12:].decode("utf-8"))
+ return {
+ "code": code,
+ "msg_length": msg_length,
+ "payload_msg": error_msg
+ }
+
+ try:
+ json_data = res[12:].decode("utf-8")
+ return {"payload_msg": json.loads(json_data)}
+ except (UnicodeDecodeError, json.JSONDecodeError):
+ return {"error": "JSON解析失败"}
+
+ except Exception:
+ return {"error": "解析响应失败"}
+
+ def _calculate_result(self, service_name, latencies, test_count):
+ """计算结果"""
+ valid_latencies = [l for l in latencies if l > 0]
+ if valid_latencies:
+ avg_latency = sum(valid_latencies) / len(valid_latencies)
+ status = f"成功({len(valid_latencies)}/{test_count}次有效)"
+ else:
+ avg_latency = 0
+ status = "失败: 所有测试均失败"
+ return {"name": service_name, "latency": avg_latency, "status": status}
+
+ def _print_results(self, test_count):
+ """打印测试结果"""
+ if not self.results:
+ print("没有有效的ASR测试结果")
+ return
+
+ print(f"\n{'='*60}")
+ print("流式ASR首词响应时间测试结果")
+ print(f"{'='*60}")
+ print(f"测试次数: 每个ASR服务测试 {test_count} 次")
+
+ # 排序结果:成功优先,按延迟升序
+ success_results = sorted(
+ [r for r in self.results if "成功" in r["status"]],
+ key=lambda x: x["latency"]
+ )
+ failed_results = [r for r in self.results if "成功" not in r["status"]]
+
+ table_data = [
+ [r["name"], f"{r['latency']:.3f}", r["status"]]
+ for r in success_results + failed_results
+ ]
+
+ print(tabulate(table_data, headers=["ASR服务", "首词延迟(秒)", "状态"], tablefmt="grid"))
+ print("\n测试说明:测量从发送请求到接收第一个识别结果的时间,取多次测试平均值")
+ print("- 超时控制: 单个请求最大等待时间为10秒")
+ print("- 错误处理: 无法连接和超时的列为网络错误")
+ print("- 排序规则: 按平均耗时从快到慢排序")
+
+ async def run(self, test_count=5):
+ """执行测试"""
+ print(f"开始流式ASR首词响应时间测试...")
+ print(f"每个ASR服务测试次数: {test_count}次")
+
+ if not self.config.get("ASR"):
+ print("配置文件中未找到ASR配置")
+ return
+
+ # 测试每种ASR服务
+ self.results = []
+
+ # 测试豆包ASR
+ if self.config["ASR"].get("DoubaoStreamASR"):
+ result = await self.test_doubao_stream_asr(test_count)
+ self.results.append(result)
+ else:
+ print("配置文件中未找到豆包流式ASR配置,跳过测试")
+
+ # 测试阿里云ASR
+ if self.config["ASR"].get("AliyunStreamASR"):
+ result = await self.test_aliyun_stream_asr(test_count)
+ self.results.append(result)
+ else:
+ print("配置文件中未找到阿里云流式ASR配置,跳过测试")
+
+ # 打印结果
+ self._print_results(test_count)
+
+async def main():
+ import argparse
+
+ parser = argparse.ArgumentParser(description="流式ASR首词响应时间测试工具")
+ parser.add_argument("--count", type=int, default=5, help="测试次数")
+
+ args = parser.parse_args()
+ await DoubaoStreamASRPerformanceTester().run(args.count)
+
+if __name__ == "__main__":
+ import os
+ import gzip
+ import opuslib_next
+ asyncio.run(main())
\ No newline at end of file
diff --git a/main/xiaozhi-server/performance_tester/performance_tester_stream_tts.py b/main/xiaozhi-server/performance_tester/performance_tester_stream_tts.py
new file mode 100644
index 00000000..309ad339
--- /dev/null
+++ b/main/xiaozhi-server/performance_tester/performance_tester_stream_tts.py
@@ -0,0 +1,388 @@
+import asyncio
+import time
+import json
+import uuid
+import aiohttp
+import websockets
+from tabulate import tabulate
+from config.settings import load_config
+
+description = "流式TTS语音合成首词耗时测试"
+class StreamTTSPerformanceTester:
+ def __init__(self):
+ self.config = load_config()
+ self.test_texts = [
+ "你好,这是一句话。"
+ ]
+ self.results = []
+
+ async def test_aliyun_tts(self, text=None, test_count=5):
+ """测试阿里云流式TTS首词延迟(测试多次取平均)"""
+ text = text or self.test_texts[0]
+ latencies = []
+
+ for i in range(test_count):
+ try:
+ tts_config = self.config["TTS"]["AliyunStreamTTS"]
+ appkey = tts_config["appkey"]
+ token = tts_config["token"]
+ voice = tts_config["voice"]
+ host = tts_config["host"]
+ ws_url = f"wss://{host}/ws/v1"
+
+ start_time = time.time()
+ async with websockets.connect(ws_url, extra_headers={"X-NLS-Token": token}) as ws:
+ task_id = str(uuid.uuid4())
+ message_id = str(uuid.uuid4())
+
+ start_request = {
+ "header": {
+ "message_id": message_id,
+ "task_id": task_id,
+ "namespace": "FlowingSpeechSynthesizer",
+ "name": "StartSynthesis",
+ "appkey": appkey,
+ },
+ "payload": {
+ "voice": voice,
+ "format": "pcm",
+ "sample_rate": 16000,
+ "volume": 50,
+ "speech_rate": 0,
+ "pitch_rate": 0,
+ }
+ }
+ await ws.send(json.dumps(start_request))
+
+ start_response = json.loads(await ws.recv())
+ if start_response["header"]["name"] != "SynthesisStarted":
+ raise Exception("启动合成失败")
+
+ run_request = {
+ "header": {
+ "message_id": str(uuid.uuid4()),
+ "task_id": task_id,
+ "namespace": "FlowingSpeechSynthesizer",
+ "name": "RunSynthesis",
+ "appkey": appkey,
+ },
+ "payload": {"text": text}
+ }
+ await ws.send(json.dumps(run_request))
+
+ while True:
+ response = await ws.recv()
+ if isinstance(response, bytes):
+ latency = time.time() - start_time
+ latencies.append(latency)
+ break
+ elif isinstance(response, str):
+ data = json.loads(response)
+ if data["header"]["name"] == "TaskFailed":
+ raise Exception(f"合成失败: {data['payload']['error_info']}")
+
+ except Exception as e:
+ latencies.append(0)
+
+ return self._calculate_result("阿里云TTS", latencies, test_count)
+
+ async def test_doubao_tts(self, text=None, test_count=5):
+ """测试火山引擎流式TTS首词延迟(测试多次取平均)"""
+ text = text or self.test_texts[0]
+ latencies = []
+
+ for i in range(test_count):
+ try:
+ tts_config = self.config["TTS"]["HuoshanDoubleStreamTTS"]
+ ws_url = tts_config["ws_url"]
+ app_id = tts_config["appid"]
+ access_token = tts_config["access_token"]
+ resource_id = tts_config["resource_id"]
+ speaker = tts_config["speaker"]
+
+ start_time = time.time()
+ ws_header = {
+ "X-Api-App-Key": app_id,
+ "X-Api-Access-Key": access_token,
+ "X-Api-Resource-Id": resource_id,
+ "X-Api-Connect-Id": str(uuid.uuid4()),
+ }
+ async with websockets.connect(ws_url, additional_headers=ws_header, max_size=1000000000) as ws:
+ session_id = uuid.uuid4().hex
+
+ # 发送会话启动请求
+ header = bytes([
+ (0b0001 << 4) | 0b0001,
+ 0b0001 << 4 | 0b100,
+ 0b0001 << 4 | 0b0000,
+ 0
+ ])
+ optional = bytearray()
+ optional.extend((1).to_bytes(4, "big", signed=True))
+ session_id_bytes = session_id.encode()
+ optional.extend(len(session_id_bytes).to_bytes(4, "big", signed=True))
+ optional.extend(session_id_bytes)
+ payload = json.dumps({"speaker": speaker}).encode()
+ await ws.send(header + optional + len(payload).to_bytes(4, "big", signed=True) + payload)
+
+ # 发送文本
+ header = bytes([
+ (0b0001 << 4) | 0b0001,
+ 0b0001 << 4 | 0b100,
+ 0b0001 << 4 | 0b0000,
+ 0
+ ])
+ optional = bytearray()
+ optional.extend((200).to_bytes(4, "big", signed=True))
+ session_id_bytes = session_id.encode()
+ optional.extend(len(session_id_bytes).to_bytes(4, "big", signed=True))
+ optional.extend(session_id_bytes)
+ payload = json.dumps({"text": text, "speaker": speaker}).encode()
+ await ws.send(header + optional + len(payload).to_bytes(4, "big", signed=True) + payload)
+
+ first_chunk = await ws.recv()
+ latency = time.time() - start_time
+ latencies.append(latency)
+
+ except Exception as e:
+ latencies.append(0)
+
+ return self._calculate_result("火山引擎TTS", latencies, test_count)
+
+ async def test_paddlespeech_tts(self, text=None, test_count=5):
+ """测试PaddleSpeech流式TTS首词延迟(测试多次取平均)"""
+ text = text or self.test_texts[0]
+ latencies = []
+
+ for i in range(test_count):
+ try:
+ tts_config = self.config["TTS"]["PaddleSpeechTTS"]
+ tts_url = tts_config["url"]
+ spk_id = tts_config["spk_id"]
+ speed = tts_config["speed"]
+ volume = tts_config["volume"]
+
+ start_time = time.time()
+ async with websockets.connect(tts_url) as ws:
+ # 发送开始请求
+ await ws.send(json.dumps({
+ "task": "tts",
+ "signal": "start"
+ }))
+
+ start_response = json.loads(await ws.recv())
+ if start_response.get("status") != 0:
+ raise Exception("连接失败")
+
+ # 发送文本数据
+ await ws.send(json.dumps({
+ "text": text,
+ "spk_id": spk_id,
+ "speed": speed,
+ "volume": volume
+ }))
+
+ # 接收第一个数据块
+ first_chunk = await ws.recv()
+ latency = time.time() - start_time
+ latencies.append(latency)
+
+ # 发送结束请求
+ end_request = {
+ "task": "tts",
+ "signal": "end"
+ }
+ await ws.send(json.dumps(end_request))
+
+ # 确保连接正常关闭
+ try:
+ await ws.recv()
+ except websockets.exceptions.ConnectionClosedOK:
+ pass
+
+ except Exception as e:
+ latencies.append(0)
+
+ return self._calculate_result("PaddleSpeechTTS", latencies, test_count)
+
+ async def test_indexstream_tts(self, text=None, test_count=5):
+ """测试IndexStream流式TTS首词延迟(测试多次取平均)"""
+ text = text or self.test_texts[0]
+ latencies = []
+
+ for i in range(test_count):
+ try:
+ tts_config = self.config["TTS"]["IndexStreamTTS"]
+ api_url = tts_config.get("api_url")
+ voice = tts_config.get("voice")
+
+ start_time = time.time()
+
+ async with aiohttp.ClientSession() as session:
+ payload = {"text": text, "character": voice}
+ async with session.post(api_url, json=payload, timeout=10) as resp:
+ if resp.status != 200:
+ raise Exception(f"请求失败: {resp.status}, {await resp.text()}")
+
+ async for chunk in resp.content.iter_any():
+ data = chunk[0] if isinstance(chunk, (list, tuple)) else chunk
+ if not data:
+ continue
+
+ latency = time.time() - start_time
+ latencies.append(latency)
+ resp.close()
+ break
+ else:
+ latencies.append(0)
+
+ except Exception as e:
+ latencies.append(0)
+
+ return self._calculate_result("IndexStreamTTS", latencies, test_count)
+
+ async def test_linkerai_tts(self, text=None, test_count=5):
+ """测试Linkerai流式TTS首词延迟(测试多次取平均)"""
+ text = text or self.test_texts[0]
+ latencies = []
+
+ for i in range(test_count):
+ try:
+ tts_config = self.config["TTS"]["LinkeraiTTS"]
+ api_url = tts_config["api_url"]
+ access_token = tts_config["access_token"]
+ voice = tts_config["voice"]
+
+ start_time = time.time()
+ async with aiohttp.ClientSession() as session:
+ params = {
+ "tts_text": text,
+ "spk_id": voice,
+ "frame_durition": 60,
+ "stream": "true",
+ "target_sr": 16000,
+ "audio_format": "pcm",
+ "instruct_text": "请生成一段自然流畅的语音",
+ }
+ headers = {
+ "Authorization": f"Bearer {access_token}",
+ "Content-Type": "application/json",
+ }
+
+ async with session.get(api_url, params=params, headers=headers, timeout=10) as resp:
+ if resp.status != 200:
+ raise Exception(f"请求失败: {resp.status}, {await resp.text()}")
+
+ # 接收第一个数据块
+ async for _ in resp.content.iter_any():
+ latency = time.time() - start_time
+ latencies.append(latency)
+ break
+ else:
+ latencies.append(0)
+
+ except Exception as e:
+ latencies.append(0)
+
+ return self._calculate_result("LinkeraiTTS", latencies, test_count)
+
+
+ def _calculate_result(self, service_name, latencies, test_count):
+ """计算测试结果"""
+ valid_latencies = [l for l in latencies if l > 0]
+ if valid_latencies:
+ avg_latency = sum(valid_latencies) / len(valid_latencies)
+ status = f"成功({len(valid_latencies)}/{test_count}次有效)"
+ else:
+ avg_latency = 0
+ status = "失败: 所有测试均失败"
+ return {"name": service_name, "latency": avg_latency, "status": status}
+
+ def _print_results(self, test_text, test_count):
+ """打印测试结果"""
+ if not self.results:
+ print("没有有效的TTS测试结果")
+ return
+
+ print(f"\n{'='*60}")
+ print("流式TTS首词延迟测试结果")
+ print(f"{'='*60}")
+ print(f"测试文本: {test_text}")
+ print(f"测试次数: 每个TTS服务测试 {test_count} 次")
+
+ # 排序结果:成功优先,按延迟升序
+ success_results = sorted(
+ [r for r in self.results if "成功" in r["status"]],
+ key=lambda x: x["latency"]
+ )
+ failed_results = [r for r in self.results if "成功" not in r["status"]]
+
+ table_data = [
+ [r["name"], f"{r['latency']:.3f}", r["status"]]
+ for r in success_results + failed_results
+ ]
+
+ print(tabulate(table_data, headers=["TTS服务", "首词延迟(秒)", "状态"], tablefmt="grid"))
+ print("\n测试说明:测量从发送请求到接收第一个音频数据块的时间,取多次测试平均值")
+ print("- 超时控制: 单个请求最大等待时间为10秒")
+ print("- 错误处理: 无法连接和超时的列为网络错误")
+ print("- 排序规则: 按平均耗时从快到慢排序")
+
+
+ async def run(self, test_text=None, test_count=5):
+ """执行测试
+
+ Args:
+ test_text: 要测试的文本,如果为None则使用默认文本
+ test_count: 每个TTS服务的测试次数
+ """
+ test_text = test_text or self.test_texts[0]
+ print(f"开始流式TTS首词延迟测试...")
+ print(f"测试文本: {test_text}")
+ print(f"每个TTS服务测试次数: {test_count}次")
+
+ if not self.config.get("TTS"):
+ print("配置文件中未找到TTS配置")
+ return
+
+ # 测试每种TTS服务
+ self.results = []
+
+ # 测试阿里云TTS
+ result = await self.test_aliyun_tts(test_text, test_count)
+ self.results.append(result)
+
+ # 测试火山引擎TTS
+ result = await self.test_doubao_tts(test_text, test_count)
+ self.results.append(result)
+
+ # 测试PaddleSpeech TTS
+ result = await self.test_paddlespeech_tts(test_text, test_count)
+ self.results.append(result)
+
+ # 测试Linkerai TTS
+ result = await self.test_linkerai_tts(test_text, test_count)
+ self.results.append(result)
+
+ # 测试IndexStreamTTS
+ result = await self.test_indexstream_tts(test_text, test_count)
+ self.results.append(result)
+
+ # 打印结果
+ self._print_results(test_text, test_count)
+
+
+async def main():
+ import argparse
+
+ parser = argparse.ArgumentParser(description="流式TTS首词延迟测试工具")
+ parser.add_argument("--text", help="要测试的文本内容")
+ parser.add_argument("--count", type=int, default=5, help="每个TTS服务的测试次数")
+
+ args = parser.parse_args()
+ await StreamTTSPerformanceTester().run(args.text, args.count)
+
+
+if __name__ == "__main__":
+ import asyncio
+ asyncio.run(main())
\ No newline at end of file
diff --git a/main/xiaozhi-server/performance_tester/performance_tester_tts.py b/main/xiaozhi-server/performance_tester/performance_tester_tts.py
index 4702c5a9..61e49682 100644
--- a/main/xiaozhi-server/performance_tester/performance_tester_tts.py
+++ b/main/xiaozhi-server/performance_tester/performance_tester_tts.py
@@ -86,22 +86,67 @@ class TTSPerformanceTester:
print("没有有效的TTS测试结果")
return
- table = []
+ headers = ["TTS模块", "平均耗时(秒)", "测试句子数", "状态"]
+ table_data = []
+
+ # 收集所有数据并分类
+ valid_results = []
+ error_results = []
+
for name, data in self.results.items():
if data["errors"] == 0:
- table.append(
- [name, f"{data['avg_time']:.3f}秒/句", len(self.test_sentences[:3])]
- )
+ # 正常结果
+ avg_time = f"{data['avg_time']:.3f}"
+ test_count = len(self.test_sentences[:3])
+ status = "✅ 正常"
+
+ # 保存用于排序的值
+ valid_results.append({
+ "name": name,
+ "avg_time": avg_time,
+ "test_count": test_count,
+ "status": status,
+ "sort_key": data['avg_time']
+ })
+ else:
+ # 错误结果
+ avg_time = "-"
+ test_count = "0/3"
+
+ # 默认错误类型为网络错误
+ error_type = "网络错误"
+ status = f"❌ {error_type}"
+
+ error_results.append([name, avg_time, test_count, status])
+
+ # 按平均耗时升序排序
+ valid_results.sort(key=lambda x: x["sort_key"])
+
+ # 将排序后的有效结果转换为表格数据
+ for result in valid_results:
+ table_data.append([
+ result["name"],
+ result["avg_time"],
+ result["test_count"],
+ result["status"]
+ ])
+
+ # 将错误结果添加到表格数据末尾
+ table_data.extend(error_results)
print("\nTTS性能测试结果:")
print(
tabulate(
- table,
- headers=["TTS模块", "平均耗时", "测试句子数"],
- tablefmt="github",
- colalign=("left", "right", "right"),
+ table_data,
+ headers=headers,
+ tablefmt="grid",
+ colalign=("left", "right", "right", "left"),
)
)
+ print("\n测试说明:")
+ print("- 超时控制: 单个请求最大等待时间为10秒")
+ print("- 错误处理: 无法连接和超时的列为网络错误")
+ print("- 排序规则: 按平均耗时从快到慢排序")
async def run(self):
"""执行测试"""
@@ -119,10 +164,9 @@ class TTSPerformanceTester:
# 并发执行测试
results = await asyncio.gather(*tasks)
- # 保存有效结果
+ # 保存所有结果,包括错误
for result in results:
- if result["errors"] == 0:
- self.results[result["name"]] = result
+ self.results[result["name"]] = result
# 打印结果
self._print_results()
diff --git a/main/xiaozhi-server/plugins_func/functions/get_news_from_newsnow.py b/main/xiaozhi-server/plugins_func/functions/get_news_from_newsnow.py
index 54641106..1d60aefd 100644
--- a/main/xiaozhi-server/plugins_func/functions/get_news_from_newsnow.py
+++ b/main/xiaozhi-server/plugins_func/functions/get_news_from_newsnow.py
@@ -125,7 +125,8 @@ def fetch_news_from_api(conn, source="thepaper"):
]["get_news_from_newsnow"].get("url"):
api_url = conn.config["plugins"]["get_news_from_newsnow"]["url"] + source
- response = requests.get(api_url, timeout=10)
+ headers = {"User-Agent": "Mozilla/5.0"}
+ response = requests.get(api_url, headers=headers, timeout=10)
response.raise_for_status()
data = response.json()
@@ -144,7 +145,8 @@ def fetch_news_from_api(conn, source="thepaper"):
def fetch_news_detail(url):
"""获取新闻详情页内容并使用MarkItDown清理HTML"""
try:
- response = requests.get(url, timeout=10)
+ headers = {"User-Agent": "Mozilla/5.0"}
+ response = requests.get(url, headers=headers, timeout=10)
response.raise_for_status()
# 使用MarkItDown清理HTML内容
diff --git a/main/xiaozhi-server/plugins_func/functions/get_weather.py b/main/xiaozhi-server/plugins_func/functions/get_weather.py
index a3af6c03..38770a3f 100644
--- a/main/xiaozhi-server/plugins_func/functions/get_weather.py
+++ b/main/xiaozhi-server/plugins_func/functions/get_weather.py
@@ -110,6 +110,11 @@ WEATHER_CODE_MAP = {
def fetch_city_info(location, api_key, api_host):
url = f"https://{api_host}/geo/v2/city/lookup?key={api_key}&location={location}&lang=zh"
response = requests.get(url, headers=HEADERS).json()
+ if response.get("error") is not None:
+ logger.bind(tag=TAG).error(
+ f"获取天气失败,原因:{response.get('error', {}).get('detail')}"
+ )
+ return None
return response.get("location", [])[0] if response.get("location") else None
diff --git a/main/xiaozhi-server/plugins_func/functions/hass_get_state.py b/main/xiaozhi-server/plugins_func/functions/hass_get_state.py
index 94478e35..4efb12b8 100644
--- a/main/xiaozhi-server/plugins_func/functions/hass_get_state.py
+++ b/main/xiaozhi-server/plugins_func/functions/hass_get_state.py
@@ -29,12 +29,7 @@ hass_get_state_function_desc = {
@register_function("hass_get_state", hass_get_state_function_desc, ToolType.SYSTEM_CTL)
def hass_get_state(conn, entity_id=""):
try:
-
- future = asyncio.run_coroutine_threadsafe(
- handle_hass_get_state(conn, entity_id), conn.loop
- )
- # 添加10秒超时
- ha_response = future.result(timeout=10)
+ ha_response = handle_hass_get_state(conn, entity_id)
return ActionResponse(Action.REQLLM, ha_response, None)
except asyncio.TimeoutError:
logger.bind(tag=TAG).error("获取Home Assistant状态超时")
@@ -45,13 +40,13 @@ def hass_get_state(conn, entity_id=""):
return ActionResponse(Action.ERROR, error_msg, None)
-async def handle_hass_get_state(conn, entity_id):
+def handle_hass_get_state(conn, entity_id):
ha_config = initialize_hass_handler(conn)
api_key = ha_config.get("api_key")
base_url = ha_config.get("base_url")
url = f"{base_url}/api/states/{entity_id}"
headers = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"}
- response = requests.get(url, headers=headers)
+ response = requests.get(url, headers=headers, timeout=5)
if response.status_code == 200:
responsetext = "设备状态:" + response.json()["state"] + " "
logger.bind(tag=TAG).info(f"api返回内容: {response.json()}")
diff --git a/main/xiaozhi-server/plugins_func/functions/hass_set_state.py b/main/xiaozhi-server/plugins_func/functions/hass_set_state.py
index c1fcdaa6..3addc823 100644
--- a/main/xiaozhi-server/plugins_func/functions/hass_set_state.py
+++ b/main/xiaozhi-server/plugins_func/functions/hass_set_state.py
@@ -54,11 +54,7 @@ def hass_set_state(conn, entity_id="", state=None):
if state is None:
state = {}
try:
- future = asyncio.run_coroutine_threadsafe(
- handle_hass_set_state(conn, entity_id, state), conn.loop
- )
- # 添加10秒超时
- ha_response = future.result(timeout=10)
+ ha_response = handle_hass_set_state(conn, entity_id, state)
return ActionResponse(Action.REQLLM, ha_response, None)
except asyncio.TimeoutError:
logger.bind(tag=TAG).error("设置Home Assistant状态超时")
@@ -69,7 +65,7 @@ def hass_set_state(conn, entity_id="", state=None):
return ActionResponse(Action.ERROR, error_msg, None)
-async def handle_hass_set_state(conn, entity_id, state):
+def handle_hass_set_state(conn, entity_id, state):
ha_config = initialize_hass_handler(conn)
api_key = ha_config.get("api_key")
base_url = ha_config.get("base_url")
@@ -169,7 +165,7 @@ async def handle_hass_set_state(conn, entity_id, state):
data = {"entity_id": entity_id, arg: value}
url = f"{base_url}/api/services/{domain}/{action}"
headers = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"}
- response = requests.post(url, headers=headers, json=data)
+ response = requests.post(url, headers=headers, json=data, timeout=5) # 设置5秒超时
logger.bind(tag=TAG).info(
f"设置状态:{description},url:{url},return_code:{response.status_code}"
)
diff --git a/main/xiaozhi-server/test/js/StreamingContext.js b/main/xiaozhi-server/test/js/StreamingContext.js
new file mode 100644
index 00000000..576d9761
--- /dev/null
+++ b/main/xiaozhi-server/test/js/StreamingContext.js
@@ -0,0 +1,149 @@
+import BlockingQueue from './utils/BlockingQueue.js';
+import { log } from './utils/logger.js';
+
+// 音频流播放上下文类
+export class StreamingContext {
+ constructor(opusDecoder, audioContext, sampleRate, channels, minAudioDuration) {
+ this.opusDecoder = opusDecoder;
+ this.audioContext = audioContext;
+
+ // 音频参数
+ this.sampleRate = sampleRate;
+ this.channels = channels;
+ this.minAudioDuration = minAudioDuration;
+
+ // 初始化队列和状态
+ this.queue = []; // 已解码的PCM队列。正在播放
+ this.activeQueue = new BlockingQueue(); // 已解码的PCM队列。准备播放
+ this.pendingAudioBufferQueue = []; // 待处理的缓存队列
+ this.audioBufferQueue = new BlockingQueue(); // 缓存队列
+ this.playing = false; // 是否正在播放
+ this.endOfStream = false; // 是否收到结束信号
+ this.source = null; // 当前音频源
+ this.totalSamples = 0; // 累积的总样本数
+ this.lastPlayTime = 0; // 上次播放的时间戳
+ }
+
+ // 缓存音频数组
+ pushAudioBuffer(item) {
+ this.audioBufferQueue.enqueue(...item);
+ }
+
+ // 获取需要处理缓存队列,单线程:在audioBufferQueue一直更新的状态下不会出现安全问题
+ async getPendingAudioBufferQueue() {
+ // 原子交换 + 清空
+ [this.pendingAudioBufferQueue, this.audioBufferQueue] = [await this.audioBufferQueue.dequeue(), new BlockingQueue()];
+ }
+
+ // 获取正在播放已解码的PCM队列,单线程:在activeQueue一直更新的状态下不会出现安全问题
+ async getQueue(minSamples) {
+ let TepArray = [];
+ const num = minSamples - this.queue.length > 0 ? minSamples - this.queue.length : 1;
+ // 原子交换 + 清空
+ [TepArray, this.activeQueue] = [await this.activeQueue.dequeue(num), new BlockingQueue()];
+ this.queue.push(...TepArray);
+ }
+
+ // 将Int16音频数据转换为Float32音频数据
+ convertInt16ToFloat32(int16Data) {
+ const float32Data = new Float32Array(int16Data.length);
+ for (let i = 0; i < int16Data.length; i++) {
+ // 将[-32768,32767]范围转换为[-1,1]
+ float32Data[i] = int16Data[i] / (int16Data[i] < 0 ? 0x8000 : 0x7FFF);
+ }
+ return float32Data;
+ }
+
+ // 将Opus数据解码为PCM
+ async decodeOpusFrames() {
+ if (!this.opusDecoder) {
+ log('Opus解码器未初始化,无法解码', 'error');
+ return;
+ } else {
+ log('Opus解码器启动', 'info');
+ }
+
+ while (true) {
+ let decodedSamples = [];
+ for (const frame of this.pendingAudioBufferQueue) {
+ try {
+ // 使用Opus解码器解码
+ const frameData = this.opusDecoder.decode(frame);
+ if (frameData && frameData.length > 0) {
+ // 转换为Float32
+ const floatData = this.convertInt16ToFloat32(frameData);
+ // 使用循环替代展开运算符
+ for (let i = 0; i < floatData.length; i++) {
+ decodedSamples.push(floatData[i]);
+ }
+ }
+ } catch (error) {
+ log("Opus解码失败: " + error.message, 'error');
+ }
+ }
+
+ if (decodedSamples.length > 0) {
+ // 使用循环替代展开运算符
+ for (let i = 0; i < decodedSamples.length; i++) {
+ this.activeQueue.enqueue(decodedSamples[i]);
+ }
+ this.totalSamples += decodedSamples.length;
+ } else {
+ log('没有成功解码的样本', 'warning');
+ }
+ await this.getPendingAudioBufferQueue();
+ }
+ }
+
+ // 开始播放音频
+ async startPlaying() {
+ while (true) {
+ // 如果累积了至少0.3秒的音频,开始播放
+ const minSamples = this.sampleRate * this.minAudioDuration * 3;
+ if (!this.playing && this.queue.length < minSamples) {
+ await this.getQueue(minSamples);
+ }
+ this.playing = true;
+ while (this.playing && this.queue.length) {
+ // 创建新的音频缓冲区
+ const minPlaySamples = Math.min(this.queue.length, this.sampleRate);
+ const currentSamples = this.queue.splice(0, minPlaySamples);
+
+ const audioBuffer = this.audioContext.createBuffer(this.channels, currentSamples.length, this.sampleRate);
+ audioBuffer.copyToChannel(new Float32Array(currentSamples), 0);
+
+ // 创建音频源
+ this.source = this.audioContext.createBufferSource();
+ this.source.buffer = audioBuffer;
+
+ // 创建增益节点用于平滑过渡
+ const gainNode = this.audioContext.createGain();
+
+ // 应用淡入淡出效果避免爆音
+ const fadeDuration = 0.02; // 20毫秒
+ gainNode.gain.setValueAtTime(0, this.audioContext.currentTime);
+ gainNode.gain.linearRampToValueAtTime(1, this.audioContext.currentTime + fadeDuration);
+
+ const duration = audioBuffer.duration;
+ if (duration > fadeDuration * 2) {
+ gainNode.gain.setValueAtTime(1, this.audioContext.currentTime + duration - fadeDuration);
+ gainNode.gain.linearRampToValueAtTime(0, this.audioContext.currentTime + duration);
+ }
+
+ // 连接节点并开始播放
+ this.source.connect(gainNode);
+ gainNode.connect(this.audioContext.destination);
+
+ this.lastPlayTime = this.audioContext.currentTime;
+ log(`开始播放 ${currentSamples.length} 个样本,约 ${(currentSamples.length / this.sampleRate).toFixed(2)} 秒`, 'info');
+ this.source.start();
+ }
+ await this.getQueue(minSamples);
+ }
+ }
+}
+
+// 创建streamingContext实例的工厂函数
+export function createStreamingContext(opusDecoder, audioContext, sampleRate, channels, minAudioDuration) {
+ return new StreamingContext(opusDecoder, audioContext, sampleRate, channels, minAudioDuration);
+}
\ No newline at end of file
diff --git a/main/xiaozhi-server/test/test_page.html b/main/xiaozhi-server/test/test_page.html
index 6f8fdba9..547f2b31 100644
--- a/main/xiaozhi-server/test/test_page.html
+++ b/main/xiaozhi-server/test/test_page.html
@@ -181,6 +181,7 @@
import { checkOpusLoaded, initOpusEncoder } from './js/opus.js';
import { addMessage } from './js/document.js'
import BlockingQueue from './js/utils/BlockingQueue.js'
+ import { createStreamingContext } from './js/StreamingContext.js'
// 需要加载的脚本列表 - 移除Opus依赖
const scriptFiles = [];
@@ -230,6 +231,16 @@
const conversationDiv = document.getElementById('conversation');
const logContainer = document.getElementById('logContainer');
+ function getAudioContextInstance() {
+ if (!audioContext) {
+ audioContext = new (window.AudioContext || window.webkitAudioContext)({
+ sampleRate: SAMPLE_RATE,
+ latencyHint: 'interactive'
+ });
+ log('创建音频上下文,采样率: ' + SAMPLE_RATE + 'Hz', 'debug');
+ }
+ return audioContext;
+ }
// 初始化可视化器
function initVisualizer() {
@@ -306,12 +317,7 @@
// 确保Opus解码器已初始化
try {
// 确保音频上下文存在
- if (!audioContext) {
- audioContext = new (window.AudioContext || window.webkitAudioContext)({
- sampleRate: SAMPLE_RATE
- });
- log('创建音频上下文,采样率: ' + SAMPLE_RATE + 'Hz', 'debug');
- }
+ audioContext = getAudioContextInstance();
// 确保解码器已初始化
if (!opusDecoder) {
@@ -331,125 +337,7 @@
// 创建流式播放上下文
if (!streamingContext) {
- streamingContext = {
- queue: [], // 已解码的PCM队列。正在播放
- activeQueue: new BlockingQueue(), // 已解码的PCM队列。准备播放
- pendingAudioBufferQueue: [], // 待处理的缓存队列
- audioBufferQueue: new BlockingQueue(), // 缓存队列
- playing: false, // 是否正在播放
- endOfStream: false, // 是否收到结束信号
- source: null, // 当前音频源
- totalSamples: 0, // 累积的总样本数
- lastPlayTime: 0, // 上次播放的时间戳
-
-
- // 缓存音频数组
- pushAudioBuffer: function (item) {
- this.audioBufferQueue.enqueue(...item)
- },
-
- // 获取需要处理缓存队列,单线程:在audioBufferQueue一直更新的状态下不会出现安全问题
- getPendingAudioBufferQueue: async function () {
- // 原子交换 + 清空
- [this.pendingAudioBufferQueue, this.audioBufferQueue] = [await this.audioBufferQueue.dequeue(), new BlockingQueue()];
-
- },
- // 获取正在播放已解码的PCM队列,单线程:在activeQueue一直更新的状态下不会出现安全问题
- getQueue: async function (minSamples) {
- let TepArray = []
- const num = minSamples - this.queue.length > 0 ? minSamples - this.queue.length : 1;
- // 原子交换 + 清空
- [TepArray, this.activeQueue] = [await this.activeQueue.dequeue(num), new BlockingQueue()];
- this.queue.push(...TepArray)
- },
- // 将Opus数据解码为PCM
- decodeOpusFrames: async function () {
- if (!opusDecoder) {
- log('Opus解码器未初始化,无法解码', 'error');
- return;
- } else {
- log('Opus解码器启动', 'info');
- }
-
- while (true) {
- let decodedSamples = [];
- for (const frame of this.pendingAudioBufferQueue) {
- try {
- // 使用Opus解码器解码
- const frameData = opusDecoder.decode(frame);
- if (frameData && frameData.length > 0) {
- // 转换为Float32
- const floatData = convertInt16ToFloat32(frameData);
- // 使用循环替代展开运算符
- for (let i = 0; i < floatData.length; i++) {
- decodedSamples.push(floatData[i]);
- }
- }
- } catch (error) {
- log("Opus解码失败: " + error.message, 'error');
- }
- }
-
- if (decodedSamples.length > 0) {
- // 使用循环替代展开运算符
- for (let i = 0; i < decodedSamples.length; i++) {
- this.activeQueue.enqueue(decodedSamples[i]);
- }
- this.totalSamples += decodedSamples.length;
- } else {
- log('没有成功解码的样本', 'warning');
- }
- await this.getPendingAudioBufferQueue();
- }
- },
-
- // 开始播放音频
- startPlaying: async function () {
- while (true) {
- // 如果累积了至少0.3秒的音频,开始播放
- const minSamples = SAMPLE_RATE * MIN_AUDIO_DURATION * 3;
- if (!this.playing && this.queue.length < minSamples) {
- await this.getQueue(minSamples)
- }
- this.playing = true;
- while (this.playing && this.queue.length) {
- // 创建新的音频缓冲区
- const minPlaySamples = Math.min(this.queue.length, SAMPLE_RATE);
- const currentSamples = this.queue.splice(0, minPlaySamples);
-
- const audioBuffer = audioContext.createBuffer(CHANNELS, currentSamples.length, SAMPLE_RATE);
- audioBuffer.copyToChannel(new Float32Array(currentSamples), 0);
-
- // 创建音频源
- this.source = audioContext.createBufferSource();
- this.source.buffer = audioBuffer;
-
- // 创建增益节点用于平滑过渡
- const gainNode = audioContext.createGain();
-
- // 应用淡入淡出效果避免爆音
- const fadeDuration = 0.02; // 20毫秒
- gainNode.gain.setValueAtTime(0, audioContext.currentTime);
- gainNode.gain.linearRampToValueAtTime(1, audioContext.currentTime + fadeDuration);
-
- const duration = audioBuffer.duration;
- if (duration > fadeDuration * 2) {
- gainNode.gain.setValueAtTime(1, audioContext.currentTime + duration - fadeDuration);
- gainNode.gain.linearRampToValueAtTime(0, audioContext.currentTime + duration);
- }
-
- // 连接节点并开始播放
- this.source.connect(gainNode);
- gainNode.connect(audioContext.destination);
-
- this.lastPlayTime = audioContext.currentTime;
- log(`开始播放 ${currentSamples.length} 个样本,约 ${(currentSamples.length / SAMPLE_RATE).toFixed(2)} 秒`, 'info');
- this.source.start();
- }
- await this.getQueue(minSamples)
- }
- }
- };
+ streamingContext = createStreamingContext(opusDecoder, audioContext, SAMPLE_RATE, CHANNELS, MIN_AUDIO_DURATION);
}
streamingContext.decodeOpusFrames();
@@ -462,15 +350,7 @@
}
}
- // 将Int16音频数据转换为Float32音频数据
- function convertInt16ToFloat32(int16Data) {
- const float32Data = new Float32Array(int16Data.length);
- for (let i = 0; i < int16Data.length; i++) {
- // 将[-32768,32767]范围转换为[-1,1]
- float32Data[i] = int16Data[i] / (int16Data[i] < 0 ? 0x8000 : 0x7FFF);
- }
- return float32Data;
- }
+
// 初始化Opus解码器 - 确保完全初始化完成后才返回
async function initOpusDecoder() {
@@ -617,10 +497,7 @@
log('已获取麦克风访问权限', 'success');
// 创建音频上下文
- audioContext = new (window.AudioContext || window.webkitAudioContext)({
- sampleRate: 16000, // 确保采样率与服务器期望的一致
- latencyHint: 'interactive'
- });
+ audioContext = getAudioContextInstance();
const source = audioContext.createMediaStreamSource(stream);
// 获取实际音频轨道设置
@@ -1416,12 +1293,7 @@
// 创建音频处理器
async function createAudioProcessor() {
- if (!audioContext) {
- audioContext = new (window.AudioContext || window.webkitAudioContext)({
- sampleRate: 16000,
- latencyHint: 'interactive'
- });
- }
+ audioContext = getAudioContextInstance();
try {
// 检查是否支持AudioWorklet
@@ -1616,12 +1488,7 @@
});
// 创建音频上下文和分析器
- if (!audioContext) {
- audioContext = new (window.AudioContext || window.webkitAudioContext)({
- sampleRate: 16000,
- latencyHint: 'interactive'
- });
- }
+ audioContext = getAudioContextInstance();
// 创建音频处理器
const processorResult = await createAudioProcessor();