Merge branch 'main' into py_MinmaxStreamTTS_test

This commit is contained in:
Sakura-RanChen
2025-09-05 14:53:02 +08:00
64 changed files with 2931 additions and 883 deletions
+12 -1
View File
@@ -31,6 +31,9 @@ jobs:
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@v3
with:
driver-opts: |
network=host
- name: Login to GitHub Container Registry
uses: docker/login-action@v3
@@ -60,6 +63,10 @@ jobs:
tags: |
${{ env.IS_VERSION == 'true' && format('ghcr.io/{0}:server_{1},ghcr.io/{0}:server_latest', github.repository, env.VERSION) || format('ghcr.io/{0}:server_latest', github.repository) }}
platforms: linux/amd64,linux/arm64
cache-from: type=gha
cache-to: type=gha,mode=max
build-args: |
BUILDKIT_PROGRESS=plain
# 构建 manager-api 镜像
- name: Build and push manager-web
@@ -70,4 +77,8 @@ jobs:
push: true
tags: |
${{ env.IS_VERSION == 'true' && format('ghcr.io/{0}:web_{1},ghcr.io/{0}:web_latest', github.repository, env.VERSION) || format('ghcr.io/{0}:web_latest', github.repository) }}
platforms: linux/amd64,linux/arm64
platforms: linux/amd64,linux/arm64
cache-from: type=gha
cache-to: type=gha,mode=max
build-args: |
BUILDKIT_PROGRESS=plain
+9 -2
View File
@@ -3,10 +3,17 @@ FROM python:3.10-slim AS builder
WORKDIR /app
# 配置pip使用国内镜像源(阿里云)并设置超时和重试
RUN pip config set global.index-url https://mirrors.aliyun.com/pypi/simple/ && \
pip config set global.trusted-host mirrors.aliyun.com && \
pip config set global.timeout 120 && \
pip config set install.retries 5
COPY main/xiaozhi-server/requirements.txt .
# 安装Python依赖
RUN pip install --no-cache-dir -r requirements.txt
# 安装Python依赖,使用并行下载
RUN pip install --no-cache-dir --upgrade pip setuptools wheel && \
pip install --no-cache-dir -r requirements.txt --default-timeout=120 --retries 5
# 第二阶段:生产镜像
FROM python:3.10-slim
+3 -3
View File
@@ -18,12 +18,12 @@ FROM bellsoft/liberica-runtime-container:jre-21-glibc
# 安装Nginx和字体库
RUN apk update && \
apk add --no-cache nginx bash && \
apk add --no-cache fontconfig ttf-dejavu msttcorefonts-installer && \
apk add --no-cache nginx bash fontconfig ttf-dejavu && \
apk add --no-cache --repository=http://dl-cdn.alpinelinux.org/alpine/edge/testing/ msttcorefonts-installer || true && \
rm -rf /var/cache/apk/*
# 更新字体缓存
RUN printf 'YES\n' | update-ms-fonts && fc-cache -f -v
RUN (printf 'YES\n' | update-ms-fonts || true) && fc-cache -f -v
# 配置Nginx
COPY docs/docker/nginx.conf /etc/nginx/nginx.conf
+13 -6
View File
@@ -38,9 +38,9 @@
</p>
<p align="center">
By Professor Siyuan Liu Research and Development Group ( South China University of Technology)
Spearheaded by Professor Siyuan Liu's Team (South China University of Technology)
</br>
刘思源教授团队研发(华南理工大学)
刘思源教授团队主导研发(华南理工大学)
</br>
<img src="./docs/images/hnlg.jpg" alt="华南理工大学" width="50%">
</p>
@@ -94,7 +94,7 @@ By Professor Siyuan Liu Research and Development Group ( South China University
</tr>
<tr>
<td>
<a href="https://www.bilibili.com/video/BV1Vy96YCE3R" target="_blank">
<a href="https://www.bilibili.com/video/BV1vchQzaEse" target="_blank">
<picture>
<img alt="自定义音色" src="docs/images/demo6.png" />
</picture>
@@ -197,6 +197,7 @@ By Professor Siyuan Liu Research and Development Group ( South China University
```
智控台地址: https://2662r3426b.vicp.fun
智控台(h5版): https://2662r3426b.vicp.fun/h5/index.html
服务测试工具: https://2662r3426b.vicp.fun/test/
OTA接口地址: https://2662r3426b.vicp.fun/xiaozhi/ota/
@@ -282,6 +283,8 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/
| dify 接口调用 | Dify | - |
| fastgpt 接口调用 | Fastgpt | - |
| coze 接口调用 | Coze | - |
| xinference 接口调用 | Xinference | - |
| homeassistant 接口调用 | HomeAssistant | - |
实际上,任何支持 openai 接口调用的 LLM 均可接入使用。
@@ -301,8 +304,8 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/
| 使用方式 | 支持平台 | 免费平台 |
|:---:|:---:|:---:|
| 接口调用 | EdgeTTS、火山引擎豆包TTS、腾讯云、阿里云TTS、CosyVoiceSiliconflow、TTS302AI、CozeCnTTS、GizwitsTTS、ACGNTTS、OpenAITTS、灵犀流式TTS | 灵犀流式TTS、EdgeTTS、CosyVoiceSiliconflow(部分) |
| 本地服务 | FishSpeech、GPT_SOVITS_V2、GPT_SOVITS_V3、MinimaxTTS | FishSpeech、GPT_SOVITS_V2、GPT_SOVITS_V3、MinimaxTTS |
| 接口调用 | EdgeTTS、火山引擎豆包TTS、腾讯云、阿里云TTS、阿里云流式TTS、CosyVoiceSiliconflow、TTS302AI、CozeCnTTS、GizwitsTTS、ACGNTTS、OpenAITTS、灵犀流式TTS、MinimaxTTS、火山双流式TTS | 灵犀流式TTS、EdgeTTS、CosyVoiceSiliconflow(部分) |
| 本地服务 | FishSpeech、GPT_SOVITS_V2、GPT_SOVITS_V3、Index-TTS、PaddleSpeech | Index-TTS、PaddleSpeech、FishSpeech、GPT_SOVITS_V2、GPT_SOVITS_V3 |
---
@@ -319,7 +322,7 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/
| 使用方式 | 支持平台 | 免费平台 |
|:---:|:---:|:---:|
| 本地使用 | FunASR、SherpaASR | FunASR、SherpaASR |
| 接口调用 | DoubaoASR、FunASRServer、TencentASR、AliyunASR | FunASRServer |
| 接口调用 | DoubaoASR、Doubao流式ASR、FunASRServer、TencentASR、AliyunASR、Aliyun流式ASR、百度ASR、OpenAI ASR | FunASRServer |
---
@@ -337,6 +340,7 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/
|:------:|:---------------:|:----:|:---------:|:--:|
| Memory | mem0ai | 接口调用 | 1000次/月额度 | |
| Memory | mem_local_short | 本地总结 | 免费 | |
| Memory | nomem | 无记忆模式 | 免费 | |
---
@@ -346,6 +350,7 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/
|:------:|:-------------:|:----:|:-------:|:---------------------:|
| Intent | intent_llm | 接口调用 | 根据LLM收费 | 通过大模型识别意图,通用性强 |
| Intent | function_call | 接口调用 | 根据LLM收费 | 通过大模型函数调用完成意图,速度快,效果好 |
| Intent | nointent | 无意图模式 | 免费 | 不进行意图识别,直接返回对话结果 |
---
@@ -356,8 +361,10 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/
| <img src="./docs/images/logo_bailing.png" width="160"> | [百聆语音对话机器人](https://github.com/wwbin2017/bailing) | 本项目受[百聆语音对话机器人](https://github.com/wwbin2017/bailing)启发,并在其基础上实现 |
| <img src="./docs/images/logo_tenclass.png" width="160"> | [十方融海](https://www.tenclass.com/) | 感谢[十方融海](https://www.tenclass.com/)为小智生态制定了标准的通讯协议、多设备兼容性方案及高并发场景实践示范;为本项目提供了全链路技术文档支持 |
| <img src="./docs/images/logo_xuanfeng.png" width="160"> | [玄凤科技](https://github.com/Eric0308) | 感谢[玄凤科技](https://github.com/Eric0308)贡献函数调用框架、MCP通信协议及插件化调用机制的实现代码,通过标准化的指令调度体系与动态扩展能力,显著提升了前端设备(IoT)的交互效率和功能延展性 |
| <img src="./docs/images/logo_junsen.png" width="160"> | [huangjunsen](https://github.com/huangjunsen0406) | 感谢[huangjunsen](https://github.com/huangjunsen0406) 贡献`智控台移动端`模块,实现了跨平台移动设备的高效控制与实时交互,大幅提升了系统在移动场景下的操作便捷性和管理效率 |
| <img src="./docs/images/logo_huiyuan.png" width="160"> | [汇远设计](http://ui.kwd988.net/) | 感谢[汇远设计](http://ui.kwd988.net/)为本项目提供专业视觉解决方案,用其服务超千家企业的设计实战经验,赋能本项目产品用户体验 |
| <img src="./docs/images/logo_qinren.png" width="160"> | [西安勤人信息科技](https://www.029app.com/) | 感谢[西安勤人信息科技](https://www.029app.com/)深化本项目视觉体系,确保整体设计风格在多场景应用中的一致性和扩展性 |
| <img src="./docs/images/logo_contributors.png" width="160"> | [代码贡献者](https://github.com/xinnan-tech/xiaozhi-esp32-server/graphs/contributors) | 感谢[所有代码贡献者](https://github.com/xinnan-tech/xiaozhi-esp32-server/graphs/contributors)贡献者,你们的付出让项目更加健壮和强大。 |
<a href="https://star-history.com/#xinnan-tech/xiaozhi-esp32-server&Date">
+6 -3
View File
@@ -38,9 +38,9 @@ Supports MCP endpoints and voiceprint recognition
</p>
<p align="center">
By Professor Siyuan Liu Research and Development Group (South China University of Technology)
Spearheaded by Professor Siyuan Liu's Team (South China University of Technology)
</br>
刘思源教授团队研发(华南理工大学)
刘思源教授团队主导研发(华南理工大学)
</br>
<img src="./docs/images/hnlg.jpg" alt="South China University of Technology" width="50%">
</p>
@@ -93,7 +93,7 @@ Want to see the usage effects? Click the videos below 🎥
</tr>
<tr>
<td>
<a href="https://www.bilibili.com/video/BV1Vy96YCE3R" target="_blank">
<a href="https://www.bilibili.com/video/BV1vchQzaEse" target="_blank">
<picture>
<img alt="Custom voice timbre" src="docs/images/demo6.png" />
</picture>
@@ -194,6 +194,7 @@ This project provides two deployment methods. Please choose based on your specif
```
Intelligent Control Console Address: https://2662r3426b.vicp.fun
Intelligent Control Console Address (H5): https://2662r3426b.vicp.fun/h5/index.html
Service Test Tool: https://2662r3426b.vicp.fun/test/
OTA Interface Address: https://2662r3426b.vicp.fun/xiaozhi/ota/
@@ -352,8 +353,10 @@ In fact, any VLLM that supports OpenAI interface calls can be integrated and use
| <img src="./docs/images/logo_bailing.png" width="160"> | [Bailing Voice Dialogue Robot](https://github.com/wwbin2017/bailing) | This project is inspired by [Bailing Voice Dialogue Robot](https://github.com/wwbin2017/bailing) and implemented on its basis |
| <img src="./docs/images/logo_tenclass.png" width="160"> | [Tenclass](https://www.tenclass.com/) | Thanks to [Tenclass](https://www.tenclass.com/) for formulating standard communication protocols, multi-device compatibility solutions, and high-concurrency scenario practice demonstrations for the Xiaozhi ecosystem; providing full-link technical documentation support for this project |
| <img src="./docs/images/logo_xuanfeng.png" width="160"> | [Xuanfeng Technology](https://github.com/Eric0308) | Thanks to [Xuanfeng Technology](https://github.com/Eric0308) for contributing function calling framework, MCP communication protocol, and plugin-based calling mechanism implementation code. Through standardized instruction scheduling system and dynamic expansion capabilities, it significantly improves the interaction efficiency and functional extensibility of frontend devices (IoT) |
| <img src="./docs/images/logo_junsen.png" width="160"> | [huangjunsen](https://github.com/huangjunsen0406) | Thanks to [huangjunsen](https://github.com/huangjunsen0406) for contributing the `Smart Control Console Mobile` module, which enables efficient control and real-time interaction across mobile devices, significantly enhancing the system's operational convenience and management efficiency in mobile scenarios. |
| <img src="./docs/images/logo_huiyuan.png" width="160"> | [Huiyuan Design](http://ui.kwd988.net/) | Thanks to [Huiyuan Design](http://ui.kwd988.net/) for providing professional visual solutions for this project, using their design practical experience serving over a thousand enterprises to empower this project's product user experience |
| <img src="./docs/images/logo_qinren.png" width="160"> | [Xi'an Qinren Information Technology](https://www.029app.com/) | Thanks to [Xi'an Qinren Information Technology](https://www.029app.com/) for deepening this project's visual system, ensuring consistency and extensibility of overall design style in multi-scenario applications |
| <img src="./docs/images/logo_contributors.png" width="160"> | [Code Contributors](https://github.com/xinnan-tech/xiaozhi-esp32-server/graphs/contributors) | Thanks to [all code contributors](https://github.com/xinnan-tech/xiaozhi-esp32-server/graphs/contributors), your efforts have made the project more robust and powerful. |
<a href="https://star-history.com/#xinnan-tech/xiaozhi-esp32-server&Date">
+6 -7
View File
@@ -46,7 +46,7 @@ cat << "EOF"
\/ \__,_||_| |_||_||_||_| \__,_| |_| \_| \__,_||_| |_||_| \__,_| \__,_|
EOF
echo -e "\e[0m" # 重置颜色
echo -e "\e[1;36m 小智服务端全量部署一键安装脚本 Ver 0.2 \e[0m\n"
echo -e "\e[1;36m 小智服务端全量部署一键安装脚本 Ver 0.2 2025年8月20日更新 \e[0m\n"
sleep 1
@@ -376,7 +376,7 @@ done
echo "服务端启动成功!正在完成配置..."
echo "正在启动服务..."
docker compose -f docker-compose_all.yml up -d
docker compose -f /opt/xiaozhi-server/docker-compose_all.yml up -d
echo "服务启动完成!"
)
@@ -402,13 +402,12 @@ fi
# 获取并显示地址信息
LOCAL_IP=$(hostname -I | awk '{print $1}')
WEBSOCKET_ADDR=$(docker logs xiaozhi-esp32-server 2>&1 | tac | grep -m 1 -E -o "ws://[^ ]+")
VISION_ADDR=$(docker logs xiaozhi-esp32-server 2>&1 | tac | grep -m 1 "视觉" | grep -m 1 -E -o "http://[^ ]+")
# 修复日志文件获取不到ws的问题,改为硬编码
whiptail --title "安装完成!" --msgbox "\
服务端相关地址如下:\n\
管理后台访问地址: http://$LOCAL_IP:8002\n\
OTA 地址: http://$LOCAL_IP:8002/xiaozhi/ota/\n\
视觉分析接口地址: $VISION_ADDR\n\
WebSocket 地址: $WEBSOCKET_ADDR\n\
\n安装完毕!感谢您的使用!\n按Enter键退出..." 16 70
视觉分析接口地址: http://$LOCAL_IP:8003/mcp/vision/explain\n\
WebSocket 地址: ws://$LOCAL_IP:8000/xiaozhi/v1/\n\
\n安装完毕!感谢您的使用!\n按Enter键退出..." 16 70
+7 -4
View File
@@ -7,11 +7,13 @@ docker镜像已支持x86架构、arm64架构的CPU,支持在国产操作系统
如果您的电脑还没安装docker,可以按照这里的教程安装:[docker安装](https://www.runoob.com/docker/ubuntu-docker-install.html)
docker 安装全模块有两种方式,你可以[1.1使用懒人脚本](#1.1 懒人脚本)(作者[@VanillaNahida](https://github.com/VanillaNahida)自动帮你下载所需的文件和配置文件,你可以使用[1.2手动部署](#1.2 手动部署)从零搭建。
docker 安装全模块有两种方式,你可以[使用懒人脚本](./Deployment_all.md#11-懒人脚本)(作者[@VanillaNahida](https://github.com/VanillaNahida)
脚本会自动帮你下载所需的文件和配置文件,你也可以使用[手动部署](./Deployment_all.md#12-手动部署)从零搭建。
### 1.1 懒人脚本
你可以使用以下命令一键安装全模块版小智服务端:
部署简便,可以参考[视频教程](https://www.bilibili.com/video/BV17bbvzHExd/) ,文字版教程如下:
> [!NOTE]
> 暂且只支持Ubuntu服务器一键部署,其他系统未尝试,可能会有一些奇怪的bug
@@ -474,7 +476,8 @@ ws://你电脑局域网的ip:8000/xiaozhi/v1/
4、[如何部署MCP接入点](./mcp-endpoint-enable.md)<br/>
5、[如何接入MCP接入点](./mcp-endpoint-integration.md)<br/>
6、[如何开启声纹识别](./voiceprint-integration.md)<br/>
10、[新闻插件源配置指南](./newsnow_plugin_config.md)<br/>
7、[新闻插件源配置指南](./newsnow_plugin_config.md)<br/>
8、[天气插件使用指南](./weather-integration.md)<br/>
## 语音克隆、本地语音部署相关教程
1、[如何部署集成index-tts本地语音](./index-stream-integration.md)<br/>
2、[如何部署集成fish-speech本地语音](./fish-speech-integration.md)<br/>
Binary file not shown.

Before

Width:  |  Height:  |  Size: 306 KiB

After

Width:  |  Height:  |  Size: 111 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 25 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 34 KiB

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