mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-29 04:13:56 +08:00
Merge branch 'main' into fix
This commit is contained in:
@@ -237,7 +237,7 @@ public interface Constant {
|
|||||||
/**
|
/**
|
||||||
* 版本号
|
* 版本号
|
||||||
*/
|
*/
|
||||||
public static final String VERSION = "0.7.6";
|
public static final String VERSION = "0.7.7";
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* 无效固件URL
|
* 无效固件URL
|
||||||
|
|||||||
@@ -56,11 +56,11 @@ function getCacheInfo() {
|
|||||||
// 验证URL格式
|
// 验证URL格式
|
||||||
function validateUrl() {
|
function validateUrl() {
|
||||||
urlError.value = ''
|
urlError.value = ''
|
||||||
|
|
||||||
if (!baseUrlInput.value) {
|
if (!baseUrlInput.value) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if (!/^https?:\/\/.+\/xiaozhi$/.test(baseUrlInput.value)) {
|
if (!/^https?:\/\/.+\/xiaozhi$/.test(baseUrlInput.value)) {
|
||||||
urlError.value = '请输入有效的服务端地址(以 http 或 https 开头,并以 /xiaozhi 结尾)'
|
urlError.value = '请输入有效的服务端地址(以 http 或 https 开头,并以 /xiaozhi 结尾)'
|
||||||
}
|
}
|
||||||
@@ -70,7 +70,7 @@ function validateUrl() {
|
|||||||
async function testServerBaseUrl() {
|
async function testServerBaseUrl() {
|
||||||
// 先清除错误信息
|
// 先清除错误信息
|
||||||
urlError.value = ''
|
urlError.value = ''
|
||||||
|
|
||||||
if (!baseUrlInput.value || !/^https?:\/\/.+\/xiaozhi$/.test(baseUrlInput.value)) {
|
if (!baseUrlInput.value || !/^https?:\/\/.+\/xiaozhi$/.test(baseUrlInput.value)) {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
@@ -113,20 +113,20 @@ async function saveServerBaseUrl() {
|
|||||||
clearAllCacheAfterUrlChange()
|
clearAllCacheAfterUrlChange()
|
||||||
|
|
||||||
uni.showModal({
|
uni.showModal({
|
||||||
title: '重启应用',
|
title: '重启应用',
|
||||||
content: '服务端地址已保存并清空缓存,是否立即重启生效?',
|
content: '服务端地址已保存并清空缓存,是否立即重启生效?',
|
||||||
confirmText: '立即重启',
|
confirmText: '立即重启',
|
||||||
cancelText: '稍后',
|
cancelText: '稍后',
|
||||||
success: (res) => {
|
success: (res) => {
|
||||||
if (res.confirm) {
|
if (res.confirm) {
|
||||||
restartApp()
|
restartApp()
|
||||||
}
|
}
|
||||||
else {
|
else {
|
||||||
toast.success('已保存,可稍后手动重启应用')
|
toast.success('已保存,可稍后手动重启应用')
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// 重置为 env 默认
|
// 重置为 env 默认
|
||||||
function resetServerBaseUrl() {
|
function resetServerBaseUrl() {
|
||||||
@@ -222,7 +222,7 @@ function showAbout() {
|
|||||||
title: `关于${import.meta.env.VITE_APP_TITLE}`,
|
title: `关于${import.meta.env.VITE_APP_TITLE}`,
|
||||||
content: `${import.meta.env.VITE_APP_TITLE}\n\n基于 Vue.js 3 + uni-app 构建的跨平台移动端管理应用,为小智ESP32智能硬件提供设备管理、智能体配置等功能。\n\n© 2025 xiaozhi-esp32-server`,
|
content: `${import.meta.env.VITE_APP_TITLE}\n\n基于 Vue.js 3 + uni-app 构建的跨平台移动端管理应用,为小智ESP32智能硬件提供设备管理、智能体配置等功能。\n\n© 2025 xiaozhi-esp32-server`,
|
||||||
title: `关于小智智控台`,
|
title: `关于小智智控台`,
|
||||||
content: `小智智控台\n\n基于 Vue.js 3 + uni-app 构建的跨平台移动端管理应用,为小智ESP32智能硬件提供设备管理、智能体配置等功能。\n\n© 2025 xiaozhi-esp32-server 0.7.6`,
|
content: `小智智控台\n\n基于 Vue.js 3 + uni-app 构建的跨平台移动端管理应用,为小智ESP32智能硬件提供设备管理、智能体配置等功能。\n\n© 2025 xiaozhi-esp32-server 0.7.7`,
|
||||||
showCancel: false,
|
showCancel: false,
|
||||||
confirmText: '确定',
|
confirmText: '确定',
|
||||||
})
|
})
|
||||||
@@ -263,17 +263,10 @@ onMounted(async () => {
|
|||||||
|
|
||||||
<view class="mb-[24rpx]">
|
<view class="mb-[24rpx]">
|
||||||
<view class="w-full rounded-[16rpx] border border-[#eeeeee] bg-[#f5f7fb] overflow-hidden">
|
<view class="w-full rounded-[16rpx] border border-[#eeeeee] bg-[#f5f7fb] overflow-hidden">
|
||||||
<wd-input
|
<wd-input v-model="baseUrlInput" type="text" clearable :maxlength="200"
|
||||||
v-model="baseUrlInput"
|
|
||||||
type="text"
|
|
||||||
clearable
|
|
||||||
:maxlength="200"
|
|
||||||
placeholder="输入服务端地址,如 https://example.com/xiaozhi"
|
placeholder="输入服务端地址,如 https://example.com/xiaozhi"
|
||||||
custom-class="!border-none !bg-transparent h-[88rpx] px-[24rpx] items-center"
|
custom-class="!border-none !bg-transparent h-[88rpx] px-[24rpx] items-center"
|
||||||
input-class="text-[28rpx] text-[#232338]"
|
input-class="text-[28rpx] text-[#232338]" @input="validateUrl" @blur="validateUrl" />
|
||||||
@input="validateUrl"
|
|
||||||
@blur="validateUrl"
|
|
||||||
/>
|
|
||||||
</view>
|
</view>
|
||||||
<text v-if="urlError" class="mt-[8rpx] block text-[24rpx] text-[#ff4d4f]">
|
<text v-if="urlError" class="mt-[8rpx] block text-[24rpx] text-[#ff4d4f]">
|
||||||
{{ urlError }}
|
{{ urlError }}
|
||||||
@@ -371,7 +364,7 @@ onMounted(async () => {
|
|||||||
|
|
||||||
<!-- 底部安全距离 -->
|
<!-- 底部安全距离 -->
|
||||||
<!-- 底部安全距离 -->
|
<!-- 底部安全距离 -->
|
||||||
<view style="height: env(safe-area-inset-bottom);" />
|
<view style="height: env(safe-area-inset-bottom);" />
|
||||||
</view>
|
</view>
|
||||||
</view>
|
</view>
|
||||||
</template>
|
</template>
|
||||||
|
|||||||
@@ -114,7 +114,10 @@ plugins:
|
|||||||
# 想稳定一点就自行申请替换,每天有1000次免费调用
|
# 想稳定一点就自行申请替换,每天有1000次免费调用
|
||||||
# 申请地址:https://console.qweather.com/#/apps/create-key/over
|
# 申请地址:https://console.qweather.com/#/apps/create-key/over
|
||||||
# 申请后通过这个链接可以找到自己的apihost:https://console.qweather.com/setting?lang=zh
|
# 申请后通过这个链接可以找到自己的apihost:https://console.qweather.com/setting?lang=zh
|
||||||
get_weather: {"api_host":"mj7p3y7naa.re.qweatherapi.com", "api_key": "a861d0d5e7bf4ee1a83d9a9e4f96d4da", "default_location": "广州" }
|
get_weather:
|
||||||
|
api_host: "mj7p3y7naa.re.qweatherapi.com"
|
||||||
|
api_key: "a861d0d5e7bf4ee1a83d9a9e4f96d4da"
|
||||||
|
default_location: "广州"
|
||||||
# 获取新闻插件的配置,这里根据需要的新闻类型传入对应的url链接,默认支持社会、科技、财经新闻
|
# 获取新闻插件的配置,这里根据需要的新闻类型传入对应的url链接,默认支持社会、科技、财经新闻
|
||||||
# 更多类型的新闻列表查看 https://www.chinanews.com.cn/rss/
|
# 更多类型的新闻列表查看 https://www.chinanews.com.cn/rss/
|
||||||
get_news_from_chinanews:
|
get_news_from_chinanews:
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ from config.config_loader import load_config
|
|||||||
from config.settings import check_config_file
|
from config.settings import check_config_file
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
|
|
||||||
SERVER_VERSION = "0.7.6"
|
SERVER_VERSION = "0.7.7"
|
||||||
_logger_initialized = False
|
_logger_initialized = False
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -4,15 +4,15 @@ import random
|
|||||||
import asyncio
|
import asyncio
|
||||||
from core.utils.dialogue import Message
|
from core.utils.dialogue import Message
|
||||||
from core.utils.util import audio_to_data
|
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.handle.sendAudioHandle import sendAudioMessage, send_stt_message
|
||||||
from core.utils.util import remove_punctuation_and_length, opus_datas_to_wav_bytes
|
from core.utils.util import remove_punctuation_and_length, opus_datas_to_wav_bytes
|
||||||
from core.providers.tts.dto.dto import ContentType, SentenceType
|
|
||||||
from core.providers.tools.device_mcp import (
|
from core.providers.tools.device_mcp import (
|
||||||
MCPClient,
|
MCPClient,
|
||||||
send_mcp_initialize_message,
|
send_mcp_initialize_message,
|
||||||
send_mcp_tools_list_request,
|
send_mcp_tools_list_request,
|
||||||
)
|
)
|
||||||
from core.utils.wakeup_word import WakeupWordsConfig
|
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
|
|
||||||
@@ -56,7 +56,16 @@ async def checkWakeupWords(conn, text):
|
|||||||
"enable_wakeup_words_response_cache"
|
"enable_wakeup_words_response_cache"
|
||||||
]
|
]
|
||||||
|
|
||||||
if not enable_wakeup_words_response_cache or not conn.tts:
|
# 等待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
|
return False
|
||||||
|
|
||||||
_, filtered_text = remove_punctuation_and_length(text)
|
_, filtered_text = remove_punctuation_and_length(text)
|
||||||
@@ -81,9 +90,10 @@ async def checkWakeupWords(conn, text):
|
|||||||
"text": "哈啰啊,我是小智啦,声音好听的台湾女孩一枚,超开心认识你耶,最近在忙啥,别忘了给我来点有趣的料哦,我超爱听八卦的啦",
|
"text": "哈啰啊,我是小智啦,声音好听的台湾女孩一枚,超开心认识你耶,最近在忙啥,别忘了给我来点有趣的料哦,我超爱听八卦的啦",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# 获取音频数据
|
||||||
|
opus_packets = audio_to_data(response.get("file_path"))
|
||||||
# 播放唤醒词回复
|
# 播放唤醒词回复
|
||||||
conn.client_abort = False
|
conn.client_abort = False
|
||||||
opus_packets, _ = audio_to_data(response.get("file_path"))
|
|
||||||
|
|
||||||
conn.logger.bind(tag=TAG).info(f"播放唤醒词回复: {response.get('text')}")
|
conn.logger.bind(tag=TAG).info(f"播放唤醒词回复: {response.get('text')}")
|
||||||
await sendAudioMessage(conn, SentenceType.FIRST, opus_packets, response.get("text"))
|
await sendAudioMessage(conn, SentenceType.FIRST, opus_packets, response.get("text"))
|
||||||
@@ -138,4 +148,4 @@ async def wakeupWordsResponse(conn):
|
|||||||
finally:
|
finally:
|
||||||
# 确保在任何情况下都释放锁
|
# 确保在任何情况下都释放锁
|
||||||
if _wakeup_response_lock.locked():
|
if _wakeup_response_lock.locked():
|
||||||
_wakeup_response_lock.release()
|
_wakeup_response_lock.release()
|
||||||
@@ -1,12 +1,12 @@
|
|||||||
import json
|
import json
|
||||||
import asyncio
|
|
||||||
import uuid
|
import uuid
|
||||||
from core.handle.sendAudioHandle import send_stt_message
|
import asyncio
|
||||||
from core.handle.helloHandle import checkWakeupWords
|
|
||||||
from core.utils.util import remove_punctuation_and_length
|
|
||||||
from core.providers.tts.dto.dto import ContentType
|
|
||||||
from core.utils.dialogue import Message
|
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 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 TTSMessageDTO, SentenceType
|
from core.providers.tts.dto.dto import TTSMessageDTO, SentenceType
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
@@ -24,9 +24,10 @@ async def handle_user_intent(conn, text):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
# 检查是否有明确的退出命令
|
# 检查是否有明确的退出命令
|
||||||
filtered_text = remove_punctuation_and_length(text)[1]
|
_, filtered_text = remove_punctuation_and_length(text)
|
||||||
if await check_direct_exit(conn, filtered_text):
|
if await check_direct_exit(conn, filtered_text):
|
||||||
return True
|
return True
|
||||||
|
|
||||||
# 检查是否是唤醒词
|
# 检查是否是唤醒词
|
||||||
if await checkWakeupWords(conn, filtered_text):
|
if await checkWakeupWords(conn, filtered_text):
|
||||||
return True
|
return True
|
||||||
|
|||||||
@@ -1,12 +1,11 @@
|
|||||||
from core.handle.sendAudioHandle import send_stt_message
|
import time
|
||||||
|
import json
|
||||||
|
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.handle.intentHandler import handle_user_intent
|
||||||
from core.utils.output_counter import check_device_output_limit
|
from core.utils.output_counter import check_device_output_limit
|
||||||
from core.handle.abortHandle import handleAbortMessage
|
from core.handle.sendAudioHandle import send_stt_message, SentenceType
|
||||||
import time
|
|
||||||
import asyncio
|
|
||||||
import json
|
|
||||||
from core.handle.sendAudioHandle import SentenceType
|
|
||||||
from core.utils.util import audio_to_data
|
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
|
|
||||||
@@ -22,7 +21,6 @@ async def handleAudioMessage(conn, audio):
|
|||||||
if not hasattr(conn, "vad_resume_task") or conn.vad_resume_task.done():
|
if not hasattr(conn, "vad_resume_task") or conn.vad_resume_task.done():
|
||||||
conn.vad_resume_task = asyncio.create_task(resume_vad_detection(conn))
|
conn.vad_resume_task = asyncio.create_task(resume_vad_detection(conn))
|
||||||
return
|
return
|
||||||
|
|
||||||
if have_voice:
|
if have_voice:
|
||||||
if conn.client_is_speaking:
|
if conn.client_is_speaking:
|
||||||
await handleAbortMessage(conn)
|
await handleAbortMessage(conn)
|
||||||
@@ -31,18 +29,16 @@ async def handleAudioMessage(conn, audio):
|
|||||||
# 接收音频
|
# 接收音频
|
||||||
await conn.asr.receive_audio(conn, audio, have_voice)
|
await conn.asr.receive_audio(conn, audio, have_voice)
|
||||||
|
|
||||||
|
|
||||||
async def resume_vad_detection(conn):
|
async def resume_vad_detection(conn):
|
||||||
# 等待2秒后恢复VAD检测
|
# 等待2秒后恢复VAD检测
|
||||||
await asyncio.sleep(1)
|
await asyncio.sleep(1)
|
||||||
conn.just_woken_up = False
|
conn.just_woken_up = False
|
||||||
|
|
||||||
|
|
||||||
async def startToChat(conn, text):
|
async def startToChat(conn, text):
|
||||||
# 检查输入是否是JSON格式(包含说话人信息)
|
# 检查输入是否是JSON格式(包含说话人信息)
|
||||||
speaker_name = None
|
speaker_name = None
|
||||||
actual_text = text
|
actual_text = text
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# 尝试解析JSON格式的输入
|
# 尝试解析JSON格式的输入
|
||||||
if text.strip().startswith('{') and text.strip().endswith('}'):
|
if text.strip().startswith('{') and text.strip().endswith('}'):
|
||||||
@@ -51,13 +47,13 @@ async def startToChat(conn, text):
|
|||||||
speaker_name = data['speaker']
|
speaker_name = data['speaker']
|
||||||
actual_text = data['content']
|
actual_text = data['content']
|
||||||
conn.logger.bind(tag=TAG).info(f"解析到说话人信息: {speaker_name}")
|
conn.logger.bind(tag=TAG).info(f"解析到说话人信息: {speaker_name}")
|
||||||
|
|
||||||
# 直接使用JSON格式的文本,不解析
|
# 直接使用JSON格式的文本,不解析
|
||||||
actual_text = text
|
actual_text = text
|
||||||
except (json.JSONDecodeError, KeyError):
|
except (json.JSONDecodeError, KeyError):
|
||||||
# 如果解析失败,继续使用原始文本
|
# 如果解析失败,继续使用原始文本
|
||||||
pass
|
pass
|
||||||
|
|
||||||
# 保存说话人信息到连接对象
|
# 保存说话人信息到连接对象
|
||||||
if speaker_name:
|
if speaker_name:
|
||||||
conn.current_speaker = speaker_name
|
conn.current_speaker = speaker_name
|
||||||
@@ -118,10 +114,12 @@ async def no_voice_close_connect(conn, have_voice):
|
|||||||
|
|
||||||
|
|
||||||
async def max_out_size(conn):
|
async def max_out_size(conn):
|
||||||
|
# 播放超出最大输出字数的提示
|
||||||
|
conn.client_abort = False
|
||||||
text = "不好意思,我现在有点事情要忙,明天这个时候我们再聊,约好了哦!明天不见不散,拜拜!"
|
text = "不好意思,我现在有点事情要忙,明天这个时候我们再聊,约好了哦!明天不见不散,拜拜!"
|
||||||
await send_stt_message(conn, text)
|
await send_stt_message(conn, text)
|
||||||
file_path = "config/assets/max_output_size.wav"
|
file_path = "config/assets/max_output_size.wav"
|
||||||
opus_packets, _ = audio_to_data(file_path)
|
opus_packets = audio_to_data(file_path)
|
||||||
conn.tts.tts_audio_queue.put((SentenceType.LAST, opus_packets, text))
|
conn.tts.tts_audio_queue.put((SentenceType.LAST, opus_packets, text))
|
||||||
conn.close_after_chat = True
|
conn.close_after_chat = True
|
||||||
|
|
||||||
@@ -140,7 +138,7 @@ async def check_bind_device(conn):
|
|||||||
|
|
||||||
# 播放提示音
|
# 播放提示音
|
||||||
music_path = "config/assets/bind_code.wav"
|
music_path = "config/assets/bind_code.wav"
|
||||||
opus_packets, _ = audio_to_data(music_path)
|
opus_packets = audio_to_data(music_path)
|
||||||
conn.tts.tts_audio_queue.put((SentenceType.FIRST, opus_packets, text))
|
conn.tts.tts_audio_queue.put((SentenceType.FIRST, opus_packets, text))
|
||||||
|
|
||||||
# 逐个播放数字
|
# 逐个播放数字
|
||||||
@@ -148,15 +146,17 @@ async def check_bind_device(conn):
|
|||||||
try:
|
try:
|
||||||
digit = conn.bind_code[i]
|
digit = conn.bind_code[i]
|
||||||
num_path = f"config/assets/bind_code/{digit}.wav"
|
num_path = f"config/assets/bind_code/{digit}.wav"
|
||||||
num_packets, _ = audio_to_data(num_path)
|
num_packets = audio_to_data(num_path)
|
||||||
conn.tts.tts_audio_queue.put((SentenceType.MIDDLE, num_packets, None))
|
conn.tts.tts_audio_queue.put((SentenceType.MIDDLE, num_packets, None))
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
conn.logger.bind(tag=TAG).error(f"播放数字音频失败: {e}")
|
conn.logger.bind(tag=TAG).error(f"播放数字音频失败: {e}")
|
||||||
continue
|
continue
|
||||||
conn.tts.tts_audio_queue.put((SentenceType.LAST, [], None))
|
conn.tts.tts_audio_queue.put((SentenceType.LAST, [], None))
|
||||||
else:
|
else:
|
||||||
|
# 播放未绑定提示
|
||||||
|
conn.client_abort = False
|
||||||
text = f"没有找到该设备的版本信息,请正确配置 OTA地址,然后重新编译固件。"
|
text = f"没有找到该设备的版本信息,请正确配置 OTA地址,然后重新编译固件。"
|
||||||
await send_stt_message(conn, text)
|
await send_stt_message(conn, text)
|
||||||
music_path = "config/assets/bind_not_found.wav"
|
music_path = "config/assets/bind_not_found.wav"
|
||||||
opus_packets, _ = audio_to_data(music_path)
|
opus_packets = audio_to_data(music_path)
|
||||||
conn.tts.tts_audio_queue.put((SentenceType.LAST, opus_packets, text))
|
conn.tts.tts_audio_queue.put((SentenceType.LAST, opus_packets, text))
|
||||||
|
|||||||
@@ -1,25 +1,26 @@
|
|||||||
import json
|
import json
|
||||||
import asyncio
|
|
||||||
import time
|
import time
|
||||||
from core.providers.tts.dto.dto import SentenceType
|
import asyncio
|
||||||
from core.utils import textUtils
|
from core.utils import textUtils
|
||||||
|
from core.utils.util import audio_to_data
|
||||||
|
from core.providers.tts.dto.dto import SentenceType
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
|
|
||||||
|
|
||||||
async def sendAudioMessage(conn, sentenceType, audios, text):
|
async def sendAudioMessage(conn, sentenceType, audios, text):
|
||||||
# 发送句子开始消息
|
|
||||||
conn.logger.bind(tag=TAG).info(f"发送音频消息: {sentenceType}, {text}")
|
|
||||||
|
|
||||||
pre_buffer = False
|
|
||||||
if conn.tts.tts_audio_first_sentence:
|
if conn.tts.tts_audio_first_sentence:
|
||||||
conn.logger.bind(tag=TAG).info(f"发送第一段语音: {text}")
|
conn.logger.bind(tag=TAG).info(f"发送第一段语音: {text}")
|
||||||
conn.tts.tts_audio_first_sentence = False
|
conn.tts.tts_audio_first_sentence = False
|
||||||
pre_buffer = True
|
await send_tts_message(conn, "start", None)
|
||||||
|
|
||||||
await send_tts_message(conn, "sentence_start", text)
|
if sentenceType == SentenceType.FIRST:
|
||||||
|
await send_tts_message(conn, "sentence_start", text)
|
||||||
|
|
||||||
await sendAudio(conn, audios, pre_buffer)
|
await sendAudio(conn, audios)
|
||||||
|
# 发送句子开始消息
|
||||||
|
if sentenceType is not SentenceType.MIDDLE:
|
||||||
|
conn.logger.bind(tag=TAG).info(f"发送音频消息: {sentenceType}, {text}")
|
||||||
|
|
||||||
# 发送结束消息(如果是最后一个文本)
|
# 发送结束消息(如果是最后一个文本)
|
||||||
if conn.llm_finish_task and sentenceType == SentenceType.LAST:
|
if conn.llm_finish_task and sentenceType == SentenceType.LAST:
|
||||||
@@ -30,45 +31,83 @@ async def sendAudioMessage(conn, sentenceType, audios, text):
|
|||||||
|
|
||||||
|
|
||||||
# 播放音频
|
# 播放音频
|
||||||
async def sendAudio(conn, audios, pre_buffer=True):
|
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:
|
if audios is None or len(audios) == 0:
|
||||||
return
|
return
|
||||||
# 流控参数优化
|
|
||||||
frame_duration = 60 # 帧时长(毫秒),匹配 Opus 编码
|
|
||||||
start_time = time.perf_counter()
|
|
||||||
play_position = 0
|
|
||||||
|
|
||||||
# 仅当第一句话时执行预缓冲
|
if isinstance(audios, bytes):
|
||||||
if pre_buffer:
|
|
||||||
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:]
|
|
||||||
else:
|
|
||||||
remaining_audios = audios
|
|
||||||
|
|
||||||
# 播放剩余音频帧
|
|
||||||
for opus_packet in remaining_audios:
|
|
||||||
if conn.client_abort:
|
if conn.client_abort:
|
||||||
break
|
return
|
||||||
|
|
||||||
# 重置没有声音的状态
|
|
||||||
conn.last_activity_time = time.time() * 1000
|
conn.last_activity_time = time.time() * 1000
|
||||||
|
|
||||||
# 计算预期发送时间
|
# 获取或初始化流控状态
|
||||||
expected_time = start_time + (play_position / 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()
|
current_time = time.perf_counter()
|
||||||
|
# 计算预期发送时间
|
||||||
|
expected_time = flow_control["start_time"] + (
|
||||||
|
flow_control["packet_count"] * frame_duration / 1000
|
||||||
|
)
|
||||||
delay = expected_time - current_time
|
delay = expected_time - current_time
|
||||||
if delay > 0:
|
if delay > 0:
|
||||||
await asyncio.sleep(delay)
|
await asyncio.sleep(delay)
|
||||||
|
|
||||||
await conn.websocket.send(opus_packet)
|
# 发送数据包
|
||||||
|
await conn.websocket.send(audios)
|
||||||
|
|
||||||
play_position += frame_duration
|
# 更新流控状态
|
||||||
|
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):
|
async def send_tts_message(conn, state, text=None):
|
||||||
"""发送 TTS 状态消息"""
|
"""发送 TTS 状态消息"""
|
||||||
|
if text is None and state == "sentence_start":
|
||||||
|
return
|
||||||
message = {"type": "tts", "state": state, "session_id": conn.session_id}
|
message = {"type": "tts", "state": state, "session_id": conn.session_id}
|
||||||
if text is not None:
|
if text is not None:
|
||||||
message["text"] = textUtils.check_emoji(text)
|
message["text"] = textUtils.check_emoji(text)
|
||||||
@@ -81,7 +120,7 @@ async def send_tts_message(conn, state, text=None):
|
|||||||
stop_tts_notify_voice = conn.config.get(
|
stop_tts_notify_voice = conn.config.get(
|
||||||
"stop_tts_notify_voice", "config/assets/tts_notify.mp3"
|
"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)
|
await sendAudio(conn, audios)
|
||||||
# 清除服务端讲话状态
|
# 清除服务端讲话状态
|
||||||
conn.clearSpeakStatus()
|
conn.clearSpeakStatus()
|
||||||
@@ -91,18 +130,17 @@ async def send_tts_message(conn, state, text=None):
|
|||||||
|
|
||||||
|
|
||||||
async def send_stt_message(conn, text):
|
async def send_stt_message(conn, text):
|
||||||
|
"""发送 STT 状态消息"""
|
||||||
end_prompt_str = conn.config.get("end_prompt", {}).get("prompt")
|
end_prompt_str = conn.config.get("end_prompt", {}).get("prompt")
|
||||||
if end_prompt_str and end_prompt_str == text:
|
if end_prompt_str and end_prompt_str == text:
|
||||||
await send_tts_message(conn, "start")
|
await send_tts_message(conn, "start")
|
||||||
return
|
return
|
||||||
|
|
||||||
"""发送 STT 状态消息"""
|
|
||||||
|
|
||||||
# 解析JSON格式,提取实际的用户说话内容
|
# 解析JSON格式,提取实际的用户说话内容
|
||||||
display_text = text
|
display_text = text
|
||||||
try:
|
try:
|
||||||
# 尝试解析JSON格式
|
# 尝试解析JSON格式
|
||||||
if text.strip().startswith('{') and text.strip().endswith('}'):
|
if text.strip().startswith("{") and text.strip().endswith("}"):
|
||||||
parsed_data = json.loads(text)
|
parsed_data = json.loads(text)
|
||||||
if isinstance(parsed_data, dict) and "content" in parsed_data:
|
if isinstance(parsed_data, dict) and "content" in parsed_data:
|
||||||
# 如果是包含说话人信息的JSON格式,只显示content部分
|
# 如果是包含说话人信息的JSON格式,只显示content部分
|
||||||
|
|||||||
@@ -1,169 +1,14 @@
|
|||||||
import json
|
from core.handle.textMessageHandlerRegistry import TextMessageHandlerRegistry
|
||||||
import time
|
from core.handle.textMessageProcessor import TextMessageProcessor
|
||||||
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.handle.sendAudioHandle import send_stt_message, send_tts_message
|
|
||||||
from core.providers.tools.device_iot import handleIotDescriptors, handleIotStatus
|
|
||||||
from core.handle.reportHandle import enqueue_asr_report
|
|
||||||
import asyncio
|
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
|
|
||||||
|
# 全局处理器注册表
|
||||||
|
message_registry = TextMessageHandlerRegistry()
|
||||||
|
|
||||||
|
# 创建全局消息处理器实例
|
||||||
|
message_processor = TextMessageProcessor(message_registry)
|
||||||
|
|
||||||
async def handleTextMessage(conn, message):
|
async def handleTextMessage(conn, message):
|
||||||
"""处理文本消息"""
|
"""处理文本消息"""
|
||||||
try:
|
await message_processor.process_message(conn, message)
|
||||||
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")
|
|
||||||
# 是否开启唤醒词回复
|
|
||||||
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)
|
|
||||||
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)
|
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -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)
|
||||||
@@ -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"]))
|
||||||
@@ -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)
|
||||||
@@ -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"])
|
||||||
|
)
|
||||||
@@ -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)
|
||||||
@@ -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
|
||||||
@@ -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())
|
||||||
@@ -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)
|
||||||
@@ -0,0 +1,11 @@
|
|||||||
|
from enum import Enum
|
||||||
|
|
||||||
|
|
||||||
|
class TextMessageType(Enum):
|
||||||
|
"""消息类型枚举"""
|
||||||
|
HELLO = "hello"
|
||||||
|
ABORT = "abort"
|
||||||
|
LISTEN = "listen"
|
||||||
|
IOT = "iot"
|
||||||
|
MCP = "mcp"
|
||||||
|
SERVER = "server"
|
||||||
@@ -1,18 +1,18 @@
|
|||||||
import os
|
import os
|
||||||
|
import io
|
||||||
import wave
|
import wave
|
||||||
import uuid
|
import uuid
|
||||||
|
import json
|
||||||
|
import time
|
||||||
import queue
|
import queue
|
||||||
import asyncio
|
import asyncio
|
||||||
import traceback
|
import traceback
|
||||||
import threading
|
import threading
|
||||||
import opuslib_next
|
import opuslib_next
|
||||||
import json
|
|
||||||
import io
|
|
||||||
import time
|
|
||||||
import concurrent.futures
|
import concurrent.futures
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from config.logger import setup_logging
|
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.receiveAudioHandle import startToChat
|
||||||
from core.handle.reportHandle import enqueue_asr_report
|
from core.handle.reportHandle import enqueue_asr_report
|
||||||
from core.utils.util import remove_punctuation_and_length
|
from core.utils.util import remove_punctuation_and_length
|
||||||
@@ -87,11 +87,9 @@ class ASRProviderBase(ABC):
|
|||||||
|
|
||||||
# 预先准备WAV数据
|
# 预先准备WAV数据
|
||||||
wav_data = None
|
wav_data = None
|
||||||
# 使用连接的声纹识别提供者
|
|
||||||
if conn.voiceprint_provider and combined_pcm_data:
|
if conn.voiceprint_provider and combined_pcm_data:
|
||||||
wav_data = self._pcm_to_wav(combined_pcm_data)
|
wav_data = self._pcm_to_wav(combined_pcm_data)
|
||||||
|
|
||||||
|
|
||||||
# 定义ASR任务
|
# 定义ASR任务
|
||||||
def run_asr():
|
def run_asr():
|
||||||
start_time = time.monotonic()
|
start_time = time.monotonic()
|
||||||
@@ -132,8 +130,6 @@ class ASRProviderBase(ABC):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
# 使用线程池执行器并行运行
|
# 使用线程池执行器并行运行
|
||||||
parallel_start_time = time.monotonic()
|
|
||||||
|
|
||||||
with concurrent.futures.ThreadPoolExecutor(max_workers=2) as thread_executor:
|
with concurrent.futures.ThreadPoolExecutor(max_workers=2) as thread_executor:
|
||||||
asr_future = thread_executor.submit(run_asr)
|
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)
|
speaker_name = results.get("voiceprint", None)
|
||||||
|
|
||||||
# 记录识别结果
|
# 记录识别结果
|
||||||
|
|||||||
@@ -167,14 +167,13 @@ class ServerMCPClient:
|
|||||||
|
|
||||||
# 建立SSEClient
|
# 建立SSEClient
|
||||||
elif "url" in self.config:
|
elif "url" in self.config:
|
||||||
|
headers = dict(self.config.get("headers", {}))
|
||||||
|
# TODO 兼容旧版本
|
||||||
if "API_ACCESS_TOKEN" in self.config:
|
if "API_ACCESS_TOKEN" in self.config:
|
||||||
headers = {
|
headers["Authorization"] = f"Bearer {self.config['API_ACCESS_TOKEN']}"
|
||||||
"Authorization": f"Bearer {self.config['API_ACCESS_TOKEN']}"
|
self.logger.bind(tag=TAG).warning(f"你正在使用旧过时的配置 API_ACCESS_TOKEN ,请在.mcp_server_settings.json中将API_ACCESS_TOKEN直接设置在headers中,例如 'Authorization': 'Bearer API_ACCESS_TOKEN'")
|
||||||
}
|
|
||||||
else:
|
|
||||||
headers = {}
|
|
||||||
sse_r, sse_w = await stack.enter_async_context(
|
sse_r, sse_w = await stack.enter_async_context(
|
||||||
sse_client(self.config["url"], headers=headers)
|
sse_client(self.config["url"], headers=headers, timeout=self.config.get("timeout", 5), sse_read_timeout=self.config.get("sse_read_timeout", 60 * 5))
|
||||||
)
|
)
|
||||||
read_stream, write_stream = sse_r, sse_w
|
read_stream, write_stream = sse_r, sse_w
|
||||||
|
|
||||||
|
|||||||
@@ -268,11 +268,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
)
|
)
|
||||||
if message.content_file and os.path.exists(message.content_file):
|
if message.content_file and os.path.exists(message.content_file):
|
||||||
# 先处理文件音频数据
|
# 先处理文件音频数据
|
||||||
file_audio = self._process_audio_file(message.content_file)
|
self._process_audio_file_stream(message.content_file, callback=lambda audio_data: self.handle_audio_file(audio_data, message.content_detail))
|
||||||
self.before_stop_play_files.append(
|
|
||||||
(file_audio, message.content_detail)
|
|
||||||
)
|
|
||||||
|
|
||||||
if message.sentence_type == SentenceType.LAST:
|
if message.sentence_type == SentenceType.LAST:
|
||||||
try:
|
try:
|
||||||
logger.bind(tag=TAG).info("开始结束TTS会话...")
|
logger.bind(tag=TAG).info("开始结束TTS会话...")
|
||||||
@@ -422,9 +418,6 @@ class TTSProvider(TTSProviderBase):
|
|||||||
|
|
||||||
async def _start_monitor_tts_response(self):
|
async def _start_monitor_tts_response(self):
|
||||||
"""监听TTS响应"""
|
"""监听TTS响应"""
|
||||||
opus_datas_cache = []
|
|
||||||
is_first_sentence = True
|
|
||||||
first_sentence_segment_count = 0 # 添加计数器
|
|
||||||
try:
|
try:
|
||||||
session_finished = False # 标记会话是否正常结束
|
session_finished = False # 标记会话是否正常结束
|
||||||
while not self.conn.stop_event.is_set():
|
while not self.conn.stop_event.is_set():
|
||||||
@@ -445,28 +438,16 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.tts_audio_queue.put(
|
self.tts_audio_queue.put(
|
||||||
(SentenceType.FIRST, [], None)
|
(SentenceType.FIRST, [], None)
|
||||||
)
|
)
|
||||||
elif event_name == "SentenceBegin":
|
|
||||||
opus_datas_cache = []
|
|
||||||
elif event_name == "SentenceEnd":
|
elif event_name == "SentenceEnd":
|
||||||
if (
|
# 发送缓存的数据
|
||||||
not is_first_sentence
|
if self.conn.tts_MessageText:
|
||||||
or first_sentence_segment_count > 10
|
logger.bind(tag=TAG).info(
|
||||||
):
|
f"句子语音生成成功: {self.conn.tts_MessageText}"
|
||||||
# 发送缓存的数据
|
)
|
||||||
if self.conn.tts_MessageText:
|
self.tts_audio_queue.put(
|
||||||
logger.bind(tag=TAG).info(
|
(SentenceType.FIRST, [], self.conn.tts_MessageText)
|
||||||
f"句子语音生成成功: {self.conn.tts_MessageText}"
|
)
|
||||||
)
|
self.conn.tts_MessageText = None
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus_datas_cache, self.conn.tts_MessageText)
|
|
||||||
)
|
|
||||||
self.conn.tts_MessageText = None
|
|
||||||
else:
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus_datas_cache, None)
|
|
||||||
)
|
|
||||||
# 第一句话结束后,将标志设置为False
|
|
||||||
is_first_sentence = False
|
|
||||||
elif event_name == "SynthesisCompleted":
|
elif event_name == "SynthesisCompleted":
|
||||||
logger.bind(tag=TAG).debug(f"会话结束~~")
|
logger.bind(tag=TAG).debug(f"会话结束~~")
|
||||||
self._process_before_stop_play_files()
|
self._process_before_stop_play_files()
|
||||||
@@ -477,22 +458,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
# 二进制消息(音频数据)
|
# 二进制消息(音频数据)
|
||||||
elif isinstance(msg, (bytes, bytearray)):
|
elif isinstance(msg, (bytes, bytearray)):
|
||||||
logger.bind(tag=TAG).debug(f"推送数据到队列里面~~")
|
logger.bind(tag=TAG).debug(f"推送数据到队列里面~~")
|
||||||
opus_datas = self.opus_encoder.encode_pcm_to_opus(msg, False)
|
self.opus_encoder.encode_pcm_to_opus_stream(msg, False, self.handle_opus)
|
||||||
logger.bind(tag=TAG).debug(
|
|
||||||
f"推送数据到队列里面帧数~~{len(opus_datas)}"
|
|
||||||
)
|
|
||||||
if is_first_sentence:
|
|
||||||
first_sentence_segment_count += 1
|
|
||||||
if first_sentence_segment_count <= 6:
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus_datas, None)
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
opus_datas_cache.extend(opus_datas)
|
|
||||||
else:
|
|
||||||
# 后续句子缓存
|
|
||||||
opus_datas_cache.extend(opus_datas)
|
|
||||||
|
|
||||||
except websockets.ConnectionClosed:
|
except websockets.ConnectionClosed:
|
||||||
logger.bind(tag=TAG).warning("WebSocket连接已关闭")
|
logger.bind(tag=TAG).warning("WebSocket连接已关闭")
|
||||||
break
|
break
|
||||||
@@ -615,11 +581,11 @@ class TTSProvider(TTSProviderBase):
|
|||||||
while not synthesis_completed:
|
while not synthesis_completed:
|
||||||
msg = await ws.recv()
|
msg = await ws.recv()
|
||||||
if isinstance(msg, (bytes, bytearray)):
|
if isinstance(msg, (bytes, bytearray)):
|
||||||
# 编码为Opus并收集
|
self.opus_encoder.encode_pcm_to_opus_stream(
|
||||||
opus_frames = self.opus_encoder.encode_pcm_to_opus(
|
msg,
|
||||||
msg, False
|
end_of_stream=False,
|
||||||
|
callback=lambda opus: audio_data.append(opus)
|
||||||
)
|
)
|
||||||
audio_data.extend(opus_frames)
|
|
||||||
elif isinstance(msg, str):
|
elif isinstance(msg, str):
|
||||||
data = json.loads(msg)
|
data = json.loads(msg)
|
||||||
header = data.get("header", {})
|
header = data.get("header", {})
|
||||||
@@ -650,4 +616,4 @@ class TTSProvider(TTSProviderBase):
|
|||||||
return audio_data
|
return audio_data
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"生成音频数据失败: {str(e)}")
|
logger.bind(tag=TAG).error(f"生成音频数据失败: {str(e)}")
|
||||||
return []
|
return []
|
||||||
@@ -1,19 +1,22 @@
|
|||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
import queue
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
|
import queue
|
||||||
import asyncio
|
import asyncio
|
||||||
import threading
|
import threading
|
||||||
|
import traceback
|
||||||
from core.utils import p3
|
from core.utils import p3
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from core.utils import textUtils
|
from core.utils import textUtils
|
||||||
|
from typing import Callable, Any
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from core.utils.util import audio_to_data, audio_bytes_to_data
|
|
||||||
from core.utils.tts import MarkdownCleaner
|
from core.utils.tts import MarkdownCleaner
|
||||||
from core.utils.output_counter import add_device_output
|
from core.utils.output_counter import add_device_output
|
||||||
from core.handle.reportHandle import enqueue_tts_report
|
from core.handle.reportHandle import enqueue_tts_report
|
||||||
from core.handle.sendAudioHandle import sendAudioMessage
|
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 (
|
from core.providers.tts.dto.dto import (
|
||||||
TTSMessageDTO,
|
TTSMessageDTO,
|
||||||
SentenceType,
|
SentenceType,
|
||||||
@@ -21,8 +24,6 @@ from core.providers.tts.dto.dto import (
|
|||||||
InterfaceType,
|
InterfaceType,
|
||||||
)
|
)
|
||||||
|
|
||||||
import traceback
|
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
|
|
||||||
@@ -31,7 +32,6 @@ class TTSProviderBase(ABC):
|
|||||||
def __init__(self, config, delete_audio_file):
|
def __init__(self, config, delete_audio_file):
|
||||||
self.interface_type = InterfaceType.NON_STREAM
|
self.interface_type = InterfaceType.NON_STREAM
|
||||||
self.conn = None
|
self.conn = None
|
||||||
self.tts_timeout = 10
|
|
||||||
self.delete_audio_file = delete_audio_file
|
self.delete_audio_file = delete_audio_file
|
||||||
self.audio_file_type = "wav"
|
self.audio_file_type = "wav"
|
||||||
self.output_file = config.get("output_dir", "tmp/")
|
self.output_file = config.get("output_dir", "tmp/")
|
||||||
@@ -50,11 +50,9 @@ class TTSProviderBase(ABC):
|
|||||||
";",
|
";",
|
||||||
";",
|
";",
|
||||||
":",
|
":",
|
||||||
"~",
|
|
||||||
)
|
)
|
||||||
self.first_sentence_punctuations = (
|
self.first_sentence_punctuations = (
|
||||||
",",
|
",",
|
||||||
"~",
|
|
||||||
"~",
|
"~",
|
||||||
"、",
|
"、",
|
||||||
",",
|
",",
|
||||||
@@ -77,6 +75,75 @@ class TTSProviderBase(ABC):
|
|||||||
f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
|
f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def handle_opus(self, opus_data: bytes):
|
||||||
|
logger.bind(tag=TAG).debug(f"推送数据到队列里面帧数~~ {len(opus_data)}")
|
||||||
|
self.tts_audio_queue.put((SentenceType.MIDDLE, opus_data, None))
|
||||||
|
|
||||||
|
def handle_audio_file(self, file_audio: bytes, text):
|
||||||
|
self.before_stop_play_files.append((file_audio, text))
|
||||||
|
|
||||||
|
def to_tts_stream(self, text, opus_handler: Callable[[bytes], None] = None) -> None:
|
||||||
|
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:
|
||||||
|
self.tts_audio_queue.put((SentenceType.FIRST, None, text))
|
||||||
|
audio_bytes_to_data_stream(
|
||||||
|
audio_bytes,
|
||||||
|
file_type=self.audio_file_type,
|
||||||
|
is_opus=True,
|
||||||
|
callback=opus_handler,
|
||||||
|
)
|
||||||
|
break
|
||||||
|
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},请检查网络或服务是否正常"
|
||||||
|
)
|
||||||
|
self.tts_audio_queue.put((SentenceType.FIRST, None, text))
|
||||||
|
self._process_audio_file_stream(tmp_file, callback=opus_handler)
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"Failed to generate TTS file: {e}")
|
||||||
|
return None
|
||||||
|
|
||||||
def to_tts(self, text):
|
def to_tts(self, text):
|
||||||
text = MarkdownCleaner.clean_markdown(text)
|
text = MarkdownCleaner.clean_markdown(text)
|
||||||
max_repeat_time = 5
|
max_repeat_time = 5
|
||||||
@@ -86,8 +153,12 @@ class TTSProviderBase(ABC):
|
|||||||
try:
|
try:
|
||||||
audio_bytes = asyncio.run(self.text_to_speak(text, None))
|
audio_bytes = asyncio.run(self.text_to_speak(text, None))
|
||||||
if audio_bytes:
|
if audio_bytes:
|
||||||
audio_datas, _ = audio_bytes_to_data(
|
audio_datas = []
|
||||||
audio_bytes, file_type=self.audio_file_type, is_opus=True
|
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
|
return audio_datas
|
||||||
else:
|
else:
|
||||||
@@ -139,13 +210,17 @@ class TTSProviderBase(ABC):
|
|||||||
async def text_to_speak(self, text, output_file):
|
async def text_to_speak(self, text, output_file):
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def audio_to_pcm_data(self, audio_file_path):
|
def audio_to_pcm_data_stream(
|
||||||
|
self, audio_file_path, callback: Callable[[Any], Any] = None
|
||||||
|
):
|
||||||
"""音频文件转换为PCM编码"""
|
"""音频文件转换为PCM编码"""
|
||||||
return audio_to_data(audio_file_path, is_opus=False)
|
return audio_to_data_stream(audio_file_path, is_opus=False, callback=callback)
|
||||||
|
|
||||||
def audio_to_opus_data(self, audio_file_path):
|
def audio_to_opus_data_stream(
|
||||||
|
self, audio_file_path, callback: Callable[[Any], Any] = None
|
||||||
|
):
|
||||||
"""音频文件转换为Opus编码"""
|
"""音频文件转换为Opus编码"""
|
||||||
return audio_to_data(audio_file_path, is_opus=True)
|
return audio_to_data_stream(audio_file_path, is_opus=True, callback=callback)
|
||||||
|
|
||||||
def tts_one_sentence(
|
def tts_one_sentence(
|
||||||
self,
|
self,
|
||||||
@@ -177,7 +252,6 @@ class TTSProviderBase(ABC):
|
|||||||
|
|
||||||
async def open_audio_channels(self, conn):
|
async def open_audio_channels(self, conn):
|
||||||
self.conn = conn
|
self.conn = conn
|
||||||
self.tts_timeout = conn.config.get("tts_timeout", 10)
|
|
||||||
# tts 消化线程
|
# tts 消化线程
|
||||||
self.tts_priority_thread = threading.Thread(
|
self.tts_priority_thread = threading.Thread(
|
||||||
target=self.tts_text_priority_thread, daemon=True
|
target=self.tts_text_priority_thread, daemon=True
|
||||||
@@ -212,30 +286,16 @@ class TTSProviderBase(ABC):
|
|||||||
self.tts_text_buff.append(message.content_detail)
|
self.tts_text_buff.append(message.content_detail)
|
||||||
segment_text = self._get_segment_text()
|
segment_text = self._get_segment_text()
|
||||||
if segment_text:
|
if segment_text:
|
||||||
if self.delete_audio_file:
|
self.to_tts_stream(segment_text, opus_handler=self.handle_opus)
|
||||||
audio_datas = self.to_tts(segment_text)
|
|
||||||
if audio_datas:
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(message.sentence_type, audio_datas, segment_text)
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
tts_file = self.to_tts(segment_text)
|
|
||||||
if tts_file:
|
|
||||||
audio_datas = self._process_audio_file(tts_file)
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(message.sentence_type, audio_datas, segment_text)
|
|
||||||
)
|
|
||||||
elif ContentType.FILE == message.content_type:
|
elif ContentType.FILE == message.content_type:
|
||||||
self._process_remaining_text()
|
self._process_remaining_text_stream(opus_handler=self.handle_opus)
|
||||||
tts_file = message.content_file
|
tts_file = message.content_file
|
||||||
if tts_file and os.path.exists(tts_file):
|
if tts_file and os.path.exists(tts_file):
|
||||||
audio_datas = self._process_audio_file(tts_file)
|
self._process_audio_file_stream(
|
||||||
self.tts_audio_queue.put(
|
tts_file, callback=self.handle_opus
|
||||||
(message.sentence_type, audio_datas, message.content_detail)
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if message.sentence_type == SentenceType.LAST:
|
if message.sentence_type == SentenceType.LAST:
|
||||||
self._process_remaining_text()
|
self._process_remaining_text_stream(opus_handler=self.handle_opus)
|
||||||
self.tts_audio_queue.put(
|
self.tts_audio_queue.put(
|
||||||
(message.sentence_type, [], message.content_detail)
|
(message.sentence_type, [], message.content_detail)
|
||||||
)
|
)
|
||||||
@@ -249,29 +309,59 @@ class TTSProviderBase(ABC):
|
|||||||
continue
|
continue
|
||||||
|
|
||||||
def _audio_play_priority_thread(self):
|
def _audio_play_priority_thread(self):
|
||||||
|
# 需要上报的文本和音频列表
|
||||||
|
enqueue_text = None
|
||||||
|
enqueue_audio = None
|
||||||
while not self.conn.stop_event.is_set():
|
while not self.conn.stop_event.is_set():
|
||||||
text = None
|
text = None
|
||||||
try:
|
try:
|
||||||
try:
|
try:
|
||||||
sentence_type, audio_datas, text = self.tts_audio_queue.get(
|
sentence_type, audio_datas, text = self.tts_audio_queue.get(
|
||||||
timeout=1
|
timeout=0.1
|
||||||
)
|
)
|
||||||
except queue.Empty:
|
except queue.Empty:
|
||||||
if self.conn.stop_event.is_set():
|
if self.conn.stop_event.is_set():
|
||||||
break
|
break
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
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) and enqueue_audio is not None:
|
||||||
|
enqueue_audio.append(audio_datas)
|
||||||
|
|
||||||
|
# 发送音频
|
||||||
future = asyncio.run_coroutine_threadsafe(
|
future = asyncio.run_coroutine_threadsafe(
|
||||||
sendAudioMessage(self.conn, sentence_type, audio_datas, text),
|
sendAudioMessage(self.conn, sentence_type, audio_datas, text),
|
||||||
self.conn.loop,
|
self.conn.loop,
|
||||||
)
|
)
|
||||||
future.result()
|
future.result()
|
||||||
|
|
||||||
|
# 记录输出和报告
|
||||||
if self.conn.max_output_size > 0 and text:
|
if self.conn.max_output_size > 0 and text:
|
||||||
add_device_output(self.conn.headers.get("device-id"), len(text))
|
add_device_output(self.conn.headers.get("device-id"), len(text))
|
||||||
enqueue_tts_report(self.conn, text, audio_datas)
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(
|
logger.bind(tag=TAG).error(f"audio_play_priority_thread: {text} {e}")
|
||||||
f"audio_play_priority priority_thread: {text} {e}"
|
|
||||||
)
|
|
||||||
|
|
||||||
async def start_session(self, session_id):
|
async def start_session(self, session_id):
|
||||||
pass
|
pass
|
||||||
@@ -323,22 +413,21 @@ class TTSProviderBase(ABC):
|
|||||||
else:
|
else:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _process_audio_file(self, tts_file):
|
def _process_audio_file_stream(
|
||||||
|
self, tts_file, callback: Callable[[Any], Any]
|
||||||
|
) -> None:
|
||||||
"""处理音频文件并转换为指定格式
|
"""处理音频文件并转换为指定格式
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
tts_file: 音频文件路径
|
tts_file: 音频文件路径
|
||||||
content_detail: 内容详情
|
callback: 文件处理函数
|
||||||
|
|
||||||
Returns:
|
|
||||||
tuple: (sentence_type, audio_datas, content_detail)
|
|
||||||
"""
|
"""
|
||||||
if tts_file.endswith(".p3"):
|
if tts_file.endswith(".p3"):
|
||||||
audio_datas, _ = p3.decode_opus_from_file(tts_file)
|
p3.decode_opus_from_file_stream(tts_file, callback=callback)
|
||||||
elif self.conn.audio_format == "pcm":
|
elif self.conn.audio_format == "pcm":
|
||||||
audio_datas, _ = self.audio_to_pcm_data(tts_file)
|
self.audio_to_pcm_data_stream(tts_file, callback=callback)
|
||||||
else:
|
else:
|
||||||
audio_datas, _ = self.audio_to_opus_data(tts_file)
|
self.audio_to_opus_data_stream(tts_file, callback=callback)
|
||||||
|
|
||||||
if (
|
if (
|
||||||
self.delete_audio_file
|
self.delete_audio_file
|
||||||
@@ -347,7 +436,6 @@ class TTSProviderBase(ABC):
|
|||||||
and tts_file.startswith(self.output_file)
|
and tts_file.startswith(self.output_file)
|
||||||
):
|
):
|
||||||
os.remove(tts_file)
|
os.remove(tts_file)
|
||||||
return audio_datas
|
|
||||||
|
|
||||||
def _process_before_stop_play_files(self):
|
def _process_before_stop_play_files(self):
|
||||||
for audio_datas, text in self.before_stop_play_files:
|
for audio_datas, text in self.before_stop_play_files:
|
||||||
@@ -355,7 +443,9 @@ class TTSProviderBase(ABC):
|
|||||||
self.before_stop_play_files.clear()
|
self.before_stop_play_files.clear()
|
||||||
self.tts_audio_queue.put((SentenceType.LAST, [], None))
|
self.tts_audio_queue.put((SentenceType.LAST, [], None))
|
||||||
|
|
||||||
def _process_remaining_text(self):
|
def _process_remaining_text_stream(
|
||||||
|
self, opus_handler: Callable[[bytes], None] = None
|
||||||
|
):
|
||||||
"""处理剩余的文本并生成语音
|
"""处理剩余的文本并生成语音
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
@@ -366,18 +456,7 @@ class TTSProviderBase(ABC):
|
|||||||
if remaining_text:
|
if remaining_text:
|
||||||
segment_text = textUtils.get_string_no_punctuation_or_emoji(remaining_text)
|
segment_text = textUtils.get_string_no_punctuation_or_emoji(remaining_text)
|
||||||
if segment_text:
|
if segment_text:
|
||||||
if self.delete_audio_file:
|
self.to_tts_stream(segment_text, opus_handler=opus_handler)
|
||||||
audio_datas = self.to_tts(segment_text)
|
|
||||||
if audio_datas:
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, audio_datas, segment_text)
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
tts_file = self.to_tts(segment_text)
|
|
||||||
audio_datas = self._process_audio_file(tts_file)
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, audio_datas, segment_text)
|
|
||||||
)
|
|
||||||
self.processed_chars += len(full_text)
|
self.processed_chars += len(full_text)
|
||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ import json
|
|||||||
import queue
|
import queue
|
||||||
import asyncio
|
import asyncio
|
||||||
import traceback
|
import traceback
|
||||||
|
from typing import Callable, Any
|
||||||
import websockets
|
import websockets
|
||||||
from core.utils.tts import MarkdownCleaner
|
from core.utils.tts import MarkdownCleaner
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
@@ -266,11 +267,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
)
|
)
|
||||||
if message.content_file and os.path.exists(message.content_file):
|
if message.content_file and os.path.exists(message.content_file):
|
||||||
# 先处理文件音频数据
|
# 先处理文件音频数据
|
||||||
file_audio = self._process_audio_file(message.content_file)
|
self._process_audio_file_stream(message.content_file, callback=lambda audio_data: self.handle_audio_file(audio_data, message.content_detail))
|
||||||
self.before_stop_play_files.append(
|
|
||||||
(file_audio, message.content_detail)
|
|
||||||
)
|
|
||||||
|
|
||||||
if message.sentence_type == SentenceType.LAST:
|
if message.sentence_type == SentenceType.LAST:
|
||||||
try:
|
try:
|
||||||
logger.bind(tag=TAG).info("开始结束TTS会话...")
|
logger.bind(tag=TAG).info("开始结束TTS会话...")
|
||||||
@@ -428,9 +425,6 @@ class TTSProvider(TTSProviderBase):
|
|||||||
|
|
||||||
async def _start_monitor_tts_response(self):
|
async def _start_monitor_tts_response(self):
|
||||||
"""监听TTS响应"""
|
"""监听TTS响应"""
|
||||||
opus_datas_cache = []
|
|
||||||
is_first_sentence = True
|
|
||||||
first_sentence_segment_count = 0 # 添加计数器
|
|
||||||
try:
|
try:
|
||||||
session_finished = False # 标记会话是否正常结束
|
session_finished = False # 标记会话是否正常结束
|
||||||
while not self.conn.stop_event.is_set():
|
while not self.conn.stop_event.is_set():
|
||||||
@@ -451,30 +445,14 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.tts_audio_queue.put(
|
self.tts_audio_queue.put(
|
||||||
(SentenceType.FIRST, [], self.tts_text)
|
(SentenceType.FIRST, [], self.tts_text)
|
||||||
)
|
)
|
||||||
opus_datas_cache = []
|
|
||||||
first_sentence_segment_count = 0 # 重置计数器
|
|
||||||
elif (
|
elif (
|
||||||
res.optional.event == EVENT_TTSResponse
|
res.optional.event == EVENT_TTSResponse
|
||||||
and res.header.message_type == AUDIO_ONLY_RESPONSE
|
and res.header.message_type == AUDIO_ONLY_RESPONSE
|
||||||
):
|
):
|
||||||
logger.bind(tag=TAG).debug(f"推送数据到队列里面~~")
|
logger.bind(tag=TAG).debug(f"推送数据到队列里面~~")
|
||||||
opus_datas = self.wav_to_opus_data_audio_raw(res.payload)
|
self.wav_to_opus_data_audio_raw_stream(res.payload, callback=self.handle_opus)
|
||||||
logger.bind(tag=TAG).debug(
|
|
||||||
f"推送数据到队列里面帧数~~{len(opus_datas)}"
|
|
||||||
)
|
|
||||||
# 优化:对于第一句话,不进行缓存,立即推送到队列
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus_datas, None)
|
|
||||||
)
|
|
||||||
elif res.optional.event == EVENT_TTSSentenceEnd:
|
elif res.optional.event == EVENT_TTSSentenceEnd:
|
||||||
logger.bind(tag=TAG).info(f"句子语音生成成功:{self.tts_text}")
|
logger.bind(tag=TAG).info(f"句子语音生成成功:{self.tts_text}")
|
||||||
# 优化:如果有缓存的数据,立即发送
|
|
||||||
if opus_datas_cache:
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus_datas_cache, None)
|
|
||||||
)
|
|
||||||
# 第一句话结束后,将标志设置为False
|
|
||||||
is_first_sentence = False
|
|
||||||
elif res.optional.event == EVENT_SessionFinished:
|
elif res.optional.event == EVENT_SessionFinished:
|
||||||
logger.bind(tag=TAG).debug(f"会话结束~~")
|
logger.bind(tag=TAG).debug(f"会话结束~~")
|
||||||
self._process_before_stop_play_files()
|
self._process_before_stop_play_files()
|
||||||
@@ -648,16 +626,13 @@ class TTSProvider(TTSProviderBase):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
def wav_to_opus_data_audio_raw(self, raw_data_var, is_end=False):
|
def wav_to_opus_data_audio_raw_stream(self, raw_data_var, is_end=False, callback: Callable[[Any], Any]=None):
|
||||||
opus_datas = self.opus_encoder.encode_pcm_to_opus(raw_data_var, is_end)
|
return self.opus_encoder.encode_pcm_to_opus_stream(raw_data_var, is_end, callback=callback)
|
||||||
return opus_datas
|
|
||||||
|
|
||||||
def to_tts(self, text: str) -> list:
|
def to_tts(self, text: str) -> list:
|
||||||
"""非流式生成音频数据,用于生成音频及测试场景
|
"""非流式生成音频数据,用于生成音频及测试场景
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
text: 要转换的文本
|
text: 要转换的文本
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
list: 音频数据列表
|
list: 音频数据列表
|
||||||
"""
|
"""
|
||||||
@@ -734,8 +709,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
res.optional.event == EVENT_TTSResponse
|
res.optional.event == EVENT_TTSResponse
|
||||||
and res.header.message_type == AUDIO_ONLY_RESPONSE
|
and res.header.message_type == AUDIO_ONLY_RESPONSE
|
||||||
):
|
):
|
||||||
opus_datas = self.wav_to_opus_data_audio_raw(res.payload)
|
self.wav_to_opus_data_audio_raw_stream(res.payload, callback=lambda opus_frame: audio_data.append(opus_frame))
|
||||||
audio_data.extend(opus_datas)
|
|
||||||
elif res.optional.event == EVENT_SessionFinished:
|
elif res.optional.event == EVENT_SessionFinished:
|
||||||
break
|
break
|
||||||
|
|
||||||
|
|||||||
@@ -1,10 +1,10 @@
|
|||||||
import os
|
import os
|
||||||
import queue
|
|
||||||
import asyncio
|
|
||||||
import traceback
|
|
||||||
import aiohttp
|
|
||||||
import requests
|
|
||||||
import time
|
import time
|
||||||
|
import queue
|
||||||
|
import aiohttp
|
||||||
|
import asyncio
|
||||||
|
import requests
|
||||||
|
import traceback
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from core.utils.tts import MarkdownCleaner
|
from core.utils.tts import MarkdownCleaner
|
||||||
from core.providers.tts.base import TTSProviderBase
|
from core.providers.tts.base import TTSProviderBase
|
||||||
@@ -27,15 +27,13 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.api_url = config.get("api_url", "http://8.138.114.124:11996/tts")
|
self.api_url = config.get("api_url", "http://8.138.114.124:11996/tts")
|
||||||
self.audio_format = "pcm"
|
self.audio_format = "pcm"
|
||||||
self.before_stop_play_files = []
|
self.before_stop_play_files = []
|
||||||
self.segment_count = 0
|
|
||||||
|
|
||||||
# 创建Opus编码器 需注意接口返回的采样率为24000
|
# 创建Opus编码器 需注意接口返回的采样率为24000
|
||||||
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
|
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
|
||||||
sample_rate=24000, channels=1, frame_size_ms=60
|
sample_rate=24000, channels=1, frame_size_ms=60
|
||||||
)
|
)
|
||||||
|
|
||||||
# 文本缓冲区和PCM缓冲区
|
# PCM缓冲区
|
||||||
self.text_buffer = ""
|
|
||||||
self.pcm_buffer = bytearray()
|
self.pcm_buffer = bytearray()
|
||||||
|
|
||||||
def tts_text_priority_thread(self):
|
def tts_text_priority_thread(self):
|
||||||
@@ -48,7 +46,6 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.tts_stop_request = False
|
self.tts_stop_request = False
|
||||||
self.processed_chars = 0
|
self.processed_chars = 0
|
||||||
self.tts_text_buff = []
|
self.tts_text_buff = []
|
||||||
self.segment_count = 0
|
|
||||||
self.before_stop_play_files.clear()
|
self.before_stop_play_files.clear()
|
||||||
elif ContentType.TEXT == message.content_type:
|
elif ContentType.TEXT == message.content_type:
|
||||||
self.tts_text_buff.append(message.content_detail)
|
self.tts_text_buff.append(message.content_detail)
|
||||||
@@ -62,14 +59,11 @@ class TTSProvider(TTSProviderBase):
|
|||||||
)
|
)
|
||||||
if message.content_file and os.path.exists(message.content_file):
|
if message.content_file and os.path.exists(message.content_file):
|
||||||
# 先处理文件音频数据
|
# 先处理文件音频数据
|
||||||
file_audio = self._process_audio_file(message.content_file)
|
self._process_audio_file_stream(message.content_file, callback=lambda audio_data: self.handle_audio_file(audio_data, message.content_detail))
|
||||||
self.before_stop_play_files.append(
|
|
||||||
(file_audio, message.content_detail)
|
|
||||||
)
|
|
||||||
|
|
||||||
if message.sentence_type == SentenceType.LAST:
|
if message.sentence_type == SentenceType.LAST:
|
||||||
# 处理剩余的文本
|
# 处理剩余的文本
|
||||||
self._process_remaining_text(True)
|
self._process_remaining_text_stream(True)
|
||||||
|
|
||||||
except queue.Empty:
|
except queue.Empty:
|
||||||
continue
|
continue
|
||||||
@@ -78,7 +72,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
f"处理TTS文本失败: {str(e)}, 类型: {type(e).__name__}, 堆栈: {traceback.format_exc()}"
|
f"处理TTS文本失败: {str(e)}, 类型: {type(e).__name__}, 堆栈: {traceback.format_exc()}"
|
||||||
)
|
)
|
||||||
|
|
||||||
def _process_remaining_text(self, is_last=False):
|
def _process_remaining_text_stream(self, is_last=False):
|
||||||
"""处理剩余的文本并生成语音
|
"""处理剩余的文本并生成语音
|
||||||
Returns:
|
Returns:
|
||||||
bool: 是否成功处理了文本
|
bool: 是否成功处理了文本
|
||||||
@@ -143,8 +137,6 @@ class TTSProvider(TTSProviderBase):
|
|||||||
return
|
return
|
||||||
|
|
||||||
self.pcm_buffer.clear()
|
self.pcm_buffer.clear()
|
||||||
opus_datas_cache = []
|
|
||||||
|
|
||||||
self.tts_audio_queue.put((SentenceType.FIRST, [], text))
|
self.tts_audio_queue.put((SentenceType.FIRST, [], text))
|
||||||
|
|
||||||
# 处理音频流数据
|
# 处理音频流数据
|
||||||
@@ -158,41 +150,22 @@ class TTSProvider(TTSProviderBase):
|
|||||||
while len(self.pcm_buffer) >= frame_bytes:
|
while len(self.pcm_buffer) >= frame_bytes:
|
||||||
frame = bytes(self.pcm_buffer[:frame_bytes])
|
frame = bytes(self.pcm_buffer[:frame_bytes])
|
||||||
del self.pcm_buffer[:frame_bytes]
|
del self.pcm_buffer[:frame_bytes]
|
||||||
opus = self.opus_encoder.encode_pcm_to_opus(
|
|
||||||
frame, end_of_stream=False
|
self.opus_encoder.encode_pcm_to_opus_stream(
|
||||||
|
frame,
|
||||||
|
end_of_stream=False,
|
||||||
|
callback=self.handle_opus
|
||||||
)
|
)
|
||||||
if opus:
|
|
||||||
if self.segment_count < 10: # 前10个片段直接发送
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus, None)
|
|
||||||
)
|
|
||||||
self.segment_count += 1
|
|
||||||
else:
|
|
||||||
opus_datas_cache.extend(opus)
|
|
||||||
|
|
||||||
# flush 剩余不足一帧的数据
|
# flush 剩余不足一帧的数据
|
||||||
if self.pcm_buffer:
|
if self.pcm_buffer:
|
||||||
opus = self.opus_encoder.encode_pcm_to_opus(
|
self.opus_encoder.encode_pcm_to_opus_stream(
|
||||||
bytes(self.pcm_buffer), end_of_stream=True
|
bytes(self.pcm_buffer),
|
||||||
|
end_of_stream=True,
|
||||||
|
callback=self.handle_opus
|
||||||
)
|
)
|
||||||
if opus:
|
|
||||||
if self.segment_count < 10: # 前10个片段直接发送
|
|
||||||
# 直接发送
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus, None)
|
|
||||||
)
|
|
||||||
self.segment_count += 1
|
|
||||||
else:
|
|
||||||
# 后续片段缓存
|
|
||||||
opus_datas_cache.extend(opus)
|
|
||||||
self.pcm_buffer.clear()
|
self.pcm_buffer.clear()
|
||||||
|
|
||||||
# 如果不是前10个片段,发送缓存的数据
|
|
||||||
if self.segment_count >= 10 and opus_datas_cache:
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus_datas_cache, None)
|
|
||||||
)
|
|
||||||
|
|
||||||
# 如果是最后一段,输出音频获取完毕
|
# 如果是最后一段,输出音频获取完毕
|
||||||
if is_last:
|
if is_last:
|
||||||
self._process_before_stop_play_files()
|
self._process_before_stop_play_files()
|
||||||
@@ -209,10 +182,8 @@ class TTSProvider(TTSProviderBase):
|
|||||||
|
|
||||||
def to_tts(self, text: str) -> list:
|
def to_tts(self, text: str) -> list:
|
||||||
"""非流式TTS处理,用于测试及保存音频文件的场景
|
"""非流式TTS处理,用于测试及保存音频文件的场景
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
text: 要转换的文本
|
text: 要转换的文本
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
list: 返回opus编码后的音频数据列表
|
list: 返回opus编码后的音频数据列表
|
||||||
"""
|
"""
|
||||||
@@ -251,14 +222,14 @@ class TTSProvider(TTSProviderBase):
|
|||||||
# 最后一帧可能不足,用0填充
|
# 最后一帧可能不足,用0填充
|
||||||
frame = frame + b"\x00" * (frame_bytes - len(frame))
|
frame = frame + b"\x00" * (frame_bytes - len(frame))
|
||||||
|
|
||||||
opus = self.opus_encoder.encode_pcm_to_opus(
|
self.opus_encoder.encode_pcm_to_opus_stream(
|
||||||
frame, end_of_stream=(i + frame_bytes >= len(pcm_data))
|
frame,
|
||||||
|
end_of_stream=(i + frame_bytes >= len(pcm_data)),
|
||||||
|
callback=lambda opus: opus_datas.append(opus)
|
||||||
)
|
)
|
||||||
if opus:
|
|
||||||
opus_datas.extend(opus)
|
|
||||||
|
|
||||||
return opus_datas
|
return opus_datas
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"TTS请求异常: {e}")
|
logger.bind(tag=TAG).error(f"TTS请求异常: {e}")
|
||||||
return []
|
return []
|
||||||
@@ -1,10 +1,10 @@
|
|||||||
import os
|
import os
|
||||||
import queue
|
|
||||||
import asyncio
|
|
||||||
import traceback
|
|
||||||
import aiohttp
|
|
||||||
import requests
|
|
||||||
import time
|
import time
|
||||||
|
import queue
|
||||||
|
import aiohttp
|
||||||
|
import asyncio
|
||||||
|
import requests
|
||||||
|
import traceback
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from core.utils.tts import MarkdownCleaner
|
from core.utils.tts import MarkdownCleaner
|
||||||
from core.providers.tts.base import TTSProviderBase
|
from core.providers.tts.base import TTSProviderBase
|
||||||
@@ -24,23 +24,15 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.api_url = config.get("api_url")
|
self.api_url = config.get("api_url")
|
||||||
self.audio_format = "pcm"
|
self.audio_format = "pcm"
|
||||||
self.before_stop_play_files = []
|
self.before_stop_play_files = []
|
||||||
self.segment_count = 0 # 添加片段计数器
|
|
||||||
|
|
||||||
# 创建Opus编码器
|
# 创建Opus编码器
|
||||||
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
|
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
|
||||||
sample_rate=16000, channels=1, frame_size_ms=60
|
sample_rate=16000, channels=1, frame_size_ms=60
|
||||||
)
|
)
|
||||||
|
|
||||||
# 添加文本缓冲区
|
|
||||||
self.text_buffer = ""
|
|
||||||
|
|
||||||
# PCM缓冲区
|
# PCM缓冲区
|
||||||
self.pcm_buffer = bytearray()
|
self.pcm_buffer = bytearray()
|
||||||
|
|
||||||
###################################################################################
|
|
||||||
# linkerai单流式TTS重写父类的方法--开始
|
|
||||||
###################################################################################
|
|
||||||
|
|
||||||
def tts_text_priority_thread(self):
|
def tts_text_priority_thread(self):
|
||||||
"""流式文本处理线程"""
|
"""流式文本处理线程"""
|
||||||
while not self.conn.stop_event.is_set():
|
while not self.conn.stop_event.is_set():
|
||||||
@@ -51,7 +43,6 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.tts_stop_request = False
|
self.tts_stop_request = False
|
||||||
self.processed_chars = 0
|
self.processed_chars = 0
|
||||||
self.tts_text_buff = []
|
self.tts_text_buff = []
|
||||||
self.segment_count = 0
|
|
||||||
self.before_stop_play_files.clear()
|
self.before_stop_play_files.clear()
|
||||||
elif ContentType.TEXT == message.content_type:
|
elif ContentType.TEXT == message.content_type:
|
||||||
self.tts_text_buff.append(message.content_detail)
|
self.tts_text_buff.append(message.content_detail)
|
||||||
@@ -65,14 +56,10 @@ class TTSProvider(TTSProviderBase):
|
|||||||
)
|
)
|
||||||
if message.content_file and os.path.exists(message.content_file):
|
if message.content_file and os.path.exists(message.content_file):
|
||||||
# 先处理文件音频数据
|
# 先处理文件音频数据
|
||||||
file_audio = self._process_audio_file(message.content_file)
|
self._process_audio_file_stream(message.content_file, callback=lambda audio_data: self.handle_audio_file(audio_data, message.content_detail))
|
||||||
self.before_stop_play_files.append(
|
|
||||||
(file_audio, message.content_detail)
|
|
||||||
)
|
|
||||||
|
|
||||||
if message.sentence_type == SentenceType.LAST:
|
if message.sentence_type == SentenceType.LAST:
|
||||||
# 处理剩余的文本
|
# 处理剩余的文本
|
||||||
self._process_remaining_text(True)
|
self._process_remaining_text_stream(True)
|
||||||
|
|
||||||
except queue.Empty:
|
except queue.Empty:
|
||||||
continue
|
continue
|
||||||
@@ -81,7 +68,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
f"处理TTS文本失败: {str(e)}, 类型: {type(e).__name__}, 堆栈: {traceback.format_exc()}"
|
f"处理TTS文本失败: {str(e)}, 类型: {type(e).__name__}, 堆栈: {traceback.format_exc()}"
|
||||||
)
|
)
|
||||||
|
|
||||||
def _process_remaining_text(self, is_last=False):
|
def _process_remaining_text_stream(self, is_last=False):
|
||||||
"""处理剩余的文本并生成语音
|
"""处理剩余的文本并生成语音
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
@@ -124,10 +111,6 @@ class TTSProvider(TTSProviderBase):
|
|||||||
finally:
|
finally:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
###################################################################################
|
|
||||||
# linkerai单流式TTS重写父类的方法--结束
|
|
||||||
###################################################################################
|
|
||||||
|
|
||||||
async def text_to_speak(self, text, is_last):
|
async def text_to_speak(self, text, is_last):
|
||||||
"""流式处理TTS音频,每句只推送一次音频列表"""
|
"""流式处理TTS音频,每句只推送一次音频列表"""
|
||||||
await self._tts_request(text, is_last)
|
await self._tts_request(text, is_last)
|
||||||
@@ -176,8 +159,6 @@ class TTSProvider(TTSProviderBase):
|
|||||||
return
|
return
|
||||||
|
|
||||||
self.pcm_buffer.clear()
|
self.pcm_buffer.clear()
|
||||||
opus_datas_cache = []
|
|
||||||
|
|
||||||
self.tts_audio_queue.put((SentenceType.FIRST, [], text))
|
self.tts_audio_queue.put((SentenceType.FIRST, [], text))
|
||||||
|
|
||||||
# 兼容 iter_chunked / iter_chunks / iter_any
|
# 兼容 iter_chunked / iter_chunks / iter_any
|
||||||
@@ -194,41 +175,21 @@ class TTSProvider(TTSProviderBase):
|
|||||||
frame = bytes(self.pcm_buffer[:frame_bytes])
|
frame = bytes(self.pcm_buffer[:frame_bytes])
|
||||||
del self.pcm_buffer[:frame_bytes]
|
del self.pcm_buffer[:frame_bytes]
|
||||||
|
|
||||||
opus = self.opus_encoder.encode_pcm_to_opus(
|
self.opus_encoder.encode_pcm_to_opus_stream(
|
||||||
frame, end_of_stream=False
|
frame,
|
||||||
|
end_of_stream=False,
|
||||||
|
callback=self.handle_opus
|
||||||
)
|
)
|
||||||
if opus:
|
|
||||||
if self.segment_count < 10: # 前10个片段直接发送
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus, None)
|
|
||||||
)
|
|
||||||
self.segment_count += 1
|
|
||||||
else:
|
|
||||||
opus_datas_cache.extend(opus)
|
|
||||||
|
|
||||||
# flush 剩余不足一帧的数据
|
# flush 剩余不足一帧的数据
|
||||||
if self.pcm_buffer:
|
if self.pcm_buffer:
|
||||||
opus = self.opus_encoder.encode_pcm_to_opus(
|
self.opus_encoder.encode_pcm_to_opus_stream(
|
||||||
bytes(self.pcm_buffer), end_of_stream=True
|
bytes(self.pcm_buffer),
|
||||||
|
end_of_stream=True,
|
||||||
|
callback=self.handle_opus
|
||||||
)
|
)
|
||||||
if opus:
|
|
||||||
if self.segment_count < 10: # 前10个片段直接发送
|
|
||||||
# 直接发送
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus, None)
|
|
||||||
)
|
|
||||||
self.segment_count += 1
|
|
||||||
else:
|
|
||||||
# 后续片段缓存
|
|
||||||
opus_datas_cache.extend(opus)
|
|
||||||
self.pcm_buffer.clear()
|
self.pcm_buffer.clear()
|
||||||
|
|
||||||
# 如果不是前10个片段,发送缓存的数据
|
|
||||||
if self.segment_count >= 10 and opus_datas_cache:
|
|
||||||
self.tts_audio_queue.put(
|
|
||||||
(SentenceType.MIDDLE, opus_datas_cache, None)
|
|
||||||
)
|
|
||||||
|
|
||||||
# 如果是最后一段,输出音频获取完毕
|
# 如果是最后一段,输出音频获取完毕
|
||||||
if is_last:
|
if is_last:
|
||||||
self._process_before_stop_play_files()
|
self._process_before_stop_play_files()
|
||||||
@@ -239,10 +200,8 @@ class TTSProvider(TTSProviderBase):
|
|||||||
|
|
||||||
def to_tts(self, text: str) -> list:
|
def to_tts(self, text: str) -> list:
|
||||||
"""非流式TTS处理,用于测试及保存音频文件的场景
|
"""非流式TTS处理,用于测试及保存音频文件的场景
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
text: 要转换的文本
|
text: 要转换的文本
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
list: 返回opus编码后的音频数据列表
|
list: 返回opus编码后的音频数据列表
|
||||||
"""
|
"""
|
||||||
@@ -295,14 +254,14 @@ class TTSProvider(TTSProviderBase):
|
|||||||
# 最后一帧可能不足,用0填充
|
# 最后一帧可能不足,用0填充
|
||||||
frame = frame + b"\x00" * (frame_bytes - len(frame))
|
frame = frame + b"\x00" * (frame_bytes - len(frame))
|
||||||
|
|
||||||
opus = self.opus_encoder.encode_pcm_to_opus(
|
self.opus_encoder.encode_pcm_to_opus_stream(
|
||||||
frame, end_of_stream=(i + frame_bytes >= len(pcm_data))
|
frame,
|
||||||
|
end_of_stream=(i + frame_bytes >= len(pcm_data)),
|
||||||
|
callback=lambda opus: opus_datas.append(opus)
|
||||||
)
|
)
|
||||||
if opus:
|
|
||||||
opus_datas.extend(opus)
|
|
||||||
|
|
||||||
return opus_datas
|
return opus_datas
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"TTS请求异常: {e}")
|
logger.bind(tag=TAG).error(f"TTS请求异常: {e}")
|
||||||
return []
|
return []
|
||||||
@@ -1,13 +1,15 @@
|
|||||||
import asyncio
|
|
||||||
import json
|
|
||||||
import base64
|
|
||||||
import aiohttp
|
|
||||||
import numpy as np
|
|
||||||
import io
|
import io
|
||||||
import wave
|
import wave
|
||||||
|
import json
|
||||||
|
import base64
|
||||||
|
import asyncio
|
||||||
import websockets
|
import websockets
|
||||||
from core.providers.tts.base import TTSProviderBase
|
import numpy as np
|
||||||
|
from datetime import datetime
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
|
from core.providers.tts.base import TTSProviderBase
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
@@ -18,11 +20,12 @@ class TTSProvider(TTSProviderBase):
|
|||||||
super().__init__(config, delete_audio_file)
|
super().__init__(config, delete_audio_file)
|
||||||
self.url = config.get("url", "ws://192.168.1.10:8092/paddlespeech/tts/streaming")
|
self.url = config.get("url", "ws://192.168.1.10:8092/paddlespeech/tts/streaming")
|
||||||
self.protocol = config.get("protocol", "websocket")
|
self.protocol = config.get("protocol", "websocket")
|
||||||
|
|
||||||
if config.get("private_voice"):
|
if config.get("private_voice"):
|
||||||
self.spk_id = int(config.get("private_voice"))
|
self.spk_id = int(config.get("private_voice"))
|
||||||
else:
|
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)
|
sample_rate = config.get("sample_rate", 24000)
|
||||||
self.sample_rate = float(sample_rate) if sample_rate else 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)
|
volume = config.get("volume", 1.0)
|
||||||
self.volume = float(volume) if volume else 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,
|
async def pcm_to_wav(self, pcm_data: bytes, sample_rate: int = 24000, num_channels: int = 1,
|
||||||
bits_per_sample: int = 16) -> bytes:
|
bits_per_sample: int = 16) -> bytes:
|
||||||
@@ -58,43 +75,9 @@ class TTSProvider(TTSProviderBase):
|
|||||||
async def text_to_speak(self, text, output_file):
|
async def text_to_speak(self, text, output_file):
|
||||||
if self.protocol == "websocket":
|
if self.protocol == "websocket":
|
||||||
return await self.text_streaming(text, output_file)
|
return await self.text_streaming(text, output_file)
|
||||||
elif self.protocol == "http":
|
|
||||||
return await self.text(text, output_file)
|
|
||||||
else:
|
else:
|
||||||
raise ValueError("Unsupported protocol. Please use 'websocket' or 'http'.")
|
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):
|
async def text_streaming(self, text, output_file):
|
||||||
try:
|
try:
|
||||||
# 使用 websockets 异步连接到 WebSocket 服务器
|
# 使用 websockets 异步连接到 WebSocket 服务器
|
||||||
@@ -151,6 +134,12 @@ class TTSProvider(TTSProviderBase):
|
|||||||
# 接收结束响应避免服务抛出异常
|
# 接收结束响应避免服务抛出异常
|
||||||
await ws.recv()
|
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:
|
if output_file:
|
||||||
with open(output_file, "wb") as file_to_save:
|
with open(output_file, "wb") as file_to_save:
|
||||||
@@ -159,4 +148,4 @@ class TTSProvider(TTSProviderBase):
|
|||||||
return wav_data
|
return wav_data
|
||||||
|
|
||||||
except Exception as e:
|
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}")
|
||||||
@@ -34,7 +34,7 @@ class VADProvider(VADProviderBase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# 至少要多少帧才算有语音
|
# 至少要多少帧才算有语音
|
||||||
self.frame_window_threshold = 1
|
self.frame_window_threshold = 3
|
||||||
|
|
||||||
def is_vad(self, conn, opus_packet):
|
def is_vad(self, conn, opus_packet):
|
||||||
try:
|
try:
|
||||||
@@ -70,9 +70,7 @@ class VADProvider(VADProviderBase):
|
|||||||
|
|
||||||
# 更新滑动窗口
|
# 更新滑动窗口
|
||||||
conn.client_voice_window.append(is_voice)
|
conn.client_voice_window.append(is_voice)
|
||||||
client_have_voice = (
|
client_have_voice = (conn.client_voice_window.count(True) >= self.frame_window_threshold)
|
||||||
conn.client_voice_window.count(True) >= self.frame_window_threshold
|
|
||||||
)
|
|
||||||
|
|
||||||
# 如果之前有声音,但本次没有声音,且与上次有声音的时间差已经超过了静默阈值,则认为已经说完一句话
|
# 如果之前有声音,但本次没有声音,且与上次有声音的时间差已经超过了静默阈值,则认为已经说完一句话
|
||||||
if conn.client_have_voice and not client_have_voice:
|
if conn.client_have_voice and not client_have_voice:
|
||||||
|
|||||||
@@ -5,12 +5,10 @@ Opus编码工具类
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
import traceback
|
import traceback
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from typing import List, Optional
|
|
||||||
from opuslib_next import Encoder
|
from opuslib_next import Encoder
|
||||||
from opuslib_next import constants
|
from opuslib_next import constants
|
||||||
|
from typing import Optional, Callable, Any
|
||||||
|
|
||||||
class OpusEncoderUtils:
|
class OpusEncoderUtils:
|
||||||
"""PCM到Opus的编码器"""
|
"""PCM到Opus的编码器"""
|
||||||
@@ -56,13 +54,14 @@ class OpusEncoderUtils:
|
|||||||
self.encoder.reset_state()
|
self.encoder.reset_state()
|
||||||
self.buffer = np.array([], dtype=np.int16)
|
self.buffer = np.array([], dtype=np.int16)
|
||||||
|
|
||||||
def encode_pcm_to_opus(self, pcm_data: bytes, end_of_stream: bool) -> List[bytes]:
|
def encode_pcm_to_opus_stream(self, pcm_data: bytes, end_of_stream: bool, callback: Callable[[Any], Any]):
|
||||||
"""
|
"""
|
||||||
将PCM数据编码为Opus格式
|
将PCM数据编码为Opus格式,以流式方式进行处理
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
pcm_data: PCM字节数据
|
pcm_data: PCM字节数据
|
||||||
end_of_stream: 是否为流的结束
|
end_of_stream: 是否为流的结束,
|
||||||
|
callback: opus处理方法
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Opus数据包列表
|
Opus数据包列表
|
||||||
@@ -76,7 +75,6 @@ class OpusEncoderUtils:
|
|||||||
# 将新数据追加到缓冲区
|
# 将新数据追加到缓冲区
|
||||||
self.buffer = np.append(self.buffer, new_samples)
|
self.buffer = np.append(self.buffer, new_samples)
|
||||||
|
|
||||||
opus_packets = []
|
|
||||||
offset = 0
|
offset = 0
|
||||||
|
|
||||||
# 处理所有完整帧
|
# 处理所有完整帧
|
||||||
@@ -84,7 +82,7 @@ class OpusEncoderUtils:
|
|||||||
frame = self.buffer[offset : offset + self.total_frame_size]
|
frame = self.buffer[offset : offset + self.total_frame_size]
|
||||||
output = self._encode(frame)
|
output = self._encode(frame)
|
||||||
if output:
|
if output:
|
||||||
opus_packets.append(output)
|
callback(output)
|
||||||
offset += self.total_frame_size
|
offset += self.total_frame_size
|
||||||
|
|
||||||
# 保留未处理的样本
|
# 保留未处理的样本
|
||||||
@@ -98,11 +96,9 @@ class OpusEncoderUtils:
|
|||||||
|
|
||||||
output = self._encode(last_frame)
|
output = self._encode(last_frame)
|
||||||
if output:
|
if output:
|
||||||
opus_packets.append(output)
|
callback(output)
|
||||||
self.buffer = np.array([], dtype=np.int16)
|
self.buffer = np.array([], dtype=np.int16)
|
||||||
|
|
||||||
return opus_packets
|
|
||||||
|
|
||||||
def _encode(self, frame: np.ndarray) -> Optional[bytes]:
|
def _encode(self, frame: np.ndarray) -> Optional[bytes]:
|
||||||
"""编码一帧音频数据"""
|
"""编码一帧音频数据"""
|
||||||
try:
|
try:
|
||||||
@@ -133,4 +129,4 @@ class OpusEncoderUtils:
|
|||||||
def close(self):
|
def close(self):
|
||||||
"""关闭编码器并释放资源"""
|
"""关闭编码器并释放资源"""
|
||||||
# opuslib没有明确的关闭方法,Python的垃圾回收会处理
|
# opuslib没有明确的关闭方法,Python的垃圾回收会处理
|
||||||
pass
|
pass
|
||||||
@@ -1,16 +1,17 @@
|
|||||||
import json
|
|
||||||
import socket
|
|
||||||
import subprocess
|
|
||||||
import re
|
import re
|
||||||
import os
|
import os
|
||||||
|
import json
|
||||||
|
import copy
|
||||||
import wave
|
import wave
|
||||||
|
import socket
|
||||||
|
import requests
|
||||||
|
import subprocess
|
||||||
|
import numpy as np
|
||||||
|
import opuslib_next
|
||||||
from io import BytesIO
|
from io import BytesIO
|
||||||
from core.utils import p3
|
from core.utils import p3
|
||||||
import numpy as np
|
|
||||||
import requests
|
|
||||||
import opuslib_next
|
|
||||||
from pydub import AudioSegment
|
from pydub import AudioSegment
|
||||||
import copy
|
from typing import Callable, Any
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
emoji_map = {
|
emoji_map = {
|
||||||
@@ -211,7 +212,7 @@ def extract_json_from_string(input_string):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
def audio_to_data(audio_file_path, is_opus=True):
|
def audio_to_data_stream(audio_file_path, is_opus=True, callback: Callable[[Any], Any]=None) -> None:
|
||||||
# 获取文件后缀名
|
# 获取文件后缀名
|
||||||
file_type = os.path.splitext(audio_file_path)[1]
|
file_type = os.path.splitext(audio_file_path)[1]
|
||||||
if file_type:
|
if file_type:
|
||||||
@@ -224,33 +225,32 @@ def audio_to_data(audio_file_path, is_opus=True):
|
|||||||
# 转换为单声道/16kHz采样率/16位小端编码(确保与编码器匹配)
|
# 转换为单声道/16kHz采样率/16位小端编码(确保与编码器匹配)
|
||||||
audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2)
|
audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2)
|
||||||
|
|
||||||
# 音频时长(秒)
|
# 获取原始PCM数据(16位小端)
|
||||||
duration = len(audio) / 1000.0
|
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位小端)
|
# 获取原始PCM数据(16位小端)
|
||||||
raw_data = audio.raw_data
|
raw_data = audio.raw_data
|
||||||
return pcm_to_data(raw_data, is_opus), duration
|
|
||||||
|
|
||||||
|
|
||||||
def audio_bytes_to_data(audio_bytes, file_type, is_opus=True):
|
|
||||||
"""
|
|
||||||
直接用音频二进制数据转为opus/pcm数据,支持wav、mp3、p3
|
|
||||||
"""
|
|
||||||
if file_type == "p3":
|
|
||||||
# 直接用p3解码
|
|
||||||
return p3.decode_opus_from_bytes(audio_bytes)
|
|
||||||
else:
|
|
||||||
# 其他格式用pydub
|
|
||||||
audio = AudioSegment.from_file(
|
|
||||||
BytesIO(audio_bytes), format=file_type, parameters=["-nostdin"]
|
|
||||||
)
|
|
||||||
audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2)
|
|
||||||
duration = len(audio) / 1000.0
|
|
||||||
raw_data = audio.raw_data
|
|
||||||
return pcm_to_data(raw_data, is_opus), duration
|
|
||||||
|
|
||||||
|
|
||||||
def pcm_to_data(raw_data, is_opus=True):
|
|
||||||
# 初始化Opus编码器
|
# 初始化Opus编码器
|
||||||
encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO)
|
encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO)
|
||||||
|
|
||||||
@@ -280,6 +280,49 @@ def pcm_to_data(raw_data, is_opus=True):
|
|||||||
|
|
||||||
return datas
|
return datas
|
||||||
|
|
||||||
|
def audio_bytes_to_data_stream(audio_bytes, file_type, is_opus, callback: Callable[[Any], Any]) -> None:
|
||||||
|
"""
|
||||||
|
直接用音频二进制数据转为opus/pcm数据,支持wav、mp3、p3
|
||||||
|
"""
|
||||||
|
if file_type == "p3":
|
||||||
|
# 直接用p3解码
|
||||||
|
return p3.decode_opus_from_bytes_stream(audio_bytes, callback)
|
||||||
|
else:
|
||||||
|
# 其他格式用pydub
|
||||||
|
audio = AudioSegment.from_file(
|
||||||
|
BytesIO(audio_bytes), format=file_type, parameters=["-nostdin"]
|
||||||
|
)
|
||||||
|
audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2)
|
||||||
|
raw_data = audio.raw_data
|
||||||
|
pcm_to_data_stream(raw_data, is_opus, callback)
|
||||||
|
|
||||||
|
|
||||||
|
def pcm_to_data_stream(raw_data, is_opus=True, callback: Callable[[Any], Any] = None):
|
||||||
|
# 初始化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
|
||||||
|
|
||||||
|
# 按帧处理所有音频数据(包括最后一帧可能补零)
|
||||||
|
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)
|
||||||
|
callback(frame_data)
|
||||||
|
else:
|
||||||
|
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):
|
def opus_datas_to_wav_bytes(opus_datas, sample_rate=16000, channels=1):
|
||||||
"""
|
"""
|
||||||
@@ -307,7 +350,6 @@ def opus_datas_to_wav_bytes(opus_datas, sample_rate=16000, channels=1):
|
|||||||
wf.writeframes(pcm_bytes)
|
wf.writeframes(pcm_bytes)
|
||||||
return wav_buffer.getvalue()
|
return wav_buffer.getvalue()
|
||||||
|
|
||||||
|
|
||||||
def check_vad_update(before_config, new_config):
|
def check_vad_update(before_config, new_config):
|
||||||
if (
|
if (
|
||||||
new_config.get("selected_module") is None
|
new_config.get("selected_module") is None
|
||||||
|
|||||||
@@ -137,4 +137,4 @@ class WakeupWordsConfig:
|
|||||||
return file_path
|
return file_path
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
print(f"生成音频文件路径失败: {e}")
|
print(f"生成音频文件路径失败: {e}")
|
||||||
raise
|
raise
|
||||||
@@ -36,6 +36,12 @@
|
|||||||
"command": "npx",
|
"command": "npx",
|
||||||
"args": ["-y", "@simonb97/server-win-cli"],
|
"args": ["-y", "@simonb97/server-win-cli"],
|
||||||
"link": "https://github.com/SimonB97/win-cli-mcp-server"
|
"link": "https://github.com/SimonB97/win-cli-mcp-server"
|
||||||
|
},
|
||||||
|
"sse-mcp-server": {
|
||||||
|
"url": "http://localhost:8080/sse",
|
||||||
|
"headers": {
|
||||||
|
"Authorization": "Bearer YOUR TOKEN"
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user