mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-23 23:53:55 +08:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
423626a6bf | ||
|
|
182acc0787 | ||
|
|
705aef462c | ||
|
|
1f092dd80e | ||
|
|
4d0ddd7ff3 | ||
|
|
2a212ae759 | ||
|
|
1d57e75d28 | ||
|
|
ac1e254eed | ||
|
|
6074431364 | ||
|
|
d5fc4d48b4 | ||
|
|
61232e7dda | ||
|
|
594b4f1d75 | ||
|
|
98174bcc16 | ||
|
|
29d90df9ea | ||
|
|
50e6e4817b | ||
|
|
c6a88e41e1 | ||
|
|
ae1c41ba82 | ||
|
|
837bb74576 | ||
|
|
1284418c18 | ||
|
|
3eb3a4502f | ||
|
|
6023af2f21 | ||
|
|
2cd90d8066 | ||
|
|
c1dfb540a3 | ||
|
|
bc44ea8757 | ||
|
|
77743ef8ba | ||
|
|
f8de052d54 | ||
|
|
cb4eb21551 | ||
|
|
8357d7abc3 | ||
|
|
5a01bfcdbf | ||
|
|
01eb416a03 | ||
|
|
dca02c1f4b | ||
|
|
2654802bb0 | ||
|
|
d7eecfdcea | ||
|
|
8df9846ad7 | ||
|
|
e43c0135a7 | ||
|
|
adf1a47945 | ||
|
|
ca1beb956a | ||
|
|
c694d3d6f1 | ||
|
|
a314d03870 | ||
|
|
f2fd3a0b7e | ||
|
|
7b693527d9 | ||
|
|
93e0d57783 | ||
|
|
18027d5d5c | ||
|
|
0ef514a68f | ||
|
|
4780f5e972 | ||
|
|
5b8a567f26 | ||
|
|
23cb7616d9 | ||
|
|
488522de6e | ||
|
|
03bcf5d2c7 | ||
|
|
fba758ea36 | ||
|
|
0268f90e9c | ||
|
|
4cc7247c37 | ||
|
|
5d94ed853d | ||
|
|
d06e297c2d | ||
|
|
6304467d3a | ||
|
|
e52601a584 | ||
|
|
74226581b7 | ||
|
|
bc28504b3e | ||
|
|
109811199d | ||
|
|
610fa4d101 | ||
|
|
0269620e66 | ||
|
|
10ebe975cf | ||
|
|
0bc609cac4 | ||
|
|
3657f6ce75 | ||
|
|
e62f72810c | ||
|
|
21553410f9 | ||
|
|
a6bc910b0f | ||
|
|
3e68667323 | ||
|
|
c0d4bbcecf | ||
|
|
6bf6159e6c | ||
|
|
fd4193daab | ||
|
|
da2077f3dd | ||
|
|
ee65032f7d | ||
|
|
2d7d75c290 | ||
|
|
237e88ba89 | ||
|
|
48223e4f39 | ||
|
|
ce776f210c | ||
|
|
fb5e25ec16 | ||
|
|
c4f2411fee | ||
|
|
3c3be950e9 | ||
|
|
b24ee0b9a6 | ||
|
|
e1d5caa0fd | ||
|
|
35ba9b0e6b | ||
|
|
66de07823d | ||
|
|
fd96e71d04 | ||
|
|
92affd6e13 | ||
|
|
3130044909 |
@@ -121,6 +121,33 @@
|
|||||||
</a>
|
</a>
|
||||||
</td>
|
</td>
|
||||||
</tr>
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td>
|
||||||
|
<a href="https://www.bilibili.com/video/BV12J7WzBEaH" target="_blank">
|
||||||
|
<picture>
|
||||||
|
<img alt="实时打断" src="docs/images/demo10.png" />
|
||||||
|
</picture>
|
||||||
|
</a>
|
||||||
|
</td>
|
||||||
|
<td>
|
||||||
|
<a href="https://www.bilibili.com/video/BV1Co76z7EvK" target="_blank">
|
||||||
|
<picture>
|
||||||
|
<img alt="拍照识物品" src="docs/images/demo12.png" />
|
||||||
|
</picture>
|
||||||
|
</a>
|
||||||
|
</td>
|
||||||
|
<td>
|
||||||
|
<a href="https://www.bilibili.com/video/BV1TJ7WzzEo6" target="_blank">
|
||||||
|
<picture>
|
||||||
|
<img alt="多指令任务" src="docs/images/demo11.png" />
|
||||||
|
</picture>
|
||||||
|
</a>
|
||||||
|
</td>
|
||||||
|
<td>
|
||||||
|
</td>
|
||||||
|
<td>
|
||||||
|
</td>
|
||||||
|
</tr>
|
||||||
</table>
|
</table>
|
||||||
|
|
||||||
---
|
---
|
||||||
@@ -143,8 +170,8 @@
|
|||||||
#### 🚀 部署方式选择
|
#### 🚀 部署方式选择
|
||||||
| 部署方式 | 特点 | 适用场景 | 部署文档 | 配置要求 | 视频教程 |
|
| 部署方式 | 特点 | 适用场景 | 部署文档 | 配置要求 | 视频教程 |
|
||||||
|---------|------|---------|---------|---------|---------|
|
|---------|------|---------|---------|---------|---------|
|
||||||
| **最简化安装** | 智能对话、IOT功能,数据存储在配置文件 | 低配置环境,无需数据库 | [Docker版](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E5%8F%AA%E8%BF%90%E8%A1%8Cserver) / [源码部署](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E5%8F%AA%E8%BF%90%E8%A1%8Cserver)| 如果使用`FunASR`要2核4G,如果全API,要2核2G | - |
|
| **最简化安装** | 智能对话、IOT、MCP、视觉感知,数据存储在配置文件 | 低配置环境,无需数据库 | [①Docker版](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E5%8F%AA%E8%BF%90%E8%A1%8Cserver) / [②源码部署](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E5%8F%AA%E8%BF%90%E8%A1%8Cserver)| 如果使用`FunASR`要2核4G,如果全API,要2核2G | - |
|
||||||
| **全模块安装** | 智能对话、IOT、OTA、智控台,数据存储在数据库 | 完整功能体验 |[Docker版](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [源码部署](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) | 如果使用`FunASR`要4核8G,如果全API,要2核4G| [本地源码视频教程](https://www.bilibili.com/video/BV1wBJhz4Ewe) |
|
| **全模块安装** | 智能对话、IOT、MCP、视觉感知、OTA、智控台,数据存储在数据库 | 完整功能体验 |[①Docker版](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [②源码部署](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) / [③源码部署自动更新教程](./docs/dev-ops-integration.md) | 如果使用`FunASR`要4核8G,如果全API,要2核4G| [本地源码启动视频教程](https://www.bilibili.com/video/BV1wBJhz4Ewe) |
|
||||||
|
|
||||||
|
|
||||||
> 💡 提示:以下是按最新代码部署后的测试平台,有需要可烧录测试,并发为6个,每天会清空数据
|
> 💡 提示:以下是按最新代码部署后的测试平台,有需要可烧录测试,并发为6个,每天会清空数据
|
||||||
@@ -159,34 +186,47 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/
|
|||||||
|
|
||||||
#### 🚩 配置说明和推荐
|
#### 🚩 配置说明和推荐
|
||||||
> [!Note]
|
> [!Note]
|
||||||
> 本项目默认的配置是`入门全免费`设置,如果想效果更优,推荐使用`全流式配置`。
|
> 本项目默认的配置是`入门全免费`设置,如果想效果更优,推荐使用`流式配置`。
|
||||||
>
|
>
|
||||||
> 本项目自`0.5.2`版本,已支持整个生命周期全流式,相比`0.5`版本以前,响应速度提升约`2.5秒`
|
> 本项目自`0.5.2`版本,已支持使用流式配置,相比`0.5`版本以前,响应速度提升约`2.5秒`
|
||||||
|
|
||||||
|
| 模块名称 | 入门全免费设置 | 流式配置 |
|
||||||
| 模块名称 | 入门全免费设置 | 全流式配置 |
|
|:---:|:---:|:---:|
|
||||||
|---------|---------|------|
|
| ASR(语音识别) | FunASR(本地) | 👍DoubaoStreamASR(火山流式语音识别) |
|
||||||
| ASR(语音识别) | FunASR(本地) | ✅DoubaoASR(火山流式语音识别) |
|
| LLM(大模型) | ChatGLMLLM(智谱glm-4-flash) | 👍DoubaoLLM(火山doubao-1-5-pro-32k-250115) |
|
||||||
| LLM(大模型) | ChatGLMLLM(智谱glm-4-flash) | ✅DoubaoLLM(火山doubao-1-5-pro-32k-250115) |
|
| VLLM(视觉大模型) | ChatGLMVLLM(智谱glm-4v-flash) | 👍QwenVLVLLM(千问qwen2.5-vl-3b-instructh) |
|
||||||
| TTS(语音合成) | EdgeTTS(微软语音) | ✅HuoshanDoubleStreamTTS(火山双流式语音合成) |
|
| TTS(语音合成) | 👍LinkeraiTTS(灵犀流式) | 👍HuoshanDoubleStreamTTS(火山双流式语音合成) |
|
||||||
| Intent(意图识别) | function_call(函数调用) | ✅function_call(函数调用) |
|
| Intent(意图识别) | function_call(函数调用) | ✅function_call(函数调用) |
|
||||||
| Memory(记忆功能) | mem_local_short(本地短期记忆) | ✅mem_local_short(本地短期记忆) |
|
| Memory(记忆功能) | mem_local_short(本地短期记忆) | ✅mem_local_short(本地短期记忆) |
|
||||||
|
|
||||||
|
#### 🔧 测试工具
|
||||||
|
本项目提供以下测试工具,帮助您验证系统和选择合适的模型:
|
||||||
|
|
||||||
|
| 工具名称 | 位置 | 使用方法 | 功能说明 |
|
||||||
|
|:---:|:---|:---:|:---:|
|
||||||
|
| 音频交互测试工具 | main》xiaozhi-server》test》test_page.html | 使用谷歌浏览器直接打开 | 测试音频播放和接收功能,验证Python端音频处理是否正常 |
|
||||||
|
| 模型响应测试工具1 | main》xiaozhi-server》performance_tester.py | 执行 `python performance_tester.py` | 测试ASR(语音识别)、LLM(大模型)、TTS(语音合成)三个核心模块的响应速度 |
|
||||||
|
| 模型响应测试工具2 | main》xiaozhi-server》performance_tester_vllm.py | 执行 `python performance_tester_vllm.py` | 测试VLLM(视觉模型)的响应速度 |
|
||||||
|
|
||||||
|
> 💡 提示:测试模型速度时,只会测试配置了密钥的模型。
|
||||||
|
|
||||||
---
|
---
|
||||||
## 功能清单 ✨
|
## 功能清单 ✨
|
||||||
### 已实现 ✅
|
### 已实现 ✅
|
||||||
|
|
||||||
| 功能模块 | 描述 |
|
| 功能模块 | 描述 |
|
||||||
|---------|------|
|
|:---:|:---|
|
||||||
| 通信协议 | 基于 `xiaozhi-esp32` 协议,通过 WebSocket 实现数据交互 |
|
| 核心服务架构 | 基于WebSocket和HTTP服务器,提供完整的控制台管理和认证系统 |
|
||||||
| 对话交互 | 支持唤醒对话、手动对话及实时打断。长时间无对话时自动休眠 |
|
| 语音交互系统 | 支持流式ASR(语音识别)、流式TTS(语音合成)、VAD(语音活动检测),支持多语言识别和语音处理 |
|
||||||
| 意图识别 | 支持使用LLM意图识别、function call函数调用,减少硬编码意图判断 |
|
| 智能对话系统 | 支持多种LLM(大语言模型),实现智能对话 |
|
||||||
| 多语言识别 | 支持国语、粤语、英语、日语、韩语(默认使用 FunASR) |
|
| 视觉感知系统 | 支持多种VLLM(视觉大模型),实现多模态交互 |
|
||||||
| LLM 模块 | 支持灵活切换 LLM 模块,默认使用 ChatGLMLLM,也可选用阿里百炼、DeepSeek、Ollama 等接口 |
|
| 意图识别系统 | 支持LLM意图识别、Function Call函数调用,提供插件化意图处理机制 |
|
||||||
| TTS 模块 | 支持 EdgeTTS(默认)、火山引擎豆包 TTS 等多种 TTS 接口,满足语音合成需求 |
|
| 记忆系统 | 支持本地短期记忆、mem0ai接口记忆,具备记忆总结功能 |
|
||||||
| 记忆功能 | 支持超长记忆、本地总结记忆、无记忆三种模式,满足不同场景需求 |
|
| IOT/MCP控制协议 | 支持设备注册管理、智能控制接口,同时支持IOT、MCP控制协议 |
|
||||||
| IOT功能 | 支持管理注册设备IOT功能,支持基于对话上下文语境下的智能物联网控制 |
|
| 管理后台 | 提供Web管理界面,支持用户管理、系统配置和设备管理 |
|
||||||
| 智控台 | 提供Web管理界面,支持智能体管理、用户管理、系统配置等功能,方便管理员和用户进行管理 |
|
| 测试工具 | 提供性能测试工具、视觉模型测试工具和音频交互测试工具 |
|
||||||
|
| 部署支持 | 支持Docker部署和本地部署,提供完整的配置文件管理 |
|
||||||
|
| 插件系统 | 支持功能插件扩展、自定义插件开发和插件热加载 |
|
||||||
|
|
||||||
### 正在开发 🚧
|
### 正在开发 🚧
|
||||||
|
|
||||||
@@ -223,11 +263,21 @@ Websocket接口地址: wss://2662r3426b.vicp.fun/xiaozhi/v1/
|
|||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
### VLLM 视觉模型
|
||||||
|
|
||||||
|
| 使用方式 | 支持平台 | 免费平台 |
|
||||||
|
|:---:|:---:|:---:|
|
||||||
|
| openai 接口调用 | 阿里百炼、智谱ChatGLMVLLM | 智谱ChatGLMVLLM |
|
||||||
|
|
||||||
|
实际上,任何支持 openai 接口调用的 VLLM 均可接入使用。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
### TTS 语音合成
|
### TTS 语音合成
|
||||||
|
|
||||||
| 使用方式 | 支持平台 | 免费平台 |
|
| 使用方式 | 支持平台 | 免费平台 |
|
||||||
|:---:|:---:|:---:|
|
|:---:|:---:|:---:|
|
||||||
| 接口调用 | EdgeTTS、火山引擎豆包TTS、腾讯云、阿里云TTS、CosyVoiceSiliconflow、TTS302AI、CozeCnTTS、GizwitsTTS、ACGNTTS、OpenAITTS | EdgeTTS、CosyVoiceSiliconflow(部分) |
|
| 接口调用 | 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 |
|
| 本地服务 | FishSpeech、GPT_SOVITS_V2、GPT_SOVITS_V3、MinimaxTTS | FishSpeech、GPT_SOVITS_V2、GPT_SOVITS_V3、MinimaxTTS |
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|||||||
+92
-54
@@ -6,14 +6,14 @@
|
|||||||
This project provides backend services for the open-source smart hardware project
|
This project provides backend services for the open-source smart hardware project
|
||||||
<a href="https://github.com/78/xiaozhi-esp32">xiaozhi-esp32</a><br/>
|
<a href="https://github.com/78/xiaozhi-esp32">xiaozhi-esp32</a><br/>
|
||||||
Implemented using Python, Java, and Vue according to the <a href="https://ccnphfhqs21z.feishu.cn/wiki/M0XiwldO9iJwHikpXD5cEx71nKh">Xiaozhi Communication Protocol</a><br/>
|
Implemented using Python, Java, and Vue according to the <a href="https://ccnphfhqs21z.feishu.cn/wiki/M0XiwldO9iJwHikpXD5cEx71nKh">Xiaozhi Communication Protocol</a><br/>
|
||||||
Helping you quickly set up your Xiaozhi server
|
Helps you quickly set up your Xiaozhi server
|
||||||
</p>
|
</p>
|
||||||
|
|
||||||
<p align="center">
|
<p align="center">
|
||||||
<a href="./README.md">中文</a>
|
<a href="./README.md">中文</a>
|
||||||
· <a href="./docs/FAQ.md">FAQ</a>
|
· <a href="./docs/FAQ.md">FAQ</a>
|
||||||
· <a href="https://github.com/xinnan-tech/xiaozhi-esp32-server/issues">Report Issues</a>
|
· <a href="https://github.com/xinnan-tech/xiaozhi-esp32-server/issues">Report Issues</a>
|
||||||
· <a href="./README_ed.md#deployment-documentation">Deployment Guide</a>
|
· <a href="./README_en.md#deployment-documentation">Deployment Guide</a>
|
||||||
· <a href="https://github.com/xinnan-tech/xiaozhi-esp32-server/releases">Release Notes</a>
|
· <a href="https://github.com/xinnan-tech/xiaozhi-esp32-server/releases">Release Notes</a>
|
||||||
</p>
|
</p>
|
||||||
<p align="center">
|
<p align="center">
|
||||||
@@ -50,7 +50,7 @@ Want to see it in action? Check out these videos 🎥
|
|||||||
<td>
|
<td>
|
||||||
<a href="https://www.bilibili.com/video/BV1FMFyejExX" target="_blank">
|
<a href="https://www.bilibili.com/video/BV1FMFyejExX" target="_blank">
|
||||||
<picture>
|
<picture>
|
||||||
<img alt="Xiaozhi esp32 connecting to custom backend model" src="docs/images/demo1.png" />
|
<img alt="Xiaozhi esp32 connecting to own backend model" src="docs/images/demo1.png" />
|
||||||
</picture>
|
</picture>
|
||||||
</a>
|
</a>
|
||||||
</td>
|
</td>
|
||||||
@@ -64,14 +64,14 @@ Want to see it in action? Check out these videos 🎥
|
|||||||
<td>
|
<td>
|
||||||
<a href="https://www.bilibili.com/video/BV12yA2egEaC" target="_blank">
|
<a href="https://www.bilibili.com/video/BV12yA2egEaC" target="_blank">
|
||||||
<picture>
|
<picture>
|
||||||
<img alt="Cantonese communication" src="docs/images/demo3.png" />
|
<img alt="Using Cantonese" src="docs/images/demo3.png" />
|
||||||
</picture>
|
</picture>
|
||||||
</a>
|
</a>
|
||||||
</td>
|
</td>
|
||||||
<td>
|
<td>
|
||||||
<a href="https://www.bilibili.com/video/BV1pNXWYGEx1" target="_blank">
|
<a href="https://www.bilibili.com/video/BV1pNXWYGEx1" target="_blank">
|
||||||
<picture>
|
<picture>
|
||||||
<img alt="Home appliance control" src="docs/images/demo5.png" />
|
<img alt="Control home appliances" src="docs/images/demo5.png" />
|
||||||
</picture>
|
</picture>
|
||||||
</a>
|
</a>
|
||||||
</td>
|
</td>
|
||||||
@@ -94,7 +94,7 @@ Want to see it in action? Check out these videos 🎥
|
|||||||
<td>
|
<td>
|
||||||
<a href="https://www.bilibili.com/video/BV1VC96Y5EMH" target="_blank">
|
<a href="https://www.bilibili.com/video/BV1VC96Y5EMH" target="_blank">
|
||||||
<picture>
|
<picture>
|
||||||
<img alt="Music playback" src="docs/images/demo7.png" />
|
<img alt="Play music" src="docs/images/demo7.png" />
|
||||||
</picture>
|
</picture>
|
||||||
</a>
|
</a>
|
||||||
</td>
|
</td>
|
||||||
@@ -108,7 +108,7 @@ Want to see it in action? Check out these videos 🎥
|
|||||||
<td>
|
<td>
|
||||||
<a href="https://www.bilibili.com/video/BV178XuYfEpi" target="_blank">
|
<a href="https://www.bilibili.com/video/BV178XuYfEpi" target="_blank">
|
||||||
<picture>
|
<picture>
|
||||||
<img alt="IOT device control" src="docs/images/demo9.png" />
|
<img alt="IOT command control" src="docs/images/demo9.png" />
|
||||||
</picture>
|
</picture>
|
||||||
</a>
|
</a>
|
||||||
</td>
|
</td>
|
||||||
@@ -120,6 +120,33 @@ Want to see it in action? Check out these videos 🎥
|
|||||||
</a>
|
</a>
|
||||||
</td>
|
</td>
|
||||||
</tr>
|
</tr>
|
||||||
|
<tr>
|
||||||
|
<td>
|
||||||
|
<a href="https://www.bilibili.com/video/BV12J7WzBEaH" target="_blank">
|
||||||
|
<picture>
|
||||||
|
<img alt="Real-time interruption" src="docs/images/demo10.png" />
|
||||||
|
</picture>
|
||||||
|
</a>
|
||||||
|
</td>
|
||||||
|
<td>
|
||||||
|
<a href="https://www.bilibili.com/video/BV1Co76z7EvK" target="_blank">
|
||||||
|
<picture>
|
||||||
|
<img alt="Photo recognition" src="docs/images/demo12.png" />
|
||||||
|
</picture>
|
||||||
|
</a>
|
||||||
|
</td>
|
||||||
|
<td>
|
||||||
|
<a href="https://www.bilibili.com/video/BV1TJ7WzzEo6" target="_blank">
|
||||||
|
<picture>
|
||||||
|
<img alt="Multi-command tasks" src="docs/images/demo11.png" />
|
||||||
|
</picture>
|
||||||
|
</a>
|
||||||
|
</td>
|
||||||
|
<td>
|
||||||
|
</td>
|
||||||
|
<td>
|
||||||
|
</td>
|
||||||
|
</tr>
|
||||||
</table>
|
</table>
|
||||||
|
|
||||||
---
|
---
|
||||||
@@ -137,91 +164,102 @@ It is recommended that users prioritize service providers with relevant business
|
|||||||
|
|
||||||

|

|
||||||
|
|
||||||
This project offers two deployment methods. Please choose based on your specific needs:
|
This project provides two deployment methods. Please choose according to your specific needs:
|
||||||
|
|
||||||
#### 🚀 Deployment Method Selection
|
#### 🚀 Deployment Method Selection
|
||||||
|
| Deployment Method | Features | Suitable Scenarios | Deployment Guide | Requirements | Video Tutorial |
|
||||||
|
|---------|------|---------|---------|---------|---------|
|
||||||
|
| **Simplified Installation** | Smart dialogue, IOT functionality, data stored in configuration files | Low-configuration environment, no database needed | [Docker Version](./docs/Deployment.md#method-1-docker-server-only) / [Source Code Deployment](./docs/Deployment.md#method-2-local-source-code-server-only) | 2 cores 4G if using `FunASR`, 2 cores 2G if using all APIs | - |
|
||||||
|
| **Full Module Installation** | Smart dialogue, IOT, OTA, Control Panel, data stored in database | Complete functionality experience | [Docker Version](./docs/Deployment_all.md#method-1-docker-full-modules) / [Source Code Deployment](./docs/Deployment_all.md#method-2-local-source-code-full-modules) | 4 cores 8G if using `FunASR`, 2 cores 4G if using all APIs | [Local Source Code Startup Video Tutorial](https://www.bilibili.com/video/BV1wBJhz4Ewe) / [Local Source Code Auto-Update Tutorial](./docs/dev-ops-integration.md) |
|
||||||
|
|
||||||
| Deployment Method | Features | Use Case | Docker Deployment Guide | Source Code Deployment Guide |
|
> 💡 Note: Below are the test platforms deployed with the latest code. You can flash and test if needed. Concurrent users: 6, data will be cleared daily
|
||||||
|---------|------|---------|---------|---------|
|
|
||||||
| **Simplified Installation** | Smart dialogue, IOT functionality, data stored in configuration files | Low-configuration environment, no database required | [Docker Server Only](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E5%8F%AA%E8%BF%90%E8%A1%8Cserver) | [Local Source Code Server Only](./docs/Deployment.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E5%8F%AA%E8%BF%90%E8%A1%8Cserver)|
|
|
||||||
| **Full Module Installation** | Smart dialogue, IOT, OTA, Control Panel, data stored in database | Complete functionality experience |[Docker Full Module](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%B8%80docker%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) | [Local Source Code Full Module](./docs/Deployment_all.md#%E6%96%B9%E5%BC%8F%E4%BA%8C%E6%9C%AC%E5%9C%B0%E6%BA%90%E7%A0%81%E8%BF%90%E8%A1%8C%E5%85%A8%E6%A8%A1%E5%9D%97) |
|
|
||||||
|
|
||||||
> 💡 Note: Below is the test platform deployed with the latest code. You can flash and test if needed. Concurrent users: 6, data cleared daily
|
|
||||||
|
|
||||||
```
|
```
|
||||||
Control Panel: https://2662r3426b.vicp.fun
|
Control Panel Address: https://2662r3426b.vicp.fun
|
||||||
|
|
||||||
Service Test Tool: https://2662r3426b.vicp.fun/test/
|
Service Test Tool: https://2662r3426b.vicp.fun/test/
|
||||||
OTA Interface: https://2662r3426b.vicp.fun/xiaozhi/ota/
|
OTA Interface Address: https://2662r3426b.vicp.fun/xiaozhi/ota/
|
||||||
Websocket Interface: wss://2662r3426b.vicp.fun/xiaozhi/v1/
|
Websocket Interface Address: wss://2662r3426b.vicp.fun/xiaozhi/v1/
|
||||||
```
|
```
|
||||||
|
|
||||||
|
#### 🚩 Configuration Description and Recommendations
|
||||||
|
> [!Note]
|
||||||
|
> The default configuration of this project is `Entry Level Free` settings. For better results, we recommend using `Full Streaming Configuration`.
|
||||||
|
>
|
||||||
|
> Since version `0.5.2`, this project supports full streaming throughout the entire lifecycle. Compared to versions before `0.5`, response speed has improved by approximately `2.5 seconds`
|
||||||
|
|
||||||
|
| Module Name | Entry Level Free Settings | Full Streaming Configuration |
|
||||||
|
|---------|---------|------|
|
||||||
|
| ASR(Speech Recognition) | FunASR(Local) | ✅DoubaoASR(Volcano Streaming Speech Recognition) |
|
||||||
|
| LLM(Large Language Model) | ChatGLMLLM(Zhipu glm-4-flash) | ✅DoubaoLLM(Volcano doubao-1-5-pro-32k-250115) |
|
||||||
|
| VLLM(Vision Large Model) | ChatGLMVLLM(Zhipu glm-4v-flash) | ✅ChatGLMVLLM(Zhipu glm-4v-flash) |
|
||||||
|
| TTS(Speech Synthesis) | EdgeTTS(Microsoft Speech) | ✅HuoshanDoubleStreamTTS(Volcano Double Streaming Speech Synthesis) |
|
||||||
|
| Intent(Intent Recognition) | function_call(Function Call) | ✅function_call(Function Call) |
|
||||||
|
| Memory(Memory Function) | mem_local_short(Local Short-term Memory) | ✅mem_local_short(Local Short-term Memory) |
|
||||||
|
|
||||||
---
|
---
|
||||||
## Feature List ✨
|
## Feature List ✨
|
||||||
|
|
||||||
### Implemented ✅
|
### Implemented ✅
|
||||||
|
|
||||||
| Feature Module | Description |
|
| Feature Module | Description |
|
||||||
|---------|------|
|
|---------|------|
|
||||||
| Communication Protocol | Based on `xiaozhi-esp32` protocol, implements data interaction through WebSocket |
|
| Communication Protocol | Based on `xiaozhi-esp32` protocol, implements data interaction through WebSocket |
|
||||||
| Dialogue Interaction | Supports wake-up dialogue, manual dialogue, and real-time interruption. Auto-sleep after long periods of inactivity |
|
| Dialogue Interaction | Supports wake-up dialogue, manual dialogue, and real-time interruption. Auto-sleep after long periods of no dialogue |
|
||||||
| Intent Recognition | Supports LLM intent recognition, function call, reducing hard-coded intent judgment |
|
| Intent Recognition | Supports LLM intent recognition, function call, reducing hard-coded intent judgment |
|
||||||
| Multi-language Recognition | Supports Mandarin, Cantonese, English, Japanese, Korean (default using FunASR) |
|
| Multi-language Recognition | Supports Mandarin, Cantonese, English, Japanese, Korean (default using FunASR) |
|
||||||
| LLM Module | Supports flexible LLM module switching, default using ChatGLMLLM, also supports Ali Bailing, DeepSeek, Ollama, etc. |
|
| LLM Module | Supports flexible LLM module switching, default using ChatGLMLLM, can also use Ali Bailian, DeepSeek, Ollama, etc. |
|
||||||
| TTS Module | Supports EdgeTTS (default), Volcano Engine Doubao TTS, and other TTS interfaces for speech synthesis |
|
| TTS Module | Supports EdgeTTS (default), Volcano Engine Doubao TTS, and other TTS interfaces |
|
||||||
| Memory Function | Supports ultra-long memory, local summary memory, and no memory modes for different scenarios |
|
| Memory Function | Supports ultra-long memory, local summary memory, and no memory modes |
|
||||||
| IOT Function | Supports managing registered device IOT functionality, intelligent IoT control based on dialogue context |
|
| IOT Function | Supports managing registered device IOT functionality, supports smart IoT control based on dialogue context |
|
||||||
| Control Panel | Provides web management interface, supports agent management, user management, system configuration, etc. |
|
| Control Panel | Provides Web management interface, supports agent management, user management, system configuration, etc. |
|
||||||
|
|
||||||
### In Development 🚧
|
### In Development 🚧
|
||||||
|
|
||||||
To learn about specific development progress, [click here](https://github.com/users/xinnan-tech/projects/3)
|
To learn about specific development progress, [click here](https://github.com/users/xinnan-tech/projects/3)
|
||||||
|
|
||||||
If you're a software developer, here's an [Open Letter to Developers](docs/contributor_open_letter.md). Welcome to join!
|
If you are a software developer, here is an [Open Letter to Developers](docs/contributor_open_letter.md). Welcome to join!
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## Product Ecosystem 👬
|
## Product Ecosystem 👬
|
||||||
Xiaozhi is an ecosystem. When using this product, you might want to check out other excellent projects in this ecosystem:
|
Xiaozhi is an ecosystem. When using this product, you might also want to check out other excellent projects in this ecosystem
|
||||||
|
|
||||||
| Project Name | Project Link | Description |
|
| Project Name | Project Address | Project Description |
|
||||||
|:---------------------|:--------|:--------|
|
|:---------------------|:--------|:--------|
|
||||||
| Xiaozhi Android Client | [xiaozhi-android-client](https://github.com/TOM88812/xiaozhi-android-client) | A Flutter-based Android and iOS voice dialogue application supporting real-time voice interaction and text dialogue |
|
| Xiaozhi Android Client | [xiaozhi-android-client](https://github.com/TOM88812/xiaozhi-android-client) | A Flutter-based Android and iOS voice dialogue application supporting real-time voice interaction and text dialogue. |
|
||||||
| Xiaozhi PC Client | [py-xiaozhi](https://github.com/Huang-junsen/py-xiaozhi) | A Python-based AI client that allows you to experience Xiaozhi AI functionality through code without physical hardware |
|
| Xiaozhi PC Client | [py-xiaozhi](https://github.com/Huang-junsen/py-xiaozhi) | This project provides a Python-based Xiaozhi AI client, allowing you to experience Xiaozhi AI's functionality through code even without physical hardware. |
|
||||||
| Xiaozhi Java Server | [xiaozhi-esp32-server-java](https://github.com/joey-zhou/xiaozhi-esp32-server-java) | A Java-based open-source project providing complete backend service solutions |
|
| Xiaozhi Java Server | [xiaozhi-esp32-server-java](https://github.com/joey-zhou/xiaozhi-esp32-server-java) | The Java version of Xiaozhi open-source backend service is a Java-based open-source project.<br/>It includes both frontend and backend services, aiming to provide users with a complete backend service solution. |
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
## Supported Platforms/Components 📋
|
## Supported Platforms/Components List 📋
|
||||||
|
|
||||||
### LLM Language Models
|
### LLM Language Models
|
||||||
|
|
||||||
| Usage Method | Supported Platforms | Free Platforms |
|
| Usage Method | Supported Platforms | Free Platforms |
|
||||||
|:---:|:---:|:---:|
|
|:---:|:---:|:---:|
|
||||||
| openai API | Ali Bailing, Volcano Engine Doubao, DeepSeek, ChatGLM, Gemini | ChatGLM, Gemini |
|
| openai interface call | Ali Bailian, Volcano Engine Doubao, DeepSeek, Zhipu ChatGLM, Gemini | Zhipu ChatGLM, Gemini |
|
||||||
| ollama API | Ollama | - |
|
| ollama interface call | Ollama | - |
|
||||||
| dify API | Dify | - |
|
| dify interface call | Dify | - |
|
||||||
| fastgpt API | Fastgpt | - |
|
| fastgpt interface call | Fastgpt | - |
|
||||||
| coze API | Coze | - |
|
| coze interface call | Coze | - |
|
||||||
|
|
||||||
Actually, any LLM supporting openai API calls can be integrated.
|
In fact, any LLM that supports openai interface calls can be integrated and used.
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
### TTS Speech Synthesis
|
### TTS Speech Synthesis
|
||||||
|
|
||||||
| Usage Method | Supported Platforms | Free Platforms |
|
| Usage Method | Supported Platforms | Free Platforms |
|
||||||
|:---:|:---:|:---:|
|
|:---:|:---:|:---:|
|
||||||
| API Calls | EdgeTTS, Volcano Engine Doubao TTS, Tencent Cloud, Aliyun TTS, CosyVoiceSiliconflow, TTS302AI, CozeCnTTS, GizwitsTTS, ACGNTTS, OpenAITTS | EdgeTTS, CosyVoiceSiliconflow(partial) |
|
| API Call | EdgeTTS, Volcano Engine Doubao TTS, Tencent Cloud, Alibaba Cloud TTS, CosyVoiceSiliconflow, TTS302AI, CozeCnTTS, GizwitsTTS, ACGNTTS, OpenAITTS | EdgeTTS, CosyVoiceSiliconflow(partial) |
|
||||||
| Local Service | FishSpeech, GPT_SOVITS_V2, GPT_SOVITS_V3, MinimaxTTS | FishSpeech, GPT_SOVITS_V2, GPT_SOVITS_V3, MinimaxTTS |
|
| Local Service | FishSpeech, GPT_SOVITS_V2, GPT_SOVITS_V3, MinimaxTTS | FishSpeech, GPT_SOVITS_V2, GPT_SOVITS_V3, MinimaxTTS |
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
### VAD Voice Activity Detection
|
### VAD Voice Activity Detection
|
||||||
|
|
||||||
| Type | Platform Name | Usage Method | Pricing | Notes |
|
| Type | Platform Name | Usage Method | Pricing Model | Notes |
|
||||||
|:---:|:---------:|:----:|:----:|:--:|
|
|:---:|:---------:|:----:|:----:|:--:|
|
||||||
| VAD | SileroVAD | Local Use | Free | |
|
| VAD | SileroVAD | Local Usage | Free | |
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
@@ -229,26 +267,26 @@ Actually, any LLM supporting openai API calls can be integrated.
|
|||||||
|
|
||||||
| Usage Method | Supported Platforms | Free Platforms |
|
| Usage Method | Supported Platforms | Free Platforms |
|
||||||
|:---:|:---:|:---:|
|
|:---:|:---:|:---:|
|
||||||
| Local Use | FunASR, SherpaASR | FunASR, SherpaASR |
|
| Local Usage | FunASR, SherpaASR | FunASR, SherpaASR |
|
||||||
| API Calls | DoubaoASR, FunASRServer, TencentASR, AliyunASR | FunASRServer |
|
| API Call | DoubaoASR, FunASRServer, TencentASR, AliyunASR | FunASRServer |
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
### Memory Storage
|
### Memory Storage
|
||||||
|
|
||||||
| Type | Platform Name | Usage Method | Pricing | Notes |
|
| Type | Platform Name | Usage Method | Pricing Model | Notes |
|
||||||
|:------:|:---------------:|:----:|:---------:|:--:|
|
|:------:|:---------------:|:----:|:---------:|:--:|
|
||||||
| Memory | mem0ai | API Calls | 1000 calls/month quota | |
|
| Memory | mem0ai | API Call | 1000 calls/month quota | |
|
||||||
| Memory | mem_local_short | Local Summary | Free | |
|
| Memory | mem_local_short | Local Summary | Free | |
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
### Intent Recognition
|
### Intent Recognition
|
||||||
|
|
||||||
| Type | Platform Name | Usage Method | Pricing | Notes |
|
| Type | Platform Name | Usage Method | Pricing Model | Notes |
|
||||||
|:------:|:-------------:|:----:|:-------:|:---------------------:|
|
|:------:|:-------------:|:----:|:-------:|:---------------------:|
|
||||||
| Intent | intent_llm | API Calls | Based on LLM pricing | Uses large model for intent recognition, highly versatile |
|
| Intent | intent_llm | API Call | Based on LLM pricing | Uses large model for intent recognition, highly versatile |
|
||||||
| Intent | function_call | API Calls | Based on LLM pricing | Uses large model function calls for intent, fast and effective |
|
| Intent | function_call | API Call | Based on LLM pricing | Uses large model function calls for intent, fast and effective |
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
@@ -257,10 +295,10 @@ Actually, any LLM supporting openai API calls can be integrated.
|
|||||||
| Logo | Project/Company | Description |
|
| Logo | Project/Company | Description |
|
||||||
|:---:|:---:|:---|
|
|:---:|:---:|:---|
|
||||||
| <img src="./docs/images/logo_bailing.png" width="160"> | [Bailing Voice Dialogue Robot](https://github.com/wwbin2017/bailing) | This project was inspired by [Bailing Voice Dialogue Robot](https://github.com/wwbin2017/bailing) and implemented based on it |
|
| <img src="./docs/images/logo_bailing.png" width="160"> | [Bailing Voice Dialogue Robot](https://github.com/wwbin2017/bailing) | This project was inspired by [Bailing Voice Dialogue Robot](https://github.com/wwbin2017/bailing) and implemented based on it |
|
||||||
| <img src="./docs/images/logo_tenclass.png" width="160"> | [Tenclass](https://www.tenclass.com/) | Thanks to [Tenclass](https://www.tenclass.com/) for developing standard communication protocols, multi-device compatibility solutions, and high-concurrency scenario practices for the Xiaozhi ecosystem; providing comprehensive technical documentation support for this project |
|
| <img src="./docs/images/logo_tenclass.png" width="160"> | [Tenclass](https://www.tenclass.com/) | Thanks to [Tenclass](https://www.tenclass.com/) for establishing standard communication protocols, multi-device compatibility solutions, and high-concurrency scenario practices for the Xiaozhi ecosystem; providing full-chain 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 call framework, MCP communication protocol, and plugin call mechanism implementation code, significantly improving front-end device (IoT) interaction efficiency and functional extensibility through standardized instruction scheduling system and dynamic expansion capabilities |
|
| <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 the function call framework, MCP communication protocol, and plugin call mechanism implementation code, significantly improving front-end device (IoT) interaction efficiency and functional extensibility through standardized instruction scheduling system and dynamic expansion capabilities |
|
||||||
| <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, empowering product user experience with their design experience serving over a thousand enterprises |
|
| <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, empowering the product user experience with their design experience serving over a thousand enterprises |
|
||||||
| <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_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 the visual system of this project, ensuring consistency and extensibility of the overall design style in multi-scenario applications |
|
||||||
|
|
||||||
|
|
||||||
<a href="https://star-history.com/#xinnan-tech/xiaozhi-esp32-server&Date">
|
<a href="https://star-history.com/#xinnan-tech/xiaozhi-esp32-server&Date">
|
||||||
@@ -270,4 +308,4 @@ Actually, any LLM supporting openai API calls can be integrated.
|
|||||||
<source media="(prefers-color-scheme: light)" srcset="https://api.star-history.com/svg?repos=xinnan-tech/xiaozhi-esp32-server&type=Date" />
|
<source media="(prefers-color-scheme: light)" srcset="https://api.star-history.com/svg?repos=xinnan-tech/xiaozhi-esp32-server&type=Date" />
|
||||||
<img alt="Star History Chart" src="https://api.star-history.com/svg?repos=xinnan-tech/xiaozhi-esp32-server&type=Date" />
|
<img alt="Star History Chart" src="https://api.star-history.com/svg?repos=xinnan-tech/xiaozhi-esp32-server&type=Date" />
|
||||||
</picture>
|
</picture>
|
||||||
</a>
|
</a>
|
||||||
|
|||||||
+5
-1
@@ -108,7 +108,11 @@ VAD:
|
|||||||
|
|
||||||
参考教程[阿里云短信集成指南](./ali-sms-integration.md)
|
参考教程[阿里云短信集成指南](./ali-sms-integration.md)
|
||||||
|
|
||||||
### 9、更多问题,可联系我们反馈 💬
|
### 9、如何开启视觉模型实现拍照识物 📷
|
||||||
|
|
||||||
|
参考教程[视觉模型使用指南](./mcp-vision-integration.md)
|
||||||
|
|
||||||
|
### 10、更多问题,可联系我们反馈 💬
|
||||||
|
|
||||||
可以在[issues](https://github.com/xinnan-tech/xiaozhi-esp32-server/issues)提交您的问题。
|
可以在[issues](https://github.com/xinnan-tech/xiaozhi-esp32-server/issues)提交您的问题。
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,236 @@
|
|||||||
|
# Technical Documentation: xiaozhi-esp32-server
|
||||||
|
|
||||||
|
**Table of Contents:**
|
||||||
|
|
||||||
|
1. [Introduction](#1-introduction)
|
||||||
|
2. [Overall Architecture](#2-overall-architecture)
|
||||||
|
3. [Component Deep Dive](#3-component-deep-dive)
|
||||||
|
* [xiaozhi-server (Python AI Engine)](#31-xiaozhi-server-python-ai-engine)
|
||||||
|
* [manager-api (Java Management Backend)](#32-manager-api-java-management-backend)
|
||||||
|
* [manager-web (Vue.js Management Frontend)](#33-manager-web-vuejs-management-frontend)
|
||||||
|
4. [Data Flow and Interaction Mechanisms](#4-data-flow-and-interaction-mechanisms)
|
||||||
|
5. [Key Features Summary](#5-key-features-summary)
|
||||||
|
6. [Deployment and Configuration Overview](#6-deployment-and-configuration-overview)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 1. Introduction
|
||||||
|
|
||||||
|
The `xiaozhi-esp32-server` project provides a comprehensive backend system designed to power intelligent voice interactions for ESP32-based smart hardware. Its primary purpose is to enable developers to quickly establish a robust server infrastructure capable of understanding natural language commands, interacting with various AI services (for speech recognition, language understanding, and speech synthesis), managing IoT devices, and offering a web-based interface for system configuration and administration. This project facilitates the creation of customizable voice assistants and smart control systems by integrating multiple cutting-edge technologies into a cohesive and extensible platform.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 2. Overall Architecture
|
||||||
|
|
||||||
|
The `xiaozhi-esp32-server` system is architected as a distributed suite of interconnected components, each with a distinct role, ensuring modularity and scalability. The primary components are:
|
||||||
|
|
||||||
|
1. **ESP32 Hardware (Client Device):** This is the physical smart hardware device that the end-user interacts with. It's responsible for:
|
||||||
|
* Capturing user's voice commands.
|
||||||
|
* Sending captured audio to the `xiaozhi-server`.
|
||||||
|
* Receiving synthesized audio responses from `xiaozhi-server` and playing them back.
|
||||||
|
* Potentially controlling other connected peripherals or IoT devices based on commands from `xiaozhi-server`.
|
||||||
|
|
||||||
|
2. **`xiaozhi-server` (Core AI Engine):** This Python-based server is the central brain for voice processing and interaction logic. Its key responsibilities include:
|
||||||
|
* Establishing real-time, bidirectional WebSocket communication with ESP32 devices.
|
||||||
|
* Receiving audio streams and performing Voice Activity Detection (VAD).
|
||||||
|
* Converting speech to text using integrated Automatic Speech Recognition (ASR) services.
|
||||||
|
* Interpreting user intent and generating responses by interacting with Large Language Models (LLMs).
|
||||||
|
* Managing dialogue context and memory.
|
||||||
|
* Converting text responses back to speech using Text-to-Speech (TTS) services.
|
||||||
|
* Executing commands, including IoT device control via a plugin system.
|
||||||
|
* Fetching its operational configuration from the `manager-api`.
|
||||||
|
|
||||||
|
3. **`manager-api` (Management Backend):** A Java Spring Boot application that provides a RESTful API for system administration and configuration. It serves as the backend for the `manager-web` frontend and a configuration source for `xiaozhi-server`. Its functions include:
|
||||||
|
* User authentication and management for the control panel.
|
||||||
|
* Registration and management of ESP32 devices.
|
||||||
|
* Storage and retrieval of system configurations (e.g., selected AI service providers, API keys, device settings) in a MySQL database.
|
||||||
|
* Providing endpoints for `xiaozhi-server` to fetch its configuration.
|
||||||
|
* Managing voice timbre settings, OTA firmware updates, and other system parameters.
|
||||||
|
* Utilizing Redis for caching to enhance performance.
|
||||||
|
|
||||||
|
4. **`manager-web` (Web Control Panel):** A Vue.js Single Page Application (SPA) that provides a graphical user interface for administrators. It allows for:
|
||||||
|
* Easy configuration of `xiaozhi-server`'s AI services and operational parameters.
|
||||||
|
* Management of users, devices, and their respective settings.
|
||||||
|
* Monitoring system status (potentially) and managing other administrative tasks.
|
||||||
|
* Interaction with all backend functionalities exposed by `manager-api`.
|
||||||
|
|
||||||
|
**High-Level Interaction Flow:**
|
||||||
|
|
||||||
|
* The **ESP32** device captures voice and communicates primarily with **`xiaozhi-server`** via WebSockets for all voice-related interactions.
|
||||||
|
* **`xiaozhi-server`** processes the voice data, interacts with various AI cloud services or local models, and sends responses back to the ESP32.
|
||||||
|
* The **`manager-web`** frontend communicates with **`manager-api`** using RESTful HTTP calls to manage and configure the entire system.
|
||||||
|
* **`xiaozhi-server`** also communicates with **`manager-api`** (via REST) to pull its latest configuration, ensuring that changes made in the web panel are reflected in its operation.
|
||||||
|
|
||||||
|
This separation of concerns allows the `xiaozhi-server` to focus on efficient real-time AI processing, while the `manager-api` and `manager-web` provide a robust and user-friendly interface for administration and setup.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 3. Component Deep Dive
|
||||||
|
|
||||||
|
### 3.1. `xiaozhi-server` (Python AI Engine)
|
||||||
|
|
||||||
|
The `xiaozhi-server` is the intelligent core of the system, responsible for processing voice interactions, interfacing with AI services, and managing communication with ESP32 devices.
|
||||||
|
|
||||||
|
* **Purpose:**
|
||||||
|
* To provide real-time processing of voice commands from ESP32 devices.
|
||||||
|
* To integrate with various AI services for Speech-to-Text (ASR), Natural Language Understanding (via Large Language Models - LLMs), Text-to-Speech (TTS), Voice Activity Detection (VAD), Intent Recognition, and Memory.
|
||||||
|
* To manage dialogue flow and context with users.
|
||||||
|
* To execute custom functions and control IoT devices based on user commands.
|
||||||
|
* To be dynamically configurable through the `manager-api`.
|
||||||
|
|
||||||
|
* **Core Technologies:**
|
||||||
|
* **Python 3:** The primary programming language.
|
||||||
|
* **Asyncio:** Python's asynchronous programming framework, crucial for handling concurrent WebSocket connections and non-blocking I/O for AI service API calls.
|
||||||
|
* **`websockets` Library:** For WebSocket server implementation.
|
||||||
|
* **HTTP Client (e.g., `aiohttp`, `httpx`):** For asynchronous HTTP requests to `manager-api` and external AI services.
|
||||||
|
* **YAML (PyYAML):** For local configuration file parsing.
|
||||||
|
|
||||||
|
* **Key Implementation Aspects:**
|
||||||
|
|
||||||
|
1. **AI Service Provider Pattern (`core/providers/`):**
|
||||||
|
* **Concept:** A flexible design for integrating AI services. Each service type (ASR, TTS, LLM, etc.) has an abstract base class defining a common interface. Concrete classes implement this interface for specific vendors or local models.
|
||||||
|
* **Benefit:** Allows easy switching of AI service backends via configuration and simplifies adding new service integrations.
|
||||||
|
* **Initialization:** `core/utils/modules_initialize.py` acts as a factory to load and instantiate configured providers.
|
||||||
|
|
||||||
|
2. **WebSocket Communication & Connection Handling (`core/websocket_server.py`, `core/connection.py`):**
|
||||||
|
* **Server Setup:** Manages WebSocket connections from ESP32 devices.
|
||||||
|
* **Connection Isolation:** Each ESP32 client gets a dedicated `ConnectionHandler` instance, isolating its session state and dialogue.
|
||||||
|
* **Dynamic Configuration Updates:** Can fetch updated configurations from `manager-api` and re-initialize AI service modules live, without a full server restart.
|
||||||
|
|
||||||
|
3. **Message Handling & Dialogue Flow (`core/handle/`):**
|
||||||
|
* Employs a modular handler pattern. The `ConnectionHandler` dispatches message processing to specialized modules based on message type or dialogue phase (e.g., `receiveAudioHandle.py` for audio input, `intentHandler.py` for NLU, `functionHandler.py` for plugin execution, `sendAudioHandle.py` for TTS output).
|
||||||
|
|
||||||
|
4. **Plugin System for Extensible Functions (`plugins_func/`):**
|
||||||
|
* **Purpose:** Allows adding custom "skills" (e.g., weather, news, Home Assistant control).
|
||||||
|
* **Mechanism:** Plugins define functions and schemas. The LLM can request execution of these functions (function calling). `loadplugins.py` and `register.py` manage plugin discovery and registration.
|
||||||
|
|
||||||
|
5. **Configuration Management (`config/`):**
|
||||||
|
* Loads settings from a local `config.yaml` and merges them with configurations fetched from `manager-api` (via `manage_api_client.py`), enabling remote dynamic configuration.
|
||||||
|
* `logger.py` sets up structured application logging.
|
||||||
|
* `config/assets/` stores predefined audio files for system notifications.
|
||||||
|
|
||||||
|
6. **Auxiliary HTTP Server (`core/http_server.py`):**
|
||||||
|
* Handles specific HTTP requests, notably for OTA firmware updates (`/xiaozhi/ota/`) and other utility endpoints.
|
||||||
|
|
||||||
|
### 3.2. `manager-api` (Java Management Backend)
|
||||||
|
|
||||||
|
The `manager-api` component is a backend server built using Java and the Spring Boot framework, serving as the administrative hub.
|
||||||
|
|
||||||
|
* **Purpose:**
|
||||||
|
* Provide a secure RESTful API for the `manager-web` frontend.
|
||||||
|
* Act as a centralized configuration provider for `xiaozhi-server`.
|
||||||
|
* Manage persistent data (users, devices, AI configurations, voice timbres, OTA firmware).
|
||||||
|
|
||||||
|
* **Core Technologies:**
|
||||||
|
* **Java 21 & Spring Boot 3:** Core language and framework.
|
||||||
|
* **Spring MVC:** For building REST controllers.
|
||||||
|
* **MyBatis-Plus:** ORM for database interaction with MySQL.
|
||||||
|
* **MySQL:** Relational database.
|
||||||
|
* **Druid:** JDBC connection pool.
|
||||||
|
* **Redis (Spring Data Redis):** For caching.
|
||||||
|
* **Apache Shiro:** Security framework for authentication and authorization.
|
||||||
|
* **Liquibase:** Database schema migration.
|
||||||
|
* **Knife4j:** OpenAPI (Swagger) API documentation.
|
||||||
|
* **Maven:** Build and dependency management.
|
||||||
|
|
||||||
|
* **Key Implementation Aspects:**
|
||||||
|
|
||||||
|
1. **Modular Architecture (`modules/` package):**
|
||||||
|
* Business logic is organized into distinct modules (e.g., `sys` for users/roles, `agent` for assistant configs, `device` for ESP32s, `config` for `xiaozhi-server` settings, `security`, `timbre`, `ota`).
|
||||||
|
* Each module typically follows a layered pattern: Controller, Service, DAO (Mapper), Entity, DTO.
|
||||||
|
|
||||||
|
2. **Layered Architecture:**
|
||||||
|
* **Controller Layer (`@RestController`):** Defines API endpoints, handles HTTP request/response.
|
||||||
|
* **Service Layer (`@Service`):** Contains business logic, transaction management.
|
||||||
|
* **Data Access Layer (MyBatis-Plus Mappers):** Interacts with the MySQL database.
|
||||||
|
|
||||||
|
3. **Common Functionalities (`common/` package):**
|
||||||
|
* Provides shared code: base classes, global configurations (Spring, MyBatis, Redis, Knife4j), custom annotations (e.g., `@LogOperation`), AOP aspects, global exception handling, utility classes, and XSS protection.
|
||||||
|
|
||||||
|
4. **Security (Apache Shiro):**
|
||||||
|
* Manages user authentication and permissions for accessing API endpoints. Configured with Shiro Realms and security filters.
|
||||||
|
|
||||||
|
5. **Database Schema Management (Liquibase):**
|
||||||
|
* Ensures consistent database structure across environments through versioned schema changes.
|
||||||
|
|
||||||
|
### 3.3. `manager-web` (Vue.js Management Frontend)
|
||||||
|
|
||||||
|
The `manager-web` is a Single Page Application (SPA) providing the administrative user interface.
|
||||||
|
|
||||||
|
* **Purpose:**
|
||||||
|
* Offer a web-based control panel for system configuration and management.
|
||||||
|
* Enable administrators to configure `xiaozhi-server`'s AI services, manage users and devices, customize voice timbres, and handle OTA updates.
|
||||||
|
|
||||||
|
* **Core Technologies:**
|
||||||
|
* **Vue.js 2 & Vue CLI:** Core JavaScript framework and build tools.
|
||||||
|
* **Vue Router:** For client-side routing within the SPA.
|
||||||
|
* **Vuex:** For centralized state management.
|
||||||
|
* **Element UI:** UI component library for a consistent look and feel.
|
||||||
|
* **SCSS:** CSS preprocessor.
|
||||||
|
* **HTTP Client (Flyio or Axios):** For API calls to `manager-api`.
|
||||||
|
* **Workbox:** For PWA features (caching, service worker).
|
||||||
|
* **Opus Libraries:** For potential in-browser audio recording/playback.
|
||||||
|
|
||||||
|
* **Key Implementation Aspects:**
|
||||||
|
|
||||||
|
1. **SPA Structure:** Single HTML page with dynamic view updates.
|
||||||
|
2. **Component-Based Architecture:** UI built from reusable Vue components (`.vue` files in `src/views/` for pages and `src/components/` for smaller elements).
|
||||||
|
3. **Client-Side Routing (`src/router/index.js`):** Maps browser URLs to view components, with route guards for authentication.
|
||||||
|
4. **State Management (`src/store/index.js`):** Vuex manages global state (user info, device lists, etc.) via state, getters, mutations, and actions (often involving API calls).
|
||||||
|
5. **API Communication (`src/apis/`):** Modularized API service files make asynchronous calls to `manager-api`.
|
||||||
|
6. **Build Process & PWA Features:** Vue CLI (Webpack) bundles assets. Workbox enables PWA features like caching.
|
||||||
|
7. **Environment Configuration (`.env` files):** Manages settings like the `manager-api` base URL for different environments.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 4. Data Flow and Interaction Mechanisms
|
||||||
|
|
||||||
|
The system uses WebSockets for real-time voice interactions and RESTful APIs for management tasks.
|
||||||
|
|
||||||
|
* **Core Voice Interaction (ESP32 <-> `xiaozhi-server` - WebSockets):**
|
||||||
|
* ESP32 connects to `xiaozhi-server` via WebSocket.
|
||||||
|
* Audio is streamed from ESP32 to server.
|
||||||
|
* Server processes audio (VAD, ASR), interacts with LLM (possibly executing plugin functions), synthesizes response via TTS.
|
||||||
|
* Synthesized audio is streamed back to ESP32.
|
||||||
|
* JSON control/status messages are also exchanged.
|
||||||
|
|
||||||
|
* **Management & Configuration (RESTful APIs - HTTP/JSON):**
|
||||||
|
* **`manager-web` -> `manager-api`:** Admin actions in the web UI trigger REST API calls to `manager-api` for managing users, devices, configurations, etc. Shiro secures these endpoints.
|
||||||
|
* **`xiaozhi-server` -> `manager-api`:** `xiaozhi-server` pulls its operational configuration from `manager-api` via REST API calls.
|
||||||
|
|
||||||
|
* **OTA Updates (Conceptual - HTTP & WebSocket):**
|
||||||
|
* Firmware uploaded via `manager-web` to `manager-api`.
|
||||||
|
* `xiaozhi-server` may notify ESP32 of updates via WebSocket.
|
||||||
|
* ESP32 downloads firmware via HTTP from an endpoint (likely on `xiaozhi-server`).
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 5. Key Features Summary
|
||||||
|
|
||||||
|
* **Modular AI Services:** Pluggable ASR, LLM, TTS, VAD, Intent, Memory.
|
||||||
|
* **Advanced Dialogue:** Real-time interruption, contextual memory, multi-language support.
|
||||||
|
* **Extensible Skills:** Plugin system for custom functions (e.g., IoT, Home Assistant).
|
||||||
|
* **Comprehensive Web Management:** UI for users, devices, AI configs, OTA, timbres.
|
||||||
|
* **Flexible Deployment:** Docker (simplified/full) and source code options.
|
||||||
|
* **Dynamic Remote Configuration:** `xiaozhi-server` updates settings from `manager-api` live.
|
||||||
|
* **Open Source (MIT License).**
|
||||||
|
* **Cost-Effective Options:** "Entry Level Free Settings" available.
|
||||||
|
* **PWA Admin Panel:** Enhanced caching and user experience.
|
||||||
|
* **API Documentation:** Knife4j for `manager-api`.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 6. Deployment and Configuration Overview
|
||||||
|
|
||||||
|
* **Deployment:**
|
||||||
|
* **Docker:** Recommended for ease. Options for `xiaozhi-server` only or full stack (all components + databases). `docker-compose.yml` files provided.
|
||||||
|
* **Source Code:** For development or custom setups, requiring manual environment setup for Python, Java/Maven, and Node.js.
|
||||||
|
|
||||||
|
* **Configuration:**
|
||||||
|
* **`xiaozhi-server`:** Uses a local `config.yaml`, but primarily pulls dynamic configurations (AI providers, API keys) from `manager-api` via its `manage_api_client.py`.
|
||||||
|
* **`manager-api`:** Configured via Spring Boot's `application.properties` or `application.yml` (database, Redis, Shiro settings).
|
||||||
|
* **`manager-web`:** Configured via `.env` files (e.g., `manager-api` URL).
|
||||||
|
* The `manager-web` UI is the primary interface for most system configurations in a full deployment.
|
||||||
|
* Predefined profiles like "Entry Level Free Settings" and "Full Streaming Configuration" guide AI service choices.
|
||||||
|
|
||||||
|
---
|
||||||
@@ -0,0 +1,450 @@
|
|||||||
|
# 技术文档:`xiaozhi-esp32-server`
|
||||||
|
|
||||||
|
**目录:**
|
||||||
|
|
||||||
|
1. [引言](#1-引言)
|
||||||
|
2. [整体架构](#2-整体架构)
|
||||||
|
3. [核心组件深度剖析](#3-核心组件深度剖析)
|
||||||
|
* [3.1. `xiaozhi-server` (核心AI引擎 - Python实现)](#31-xiaozhi-server-核心ai引擎---python实现)
|
||||||
|
* [3.2. `manager-api` (管理后端 - Java Spring Boot实现)](#32-manager-api-管理后端---java-spring-boot实现)
|
||||||
|
* [3.3. `manager-web` (Web管理前端 - Vue.js实现)](#33-manager-web-web管理前端---vuejs实现)
|
||||||
|
4. [数据流与交互机制](#4-数据流与交互机制)
|
||||||
|
5. [核心功能概要](#5-核心功能概要)
|
||||||
|
6. [部署与配置概述](#6-部署与配置概述)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 1. 引言
|
||||||
|
|
||||||
|
`xiaozhi-esp32-server` 项目是一个专为基于ESP32的智能硬件提供支持的**综合性后端系统**。其核心目标是使开发人员能够快速构建一个强大的服务器基础设施,该设施不仅能够理解自然语言指令,还能与多种AI服务(用于语音识别、自然语言理解及语音合成)进行高效交互、管理物联网(IoT)设备,并提供一个基于Web的用户界面以进行系统配置和管理。通过将多种尖端技术整合到一个高内聚且可扩展的平台中,本项目旨在简化和加速可定制化语音助手及智能控制系统的开发进程。它不仅仅是一个简单的服务器,更是一个连接硬件、AI能力与用户管理的桥梁。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 2. 整体架构
|
||||||
|
|
||||||
|
`xiaozhi-esp32-server` 系统采用了一种**分布式、多组件协作**的架构设计,确保了系统的模块化、可维护性和可扩展性。各个核心组件各司其职,协同工作。主要组件包括:
|
||||||
|
|
||||||
|
1. **ESP32 硬件 (客户端设备):**
|
||||||
|
这是终端用户直接与之交互的物理智能硬件设备。其主要职责包括:
|
||||||
|
* 捕捉用户的语音指令。
|
||||||
|
* 将捕捉到的原始音频数据安全地发送至 `xiaozhi-server` 进行处理。
|
||||||
|
* 接收来自 `xiaozhi-server` 合成的语音回复,并通过扬声器播放给用户。
|
||||||
|
* 根据从 `xiaozhi-server` 收到的指令,控制与之连接的其他外围设备或IoT设备(例如智能灯泡、传感器等)。
|
||||||
|
|
||||||
|
2. **`xiaozhi-server` (核心AI引擎 - Python实现):**
|
||||||
|
这个基于Python的服务器是整个系统的“大脑”,负责处理所有语音相关的逻辑和AI交互。其关键职责细化如下:
|
||||||
|
* 通过WebSocket协议与ESP32设备建立**稳定、低延迟的实时双向通信链路**。
|
||||||
|
* 接收来自ESP32的音频流,并利用语音活动检测(VAD)技术精确切分有效的语音片段。
|
||||||
|
* 集成并调用自动语音识别(ASR)服务(可配置本地或云端),将语音片段转换为文本。
|
||||||
|
* 通过与大型语言模型(LLM)的交互来解析用户意图、生成智能回复,并支持复杂的自然语言理解任务。
|
||||||
|
* 管理多轮对话中的上下文信息和用户记忆,以提供连贯的交互体验。
|
||||||
|
* 调用文本转语音(TTS)服务,将LLM生成的文本回复合成为自然流畅的语音。
|
||||||
|
* 通过一个灵活的**插件系统**执行自定义命令,包括对IoT设备的控制逻辑。
|
||||||
|
* 从 `manager-api` 服务获取其详细的运行时操作配置。
|
||||||
|
|
||||||
|
3. **`manager-api` (管理后端 - Java实现):**
|
||||||
|
这是一个基于Java Spring Boot框架构建的应用程序,它为整个系统的管理和配置提供了一套安全的RESTful API。它不仅是 `manager-web` 控制台的后端支撑,也是 `xiaozhi-server` 的配置数据来源。其核心功能包括:
|
||||||
|
* 为Web控制台提供用户认证(登录、权限验证)和用户账户管理功能。
|
||||||
|
* ESP32设备的注册、信息管理以及设备特定配置的维护。
|
||||||
|
* 在**MySQL数据库**中持久化存储系统配置,例如用户选择的AI服务提供商、API密钥、设备参数、插件设置等。
|
||||||
|
* 提供特定的API端点,供 `xiaozhi-server` 拉取其所需的最新配置。
|
||||||
|
* 管理TTS音色选项、处理OTA(Over-The-Air)固件更新流程及相关元数据。
|
||||||
|
* 利用 **Redis** 作为高速缓存,存储热点数据(如会话信息、频繁访问的配置),以提升API响应速度和系统整体性能。
|
||||||
|
|
||||||
|
4. **`manager-web` (Web控制面板 - Vue.js实现):**
|
||||||
|
这是一个基于Vue.js构建的单页应用(SPA),为系统管理员提供了一个图形化、用户友好的操作界面。其主要能力包括:
|
||||||
|
* 便捷地配置 `xiaozhi-server` 所使用的各项AI服务(如ASR、LLM、TTS的提供商切换、参数调整)。
|
||||||
|
* 管理平台用户账户、角色分配及权限控制。
|
||||||
|
* 管理已注册的ESP32设备及其相关设置。
|
||||||
|
* (潜在功能)监控系统运行状态、查看日志、进行故障排查等。
|
||||||
|
* 与 `manager-api` 提供的所有后端管理功能进行全面的交互。
|
||||||
|
|
||||||
|
**高层交互流程概述:**
|
||||||
|
|
||||||
|
* **语音交互主线:** **ESP32设备**捕捉到用户语音后,通过**WebSocket**将音频数据实时传输给**`xiaozhi-server`**。`xiaozhi-server`完成一系列AI处理(VAD、ASR、LLM交互、TTS)后,再通过WebSocket将合成的语音回复发送回ESP32设备进行播放。所有与语音直接相关的实时交互均在此链路完成。
|
||||||
|
* **管理配置主线:** 管理员通过浏览器访问**`manager-web`**控制台。`manager-web`通过调用**`manager-api`**提供的**RESTful HTTP接口**来执行各种管理操作(如修改配置、管理用户或设备)。数据以JSON格式在两者间传递。
|
||||||
|
* **配置同步:** **`xiaozhi-server`**在启动或特定更新机制触发时,会主动通过HTTP请求从**`manager-api`**拉取其最新的操作配置。这确保了管理员在Web界面上所做的配置更改能够及时有效地应用到核心AI引擎的运行中。
|
||||||
|
|
||||||
|
这种**前后端分离、核心服务与管理服务分离**的架构设计,使得 `xiaozhi-server`能够专注于高效的实时AI处理任务,而 `manager-api` 和 `manager-web` 则共同提供了一个功能强大且易于使用的管理和配置平台。各组件职责清晰,有利于独立开发、测试、部署和扩展。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 3. 核心组件深度剖析
|
||||||
|
|
||||||
|
### 3.1. `xiaozhi-server` (核心AI引擎 - Python实现)
|
||||||
|
|
||||||
|
`xiaozhi-server` 作为系统的智能核心,全权负责处理语音交互、对接各类AI服务以及管理与ESP32设备间的通信。其设计目标是实现高效、灵活且可扩展的语音AI处理能力。
|
||||||
|
|
||||||
|
* **核心目标:**
|
||||||
|
* 为ESP32设备提供实时的语音指令处理服务。
|
||||||
|
* 深度集成各类AI服务,包括:自动语音识别 (ASR)、大型语言模型 (LLM) 进行自然语言理解 (NLU)、文本转语音 (TTS)、语音活动检测 (VAD)、意图识别 (Intent Recognition) 及对话记忆 (Memory)。
|
||||||
|
* 精细管理用户与设备间的对话流程及上下文状态。
|
||||||
|
* 基于用户指令,通过插件化机制执行自定义函数及控制物联网 (IoT) 设备。
|
||||||
|
* 支持通过 `manager-api`进行动态配置加载与更新。
|
||||||
|
|
||||||
|
* **核心技术栈:**
|
||||||
|
* **Python 3:** 作为主要编程语言,Python以其丰富的AI/ML生态库和快速开发特性被选用。
|
||||||
|
* **Asyncio:** Python的异步编程框架,是`xiaozhi-server`高性能的关键。它被广泛用于高效处理来自大量ESP32设备的并发WebSocket连接,以及执行与外部AI服务API通信时的非阻塞I/O操作,确保服务器在高并发下的响应能力。
|
||||||
|
* **`websockets` 库:** 提供WebSocket服务器的具体实现,支持与ESP32客户端进行全双工实时通信。
|
||||||
|
* **HTTP客户端 (如 `aiohttp`, `httpx`):** 用于异步执行HTTP请求,主要目的是从`manager-api`获取配置信息,以及与云端AI服务的API进行交互。
|
||||||
|
* **YAML (通常通过 PyYAML 库):** 用于解析本地的 `config.yaml` 配置文件。
|
||||||
|
* **FFmpeg (外部依赖):** 在 `app.py` 启动时会进行检查 (`check_ffmpeg_installed()`)。FFmpeg通常用于音频处理和格式转换,例如,确保音频数据符合特定AI服务的要求或进行内部处理。
|
||||||
|
|
||||||
|
* **关键实现细节:**
|
||||||
|
|
||||||
|
1. **AI服务提供者模式 (Provider Pattern - `core/providers/`):**
|
||||||
|
* **设计思想:** 这是`xiaozhi-server`集成不同AI服务的核心设计模式,极大地增强了系统的灵活性和可扩展性。针对每一种AI服务类型(ASR, TTS, LLM, VAD, Intent, Memory, VLLM),都在其对应子目录下定义了一个抽象基类 (ABC, Abstract Base Class),例如 `core/providers/asr/base.py`。这个基类规定了该类型服务必须实现的通用接口方法(如ASR的 `async def transcribe(self, audio_chunk: bytes) -> str: pass`)。
|
||||||
|
* **具体实现:** 各种具体的AI服务提供商或本地模型的实现,则以独立的Python类形式存在(例如 `core/providers/asr/fun_local.py` 实现了本地FunASR的逻辑,`core/providers/llm/openai.py` 实现了与OpenAI GPT模型的对接)。这些具体类继承自相应的抽象基类,并实现其定义的接口。部分提供者还使用DTOs (Data Transfer Objects, 存在于各自的 `dto/` 目录) 来结构化与外部服务交换的数据。
|
||||||
|
* **优势:** 使得核心业务逻辑能够以统一的方式调用不同的AI服务,而无需关心其底层具体实现。用户可以通过配置文件轻松切换AI服务后端。添加对新AI服务的支持也变得相对简单,只需实现对应的Provider接口。
|
||||||
|
* **动态加载与初始化:** `core/utils/modules_initialize.py` 脚本扮演了工厂的角色。它在服务器启动时,或在接收到配置更新指令时,会根据配置文件中 `selected_module` 及各项服务的具体provider设置,动态地导入并实例化相应的Provider类。
|
||||||
|
|
||||||
|
2. **WebSocket通信与连接处理 (`app.py`, `core/websocket_server.py`, `core/connection.py`):**
|
||||||
|
* **服务器启动与入口 (`app.py`):**
|
||||||
|
* `app.py` 作为主入口,负责初始化应用环境(如检查FFmpeg、加载配置、设置日志)。
|
||||||
|
* 它会生成或加载一个 `auth_key` (JWT密钥),用于保护特定的HTTP接口(如视觉分析接口 `/mcp/vision/explain`)。若配置中 `manager-api.secret` 为空,则会生成一个UUID作为 `auth_key`。
|
||||||
|
* 使用 `asyncio.create_task()` 并发启动 `WebSocketServer` (监听如 `ws://0.0.0.0:8000/xiaozhi/v1/`) 和 `SimpleHttpServer` (监听如 `http://0.0.0.0:8003/xiaozhi/ota/`)。
|
||||||
|
* 包含一个 `monitor_stdin()` 协程,用于在某些环境下保持应用存活或处理终端输入。
|
||||||
|
* **WebSocket服务器核心 (`core/websocket_server.py`):**
|
||||||
|
* `WebSocketServer` 类使用 `websockets` 库监听来自ESP32设备的连接请求。
|
||||||
|
* 对于每一个成功的WebSocket连接,它都会创建一个**独立的 `ConnectionHandler` 实例** (推测定义于 `core/connection.py`)。这种每个连接一个处理程序实例的设计模式,是实现多设备状态隔离和并发处理的关键,确保每个设备的对话流程和上下文信息互不干扰。
|
||||||
|
* 该服务器还提供一个 `_http_response` 方法,允许在同一端口上对非WebSocket升级的HTTP GET请求做出简单响应(例如返回 "Server is running"),便于进行健康检查。
|
||||||
|
* **动态配置更新:** `WebSocketServer` 包含一个 `update_config()` 异步方法。此方法使用 `config_lock` (一个 `asyncio.Lock`) 保证配置更新的原子性。它调用 `get_config_from_api()` (可能在 `config_loader.py` 中实现,通过 `manage_api_client.py` 与 `manager-api` 通信) 来获取新的配置。通过 `check_vad_update()` 和 `check_asr_update()` 等辅助函数判断是否需要重新初始化特定的AI模块,避免不必要的开销。更新后的配置会用于重新调用 `initialize_modules()`,从而实现AI服务提供者的热切换。
|
||||||
|
|
||||||
|
3. **消息处理与对话流程控制 (`core/handle/` 和 `ConnectionHandler`):**
|
||||||
|
* `ConnectionHandler` (推测) 作为每个连接的控制中心,负责接收来自ESP32的消息,并根据消息类型或当前对话状态,将其分发给 `core/handle/` 目录下的相应处理模块。这种模块化的处理器设计使得 `ConnectionHandler` 逻辑更清晰,易于扩展。
|
||||||
|
* **主要处理模块及其职责:**
|
||||||
|
* `helloHandle.py`: 处理与ESP32初次连接时的握手协议、设备认证或初始化信息交换。
|
||||||
|
* `receiveAudioHandle.py`: 接收音频流数据,调用VAD Provider进行语音活动检测,并将有效的音频片段传递给ASR Provider进行识别。
|
||||||
|
* `textHandle.py` / `intentHandler.py`: 获取ASR识别出的文本后,与Intent Provider (可能利用LLM进行意图识别) 和LLM Provider交互,以理解用户意图并生成初步回复或决策。
|
||||||
|
* `functionHandler.py`: 当LLM的响应包含执行特定“函数调用”的指令时,此模块负责从插件注册表中查找并执行对应的插件函数。
|
||||||
|
* `sendAudioHandle.py`: 将LLM最终生成的文本回复交给TTS Provider合成语音,并将音频流通过WebSocket发送回ESP32。
|
||||||
|
* `abortHandle.py`: 处理来自ESP32的中断请求,例如停止当前的TTS播报。
|
||||||
|
* `iotHandle.py`, `mcpHandle.py`: 处理与IoT设备控制相关的特定指令或更复杂的模块通信协议 (MCP)。
|
||||||
|
|
||||||
|
4. **插件化功能扩展系统 (`plugins_func/`):**
|
||||||
|
* **设计目的:** 提供一种标准化的方式来扩展语音助手的功能和“技能”,而无需修改核心代码。
|
||||||
|
* **实现机制:**
|
||||||
|
* 各个具体功能以独立的Python脚本形式存在于 `plugins_func/functions/` 目录中(例如 `get_weather.py`, `hass_set_state.py` 用于Home Assistant集成)。
|
||||||
|
* `loadplugins.py` 在服务器启动时负责扫描并加载这些插件模块。
|
||||||
|
* `register.py` (或插件模块内部的特定装饰器/函数) 可能用于定义每个插件函数的元数据,包括:
|
||||||
|
* **函数名称 (Function Name):** LLM调用时使用的标识符。
|
||||||
|
* **功能描述 (Description):** 供LLM理解此函数的作用。
|
||||||
|
* **参数模式 (Parameters Schema):** 通常是一个JSON Schema,详细定义了函数所需的参数、类型、是否必需以及描述。这是LLM能够正确生成函数调用参数的关键。
|
||||||
|
* **执行流程:** 当LLM在其思考过程中决定需要调用某个外部工具或函数来获取信息或执行操作时,它会依据预先提供的函数模式生成一个结构化的“函数调用”请求。`xiaozhi-server`中的`functionHandler.py`捕获此请求,从插件注册表中找到对应的Python函数并执行,然后将执行结果返回给LLM,LLM再基于此结果生成最终给用户的自然语言回复。
|
||||||
|
|
||||||
|
5. **配置管理 (`config/`):**
|
||||||
|
* **加载机制:** `config_loader.py` (通过 `settings.py` 被调用) 负责从根目录的 `config.yaml` 文件加载基础配置。
|
||||||
|
* **远程配置与合并:** 通过 `manage_api_client.py` (使用如`aiohttp`的库与`manager-api`通信) 可以从`manager-api`服务拉取配置。远程配置通常会覆盖本地 `config.yaml` 中的同名设置,从而实现通过Web界面动态调整服务器行为。
|
||||||
|
* **日志系统:** `logger.py` 初始化应用日志系统(可能使用 `loguru` 或对标准 `logging` 模块进行封装,支持通过 `logger.bind(tag=TAG)` 添加标签,便于追踪和过滤)。
|
||||||
|
* **静态资源:** `config/assets/` 目录下存放了用于系统提示音的静态音频文件(如设备绑定提示音 `bind_code.wav`、错误提示音等)。
|
||||||
|
|
||||||
|
6. **辅助HTTP服务 (`core/http_server.py`):**
|
||||||
|
* 与WebSocket服务并行运行一个简单的HTTP服务器,用于处理特定的HTTP请求。最主要的功能是为ESP32设备提供OTA (Over-The-Air) 固件更新的下载服务 (通过 `/xiaozhi/ota/` 端点)。此外,也可能承载其他如 `/mcp/vision/explain` (视觉分析) 等工具性HTTP接口。
|
||||||
|
|
||||||
|
综上所述,`xiaozhi-server` 是一个采用现代Python异步编程模型构建的、高度模块化、配置驱动的AI应用服务器。其精心设计的Provider模式和插件架构赋予了它强大的适应性和扩展性,能够灵活接入不同的AI能力并支持日益增长的功能需求。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### 3.2. `manager-api` (管理后端 - Java Spring Boot实现)
|
||||||
|
|
||||||
|
`manager-api` 组件是使用Java和Spring Boot框架构建的强大后端服务,作为整个`xiaozhi-esp32-server`生态系统的中央行政管理和配置中枢。
|
||||||
|
|
||||||
|
* **核心目标:**
|
||||||
|
* 为`manager-web`(Vue.js前端)提供一套安全、稳定、符合RESTful规范的API接口,使得管理员能够便捷地管理用户、设备、系统配置及其他相关资源。
|
||||||
|
* 充当`xiaozhi-server`(Python核心AI引擎)的集中化配置数据提供者,允许`xiaozhi-server`实例在启动或运行时获取其最新的操作参数。
|
||||||
|
* 持久化存储关键数据,例如:用户账户信息、设备注册详情、AI服务提供商配置(包括API密钥、选定的服务模型等)、TTS音色参数,以及OTA固件版本信息等。
|
||||||
|
|
||||||
|
* **核心技术栈:**
|
||||||
|
* **Java 21:** 项目采用的JDK版本,确保了对现代Java特性的支持。
|
||||||
|
* **Spring Boot 3:** 作为核心开发框架,极大地简化了独立、生产级别的Spring应用的创建和部署。它提供了自动配置、内嵌Web服务器(默认为Tomcat)、依赖管理等关键功能。
|
||||||
|
* **Spring MVC:** Spring框架中用于构建Web应用和RESTful API的模块。
|
||||||
|
* **MyBatis-Plus:** 一个对MyBatis进行功能增强的ORM(对象关系映射)框架。它简化了数据库操作,提供了强大的CRUD(增删改查)功能、条件构造器、代码生成器等,并能很好地与Spring Boot集成。
|
||||||
|
* **MySQL:** 作为主要的后端关系型数据库,用于存储所有需要持久化的管理数据和配置信息。
|
||||||
|
* **Druid (Alibaba Druid):** 一个功能强大的JDBC连接池实现,提供了丰富的监控功能和优秀的性能,用于高效管理数据库连接。
|
||||||
|
* **Redis (通过 Spring Data Redis):** 一个高性能的内存数据结构存储,常用于实现数据缓存(例如缓存热点配置数据、用户会话信息),以显著提升API的响应速度。
|
||||||
|
* **Apache Shiro:** 一个成熟且易用的Java安全框架,负责处理应用的认证(用户身份验证)和授权(API访问权限控制)需求。
|
||||||
|
* **Liquibase:** 一个用于跟踪、管理和应用数据库 schéma(模式)变更的开源工具。它允许开发者以数据库无关的方式定义和版本化数据库结构变更。
|
||||||
|
* **Knife4j:** 一个集成了Swagger并增强了UI的API文档生成工具,专为Java MVC框架(尤其是Spring Boot)设计。它能生成美观且易于交互的API文档界面(通常通过 `/xiaozhi/doc.html` 访问)。
|
||||||
|
* **Maven:** 用于项目的构建自动化和依赖项管理。
|
||||||
|
* **Lombok:** 一个Java库,通过注解自动生成构造函数、getter/setter、equals/hashCode、toString等样板代码,减少冗余。
|
||||||
|
* **HuTool / Google Guava:** 提供大量实用工具类,简化常见编程任务。
|
||||||
|
* **Aliyun Dysmsapi:** 阿里云短信服务SDK,用于集成发送短信功能(如验证码、通知)。
|
||||||
|
|
||||||
|
* **关键实现细节:**
|
||||||
|
|
||||||
|
1. **模块化项目结构 (`modules/` 包):**
|
||||||
|
* `manager-api` 的核心业务逻辑被清晰地划分到 `src/main/java/xiaozhi/modules/` 目录下的不同模块中。这种按功能领域划分模块的方式(例如 `sys` 负责系统管理,`agent` 负责智能体配置,`device` 负责设备管理,`config` 负责为`xiaozhi-server`提供配置,`security` 负责安全,`timbre` 负责音色管理,`ota` 负责固件升级)极大地提高了代码的可维护性和可扩展性。
|
||||||
|
* **各模块内部结构:** 每个业务模块通常遵循经典的三层架构或其变体:
|
||||||
|
* **Controller (控制层):** 位于 `xiaozhi.modules.[模块名].controller`。
|
||||||
|
* **Service (服务层):** 位于 `xiaozhi.modules.[模块名].service`。
|
||||||
|
* **DAO/Mapper (数据访问层):** 位于 `xiaozhi.modules.[模块名].dao`。
|
||||||
|
* **Entity (实体类):** 位于 `xiaozhi.modules.[模块名].entity`。
|
||||||
|
* **DTO (数据传输对象):** 位于 `xiaozhi.modules.[模块名].dto`。
|
||||||
|
|
||||||
|
2. **分层架构实现:**
|
||||||
|
* **Controller层 (`@RestController`):** 这些类使用Spring MVC注解(如 `@GetMapping`, `@PostMapping` 等)来定义API的端点(endpoints)。它们负责接收HTTP请求,将请求体中的JSON数据反序列化为DTO对象,调用相应的Service层方法处理业务逻辑,最后将Service层的返回结果序列化为JSON并作为HTTP响应返回给客户端。
|
||||||
|
* **Service层 (`@Service`):** 这些类(通常是接口及其实现类的组合)封装了核心的业务规则和操作流程。它们可能会调用一个或多个DAO/Mapper对象来与数据库交互,并常常使用 `@Transactional` 注解来管理数据库事务的原子性。
|
||||||
|
* **Data Access (DAO/Mapper) 层 (MyBatis-Plus Mappers):** 这些是Java接口,继承自MyBatis-Plus提供的 `BaseMapper<Entity>` 接口。MyBatis-Plus会为这些接口自动提供标准的CRUD方法。对于更复杂的数据库查询,开发者可以通过在Mapper接口中定义方法并使用注解(如 `@Select`, `@Update`)或编写对应的XML映射文件来实现。例如,`UserMapper.selectById(userId)` 会被MyBatis-Plus自动实现。
|
||||||
|
* **Entity层 (`@TableName`, `@TableId` 等MyBatis-Plus注解):** 这些POJO(Plain Old Java Objects)类直接映射到数据库中的表结构。Lombok的 `@Data` 注解常用于自动生成getter/setter等。
|
||||||
|
* **DTO层:** 用于在各层之间,特别是Controller层与Service层之间,以及API的请求/响应体中传递数据。使用DTO有助于解耦API接口的数据结构与数据库实体的数据结构,使API更稳定。
|
||||||
|
|
||||||
|
3. **通用功能与配置 (`common/` 包):**
|
||||||
|
* `src/main/java/xiaozhi/common/` 包提供了一系列跨模块共享的通用组件和配置:
|
||||||
|
* **基类:** 如 `BaseDao`, `BaseEntity`, `BaseService`, `CrudService`,为各模块的相应组件提供通用的属性或方法。
|
||||||
|
* **全局配置:** 包括 `MybatisPlusConfig` (MyBatis-Plus的配置,如分页插件、数据权限插件等)、`RedisConfig` (Redis连接及序列化配置)、`SwaggerConfig` (Knife4j的配置)、`AsyncConfig` (异步任务执行器配置)。
|
||||||
|
* **自定义注解:** 例如 `@LogOperation` 用于通过AOP记录操作日志,`@DataFilter` 可能用于实现数据范围过滤。
|
||||||
|
* **AOP切面:** 如 `RedisAspect` 可能用于实现方法级别的缓存逻辑。
|
||||||
|
* **全局异常处理:** `RenExceptionHandler` (使用 `@ControllerAdvice` 注解) 捕获应用中抛出的特定或所有异常 (如自定义的 `RenException`),并返回统一格式的JSON错误响应给客户端。`ErrorCode` 定义了标准化的错误码。
|
||||||
|
* **工具类:** 提供了日期转换、JSON处理(Jackson)、IP地址获取、HTTP上下文操作、统一结果封装 (`Result` 类)等多种实用工具。
|
||||||
|
* **校验工具:** `ValidatorUtils` 和 `AssertUtils` 用于简化参数校验逻辑。
|
||||||
|
* **XSS防护:** `XssFilter` 等组件用于防止跨站脚本攻击。
|
||||||
|
* **MyBatis-Plus自动填充:** `FieldMetaObjectHandler` 用于在执行插入或更新数据库操作时,自动填充如 `createTime`, `updateTime` 等公共字段。
|
||||||
|
|
||||||
|
4. **安全机制 (Apache Shiro):**
|
||||||
|
* Shiro的配置(通常在 `modules/security/config/` 或 `common/config/` 下)定义了如何进行用户认证和授权。
|
||||||
|
* **Realms (域):** 自定义的Shiro Realm类负责从数据库中查询用户信息(用户名、密码、盐值)进行身份验证,以及获取用户的角色和权限信息用于授权决策。
|
||||||
|
* **Filters (过滤器):** Shiro过滤器链被应用于保护API端点,确保只有经过认证且拥有足够权限的用户才能访问特定资源。
|
||||||
|
* **Session/Token Management:** Shiro管理用户会话。对于RESTful API,可能结合OAuth2或JWT等令牌机制实现无状态认证。
|
||||||
|
|
||||||
|
5. **数据库版本控制 (Liquibase):**
|
||||||
|
* 数据库的表结构、索引、初始数据等变更,都通过Liquibase的 `changelog` 文件(通常是XML格式)进行定义和版本化管理。当应用启动时,Liquibase会自动检查并应用必要的数据库结构更新,确保开发、测试和生产环境数据库结构的一致性。
|
||||||
|
|
||||||
|
`manager-api` 通过这些精心选择的技术和设计模式,构建了一个功能全面、结构清晰、安全可靠且易于维护和扩展的Java后端服务。其模块化的设计特别适合处理具有多种管理功能需求的复杂系统。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### 3.3. `manager-web` (Web管理前端 - Vue.js实现)
|
||||||
|
|
||||||
|
`manager-web` 组件是一个采用 Vue.js 2 框架构建的单页应用 (SPA - Single Page Application)。它为系统管理员提供了一个功能丰富、交互友好的图形用户界面,用于全面管理和配置 `xiaozhi-esp32-server` 生态系统。
|
||||||
|
|
||||||
|
* **核心目标:**
|
||||||
|
* 提供一个基于Web的集中式控制面板,供管理员进行系统操作与监控。
|
||||||
|
* 实现对 `xiaozhi-server` 中AI服务提供商(ASR、LLM、TTS等)及其相关API密钥或许可配置的便捷管理。
|
||||||
|
* 支持用户账户、角色及权限的精细化管理。
|
||||||
|
* 提供ESP32设备的注册、配置及状态查看功能。
|
||||||
|
* 允许管理员自定义TTS音色、管理OTA固件更新流程、调整系统级参数及字典数据等。
|
||||||
|
* 作为 `manager-api` 所暴露各项功能的图形化交互前端。
|
||||||
|
|
||||||
|
* **核心技术栈:**
|
||||||
|
* **Vue.js 2:** 一个渐进式的JavaScript框架,用于构建用户界面。其核心特性包括声明式渲染、组件化系统、数据绑定等,非常适合构建复杂的SPA。
|
||||||
|
* **Vue CLI (`@vue/cli-service`):** Vue.js的官方命令行工具,用于项目的快速搭建、开发服务器的运行(支持热模块替换HMR)、以及生产环境构建打包(内部集成并配置了Webpack)。
|
||||||
|
* **Vue Router (`vue-router`):** Vue.js官方的路由管理器。它负责在SPA内部实现不同“页面”或视图组件之间的导航切换,而无需重新加载整个HTML页面,提供了流畅的用户体验。
|
||||||
|
* **Vuex (`vuex`):** Vue.js官方的状态管理模式和库。它充当了应用中所有组件的“中央数据存储”,用于管理全局共享状态(例如当前登录用户信息、设备列表、应用配置等),特别适用于大型复杂应用。
|
||||||
|
* **Element UI (`element-ui`):** 一个广受欢迎的基于Vue 2.0的桌面端UI组件库。它提供了大量预先设计和实现的组件(如表单、表格、对话框、导航菜单、按钮、提示等),帮助开发者快速构建出专业且一致的用户界面。
|
||||||
|
* **JavaScript (ES6+):** 前端逻辑实现的主要编程语言,利用其现代特性进行开发。
|
||||||
|
* **SCSS (Sassy CSS):** 一种CSS预处理器,它为CSS增加了变量、嵌套规则、混合(Mixin)、继承等高级特性,使得CSS代码更易于组织、维护和复用。
|
||||||
|
* **HTTP客户端 (Flyio 或 Axios 通过 `vue-axios`):** 用于在浏览器端向 `manager-api` 后端发起异步HTTP(AJAX)请求,以获取数据或提交操作。
|
||||||
|
* **Webpack:** 一个强大的模块打包工具(由Vue CLI在底层管理和配置)。它将项目中的各种资源(JavaScript文件、CSS、图片、字体等)视为模块,并将它们打包成浏览器可识别的静态文件。
|
||||||
|
* **Workbox (通过 `workbox-webpack-plugin`):** Google开发的一个库,用于简化Service Worker的编写和PWA(Progressive Web App - 渐进式Web应用)的实现。它可以帮助生成Service Worker脚本,实现资源缓存、离线访问等功能。
|
||||||
|
* **Opus库 (`opus-decoder`, `opus-recorder`):** 这些音频处理库表明前端可能具备一些直接在浏览器中处理Opus格式音频的能力,例如:用于测试麦克风输入、允许管理员录制自定义音频片段(可能用于TTS音色样本或语音指令测试),或播放在管理界面中预览的Opus编码音频。
|
||||||
|
|
||||||
|
* **关键实现细节:**
|
||||||
|
|
||||||
|
1. **单页应用 (SPA) 结构:**
|
||||||
|
* 整个前端应用加载一个主HTML文件 (`public/index.html`)。后续的所有页面切换和内容更新都在客户端由Vue Router动态完成,无需每次都从服务器请求新的HTML页面。这种模式能提供更快的页面加载速度和更流畅的交互体验。
|
||||||
|
|
||||||
|
2. **组件化架构 (Component-Based Architecture):**
|
||||||
|
* 用户界面由一系列可复用的Vue组件 (`.vue` 单文件组件) 构成,形成一个组件树。这种方式提高了代码的模块化程度、可维护性和复用性。
|
||||||
|
* **`src/main.js`:** 应用的入口JS文件。它负责创建和初始化根Vue实例,注册全局插件(如Vue Router, Vuex, Element UI),并把根Vue实例挂载到 `public/index.html` 中的某个DOM元素上(通常是 `#app`)。
|
||||||
|
* **`src/App.vue`:** 应用的根组件。它通常定义了应用的基础布局结构(如包含导航栏、侧边栏、主内容区),并通过 `<router-view></router-view>` 标签来显示当前路由匹配到的视图组件。
|
||||||
|
* **视图组件 (`src/views/`):** 这些组件代表了应用中的各个“页面”或主要功能区(例如 `Login.vue` 登录页, `DeviceManagement.vue` 设备管理页, `UserManagement.vue` 用户管理页, `ModelConfig.vue` 模型配置页)。它们通常由Vue Router直接映射。
|
||||||
|
* **可复用UI组件 (`src/components/`):** 包含了在不同视图之间共享的、更小粒度的UI组件(例如 `HeaderBar.vue` 顶部导航栏, `AddDeviceDialog.vue` 添加设备对话框, `AudioPlayer.vue` 音频播放器组件)。
|
||||||
|
|
||||||
|
3. **客户端路由 (`src/router/index.js`):**
|
||||||
|
* Vue Router在此文件中进行配置,定义了应用的路由表。每个路由规则将一个特定的URL路径映射到一个视图组件。
|
||||||
|
* 常常包含**导航守卫 (Navigation Guards)**,例如 `beforeEach` 守卫,用于在路由跳转前执行逻辑,如检查用户是否已登录,如果未登录则重定向到登录页面,从而保护需要认证才能访问的页面。
|
||||||
|
|
||||||
|
4. **状态管理 (`src/store/index.js`):**
|
||||||
|
* Vuex被用来构建一个集中的状态管理中心(Store)。这个Store包含了:
|
||||||
|
* **State:** 存储应用级别的共享数据(例如,当前登录用户的详细信息、从API获取的设备列表、系统配置等)。
|
||||||
|
* **Getters:** 类似于Vue组件中的计算属性,用于从State派生出一些状态值,方便组件使用。
|
||||||
|
* **Mutations:** **唯一**可以同步修改State中数据的方法。它们必须是同步函数。
|
||||||
|
* **Actions:** 用于处理异步操作(如API调用)或封装多个Mutation提交。Actions会调用API,获取数据后,通过 `commit` 一个或多个Mutation来更新State。
|
||||||
|
* 例如,用户登录时,一个名为 `login` 的Action可能会被调用,它会向后端API发送登录请求,成功后获取到用户信息和token,然后 `commit` 一个名为 `SET_USER_INFO` 的Mutation来更新State中的用户信息和token。
|
||||||
|
|
||||||
|
5. **API通信 (`src/apis/`):**
|
||||||
|
* 与 `manager-api` 后端的所有HTTP通信逻辑被封装在 `src/apis/` 目录下,通常会按照后端API的模块进行组织(例如 `src/apis/module/agent.js`, `src/apis/module/device.js`)。
|
||||||
|
* 每个模块导出一系列函数,每个函数对应一个具体的API请求。这些函数内部使用配置好的HTTP客户端实例 (例如,在 `src/apis/api.js` 或 `src/apis/httpRequest.js` 中统一配置Axios或Flyio实例,可能包含设置请求基地址、请求/响应拦截器等)。
|
||||||
|
* **拦截器 (Interceptors):** HTTP客户端的请求拦截器常用于在每个请求发送前自动添加认证令牌(如JWT);响应拦截器则可用于全局处理API错误(如权限不足、服务器错误)或对响应数据进行预处理。
|
||||||
|
|
||||||
|
6. **样式与资源 (`src/styles/`, `src/assets/`):**
|
||||||
|
* `Element UI` 提供了基础的组件样式。
|
||||||
|
* `src/styles/global.scss` 文件用于定义全局共享的SCSS样式、变量、混合(Mixin)等。
|
||||||
|
* Vue单文件组件内部的 `<style scoped>` 标签允许编写只作用于当前组件的局部样式。
|
||||||
|
* `src/assets/` 目录存放图片、字体等静态资源。
|
||||||
|
|
||||||
|
7. **构建与PWA特性:**
|
||||||
|
* Vue CLI通过Webpack将所有代码和资源打包成优化的静态文件,用于生产部署。
|
||||||
|
* `workbox-webpack-plugin` 的使用(体现在 `service-worker.js` 和 `registerServiceWorker.js` 文件)表明项目集成了Service Worker技术。Service Worker可以拦截网络请求,实现前端资源的智能缓存(从而加快后续访问速度),甚至在网络断开时提供一定的离线访问能力,是PWA的核心技术之一。
|
||||||
|
|
||||||
|
8. **环境配置 (`.env`系列文件):**
|
||||||
|
* 项目根目录下的 `.env` (以及 `.env.development`, `.env.production` 等) 文件用于定义环境变量。这些变量(例如 `VUE_APP_API_BASE_URL` 来指定 `manager-api` 的基础URL)可以在应用代码中通过 `process.env.VUE_APP_XXX` 的形式访问,从而允许为不同构建环境(开发、测试、生产)配置不同的参数。
|
||||||
|
|
||||||
|
`manager-web` 通过这些技术的综合运用,构建了一个功能强大、易于维护且用户体验良好的管理界面,为 `xiaozhi-esp32-server` 系统的配置和监控提供了坚实的前端支持。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 4. 数据流与交互机制
|
||||||
|
|
||||||
|
`xiaozhi-esp32-server` 系统通过各组件间定义清晰的数据流和交互协议来协同工作。主要的通信方式依赖于针对实时交互优化的WebSocket协议和适用于客户端-服务器请求的RESTful API。
|
||||||
|
|
||||||
|
**4.1.核心语音交互流程 (ESP32设备 <-> `xiaozhi-server`)**
|
||||||
|
|
||||||
|
此流程是实时的,主要通过WebSocket进行低延迟、双向的数据交换。
|
||||||
|
|
||||||
|
* **连接建立与握手:**
|
||||||
|
* ESP32设备作为客户端,主动向`xiaozhi-server`的指定端点(例如 `ws://<服务器IP>:<WebSocket端口>/xiaozhi/v1/`)发起WebSocket连接请求。
|
||||||
|
* `xiaozhi-server` (`core/websocket_server.py`) 接收连接,并为每个成功连接的ESP32设备实例化一个独立的`ConnectionHandler`对象来管理该会话的整个生命周期。
|
||||||
|
* 连接建立后,可能会执行一个初始握手流程(由`core/handle/helloHandle.py`处理),用于交换设备标识、认证信息、协议版本或基本状态。
|
||||||
|
|
||||||
|
* **音频上行传输 (ESP32 -> `xiaozhi-server`):**
|
||||||
|
* 用户对ESP32设备讲话后,设备上的麦克风捕捉原始音频数据(通常是PCM或经过压缩如Opus的格式)。
|
||||||
|
* ESP32将这些音频数据块(chunks)作为WebSocket的**二进制消息 (binary messages)** 实时推送到`xiaozhi-server`对应的`ConnectionHandler`。
|
||||||
|
* 服务器端的`core/handle/receiveAudioHandle.py`模块负责接收、缓冲并处理这些音频数据。
|
||||||
|
|
||||||
|
* **AI核心处理 (在`xiaozhi-server`内部):**
|
||||||
|
* **VAD (语音活动检测):** `receiveAudioHandle.py`利用配置的VAD提供者(如SileroVAD)分析音频流,以准确识别语音的起始和结束点,滤除静默或噪声片段。
|
||||||
|
* **ASR (自动语音识别):** 检测到的有效语音片段被送往配置的ASR提供者(本地如FunASR,或云端服务)。ASR引擎将音频信号转换为文本字符串。
|
||||||
|
* **NLU/LLM (自然语言理解/大型语言模型):** ASR输出的文本,连同从Memory提供者获取的当前对话上下文历史,以及从`plugins_func/`加载的可用函数(工具)的描述模式,一同被传递给配置的LLM提供者。
|
||||||
|
* **函数调用执行 (若LLM决策需要):** 如果LLM分析后认为需要调用外部函数(例如查询天气、控制家电),它会生成一个结构化的函数调用请求。`core/handle/functionHandler.py`接收此请求,查找并执行在`plugins_func/`中定义的相应Python函数,并将函数的执行结果返回给LLM。LLM随后基于此结果生成最终的自然语言回复。
|
||||||
|
* **回复生成:** LLM综合所有信息(用户输入、上下文、函数调用结果等)生成最终的文本回复。
|
||||||
|
* **记忆更新:** 当前轮次的交互(用户问题、LLM回复、可能的功能调用)会被Memory提供者处理,以更新对话历史,供后续交互使用。
|
||||||
|
* **TTS (文本转语音):** LLM生成的最终文本回复被送往配置的TTS提供者,后者将文本合成为语音数据流(例如MP3或WAV格式)。
|
||||||
|
|
||||||
|
* **音频下行响应 (`xiaozhi-server` -> ESP32):**
|
||||||
|
* 由TTS提供者合成的语音数据流,通过`core/handle/sendAudioHandle.py`模块,作为WebSocket的**二进制消息**实时发送回ESP32设备。
|
||||||
|
* ESP32设备接收这些音频数据块并立即通过扬声器播放给用户。
|
||||||
|
|
||||||
|
* **控制与状态消息 (双向):**
|
||||||
|
* 除了音频流,ESP32与`xiaozhi-server`之间也通过WebSocket交换**文本消息 (text messages)**,这些消息通常采用JSON格式封装。
|
||||||
|
* **ESP32 -> Server:** 设备可能发送状态报告(如网络状况、麦克风状态)、错误代码、或特定的控制命令(例如用户按键触发的“停止TTS播报”)。
|
||||||
|
* **Server -> ESP32:** 服务器可能发送控制指令给设备(如“开始监听”、“停止监听”、调整灵敏度、下发特定配置参数)。
|
||||||
|
* `core/handle/abortHandle.py`(处理中断请求)、`core/handle/reportHandle.py`(处理设备报告)等模块负责解析和响应这些控制/状态消息。
|
||||||
|
|
||||||
|
**4.2.管理与配置流程 (`manager-web` <-> `manager-api` <-> `xiaozhi-server`)**
|
||||||
|
|
||||||
|
此流程主要依赖于基于HTTP/HTTPS的RESTful API进行请求-响应式的交互。
|
||||||
|
|
||||||
|
* **管理员UI后端交互 (`manager-web` -> `manager-api`):**
|
||||||
|
* 当管理员在`manager-web`界面执行操作时(例如保存一项配置、添加一个新用户、注册一台ESP32设备):
|
||||||
|
* Vue.js前端应用 (`manager-web`) 会通过其API封装模块(位于`src/apis/module/`)向`manager-api`的对应REST API端点发起异步HTTP请求(通常是GET, POST, PUT, DELETE)。
|
||||||
|
* 请求体和响应体通常使用JSON格式。
|
||||||
|
* `manager-api`中的`@RestController`类接收这些请求。**Apache Shiro**框架会首先对请求进行认证和授权检查。
|
||||||
|
* 通过验证后,Controller将请求分发给相应的Service层处理业务逻辑。Service层可能会与MySQL数据库(通过MyBatis-Plus)交互,并可能利用Redis进行数据缓存。
|
||||||
|
* 处理完成后,`manager-api`向`manager-web`返回一个JSON格式的HTTP响应。
|
||||||
|
* `manager-web`根据响应结果更新其Vuex状态存储和用户界面显示。
|
||||||
|
|
||||||
|
* **配置同步 (`manager-api` -> `xiaozhi-server`):**
|
||||||
|
* `xiaozhi-server`的运行依赖于从`manager-api`获取的动态配置(例如当前选用的AI服务提供商及其API密钥)。
|
||||||
|
* **拉取机制 (Pull Mechanism):** `xiaozhi-server`内部的`config/manage_api_client.py`模块,在服务器启动时或通过特定更新触发器(例如`WebSocketServer.update_config()`被调用),会向`manager-api`的一个指定端点(例如由`modules/config/controller/`中的某个Controller提供)发起HTTP GET请求。
|
||||||
|
* `manager-api`响应该请求,返回`xiaozhi-server`所需的配置数据(JSON格式)。
|
||||||
|
* `xiaozhi-server`接收到配置后,会更新其内部状态,并可能重新初始化相关的AI服务模块,以使新配置生效。
|
||||||
|
|
||||||
|
* **OTA固件更新流程 (概念性描述):**
|
||||||
|
* 管理员通过`manager-web`界面上传新的ESP32固件包到`manager-api`的特定端点。
|
||||||
|
* `manager-api`将固件文件存储起来,并记录相关元数据(版本号、适用设备型号等)。
|
||||||
|
* 当管理员触发对特定设备的OTA更新时:
|
||||||
|
* `manager-api`可能会通知`xiaozhi-server`(具体通知机制可能是一个轮询检查点,或`xiaozhi-server`暴露一个接收更新通知的API,或者更松耦合的如消息队列)。
|
||||||
|
* `xiaozhi-server`随后可以通过WebSocket向目标ESP32设备发送一条包含固件下载URL的指令消息。
|
||||||
|
* ESP32设备收到指令后,通过HTTP GET请求从该URL下载固件。此URL可能指向`xiaozhi-server`自身运行的`SimpleHttpServer`所服务的路径(如`/xiaozhi/ota/`),或者在某些架构中,也可能直接指向`manager-api`或专用的文件服务器。
|
||||||
|
|
||||||
|
**4.3. 主要协议总结:**
|
||||||
|
|
||||||
|
* **WebSocket:** 被选用于ESP32与`xiaozhi-server`之间的通信链路,因为它非常适合实时、低延迟、双向的数据流传输(尤其是音频),以及异步控制消息的传递。
|
||||||
|
* **RESTful APIs (基于HTTP/HTTPS,通常使用JSON作为数据交换格式):** 这是Web服务间通信的标准方式。用于`manager-web`(客户端)与`manager-api`(服务器)之间的请求-响应交互,也用于`xiaozhi-server`(作为客户端)从`manager-api`(作为服务器)拉取配置信息。其无状态特性、广泛的库支持和易于理解的语义使其成为此类交互的理想选择。
|
||||||
|
|
||||||
|
这种多协议并用的通信策略,确保了系统内不同类型的交互需求都能得到高效和恰当的处理,兼顾了实时性和标准化的请求-响应模式。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 5. 核心功能概要
|
||||||
|
|
||||||
|
`xiaozhi-esp32-server` 系统提供了一系列丰富的功能,旨在支持开发者构建先进的语音控制应用:
|
||||||
|
|
||||||
|
1. **全面的语音交互后端:** 提供从语音捕获指导到响应生成和动作执行的端到端解决方案。
|
||||||
|
2. **模块化和可插拔的AI服务:**
|
||||||
|
* 支持广泛的ASR(自动语音识别)、LLM(大型语言模型)、TTS(文本转语音)、VAD(语音活动检测)、意图识别和记忆提供商。
|
||||||
|
* 允许动态选择和配置这些服务(包括基于云的API和本地模型),以平衡成本、性能、隐私和语言需求。
|
||||||
|
3. **高级对话管理:**
|
||||||
|
* 支持自然交互,具有唤醒词启动对话、手动(按键说话式)对话以及对系统响应的实时打断等功能。
|
||||||
|
* 包含上下文记忆,以在多轮对话中保持连贯性。
|
||||||
|
* 在一段时间不活动后具有自动休眠模式。
|
||||||
|
4. **多语言能力:**
|
||||||
|
* 支持多种语言的识别和合成,包括普通话、粤语、英语、日语和韩语(具体取决于所选的ASR/LLM/TTS提供商)。
|
||||||
|
5. **通过插件实现的可扩展功能:**
|
||||||
|
* 强大的插件系统允许开发人员添加自定义“技能”或函数(例如,获取天气、控制智能家居设备、访问新闻)。
|
||||||
|
* 这些函数可以由LLM使用其函数调用能力,根据提供的模式来触发。
|
||||||
|
* 内置对Home Assistant集成的支持。
|
||||||
|
6. **物联网设备控制:**
|
||||||
|
* 设计用于通过语音命令管理和控制智能家居设备及其他物联网硬件,并利用插件系统。
|
||||||
|
7. **基于Web的管理控制台 (`manager-web` & `manager-api`):**
|
||||||
|
* 提供全面的图形界面,用于:
|
||||||
|
* 系统配置(AI服务选择、API密钥、操作参数)。
|
||||||
|
* 基于角色的访问控制的用户管理。
|
||||||
|
* ESP32设备注册和管理。
|
||||||
|
* 语音音色/TTS语音定制。
|
||||||
|
* ESP32设备的OTA(空中下载)固件更新管理。
|
||||||
|
* 系统参数和字典的管理。
|
||||||
|
8. **灵活的部署选项:**
|
||||||
|
* 支持通过Docker容器(用于简化的仅服务器或全栈设置)和直接从源代码部署,以适应各种环境和用户专业知识。
|
||||||
|
9. **动态远程配置:**
|
||||||
|
* `xiaozhi-server`可以从`manager-api`获取其配置,允许实时更新AI提供商和设置,而无需重新启动服务器。
|
||||||
|
10. **开源和社区驱动:**
|
||||||
|
* 根据MIT许可证授权,鼓励透明、协作和社区贡献。
|
||||||
|
11. **经济高效的解决方案:**
|
||||||
|
* 提供“入门全免费设置”路径,利用AI服务的免费套餐或本地模型,使其易于进行实验和个人项目。
|
||||||
|
12. **渐进式Web应用 (PWA) 特性:**
|
||||||
|
* `manager-web`控制面板包含Service Worker集成,以增强缓存和潜在的离线访问能力。
|
||||||
|
13. **详细的API文档:**
|
||||||
|
* `manager-api`通过Knife4j提供OpenAPI (Swagger) 文档,以便清晰理解和测试其RESTful端点。
|
||||||
|
|
||||||
|
这些功能共同使`xiaozhi-esp32-server`成为一个强大、适应性强且用户友好的平台,用于构建复杂的语音交互应用程序。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 6. 部署与配置概述
|
||||||
|
|
||||||
|
`xiaozhi-esp32-server`系统在设计上充分考虑了灵活性,提供了多种部署方法和全面的配置选项,以适应不同的使用场景和需求。
|
||||||
|
|
||||||
|
**部署选项:**
|
||||||
|
|
||||||
|
项目可以通过多种方式部署,主要包括使用Docker简化安装过程,或直接从源代码部署以获得更大的控制权和进行开发。
|
||||||
|
|
||||||
|
1. **基于Docker的部署:**
|
||||||
|
* **简化安装 (仅`xiaozhi-server`):** 此选项仅部署核心的基于Python的`xiaozhi-server`。它适用于主要需要语音AI处理能力和IoT控制,而不需要完整Web管理界面和数据库支持功能(如OTA)的用户。在此模式下,配置通常通过本地文件(`config.yaml`)管理,但如果需要,仍可将其指向一个已存在的`manager-api`实例。
|
||||||
|
* **全模块安装 (所有组件):** 此方案部署所有核心组件:`xiaozhi-server`、基于Java的`manager-api`、以及基于Vue.js的`manager-web`,同时还包括所需的数据库服务(MySQL和Redis)。这提供了完整的系统体验,包括用于全面配置和管理的Web控制面板。
|
||||||
|
* 项目为每个服务提供了`Dockerfile`定义,并使用`docker-compose.yml`文件(例如`docker-compose.yml`用于基础版,`docker-compose_all.yml`用于全功能版)来编排和管理多容器的部署。此外,还可能提供一个`docker-setup.sh`脚本来辅助自动化部分Docker环境的搭建工作。
|
||||||
|
|
||||||
|
2. **源代码部署:**
|
||||||
|
* 这种方法需要为每个组件手动设置相应的开发环境:Python环境用于`xiaozhi-server`,Java/Maven环境用于`manager-api`,Node.js/Vue CLI环境用于`manager-web`。
|
||||||
|
* 对于全模块安装,还需要手动安装和配置MySQL及Redis数据库服务。
|
||||||
|
* 这种方式通常用于项目开发、深度定制、调试,或者在对环境有特殊要求的生产场景中。
|
||||||
|
|
||||||
|
**配置管理:**
|
||||||
|
|
||||||
|
配置是定制系统行为的关键,尤其是在选择AI服务提供商和管理API密钥方面。
|
||||||
|
|
||||||
|
1. **`xiaozhi-server` 配置:**
|
||||||
|
* **本地`config.yaml`:** 位于`xiaozhi-server`根目录下的一个主要的YAML格式配置文件。它定义了服务器端口、选定的AI服务提供商(ASR、LLM、TTS、VAD、意图识别、记忆模块等)、它们各自的API密钥或模型路径、插件配置以及日志级别等。
|
||||||
|
* **通过`manager-api`进行远程配置:** `xiaozhi-server`被设计为可以从`manager-api`获取其运行配置。从`manager-api`获取的设置通常会覆盖本地`config.yaml`中的同名设置。这带来了两大好处:
|
||||||
|
* **集中管理:** 所有配置都可以通过`manager-web`界面进行统一管理。
|
||||||
|
* **动态更新:** `xiaozhi-server`可以刷新其配置并重新初始化AI模块,而无需完全重启服务。
|
||||||
|
* `xiaozhi-server`中的`config/config_loader.py`和`config/manage_api_client.py`负责处理配置的加载、合并及从`manager-api`拉取的逻辑。
|
||||||
|
|
||||||
|
2. **`manager-api` 配置:**
|
||||||
|
* 作为一个Spring Boot应用,其配置主要通过位于`src/main/resources`目录下的`application.properties`或`application.yml`文件进行管理。
|
||||||
|
* 关键配置项包括:数据库连接信息(MySQL的URL、用户名、密码)、Redis服务器地址和端口、应用服务端口(默认为8002)、Apache Shiro安全相关的设置,以及任何集成的第三方服务(如阿里云短信)的配置参数。
|
||||||
|
|
||||||
|
3. **`manager-web` 配置:**
|
||||||
|
* Vue.js前端应用的环境特定设置通过项目根目录下的`.env`系列文件(例如`.env`, `.env.development`, `.env.production`)进行管理。
|
||||||
|
* 这里最关键的配置通常是`manager-api`后端的API基础URL地址 (例如 `VUE_APP_API_BASE_URL`),前端应用将向此地址发送所有API请求。
|
||||||
|
|
||||||
|
4. **预定义的配置方案:**
|
||||||
|
* 项目文档(通常是README)中会推荐一些常见的配置组合,例如:
|
||||||
|
* **“入门全免费设置”:** 该方案旨在利用云AI服务的免费套餐额度或完全免费的本地模型,以最大程度地降低用户的初始使用成本和运营费用。
|
||||||
|
* **“全流式配置”:** 该方案优先考虑系统的响应速度和交互的流畅性,通常会选用支持流式处理的(可能付费的)AI服务。
|
||||||
|
* 这些预定义方案为用户在`xiaozhi-server`中配置AI服务提供商(通过`manager-web`界面或直接修改`config.yaml`)提供了指导。
|
||||||
|
|
||||||
|
在全模块部署的情况下,推荐使用`manager-web`控制面板作为大多数配置任务的主要操作界面,因为它提供了一种用户友好的方式来管理由`manager-api`持久化并最终由`xiaozhi-server`使用的各项设置。
|
||||||
|
|
||||||
|
---
|
||||||
@@ -0,0 +1,163 @@
|
|||||||
|
# 全模块源码部署自动升级方法
|
||||||
|
|
||||||
|
本教程是方便全模块源码部署的爱好者,如何通过自动命令,自动拉取源码,自动编译,自动启动端口运行。实现最高效率的升级系统。
|
||||||
|
|
||||||
|
本项目的测试平台`https://2662r3426b.vicp.fun`,从开放以来就使用了该方法,效果良好。
|
||||||
|
|
||||||
|
# 开始条件
|
||||||
|
- 你的电脑/服务器是linux操作系统
|
||||||
|
- 你已经跑通了整个流程
|
||||||
|
- 你喜欢跟进最新功能,但是觉得每次手动部署有点麻烦,期待有一个自动更新的方法
|
||||||
|
|
||||||
|
第二个条件必须满足,因为本教程所涉及的某些文件,JDK、Node.js环境、Conda环境等,是需要你跑通整个流程才有的,如果你没有跑通,当我讲到某个文件的时候,你可能就不知道什么意思。
|
||||||
|
|
||||||
|
# 教程效果
|
||||||
|
- 解决国内不能拉取最新项目源码问题
|
||||||
|
- 自动拉取代码编译前端文件
|
||||||
|
- 自动拉取代码编译java文件,自动杀掉8002端口,自动启动8002端口
|
||||||
|
- 自动拉取python代码,自动杀掉8000端口,自动启动8000端口
|
||||||
|
|
||||||
|
# 第一步 选好你的项目目录
|
||||||
|
|
||||||
|
例如,我规划了我的项目目录是,这是一个新建的空白的目录,如果你不想出错,可以和我一样
|
||||||
|
```
|
||||||
|
/home/system/xiaozhi
|
||||||
|
```
|
||||||
|
|
||||||
|
# 第二步 克隆本项目
|
||||||
|
此刻,先要执行第一句话,拉取源码,这句命令适用于国内网络的服务器和电脑,无需翻墙
|
||||||
|
|
||||||
|
```
|
||||||
|
cd /home/system/xiaozhi
|
||||||
|
git clone https://ghproxy.net/https://github.com/xinnan-tech/xiaozhi-esp32-server.git
|
||||||
|
```
|
||||||
|
|
||||||
|
执行完后,你的项目目录会多了一个文件夹`xiaozhi-esp32-server`,这个就是项目的源码
|
||||||
|
|
||||||
|
# 第三步 复制基础的文件
|
||||||
|
|
||||||
|
如果你之前已经跑通了整个流程,对funasr的模型文件`xiaozhi-server/models/SenseVoiceSmall/model.pt`和你的私有配置文件`xiaozhi-server/data/.config.yaml`这两个文件不会陌生。
|
||||||
|
|
||||||
|
此刻你需要把`model.pt`文件复制到新的目录去,你可以这样
|
||||||
|
```
|
||||||
|
cp 你原来的.config.yaml完整路径 /home/system/xiaozhi/xiaozhi-esp32-server/main/xiaozhi-server/data/.config.yaml
|
||||||
|
cp 你原来的model.pt完整路径 /home/system/xiaozhi/xiaozhi-esp32-server/main/xiaozhi-server/models/SenseVoiceSmall/model.pt
|
||||||
|
```
|
||||||
|
|
||||||
|
# 第四步 建立三个自动编译文件
|
||||||
|
|
||||||
|
## 4.1 自动编译mananger-web模块
|
||||||
|
在`/home/system/xiaozhi/`目录下,创建名字为`update_8001.sh`的文件,内容如下
|
||||||
|
|
||||||
|
```
|
||||||
|
cd /home/system/xiaozhi/xiaozhi-esp32-server
|
||||||
|
git fetch --all
|
||||||
|
git reset --hard
|
||||||
|
git pull origin main
|
||||||
|
|
||||||
|
|
||||||
|
cd /home/system/xiaozhi/xiaozhi-esp32-server/main/manager-web
|
||||||
|
npm install
|
||||||
|
npm run build
|
||||||
|
rm -rf /home/system/xiaozhi/manager-web
|
||||||
|
mv /home/system/xiaozhi/xiaozhi-esp32-server/main/manager-web/dist /home/system/xiaozhi/manager-web
|
||||||
|
```
|
||||||
|
|
||||||
|
保存好后执行赋权命令
|
||||||
|
```
|
||||||
|
chmod 777 update_8001.sh
|
||||||
|
```
|
||||||
|
执行完后,继续往下
|
||||||
|
|
||||||
|
## 4.2 自动编译运行manager-api模块
|
||||||
|
在`/home/system/xiaozhi/`目录下,创建名字为`update_8002.sh`的文件,内容如下
|
||||||
|
|
||||||
|
```
|
||||||
|
cd /home/system/xiaozhi/xiaozhi-esp32-server
|
||||||
|
git pull origin main
|
||||||
|
|
||||||
|
|
||||||
|
cd /home/system/xiaozhi/xiaozhi-esp32-server/main/manager-api
|
||||||
|
rm -rf target
|
||||||
|
mvn clean package -Dmaven.test.skip=true
|
||||||
|
cd /home/system/xiaozhi/
|
||||||
|
|
||||||
|
# 查找占用8002端口的进程号
|
||||||
|
PID=$(sudo netstat -tulnp | grep 8002 | awk '{print $7}' | cut -d'/' -f1)
|
||||||
|
|
||||||
|
rm -rf /home/system/xiaozhi/xiaozhi-esp32-api.jar
|
||||||
|
mv /home/system/xiaozhi/xiaozhi-esp32-server/main/manager-api/target/xiaozhi-esp32-api.jar /home/system/xiaozhi/xiaozhi-esp32-api.jar
|
||||||
|
|
||||||
|
# 检查是否找到进程号
|
||||||
|
if [ -z "$PID" ]; then
|
||||||
|
echo "没有找到占用8002端口的进程"
|
||||||
|
else
|
||||||
|
echo "找到占用8002端口的进程,进程号为: $PID"
|
||||||
|
# 杀掉进程
|
||||||
|
kill -9 $PID
|
||||||
|
kill -9 $PID
|
||||||
|
echo "已杀掉进程 $PID"
|
||||||
|
fi
|
||||||
|
|
||||||
|
nohup java -jar xiaozhi-esp32-api.jar --spring.profiles.active=dev &
|
||||||
|
```
|
||||||
|
|
||||||
|
保存好后执行赋权命令
|
||||||
|
```
|
||||||
|
chmod 777 update_8002.sh
|
||||||
|
```
|
||||||
|
执行完后,继续往下
|
||||||
|
|
||||||
|
## 4.3 自动编译运行Python项目
|
||||||
|
在`/home/system/xiaozhi/`目录下,创建名字为`update_8000.sh`的文件,内容如下
|
||||||
|
|
||||||
|
```
|
||||||
|
cd /home/system/xiaozhi/xiaozhi-esp32-server
|
||||||
|
git pull origin main
|
||||||
|
|
||||||
|
# 查找占用8000端口的进程号
|
||||||
|
PID=$(sudo netstat -tulnp | grep 8000 | awk '{print $7}' | cut -d'/' -f1)
|
||||||
|
|
||||||
|
# 检查是否找到进程号
|
||||||
|
if [ -z "$PID" ]; then
|
||||||
|
echo "没有找到占用8000端口的进程"
|
||||||
|
else
|
||||||
|
echo "找到占用8000端口的进程,进程号为: $PID"
|
||||||
|
# 杀掉进程
|
||||||
|
kill -9 $PID
|
||||||
|
kill -9 $PID
|
||||||
|
echo "已杀掉进程 $PID"
|
||||||
|
fi
|
||||||
|
cd main/xiaozhi-server
|
||||||
|
pip install -r requirements.txt
|
||||||
|
nohup python app.py >/dev/null &
|
||||||
|
```
|
||||||
|
|
||||||
|
保存好后执行赋权命令
|
||||||
|
```
|
||||||
|
chmod 777 update_8000.sh
|
||||||
|
```
|
||||||
|
执行完后,继续往下
|
||||||
|
|
||||||
|
# 日常更新
|
||||||
|
|
||||||
|
以上的脚本都建立好后,日常更新,我们只要依次执行以下命令就可以做到自动更新和启动
|
||||||
|
|
||||||
|
```
|
||||||
|
# 进入pyhton环境
|
||||||
|
conda activate xiaozhi-esp32-server
|
||||||
|
cd /home/system/xiaozhi
|
||||||
|
# 更新并启动Java程序
|
||||||
|
./update_8001.sh
|
||||||
|
# 更新web程序
|
||||||
|
./update_8002.sh
|
||||||
|
# 更新并启动python程序
|
||||||
|
./update_8000.sh
|
||||||
|
# 查看Java日志
|
||||||
|
tail -f nohup.out
|
||||||
|
# 查看Python日志
|
||||||
|
tail -f /home/system/xiaozhi/xiaozhi-esp32-server/main/xiaozhi-server/tmp/server.log
|
||||||
|
```
|
||||||
|
|
||||||
|
# 注意事项
|
||||||
|
测试平台`https://2662r3426b.vicp.fun`,是使用nginx做了反向代理。nginx.conf详细配置可以[参考这里](https://github.com/xinnan-tech/xiaozhi-esp32-server/issues/791)
|
||||||
@@ -202,7 +202,15 @@ change_role;get_weather;get_news;play_music;hass_get_state;hass_set_state
|
|||||||
|
|
||||||
#### 2. 配置小智开源服务端MCP配置信息
|
#### 2. 配置小智开源服务端MCP配置信息
|
||||||
|
|
||||||
切换到小智开源服务端`xiaozhi-esp32-server`的`mcp_server_settings.json`文件内,在`"mcpServers"`的括号内添加以下内容:
|
|
||||||
|
进入`data`目录,找到`.mcp_server_settings.json`文件。
|
||||||
|
|
||||||
|
如果你的`data`目录下没有`.mcp_server_settings.json`文件,
|
||||||
|
- 请把在`xiaozhi-server`文件夹根目录的`mcp_server_settings.json`文件复制到`data`目录下,并重命名为`.mcp_server_settings.json`
|
||||||
|
- 或[下载这个文件](https://github.com/xinnan-tech/xiaozhi-esp32-server/blob/main/main/xiaozhi-server/mcp_server_settings.json),下载到`data`目录下,并重命名为`.mcp_server_settings.json`
|
||||||
|
|
||||||
|
|
||||||
|
修改`"mcpServers"`里的这部分的内容:
|
||||||
|
|
||||||
```json
|
```json
|
||||||
"Home Assistant": {
|
"Home Assistant": {
|
||||||
|
|||||||
Binary file not shown.
|
After Width: | Height: | Size: 386 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 518 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 361 KiB |
@@ -0,0 +1,171 @@
|
|||||||
|
# 视觉模型使用指南
|
||||||
|
本教程分为两部分:
|
||||||
|
- 第一部分:单模块运行xiaozhi-server开启视觉模型
|
||||||
|
- 第二部分:全模块运行时,如何开启视觉模型
|
||||||
|
|
||||||
|
开启视觉模型前,你需要准备三件事:
|
||||||
|
- 你需要准备一台带摄像头的设备,而且这台设备已经在虾哥仓库里,实现了调用摄像头功能。例如`立创·实战派ESP32-S3开发板`
|
||||||
|
- 你设备固件的版本升级到1.6.6及以上
|
||||||
|
- 你已经成功跑通基础对话模块
|
||||||
|
|
||||||
|
## 单模块运行xiaozhi-server开启视觉模型
|
||||||
|
|
||||||
|
### 第一步确认网络
|
||||||
|
由于视觉模型会默认启动8003端口。
|
||||||
|
|
||||||
|
如果你是docker运行,请确认一下你的`docker-compose.yml`是否放了`8003`端口,如果没有就更新最新的`docker-compose.yml`文件
|
||||||
|
|
||||||
|
如果你是源码运行,确认防火墙是否放行`8003`端口
|
||||||
|
|
||||||
|
### 第二步选择你的视觉模型
|
||||||
|
打开你的`data/.config.yaml`文件,设置你的`selected_module.VLLM`设置为某个视觉模型。目前我们已经支持`openai`类型接口的视觉模型。`ChatGLMVLLM`就是其中一款兼容`openai`的模型。
|
||||||
|
|
||||||
|
```
|
||||||
|
selected_module:
|
||||||
|
VAD: ..
|
||||||
|
ASR: ..
|
||||||
|
LLM: ..
|
||||||
|
VLLM: ChatGLMVLLM
|
||||||
|
TTS: ..
|
||||||
|
Memory: ..
|
||||||
|
Intent: ..
|
||||||
|
```
|
||||||
|
|
||||||
|
假设我们使用`ChatGLMVLLM`作为视觉模型,那我们需要先登录[智谱AI](https://bigmodel.cn/usercenter/proj-mgmt/apikeys)网站,申请密钥。如果你之前已经申请过了密钥,可以复用这个密钥。
|
||||||
|
|
||||||
|
在你的配置文件中,增加这个配置,如果已经有了这个配置,就设置好你的api_key。
|
||||||
|
|
||||||
|
```
|
||||||
|
VLLM:
|
||||||
|
ChatGLMVLLM:
|
||||||
|
api_key: 你的api_key
|
||||||
|
```
|
||||||
|
|
||||||
|
### 第三步启动xiaozhi-server服务
|
||||||
|
如果你是源码,就输入命令启动
|
||||||
|
```
|
||||||
|
python app.py
|
||||||
|
```
|
||||||
|
如果你是docker运行,就重启容器
|
||||||
|
```
|
||||||
|
docker restart xiaozhi-esp32-server
|
||||||
|
```
|
||||||
|
|
||||||
|
启动后会输出以下内容的日志。
|
||||||
|
|
||||||
|
```
|
||||||
|
2025-06-01 **** - OTA接口是 http://192.168.4.7:8003/xiaozhi/ota/
|
||||||
|
2025-06-01 **** - 视觉分析接口是 http://192.168.4.7:8003/mcp/vision/explain
|
||||||
|
2025-06-01 **** - Websocket地址是 ws://192.168.4.7:8000/xiaozhi/v1/
|
||||||
|
2025-06-01 **** - =======上面的地址是websocket协议地址,请勿用浏览器访问=======
|
||||||
|
2025-06-01 **** - 如想测试websocket请用谷歌浏览器打开test目录下的test_page.html
|
||||||
|
2025-06-01 **** - =============================================================
|
||||||
|
```
|
||||||
|
|
||||||
|
启动后,使用使用浏览器打开日志里`视觉分析接口`连接。看看输出了什么?如果你是linux,没有浏览器,你可以执行这个命令:
|
||||||
|
```
|
||||||
|
curl -i 你的视觉分析接口
|
||||||
|
```
|
||||||
|
|
||||||
|
正常来说会这样显示
|
||||||
|
```
|
||||||
|
MCP Vision 接口运行正常,视觉解释接口地址是:http://xxxx:8003/mcp/vision/explain
|
||||||
|
```
|
||||||
|
|
||||||
|
请注意,如果你是公网部署,或者docker部署,一定要改一下你的`data/.config.yaml`里这个配置
|
||||||
|
```
|
||||||
|
server:
|
||||||
|
vision_explain: http://你的ip或者域名:端口号/mcp/vision/explain
|
||||||
|
```
|
||||||
|
|
||||||
|
为什么呢?因为视觉解释接口需要下发到设备,如果你的地址是局域网地址,或者是docker内部地址,设备是无法访问的。
|
||||||
|
|
||||||
|
假设你的公网地址是`111.111.111.111`,那么`vision_explain`应该这么配
|
||||||
|
|
||||||
|
```
|
||||||
|
server:
|
||||||
|
vision_explain: http://111.111.111.111:8003/mcp/vision/explain
|
||||||
|
```
|
||||||
|
|
||||||
|
如果你的MCP Vision 接口运行正常,且你也试着用浏览器访问正常打开下发的`视觉解释接口地址`,请继续下一步
|
||||||
|
|
||||||
|
### 第四步 设备唤醒开启
|
||||||
|
|
||||||
|
对设备说“请打开摄像头,说你你看到了什么”
|
||||||
|
|
||||||
|
留意xiaozhi-server的日志输出,看看有没有报错。
|
||||||
|
|
||||||
|
|
||||||
|
## 全模块运行时,如何开启视觉模型
|
||||||
|
|
||||||
|
### 第一步 确认网络
|
||||||
|
由于视觉模型会默认启动8003端口。
|
||||||
|
|
||||||
|
如果你是docker运行,请确认一下你的`docker-compose_all.yml`是否映射了`8003`端口,如果没有就更新最新的`docker-compose_all.yml`文件
|
||||||
|
|
||||||
|
如果你是源码运行,确认防火墙是否放行`8003`端口
|
||||||
|
|
||||||
|
### 第二步 确认你配置文件
|
||||||
|
|
||||||
|
打开你的`data/.config.yaml`文件,确认一下你的配置文件的结构,是否和`data/config_from_api.yaml`一样。如果不一样,或缺少某项,请补齐。
|
||||||
|
|
||||||
|
### 第三步 配置视觉模型密钥
|
||||||
|
|
||||||
|
那我们需要先登录[智谱AI](https://bigmodel.cn/usercenter/proj-mgmt/apikeys)网站,申请密钥。如果你之前已经申请过了密钥,可以复用这个密钥。
|
||||||
|
|
||||||
|
登录`智控台`,顶部菜单点击`模型配置`,在左侧栏点击`视觉打语言模型`,找到`VLLM_ChatGLMVLLM`,点击修改按钮,在弹框中,在`API密钥`输入你密钥,点击保存。
|
||||||
|
|
||||||
|
保存成功后,去到你需要测试的智能体哪里,点击`配置角色`,在打开的内容里,查看`视觉大语言模型(VLLM)`是否选择了刚才的视觉模型。点击保存。
|
||||||
|
|
||||||
|
### 第三步 启动xiaozhi-server模块
|
||||||
|
如果你是源码,就输入命令启动
|
||||||
|
```
|
||||||
|
python app.py
|
||||||
|
```
|
||||||
|
如果你是docker运行,就重启容器
|
||||||
|
```
|
||||||
|
docker restart xiaozhi-esp32-server
|
||||||
|
```
|
||||||
|
|
||||||
|
启动后会输出以下内容的日志。
|
||||||
|
|
||||||
|
```
|
||||||
|
2025-06-01 **** - 视觉分析接口是 http://192.168.4.7:8003/mcp/vision/explain
|
||||||
|
2025-06-01 **** - Websocket地址是 ws://192.168.4.7:8000/xiaozhi/v1/
|
||||||
|
2025-06-01 **** - =======上面的地址是websocket协议地址,请勿用浏览器访问=======
|
||||||
|
2025-06-01 **** - 如想测试websocket请用谷歌浏览器打开test目录下的test_page.html
|
||||||
|
2025-06-01 **** - =============================================================
|
||||||
|
```
|
||||||
|
|
||||||
|
启动后,使用使用浏览器打开日志里`视觉分析接口`连接。看看输出了什么?如果你是linux,没有浏览器,你可以执行这个命令:
|
||||||
|
```
|
||||||
|
curl -i 你的视觉分析接口
|
||||||
|
```
|
||||||
|
|
||||||
|
正常来说会这样显示
|
||||||
|
```
|
||||||
|
MCP Vision 接口运行正常,视觉解释接口地址是:http://xxxx:8003/mcp/vision/explain
|
||||||
|
```
|
||||||
|
|
||||||
|
请注意,如果你是公网部署,或者docker部署,一定要改一下你的`data/.config.yaml`里这个配置
|
||||||
|
```
|
||||||
|
server:
|
||||||
|
vision_explain: http://你的ip或者域名:端口号/mcp/vision/explain
|
||||||
|
```
|
||||||
|
|
||||||
|
为什么呢?因为视觉解释接口需要下发到设备,如果你的地址是局域网地址,或者是docker内部地址,设备是无法访问的。
|
||||||
|
|
||||||
|
假设你的公网地址是`111.111.111.111`,那么`vision_explain`应该这么配
|
||||||
|
|
||||||
|
```
|
||||||
|
server:
|
||||||
|
vision_explain: http://111.111.111.111:8003/mcp/vision/explain
|
||||||
|
```
|
||||||
|
|
||||||
|
如果你的MCP Vision 接口运行正常,且你也试着用浏览器访问正常打开下发的`视觉解释接口地址`,请继续下一步
|
||||||
|
|
||||||
|
### 第四步 设备唤醒开启
|
||||||
|
|
||||||
|
对设备说“请打开摄像头,说你你看到了什么”
|
||||||
|
|
||||||
|
留意xiaozhi-server的日志输出,看看有没有报错。
|
||||||
@@ -227,7 +227,7 @@ public interface Constant {
|
|||||||
/**
|
/**
|
||||||
* 版本号
|
* 版本号
|
||||||
*/
|
*/
|
||||||
public static final String VERSION = "0.5.2";
|
public static final String VERSION = "0.5.5";
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* 无效固件URL
|
* 无效固件URL
|
||||||
|
|||||||
@@ -171,6 +171,9 @@ public class AgentController {
|
|||||||
if (dto.getLlmModelId() != null) {
|
if (dto.getLlmModelId() != null) {
|
||||||
existingEntity.setLlmModelId(dto.getLlmModelId());
|
existingEntity.setLlmModelId(dto.getLlmModelId());
|
||||||
}
|
}
|
||||||
|
if (dto.getVllmModelId() != null) {
|
||||||
|
existingEntity.setVllmModelId(dto.getVllmModelId());
|
||||||
|
}
|
||||||
if (dto.getTtsModelId() != null) {
|
if (dto.getTtsModelId() != null) {
|
||||||
existingEntity.setTtsModelId(dto.getTtsModelId());
|
existingEntity.setTtsModelId(dto.getTtsModelId());
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -27,6 +27,9 @@ public class AgentDTO {
|
|||||||
@Schema(description = "大语言模型名称", example = "llm_model_01")
|
@Schema(description = "大语言模型名称", example = "llm_model_01")
|
||||||
private String llmModelName;
|
private String llmModelName;
|
||||||
|
|
||||||
|
@Schema(description = "视觉模型名称", example = "vllm_model_01")
|
||||||
|
private String vllmModelName;
|
||||||
|
|
||||||
@Schema(description = "记忆模型ID", example = "mem_model_01")
|
@Schema(description = "记忆模型ID", example = "mem_model_01")
|
||||||
private String memModelId;
|
private String memModelId;
|
||||||
|
|
||||||
|
|||||||
@@ -30,6 +30,9 @@ public class AgentUpdateDTO implements Serializable {
|
|||||||
@Schema(description = "大语言模型标识", example = "llm_model_02", required = false)
|
@Schema(description = "大语言模型标识", example = "llm_model_02", required = false)
|
||||||
private String llmModelId;
|
private String llmModelId;
|
||||||
|
|
||||||
|
@Schema(description = "VLLM模型标识", example = "vllm_model_02", required = false)
|
||||||
|
private String vllmModelId;
|
||||||
|
|
||||||
@Schema(description = "语音合成模型标识", example = "tts_model_02", required = false)
|
@Schema(description = "语音合成模型标识", example = "tts_model_02", required = false)
|
||||||
private String ttsModelId;
|
private String ttsModelId;
|
||||||
|
|
||||||
|
|||||||
@@ -36,6 +36,9 @@ public class AgentEntity {
|
|||||||
@Schema(description = "大语言模型标识")
|
@Schema(description = "大语言模型标识")
|
||||||
private String llmModelId;
|
private String llmModelId;
|
||||||
|
|
||||||
|
@Schema(description = "VLLM模型标识")
|
||||||
|
private String vllmModelId;
|
||||||
|
|
||||||
@Schema(description = "语音合成模型标识")
|
@Schema(description = "语音合成模型标识")
|
||||||
private String ttsModelId;
|
private String ttsModelId;
|
||||||
|
|
||||||
|
|||||||
@@ -49,6 +49,11 @@ public class AgentTemplateEntity implements Serializable {
|
|||||||
*/
|
*/
|
||||||
private String llmModelId;
|
private String llmModelId;
|
||||||
|
|
||||||
|
/**
|
||||||
|
* VLLM模型标识
|
||||||
|
*/
|
||||||
|
private String vllmModelId;
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* 语音合成模型标识
|
* 语音合成模型标识
|
||||||
*/
|
*/
|
||||||
|
|||||||
+3
@@ -102,6 +102,9 @@ public class AgentServiceImpl extends BaseServiceImpl<AgentDao, AgentEntity> imp
|
|||||||
// 获取 LLM 模型名称
|
// 获取 LLM 模型名称
|
||||||
dto.setLlmModelName(modelConfigService.getModelNameById(agent.getLlmModelId()));
|
dto.setLlmModelName(modelConfigService.getModelNameById(agent.getLlmModelId()));
|
||||||
|
|
||||||
|
// 获取 VLLM 模型名称
|
||||||
|
dto.setVllmModelName(modelConfigService.getModelNameById(agent.getVllmModelId()));
|
||||||
|
|
||||||
// 获取记忆模型名称
|
// 获取记忆模型名称
|
||||||
dto.setMemModelId(agent.getMemModelId());
|
dto.setMemModelId(agent.getMemModelId());
|
||||||
|
|
||||||
|
|||||||
+5
-2
@@ -72,6 +72,7 @@ public class ConfigServiceImpl implements ConfigService {
|
|||||||
null,
|
null,
|
||||||
null,
|
null,
|
||||||
null,
|
null,
|
||||||
|
null,
|
||||||
result,
|
result,
|
||||||
isCache);
|
isCache);
|
||||||
|
|
||||||
@@ -140,6 +141,7 @@ public class ConfigServiceImpl implements ConfigService {
|
|||||||
agent.getVadModelId(),
|
agent.getVadModelId(),
|
||||||
agent.getAsrModelId(),
|
agent.getAsrModelId(),
|
||||||
agent.getLlmModelId(),
|
agent.getLlmModelId(),
|
||||||
|
agent.getVllmModelId(),
|
||||||
agent.getTtsModelId(),
|
agent.getTtsModelId(),
|
||||||
agent.getMemModelId(),
|
agent.getMemModelId(),
|
||||||
agent.getIntentModelId(),
|
agent.getIntentModelId(),
|
||||||
@@ -241,6 +243,7 @@ public class ConfigServiceImpl implements ConfigService {
|
|||||||
String vadModelId,
|
String vadModelId,
|
||||||
String asrModelId,
|
String asrModelId,
|
||||||
String llmModelId,
|
String llmModelId,
|
||||||
|
String vllmModelId,
|
||||||
String ttsModelId,
|
String ttsModelId,
|
||||||
String memModelId,
|
String memModelId,
|
||||||
String intentModelId,
|
String intentModelId,
|
||||||
@@ -248,8 +251,8 @@ public class ConfigServiceImpl implements ConfigService {
|
|||||||
boolean isCache) {
|
boolean isCache) {
|
||||||
Map<String, String> selectedModule = new HashMap<>();
|
Map<String, String> selectedModule = new HashMap<>();
|
||||||
|
|
||||||
String[] modelTypes = { "VAD", "ASR", "TTS", "Memory", "Intent", "LLM" };
|
String[] modelTypes = { "VAD", "ASR", "TTS", "Memory", "Intent", "LLM", "VLLM" };
|
||||||
String[] modelIds = { vadModelId, asrModelId, ttsModelId, memModelId, intentModelId, llmModelId };
|
String[] modelIds = { vadModelId, asrModelId, ttsModelId, memModelId, intentModelId, llmModelId, vllmModelId };
|
||||||
String intentLLMModelId = null;
|
String intentLLMModelId = null;
|
||||||
String memLocalShortLLMModelId = null;
|
String memLocalShortLLMModelId = null;
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,29 @@
|
|||||||
|
-- VLLM模型供应器
|
||||||
|
delete from `ai_model_provider` where id = 'SYSTEM_VLLM_openai';
|
||||||
|
INSERT INTO `ai_model_provider` (`id`, `model_type`, `provider_code`, `name`, `fields`, `sort`, `creator`, `create_date`, `updater`, `update_date`) VALUES
|
||||||
|
('SYSTEM_VLLM_openai', 'VLLM', 'openai', 'OpenAI接口', '[{"key":"base_url","label":"基础URL","type":"string"},{"key":"model_name","label":"模型名称","type":"string"},{"key":"api_key","label":"API密钥","type":"string"}]', 9, 1, NOW(), 1, NOW());
|
||||||
|
|
||||||
|
-- VLLM模型配置
|
||||||
|
delete from `ai_model_config` where id = 'VLLM_ChatGLMVLLM';
|
||||||
|
INSERT INTO `ai_model_config` VALUES ('VLLM_ChatGLMVLLM', 'VLLM', 'ChatGLMVLLM', '智谱视觉AI', 1, 1, '{\"type\": \"openai\", \"model_name\": \"glm-4v-flash\", \"base_url\": \"https://open.bigmodel.cn/api/paas/v4/\", \"api_key\": \"你的api_key\"}', NULL, NULL, 1, NULL, NULL, NULL, NULL);
|
||||||
|
|
||||||
|
-- 更新文档
|
||||||
|
UPDATE `ai_model_config` SET
|
||||||
|
`doc_link` = 'https://bigmodel.cn/usercenter/proj-mgmt/apikeys',
|
||||||
|
`remark` = '智谱视觉AI配置说明:
|
||||||
|
1. 访问 https://bigmodel.cn/usercenter/proj-mgmt/apikeys
|
||||||
|
2. 注册并获取API密钥
|
||||||
|
3. 填入配置文件中' WHERE `id` = 'VLLM_ChatGLMVLLM';
|
||||||
|
|
||||||
|
|
||||||
|
-- 添加参数
|
||||||
|
INSERT INTO `sys_params` (id, param_code, param_value, value_type, param_type, remark) VALUES (113, 'server.http_port', '8003', 'number', 1, 'http服务的端口,用于启动视觉分析接口');
|
||||||
|
INSERT INTO `sys_params` (id, param_code, param_value, value_type, param_type, remark) VALUES (114, 'server.vision_explain', 'null', 'string', 1, '视觉分析接口地址,用于下发到设备,多个用;分隔');
|
||||||
|
|
||||||
|
-- 智能体表增加VLLM模型配置
|
||||||
|
ALTER TABLE `ai_agent`
|
||||||
|
ADD COLUMN `vllm_model_id` varchar(32) NULL DEFAULT 'VLLM_ChatGLMVLLM' COMMENT '视觉模型标识' AFTER `llm_model_id`;
|
||||||
|
|
||||||
|
-- 智能体模版表增加VLLM模型配置
|
||||||
|
ALTER TABLE `ai_agent_template`
|
||||||
|
ADD COLUMN `vllm_model_id` varchar(32) NULL DEFAULT 'VLLM_ChatGLMVLLM' COMMENT '视觉模型标识' AFTER `llm_model_id`;
|
||||||
@@ -0,0 +1,45 @@
|
|||||||
|
-- VLLM模型供应器
|
||||||
|
delete from `ai_model_provider` where id = 'SYSTEM_ASR_DoubaoStreamASR';
|
||||||
|
INSERT INTO `ai_model_provider` (`id`, `model_type`, `provider_code`, `name`, `fields`, `sort`, `creator`, `create_date`, `updater`, `update_date`) VALUES
|
||||||
|
('SYSTEM_ASR_DoubaoStreamASR', 'ASR', 'doubao_stream', '火山引擎语音识别(流式)', '[{"key":"appid","label":"应用ID","type":"string"},{"key":"access_token","label":"访问令牌","type":"string"},{"key":"cluster","label":"集群","type":"string"},{"key":"boosting_table_name","label":"热词文件名称","type":"string"},{"key":"correct_table_name","label":"替换词文件名称","type":"string"},{"key":"output_dir","label":"输出目录","type":"string"}]', 3, 1, NOW(), 1, NOW());
|
||||||
|
|
||||||
|
|
||||||
|
-- VLLM模型配置
|
||||||
|
delete from `ai_model_config` where id = 'ASR_DoubaoStreamASR';
|
||||||
|
INSERT INTO `ai_model_config` VALUES ('ASR_DoubaoStreamASR', 'ASR', 'DoubaoStreamASR', '豆包语音识别(流式)', 0, 1, '{\"type\": \"doubao_stream\", \"appid\": \"\", \"access_token\": \"\", \"cluster\": \"volcengine_input_common\", \"output_dir\": \"tmp/\"}', NULL, NULL, 3, NULL, NULL, NULL, NULL);
|
||||||
|
|
||||||
|
|
||||||
|
-- 更新豆包ASR配置说明
|
||||||
|
UPDATE `ai_model_config` SET
|
||||||
|
`doc_link` = 'https://console.volcengine.com/speech/app',
|
||||||
|
`remark` = '豆包ASR配置说明:
|
||||||
|
1. 豆包ASR和豆包(流式)ASR的区别是:豆包ASR是按次收费,豆包(流式)ASR是按时收费
|
||||||
|
2. 一般来说按次收费的更便宜,但是豆包(流式)ASR使用了大模型技术,效果更好
|
||||||
|
3. 需要在火山引擎控制台创建应用并获取appid和access_token
|
||||||
|
4. 支持中文语音识别
|
||||||
|
5. 需要网络连接
|
||||||
|
6. 输出文件保存在tmp/目录
|
||||||
|
申请步骤:
|
||||||
|
1. 访问 https://console.volcengine.com/speech/app
|
||||||
|
2. 创建新应用
|
||||||
|
3. 获取appid和access_token
|
||||||
|
4. 填入配置文件中
|
||||||
|
如需设置热词,请参考:https://www.volcengine.com/docs/6561/155738
|
||||||
|
' WHERE `id` = 'ASR_DoubaoASR';
|
||||||
|
|
||||||
|
UPDATE `ai_model_config` SET
|
||||||
|
`doc_link` = 'https://console.volcengine.com/speech/app',
|
||||||
|
`remark` = '豆包ASR配置说明:
|
||||||
|
1. 豆包ASR和豆包(流式)ASR的区别是:豆包ASR是按次收费,豆包(流式)ASR是按时收费
|
||||||
|
2. 一般来说按次收费的更便宜,但是豆包(流式)ASR使用了大模型技术,效果更好
|
||||||
|
3. 需要在火山引擎控制台创建应用并获取appid和access_token
|
||||||
|
4. 支持中文语音识别
|
||||||
|
5. 需要网络连接
|
||||||
|
6. 输出文件保存在tmp/目录
|
||||||
|
申请步骤:
|
||||||
|
1. 访问 https://console.volcengine.com/speech/app
|
||||||
|
2. 创建新应用
|
||||||
|
3. 获取appid和access_token
|
||||||
|
4. 填入配置文件中
|
||||||
|
如需设置热词,请参考:https://www.volcengine.com/docs/6561/155738
|
||||||
|
' WHERE `id` = 'ASR_DoubaoStreamASR';
|
||||||
@@ -0,0 +1,14 @@
|
|||||||
|
-- VLLM模型配置
|
||||||
|
delete from `ai_model_config` where id = 'VLLM_QwenVLVLLM';
|
||||||
|
INSERT INTO `ai_model_config` VALUES ('VLLM_QwenVLVLLM', 'VLLM', 'QwenVLVLLM', '千问视觉模型', 0, 1, '{\"type\": \"openai\", \"model_name\": \"qwen2.5-vl-3b-instruct\", \"base_url\": \"https://dashscope.aliyuncs.com/compatible-mode/v1\", \"api_key\": \"你的api_key\"}', NULL, NULL, 2, NULL, NULL, NULL, NULL);
|
||||||
|
|
||||||
|
-- 更新文档
|
||||||
|
UPDATE `ai_model_config` SET
|
||||||
|
`doc_link` = 'https://bailian.console.aliyun.com/?tab=api#/api/?type=model&url=https%3A%2F%2Fhelp.aliyun.com%2Fdocument_detail%2F2845564.html&renderType=iframe',
|
||||||
|
`remark` = '千问视觉模型配置说明:
|
||||||
|
1. 访问 https://bailian.console.aliyun.com/?tab=model#/api-key
|
||||||
|
2. 注册并获取API密钥
|
||||||
|
3. 填入配置文件中' WHERE `id` = 'VLLM_QwenVLVLLM';
|
||||||
|
|
||||||
|
-- 删除参数,这两个参数已挪至python配置文件
|
||||||
|
delete from `sys_params` where id in (113,114);
|
||||||
@@ -0,0 +1,20 @@
|
|||||||
|
-- 增加LinkeraiTTS供应器和模型配置
|
||||||
|
delete from `ai_model_provider` where id = 'SYSTEM_TTS_LinkeraiTTS';
|
||||||
|
INSERT INTO `ai_model_provider` (`id`, `model_type`, `provider_code`, `name`, `fields`, `sort`, `creator`, `create_date`, `updater`, `update_date`) VALUES
|
||||||
|
('SYSTEM_TTS_LinkeraiTTS', 'TTS', 'linkerai', 'Linkerai语音合成', '[{"key":"api_url","label":"API地址","type":"string"},{"key":"audio_format","label":"音频格式","type":"string"},{"key":"access_token","label":"访问令牌","type":"string"},{"key":"voice","label":"默认音色","type":"string"}]', 14, 1, NOW(), 1, NOW());
|
||||||
|
|
||||||
|
delete from `ai_model_config` where id = 'TTS_LinkeraiTTS';
|
||||||
|
INSERT INTO `ai_model_config` VALUES ('TTS_LinkeraiTTS', 'TTS', 'LinkeraiTTS', 'Linkerai语音合成', 0, 1, '{\"type\": \"linkerai\", \"api_url\": \"https://tts.linkerai.cn/tts\", \"audio_format\": \"pcm\", \"access_token\": \"U4YdYXVfpwWnk2t5Gp822zWPCuORyeJL\", \"voice\": \"OUeAo1mhq6IBExi\"}', NULL, NULL, 17, NULL, NULL, NULL, NULL);
|
||||||
|
|
||||||
|
-- LinkeraiTTS模型配置说明文档
|
||||||
|
UPDATE `ai_model_config` SET
|
||||||
|
`doc_link` = 'https://tts.linkerai.cn/docs',
|
||||||
|
`remark` = 'Linkerai语音合成服务配置说明:
|
||||||
|
1. 访问 https://linkerai.cn 注册并获取访问令牌
|
||||||
|
2. 默认的access_token供测试使用,请勿用于商业用途
|
||||||
|
3. 支持声音克隆功能,可自行上传音频,填入voice参数
|
||||||
|
4. 如果voice参数为空,将使用默认声音' WHERE `id` = 'TTS_LinkeraiTTS';
|
||||||
|
|
||||||
|
|
||||||
|
delete from `ai_tts_voice` where tts_model_id = 'TTS_LinkeraiTTS';
|
||||||
|
INSERT INTO `ai_tts_voice` VALUES ('TTS_LinkeraiTTS_0001', 'TTS_LinkeraiTTS', '芷若', 'OUeAo1mhq6IBExi', '中文', NULL, NULL, 1, NULL, NULL, NULL, NULL);
|
||||||
@@ -0,0 +1,3 @@
|
|||||||
|
-- 智控台开启唤醒词加速
|
||||||
|
update `sys_params` set param_value = '你好小智;你好小志;小爱同学;你好小鑫;你好小新;小美同学;小龙小龙;喵喵同学;小滨小滨;小冰小冰;嘿你好呀' where param_code = 'wakeup_words';
|
||||||
|
update `sys_params` set param_value = 'true' where param_code = 'enable_wakeup_words_response_cache';
|
||||||
@@ -169,4 +169,39 @@ databaseChangeLog:
|
|||||||
changes:
|
changes:
|
||||||
- sqlFile:
|
- sqlFile:
|
||||||
encoding: utf8
|
encoding: utf8
|
||||||
path: classpath:db/changelog/202505271414.sql
|
path: classpath:db/changelog/202505271414.sql
|
||||||
|
- changeSet:
|
||||||
|
id: 202506010920
|
||||||
|
author: hrz
|
||||||
|
changes:
|
||||||
|
- sqlFile:
|
||||||
|
encoding: utf8
|
||||||
|
path: classpath:db/changelog/202506010920.sql
|
||||||
|
- changeSet:
|
||||||
|
id: 202506031639
|
||||||
|
author: hrz
|
||||||
|
changes:
|
||||||
|
- sqlFile:
|
||||||
|
encoding: utf8
|
||||||
|
path: classpath:db/changelog/202506031639.sql
|
||||||
|
- changeSet:
|
||||||
|
id: 202506032232
|
||||||
|
author: hrz
|
||||||
|
changes:
|
||||||
|
- sqlFile:
|
||||||
|
encoding: utf8
|
||||||
|
path: classpath:db/changelog/202506032232.sql
|
||||||
|
- changeSet:
|
||||||
|
id: 202506051538
|
||||||
|
author: hrz
|
||||||
|
changes:
|
||||||
|
- sqlFile:
|
||||||
|
encoding: utf8
|
||||||
|
path: classpath:db/changelog/202506051538.sql
|
||||||
|
- changeSet:
|
||||||
|
id: 202506080955
|
||||||
|
author: hrz
|
||||||
|
changes:
|
||||||
|
- sqlFile:
|
||||||
|
encoding: utf8
|
||||||
|
path: classpath:db/changelog/202506080955.sql
|
||||||
@@ -14,10 +14,10 @@
|
|||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
<div class="device-name">
|
<div class="device-name">
|
||||||
设备型号:{{ device.ttsModelName }}
|
语言模型:{{ device.llmModelName }}
|
||||||
</div>
|
</div>
|
||||||
<div class="device-name">
|
<div class="device-name">
|
||||||
音色模型:{{ device.ttsVoiceName }}
|
音色模型:{{ device.ttsModelName }} ({{ device.ttsVoiceName }})
|
||||||
</div>
|
</div>
|
||||||
<div style="display: flex;gap: 10px;align-items: center;">
|
<div style="display: flex;gap: 10px;align-items: center;">
|
||||||
<div class="settings-btn" @click="handleConfigure">
|
<div class="settings-btn" @click="handleConfigure">
|
||||||
|
|||||||
@@ -30,6 +30,9 @@
|
|||||||
<el-menu-item index="llm">
|
<el-menu-item index="llm">
|
||||||
<span class="menu-text">大语言模型</span>
|
<span class="menu-text">大语言模型</span>
|
||||||
</el-menu-item>
|
</el-menu-item>
|
||||||
|
<el-menu-item index="vllm">
|
||||||
|
<span class="menu-text">视觉大模型</span>
|
||||||
|
</el-menu-item>
|
||||||
<el-menu-item index="intent">
|
<el-menu-item index="intent">
|
||||||
<span class="menu-text">意图识别</span>
|
<span class="menu-text">意图识别</span>
|
||||||
</el-menu-item>
|
</el-menu-item>
|
||||||
@@ -173,6 +176,7 @@ export default {
|
|||||||
vad: '语言活动检测模型(VAD)',
|
vad: '语言活动检测模型(VAD)',
|
||||||
asr: '语音识别模型(ASR)',
|
asr: '语音识别模型(ASR)',
|
||||||
llm: '大语言模型(LLM)',
|
llm: '大语言模型(LLM)',
|
||||||
|
vllm: '视觉大模型(VLLM)',
|
||||||
intent: '意图识别模型(Intent)',
|
intent: '意图识别模型(Intent)',
|
||||||
tts: '语音合成模型(TTS)',
|
tts: '语音合成模型(TTS)',
|
||||||
memory: '记忆模型(Memory)'
|
memory: '记忆模型(Memory)'
|
||||||
|
|||||||
@@ -64,7 +64,27 @@
|
|||||||
</el-form-item>
|
</el-form-item>
|
||||||
</div>
|
</div>
|
||||||
<div class="form-column">
|
<div class="form-column">
|
||||||
<el-form-item v-for="(model, index) in models" :key="`model-${index}`" :label="model.label"
|
<div class="model-row">
|
||||||
|
<el-form-item label="语音活动检测(VAD)" class="model-item">
|
||||||
|
<div class="model-select-wrapper">
|
||||||
|
<el-select v-model="form.model.vadModelId" filterable placeholder="请选择" class="form-select"
|
||||||
|
@change="handleModelChange('VAD', $event)">
|
||||||
|
<el-option v-for="(item, optionIndex) in modelOptions['VAD']"
|
||||||
|
:key="`option-vad-${optionIndex}`" :label="item.label" :value="item.value" />
|
||||||
|
</el-select>
|
||||||
|
</div>
|
||||||
|
</el-form-item>
|
||||||
|
<el-form-item label="语音识别(ASR)" class="model-item">
|
||||||
|
<div class="model-select-wrapper">
|
||||||
|
<el-select v-model="form.model.asrModelId" filterable placeholder="请选择" class="form-select"
|
||||||
|
@change="handleModelChange('ASR', $event)">
|
||||||
|
<el-option v-for="(item, optionIndex) in modelOptions['ASR']"
|
||||||
|
:key="`option-asr-${optionIndex}`" :label="item.label" :value="item.value" />
|
||||||
|
</el-select>
|
||||||
|
</div>
|
||||||
|
</el-form-item>
|
||||||
|
</div>
|
||||||
|
<el-form-item v-for="(model, index) in models.slice(2)" :key="`model-${index}`" :label="model.label"
|
||||||
class="model-item">
|
class="model-item">
|
||||||
<div class="model-select-wrapper">
|
<div class="model-select-wrapper">
|
||||||
<el-select v-model="form.model[model.key]" filterable placeholder="请选择" class="form-select"
|
<el-select v-model="form.model[model.key]" filterable placeholder="请选择" class="form-select"
|
||||||
@@ -148,6 +168,7 @@ export default {
|
|||||||
vadModelId: "",
|
vadModelId: "",
|
||||||
asrModelId: "",
|
asrModelId: "",
|
||||||
llmModelId: "",
|
llmModelId: "",
|
||||||
|
vllmModelId: "",
|
||||||
memModelId: "",
|
memModelId: "",
|
||||||
intentModelId: "",
|
intentModelId: "",
|
||||||
}
|
}
|
||||||
@@ -156,6 +177,7 @@ export default {
|
|||||||
{ label: '语音活动检测(VAD)', key: 'vadModelId', type: 'VAD' },
|
{ label: '语音活动检测(VAD)', key: 'vadModelId', type: 'VAD' },
|
||||||
{ label: '语音识别(ASR)', key: 'asrModelId', type: 'ASR' },
|
{ label: '语音识别(ASR)', key: 'asrModelId', type: 'ASR' },
|
||||||
{ label: '大语言模型(LLM)', key: 'llmModelId', type: 'LLM' },
|
{ label: '大语言模型(LLM)', key: 'llmModelId', type: 'LLM' },
|
||||||
|
{ label: '视觉大模型(VLLM)', key: 'vllmModelId', type: 'VLLM' },
|
||||||
{ label: '意图识别(Intent)', key: 'intentModelId', type: 'Intent' },
|
{ label: '意图识别(Intent)', key: 'intentModelId', type: 'Intent' },
|
||||||
{ label: '记忆(Memory)', key: 'memModelId', type: 'Memory' },
|
{ label: '记忆(Memory)', key: 'memModelId', type: 'Memory' },
|
||||||
{ label: '语音合成(TTS)', key: 'ttsModelId', type: 'TTS' },
|
{ label: '语音合成(TTS)', key: 'ttsModelId', type: 'TTS' },
|
||||||
@@ -189,6 +211,7 @@ export default {
|
|||||||
asrModelId: this.form.model.asrModelId,
|
asrModelId: this.form.model.asrModelId,
|
||||||
vadModelId: this.form.model.vadModelId,
|
vadModelId: this.form.model.vadModelId,
|
||||||
llmModelId: this.form.model.llmModelId,
|
llmModelId: this.form.model.llmModelId,
|
||||||
|
vllmModelId: this.form.model.vllmModelId,
|
||||||
ttsModelId: this.form.model.ttsModelId,
|
ttsModelId: this.form.model.ttsModelId,
|
||||||
ttsVoiceId: this.form.ttsVoiceId,
|
ttsVoiceId: this.form.ttsVoiceId,
|
||||||
chatHistoryConf: this.form.chatHistoryConf,
|
chatHistoryConf: this.form.chatHistoryConf,
|
||||||
@@ -236,6 +259,7 @@ export default {
|
|||||||
vadModelId: "",
|
vadModelId: "",
|
||||||
asrModelId: "",
|
asrModelId: "",
|
||||||
llmModelId: "",
|
llmModelId: "",
|
||||||
|
vllmModelId: "",
|
||||||
memModelId: "",
|
memModelId: "",
|
||||||
intentModelId: "",
|
intentModelId: "",
|
||||||
}
|
}
|
||||||
@@ -289,6 +313,7 @@ export default {
|
|||||||
vadModelId: templateData.vadModelId || this.form.model.vadModelId,
|
vadModelId: templateData.vadModelId || this.form.model.vadModelId,
|
||||||
asrModelId: templateData.asrModelId || this.form.model.asrModelId,
|
asrModelId: templateData.asrModelId || this.form.model.asrModelId,
|
||||||
llmModelId: templateData.llmModelId || this.form.model.llmModelId,
|
llmModelId: templateData.llmModelId || this.form.model.llmModelId,
|
||||||
|
vllmModelId: templateData.vllmModelId || this.form.model.vllmModelId,
|
||||||
memModelId: templateData.memModelId || this.form.model.memModelId,
|
memModelId: templateData.memModelId || this.form.model.memModelId,
|
||||||
intentModelId: templateData.intentModelId || this.form.model.intentModelId
|
intentModelId: templateData.intentModelId || this.form.model.intentModelId
|
||||||
}
|
}
|
||||||
@@ -305,6 +330,7 @@ export default {
|
|||||||
vadModelId: data.data.vadModelId,
|
vadModelId: data.data.vadModelId,
|
||||||
asrModelId: data.data.asrModelId,
|
asrModelId: data.data.asrModelId,
|
||||||
llmModelId: data.data.llmModelId,
|
llmModelId: data.data.llmModelId,
|
||||||
|
vllmModelId: data.data.vllmModelId,
|
||||||
memModelId: data.data.memModelId,
|
memModelId: data.data.memModelId,
|
||||||
intentModelId: data.data.intentModelId
|
intentModelId: data.data.intentModelId
|
||||||
}
|
}
|
||||||
@@ -587,6 +613,25 @@ export default {
|
|||||||
width: 100%;
|
width: 100%;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
.model-row {
|
||||||
|
display: flex;
|
||||||
|
gap: 20px;
|
||||||
|
margin-bottom: 6px;
|
||||||
|
}
|
||||||
|
|
||||||
|
.model-row .model-item {
|
||||||
|
flex: 1;
|
||||||
|
margin-bottom: 0;
|
||||||
|
}
|
||||||
|
|
||||||
|
.model-row .el-form-item__label {
|
||||||
|
font-size: 12px !important;
|
||||||
|
color: #3d4566 !important;
|
||||||
|
font-weight: 400;
|
||||||
|
line-height: 22px;
|
||||||
|
padding-bottom: 2px;
|
||||||
|
}
|
||||||
|
|
||||||
.function-icons {
|
.function-icons {
|
||||||
display: flex;
|
display: flex;
|
||||||
align-items: center;
|
align-items: center;
|
||||||
|
|||||||
@@ -1,11 +1,12 @@
|
|||||||
import sys
|
import sys
|
||||||
|
import uuid
|
||||||
import signal
|
import signal
|
||||||
import asyncio
|
import asyncio
|
||||||
from aioconsole import ainput
|
from aioconsole import ainput
|
||||||
from config.settings import load_config
|
from config.settings import load_config
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from core.utils.util import get_local_ip
|
from core.utils.util import get_local_ip
|
||||||
from core.ota_server import SimpleOtaServer
|
from core.http_server import SimpleHttpServer
|
||||||
from core.websocket_server import WebSocketServer
|
from core.websocket_server import WebSocketServer
|
||||||
from core.utils.util import check_ffmpeg_installed
|
from core.utils.util import check_ffmpeg_installed
|
||||||
|
|
||||||
@@ -45,25 +46,37 @@ async def main():
|
|||||||
check_ffmpeg_installed()
|
check_ffmpeg_installed()
|
||||||
config = load_config()
|
config = load_config()
|
||||||
|
|
||||||
|
# 默认使用manager-api的secret作为auth_key
|
||||||
|
# 如果secret为空,则生成随机密钥
|
||||||
|
# auth_key用于jwt认证,比如视觉分析接口的jwt认证
|
||||||
|
auth_key = config.get("manager-api", {}).get("secret", "")
|
||||||
|
if not auth_key or len(auth_key) == 0 or "你" in auth_key:
|
||||||
|
auth_key = str(uuid.uuid4().hex)
|
||||||
|
config["server"]["auth_key"] = auth_key
|
||||||
|
|
||||||
# 添加 stdin 监控任务
|
# 添加 stdin 监控任务
|
||||||
stdin_task = asyncio.create_task(monitor_stdin())
|
stdin_task = asyncio.create_task(monitor_stdin())
|
||||||
|
|
||||||
# 启动 WebSocket 服务器
|
# 启动 WebSocket 服务器
|
||||||
ws_server = WebSocketServer(config)
|
ws_server = WebSocketServer(config)
|
||||||
ws_task = asyncio.create_task(ws_server.start())
|
ws_task = asyncio.create_task(ws_server.start())
|
||||||
ota_task = None
|
# 启动 Simple http 服务器
|
||||||
|
ota_server = SimpleHttpServer(config)
|
||||||
|
ota_task = asyncio.create_task(ota_server.start())
|
||||||
|
|
||||||
read_config_from_api = config.get("read_config_from_api", False)
|
read_config_from_api = config.get("read_config_from_api", False)
|
||||||
|
port = int(config["server"].get("http_port", 8003))
|
||||||
if not read_config_from_api:
|
if not read_config_from_api:
|
||||||
# 启动 Simple OTA 服务器
|
|
||||||
ota_server = SimpleOtaServer(config)
|
|
||||||
ota_task = asyncio.create_task(ota_server.start())
|
|
||||||
|
|
||||||
logger.bind(tag=TAG).info(
|
logger.bind(tag=TAG).info(
|
||||||
"OTA接口是\t\thttp://{}:{}/xiaozhi/ota/",
|
"OTA接口是\t\thttp://{}:{}/xiaozhi/ota/",
|
||||||
get_local_ip(),
|
get_local_ip(),
|
||||||
config["server"]["ota_port"],
|
port,
|
||||||
)
|
)
|
||||||
|
logger.bind(tag=TAG).info(
|
||||||
|
"视觉分析接口是\thttp://{}:{}/mcp/vision/explain",
|
||||||
|
get_local_ip(),
|
||||||
|
port,
|
||||||
|
)
|
||||||
|
|
||||||
# 获取WebSocket配置,使用安全的默认值
|
# 获取WebSocket配置,使用安全的默认值
|
||||||
websocket_port = 8000
|
websocket_port = 8000
|
||||||
|
|||||||
@@ -2,6 +2,7 @@
|
|||||||
# 然后你想修改覆盖修改什么配置,就修改【.config.yaml】文件,而不是修改【config.yaml】文件
|
# 然后你想修改覆盖修改什么配置,就修改【.config.yaml】文件,而不是修改【config.yaml】文件
|
||||||
# 系统会优先读取【data/.config.yaml】文件的配置,如果【.config.yaml】文件里的配置不存在,系统会自动去读取【config.yaml】文件的配置。
|
# 系统会优先读取【data/.config.yaml】文件的配置,如果【.config.yaml】文件里的配置不存在,系统会自动去读取【config.yaml】文件的配置。
|
||||||
# 这样做,可以最简化配置,保护您的密钥安全。
|
# 这样做,可以最简化配置,保护您的密钥安全。
|
||||||
|
# 如果你使用了智控台,那么以下所有配置,都不会生效,请在智控台中修改配置
|
||||||
|
|
||||||
# #####################################################################################
|
# #####################################################################################
|
||||||
# #############################以下是服务器基本运行配置####################################
|
# #############################以下是服务器基本运行配置####################################
|
||||||
@@ -9,13 +10,21 @@ server:
|
|||||||
# 服务器监听地址和端口(Server listening address and port)
|
# 服务器监听地址和端口(Server listening address and port)
|
||||||
ip: 0.0.0.0
|
ip: 0.0.0.0
|
||||||
port: 8000
|
port: 8000
|
||||||
# OTA接口的端口号
|
# http服务的端口,用于简单OTA接口(单服务部署),以及视觉分析接口
|
||||||
ota_port: 8002
|
http_port: 8003
|
||||||
# 这个websocket配置是指ota接口向设备发送的websocket地址
|
# 这个websocket配置是指ota接口向设备发送的websocket地址
|
||||||
# 如果按默认的写法,ota接口会自动生成websocket地址。这个地址你可以直接用浏览器访问ota接口确认一下
|
# 如果按默认的写法,ota接口会自动生成websocket地址,并输出在启动日志里,这个地址你可以直接用浏览器访问ota接口确认一下
|
||||||
# 当你使用docker部署或使用公网部署(使用ssl、域名)时,不一定准确
|
# 当你使用docker部署或使用公网部署(使用ssl、域名)时,不一定准确
|
||||||
# 所以如果你使用docker部署或使用公网部署时,请设置正确的websocket地址
|
# 所以如果你使用docker部署时,将websocket设置成局域网地址
|
||||||
|
# 如果你使用公网部署时,将vwebsocket设置成公网地址
|
||||||
websocket: ws://你的ip或者域名:端口号/xiaozhi/v1/
|
websocket: ws://你的ip或者域名:端口号/xiaozhi/v1/
|
||||||
|
# 视觉分析接口地址
|
||||||
|
# 向设备发送的视觉分析的接口地址
|
||||||
|
# 如果按下面默认的写法,系统会自动生成视觉识别地址,并输出在启动日志里,这个地址你可以直接用浏览器访问确认一下
|
||||||
|
# 当你使用docker部署或使用公网部署(使用ssl、域名)时,不一定准确
|
||||||
|
# 所以如果你使用docker部署时,将vision_explain设置成局域网地址
|
||||||
|
# 如果你使用公网部署时,将vision_explain设置成公网地址
|
||||||
|
vision_explain: http://你的ip或者域名:端口号/mcp/vision/explain
|
||||||
# OTA返回信息时区偏移量
|
# OTA返回信息时区偏移量
|
||||||
timezone_offset: +8
|
timezone_offset: +8
|
||||||
# 认证配置
|
# 认证配置
|
||||||
@@ -160,6 +169,8 @@ selected_module:
|
|||||||
ASR: FunASR
|
ASR: FunASR
|
||||||
# 将根据配置名称对应的type调用实际的LLM适配器
|
# 将根据配置名称对应的type调用实际的LLM适配器
|
||||||
LLM: ChatGLMLLM
|
LLM: ChatGLMLLM
|
||||||
|
# 视觉语言大模型
|
||||||
|
VLLM: ChatGLMVLLM
|
||||||
# TTS将根据配置名称对应的type调用实际的TTS适配器
|
# TTS将根据配置名称对应的type调用实际的TTS适配器
|
||||||
TTS: EdgeTTS
|
TTS: EdgeTTS
|
||||||
# 记忆模块,默认不开启记忆;如果想使用超长记忆,推荐使用mem0ai;如果注重隐私,请使用本地的mem_local_short
|
# 记忆模块,默认不开启记忆;如果想使用超长记忆,推荐使用mem0ai;如果注重隐私,请使用本地的mem_local_short
|
||||||
@@ -253,6 +264,8 @@ ASR:
|
|||||||
DoubaoASR:
|
DoubaoASR:
|
||||||
# 可以在这里申请相关Key等信息
|
# 可以在这里申请相关Key等信息
|
||||||
# https://console.volcengine.com/speech/app
|
# https://console.volcengine.com/speech/app
|
||||||
|
# DoubaoASR和DoubaoStreamASR的区别是:DoubaoASR是按次收费,DoubaoStreamASR是按时收费
|
||||||
|
# 一般来说按次收费的更便宜,但是DoubaoStreamASR使用了大模型技术,效果更好
|
||||||
type: doubao
|
type: doubao
|
||||||
appid: 你的火山引擎语音合成服务appid
|
appid: 你的火山引擎语音合成服务appid
|
||||||
access_token: 你的火山引擎语音合成服务access_token
|
access_token: 你的火山引擎语音合成服务access_token
|
||||||
@@ -261,6 +274,20 @@ ASR:
|
|||||||
boosting_table_name: (选填)你的热词文件名称
|
boosting_table_name: (选填)你的热词文件名称
|
||||||
correct_table_name: (选填)你的替换词文件名称
|
correct_table_name: (选填)你的替换词文件名称
|
||||||
output_dir: tmp/
|
output_dir: tmp/
|
||||||
|
DoubaoStreamASR:
|
||||||
|
# 可以在这里申请相关Key等信息
|
||||||
|
# https://console.volcengine.com/speech/app
|
||||||
|
# DoubaoASR和DoubaoStreamASR的区别是:DoubaoASR是按次收费,DoubaoStreamASR是按时收费
|
||||||
|
# 开通地址https://console.volcengine.com/speech/service/10011
|
||||||
|
# 一般来说按次收费的更便宜,但是DoubaoStreamASR使用了大模型技术,效果更好
|
||||||
|
type: doubao_stream
|
||||||
|
appid: 你的火山引擎语音合成服务appid
|
||||||
|
access_token: 你的火山引擎语音合成服务access_token
|
||||||
|
cluster: volcengine_input_common
|
||||||
|
# 热词、替换词使用流程:https://www.volcengine.com/docs/6561/155738
|
||||||
|
boosting_table_name: (选填)你的热词文件名称
|
||||||
|
correct_table_name: (选填)你的替换词文件名称
|
||||||
|
output_dir: tmp/
|
||||||
TencentASR:
|
TencentASR:
|
||||||
# token申请地址:https://console.cloud.tencent.com/cam/capi
|
# token申请地址:https://console.cloud.tencent.com/cam/capi
|
||||||
# 免费领取资源:https://console.cloud.tencent.com/asr/resourcebundle
|
# 免费领取资源:https://console.cloud.tencent.com/asr/resourcebundle
|
||||||
@@ -297,7 +324,7 @@ VAD:
|
|||||||
type: silero
|
type: silero
|
||||||
threshold: 0.5
|
threshold: 0.5
|
||||||
model_dir: models/snakers4_silero-vad
|
model_dir: models/snakers4_silero-vad
|
||||||
min_silence_duration_ms: 700 # 如果说话停顿比较长,可以把这个值设置大一些
|
min_silence_duration_ms: 200 # 如果说话停顿比较长,可以把这个值设置大一些
|
||||||
|
|
||||||
LLM:
|
LLM:
|
||||||
# 所有openai类型均可以修改超参,以AliLLM为例
|
# 所有openai类型均可以修改超参,以AliLLM为例
|
||||||
@@ -432,6 +459,21 @@ LLM:
|
|||||||
# Xinference服务地址和模型名称
|
# Xinference服务地址和模型名称
|
||||||
model_name: qwen2.5:3b-AWQ # 使用的小模型名称,用于意图识别
|
model_name: qwen2.5:3b-AWQ # 使用的小模型名称,用于意图识别
|
||||||
base_url: http://localhost:9997 # Xinference服务地址
|
base_url: http://localhost:9997 # Xinference服务地址
|
||||||
|
# VLLM配置(视觉语言大模型)
|
||||||
|
VLLM:
|
||||||
|
ChatGLMVLLM:
|
||||||
|
type: openai
|
||||||
|
# glm-4v-flash是智谱免费AI的视觉模型,需要先在智谱AI平台创建API密钥并获取api_key
|
||||||
|
# 可在这里找到你的api key https://bigmodel.cn/usercenter/proj-mgmt/apikeys
|
||||||
|
model_name: glm-4v-flash # 智谱AI的视觉模型
|
||||||
|
url: https://open.bigmodel.cn/api/paas/v4/
|
||||||
|
api_key: 你的api_key
|
||||||
|
QwenVLVLLM:
|
||||||
|
type: openai
|
||||||
|
model_name: qwen2.5-vl-3b-instruct
|
||||||
|
url: https://dashscope.aliyuncs.com/compatible-mode/v1
|
||||||
|
# 可在这里找到你的api key https://bailian.console.aliyun.com/?apiKey=1#/api-key
|
||||||
|
api_key: 你的api_key
|
||||||
TTS:
|
TTS:
|
||||||
# 当前支持的type为edge、doubao,可自行适配
|
# 当前支持的type为edge、doubao,可自行适配
|
||||||
EdgeTTS:
|
EdgeTTS:
|
||||||
@@ -713,4 +755,15 @@ TTS:
|
|||||||
headers: # 自定义请求头
|
headers: # 自定义请求头
|
||||||
# Authorization: Bearer xxxx
|
# Authorization: Bearer xxxx
|
||||||
format: mp3 # 接口返回的音频格式
|
format: mp3 # 接口返回的音频格式
|
||||||
|
output_dir: tmp/
|
||||||
|
LinkeraiTTS:
|
||||||
|
type: linkerai
|
||||||
|
api_url: https://tts.linkerai.cn/tts
|
||||||
|
audio_format: "pcm"
|
||||||
|
# 默认的access_token供大家测试时免费使用的,此access_token请勿用于商业用途
|
||||||
|
# 如果效果不错,可自行申请token,申请地址:https://linkerai.cn
|
||||||
|
# 各参数意义见开发文档:https://tts.linkerai.cn/docs
|
||||||
|
# 支持声音克隆,可自行上传音频,填入voice参数,voice参数为空时,使用默认声音
|
||||||
|
access_token: "U4YdYXVfpwWnk2t5Gp822zWPCuORyeJL"
|
||||||
|
voice: "OUeAo1mhq6IBExi"
|
||||||
output_dir: tmp/
|
output_dir: tmp/
|
||||||
@@ -59,10 +59,14 @@ def get_config_from_api(config):
|
|||||||
"url": config["manager-api"].get("url", ""),
|
"url": config["manager-api"].get("url", ""),
|
||||||
"secret": config["manager-api"].get("secret", ""),
|
"secret": config["manager-api"].get("secret", ""),
|
||||||
}
|
}
|
||||||
|
# server的配置以本地为准
|
||||||
if config.get("server"):
|
if config.get("server"):
|
||||||
config_data["server"] = {
|
config_data["server"] = {
|
||||||
"ip": config["server"].get("ip", ""),
|
"ip": config["server"].get("ip", ""),
|
||||||
"port": config["server"].get("port", ""),
|
"port": config["server"].get("port", ""),
|
||||||
|
"http_port": config["server"].get("http_port", ""),
|
||||||
|
"vision_explain": config["server"].get("vision_explain", ""),
|
||||||
|
"auth_key": config["server"].get("auth_key", ""),
|
||||||
}
|
}
|
||||||
return config_data
|
return config_data
|
||||||
|
|
||||||
|
|||||||
@@ -3,8 +3,9 @@ import sys
|
|||||||
from loguru import logger
|
from loguru import logger
|
||||||
from config.config_loader import load_config
|
from config.config_loader import load_config
|
||||||
from config.settings import check_config_file
|
from config.settings import check_config_file
|
||||||
|
from datetime import datetime
|
||||||
|
|
||||||
SERVER_VERSION = "0.5.2"
|
SERVER_VERSION = "0.5.5"
|
||||||
_logger_initialized = False
|
_logger_initialized = False
|
||||||
|
|
||||||
|
|
||||||
@@ -59,7 +60,7 @@ def setup_logging():
|
|||||||
)
|
)
|
||||||
log_format_file = log_config.get(
|
log_format_file = log_config.get(
|
||||||
"log_format_file",
|
"log_format_file",
|
||||||
"{time:YYYY-MM-DD HH:mm:ss} - {version_{extra[selected_module]}} - {name} - {level} - {extra[tag]} - {message}",
|
"{time:YYYY-MM-DD HH:mm:ss} - {version}_{extra[selected_module]} - {name} - {level} - {extra[tag]} - {message}",
|
||||||
)
|
)
|
||||||
selected_module_str = logger._core.extra["selected_module"]
|
selected_module_str = logger._core.extra["selected_module"]
|
||||||
|
|
||||||
@@ -84,12 +85,23 @@ def setup_logging():
|
|||||||
# 输出到控制台
|
# 输出到控制台
|
||||||
logger.add(sys.stdout, format=log_format, level=log_level, filter=formatter)
|
logger.add(sys.stdout, format=log_format, level=log_level, filter=formatter)
|
||||||
|
|
||||||
# 输出到文件
|
# 输出到文件 - 统一目录,按大小轮转
|
||||||
|
# 日志文件完整路径
|
||||||
|
log_file_path = os.path.join(log_dir, log_file)
|
||||||
|
|
||||||
|
# 添加日志处理器
|
||||||
logger.add(
|
logger.add(
|
||||||
os.path.join(log_dir, log_file),
|
log_file_path,
|
||||||
format=log_format_file,
|
format=log_format_file,
|
||||||
level=log_level,
|
level=log_level,
|
||||||
filter=formatter,
|
filter=formatter,
|
||||||
|
rotation="10 MB", # 每个文件最大10MB
|
||||||
|
retention="30 days", # 保留30天
|
||||||
|
compression=None,
|
||||||
|
encoding="utf-8",
|
||||||
|
enqueue=True, # 异步安全
|
||||||
|
backtrace=True,
|
||||||
|
diagnose=True,
|
||||||
)
|
)
|
||||||
_logger_initialized = True # 标记为已初始化
|
_logger_initialized = True # 标记为已初始化
|
||||||
|
|
||||||
@@ -116,7 +128,7 @@ def update_module_string(selected_module_str):
|
|||||||
)
|
)
|
||||||
log_format_file = log_config.get(
|
log_format_file = log_config.get(
|
||||||
"log_format_file",
|
"log_format_file",
|
||||||
"{time:YYYY-MM-DD HH:mm:ss} - {version_{extra[selected_module]}} - {name} - {level} - {extra[tag]} - {message}",
|
"{time:YYYY-MM-DD HH:mm:ss} - {version}_{extra[selected_module]} - {name} - {level} - {extra[tag]} - {message}",
|
||||||
)
|
)
|
||||||
|
|
||||||
log_format = log_format.replace("{version}", SERVER_VERSION)
|
log_format = log_format.replace("{version}", SERVER_VERSION)
|
||||||
@@ -133,14 +145,26 @@ def update_module_string(selected_module_str):
|
|||||||
level=log_config.get("log_level", "INFO"),
|
level=log_config.get("log_level", "INFO"),
|
||||||
filter=formatter,
|
filter=formatter,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# 更新文件日志配置 - 统一目录,按大小轮转
|
||||||
|
log_dir = log_config.get("log_dir", "tmp")
|
||||||
|
log_file = log_config.get("log_file", "server.log")
|
||||||
|
|
||||||
|
# 日志文件完整路径
|
||||||
|
log_file_path = os.path.join(log_dir, log_file)
|
||||||
|
|
||||||
logger.add(
|
logger.add(
|
||||||
os.path.join(
|
log_file_path,
|
||||||
log_config.get("log_dir", "tmp"),
|
|
||||||
log_config.get("log_file", "server.log"),
|
|
||||||
),
|
|
||||||
format=log_format_file,
|
format=log_format_file,
|
||||||
level=log_config.get("log_level", "INFO"),
|
level=log_config.get("log_level", "INFO"),
|
||||||
filter=formatter,
|
filter=formatter,
|
||||||
|
rotation="10 MB", # 每个文件最大10MB
|
||||||
|
retention="30 days", # 保留30天
|
||||||
|
compression=None,
|
||||||
|
encoding="utf-8",
|
||||||
|
enqueue=True, # 异步安全
|
||||||
|
backtrace=True,
|
||||||
|
diagnose=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
|||||||
@@ -8,6 +8,15 @@
|
|||||||
server:
|
server:
|
||||||
ip: 0.0.0.0
|
ip: 0.0.0.0
|
||||||
port: 8000
|
port: 8000
|
||||||
|
# http服务的端口,用于视觉分析接口
|
||||||
|
http_port: 8003
|
||||||
|
# 视觉分析接口地址
|
||||||
|
# 向设备发送的视觉分析的接口地址
|
||||||
|
# 如果按下面默认的写法,系统会自动生成视觉识别地址,并输出在启动日志里,这个地址你可以直接用浏览器访问确认一下
|
||||||
|
# 当你使用docker部署或使用公网部署(使用ssl、域名)时,不一定准确
|
||||||
|
# 所以如果你使用docker部署时,将vision_explain设置成局域网地址
|
||||||
|
# 如果你使用公网部署时,将vision_explain设置成公网地址
|
||||||
|
vision_explain: http://你的ip或者域名:端口号/mcp/vision/explain
|
||||||
manager-api:
|
manager-api:
|
||||||
# 你的manager-api的地址,最好使用局域网ip
|
# 你的manager-api的地址,最好使用局域网ip
|
||||||
url: http://127.0.0.1:8002/xiaozhi
|
url: http://127.0.0.1:8002/xiaozhi
|
||||||
|
|||||||
@@ -0,0 +1,16 @@
|
|||||||
|
from aiohttp import web
|
||||||
|
from config.logger import setup_logging
|
||||||
|
|
||||||
|
|
||||||
|
class BaseHandler:
|
||||||
|
def __init__(self, config: dict):
|
||||||
|
self.config = config
|
||||||
|
self.logger = setup_logging()
|
||||||
|
|
||||||
|
def _add_cors_headers(self, response):
|
||||||
|
"""添加CORS头信息"""
|
||||||
|
response.headers["Access-Control-Allow-Headers"] = (
|
||||||
|
"client-id, content-type, device-id"
|
||||||
|
)
|
||||||
|
response.headers["Access-Control-Allow-Credentials"] = "true"
|
||||||
|
response.headers["Access-Control-Allow-Origin"] = "*"
|
||||||
+11
-52
@@ -1,18 +1,15 @@
|
|||||||
import json
|
import json
|
||||||
import time
|
import time
|
||||||
import asyncio
|
|
||||||
from aiohttp import web
|
from aiohttp import web
|
||||||
from config.logger import setup_logging
|
|
||||||
from core.utils.util import get_local_ip
|
from core.utils.util import get_local_ip
|
||||||
from core.utils.modules_initialize import initialize_modules
|
from core.api.base_handler import BaseHandler
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
|
|
||||||
|
|
||||||
class SimpleOtaServer:
|
class OTAHandler(BaseHandler):
|
||||||
def __init__(self, config: dict):
|
def __init__(self, config: dict):
|
||||||
self.config = config
|
super().__init__(config)
|
||||||
self.logger = setup_logging()
|
|
||||||
|
|
||||||
def _get_websocket_url(self, local_ip: str, port: int) -> str:
|
def _get_websocket_url(self, local_ip: str, port: int) -> str:
|
||||||
"""获取websocket地址
|
"""获取websocket地址
|
||||||
@@ -25,41 +22,15 @@ class SimpleOtaServer:
|
|||||||
str: websocket地址
|
str: websocket地址
|
||||||
"""
|
"""
|
||||||
server_config = self.config["server"]
|
server_config = self.config["server"]
|
||||||
websocket_config = server_config.get("websocket")
|
websocket_config = server_config.get("websocket", "")
|
||||||
|
|
||||||
if websocket_config and "你" not in websocket_config:
|
if "你的" not in websocket_config:
|
||||||
return websocket_config
|
return websocket_config
|
||||||
else:
|
else:
|
||||||
return f"ws://{local_ip}:{port}/xiaozhi/v1/"
|
return f"ws://{local_ip}:{port}/xiaozhi/v1/"
|
||||||
|
|
||||||
async def start(self):
|
async def handle_post(self, request):
|
||||||
server_config = self.config["server"]
|
"""处理 OTA POST 请求"""
|
||||||
host = server_config.get("ip", "0.0.0.0")
|
|
||||||
port = int(server_config.get("ota_port"))
|
|
||||||
|
|
||||||
if port:
|
|
||||||
app = web.Application()
|
|
||||||
# 添加路由
|
|
||||||
app.add_routes(
|
|
||||||
[
|
|
||||||
web.get("/xiaozhi/ota/", self._handle_ota_get_request),
|
|
||||||
web.post("/xiaozhi/ota/", self._handle_ota_request),
|
|
||||||
web.options("/xiaozhi/ota/", self._handle_ota_request),
|
|
||||||
]
|
|
||||||
)
|
|
||||||
|
|
||||||
# 运行服务
|
|
||||||
runner = web.AppRunner(app)
|
|
||||||
await runner.setup()
|
|
||||||
site = web.TCPSite(runner, host, port)
|
|
||||||
await site.start()
|
|
||||||
|
|
||||||
# 保持服务运行
|
|
||||||
while True:
|
|
||||||
await asyncio.sleep(3600) # 每隔 1 小时检查一次
|
|
||||||
|
|
||||||
async def _handle_ota_request(self, request):
|
|
||||||
"""处理 /xiaozhi/ota/ 的 POST 请求"""
|
|
||||||
try:
|
try:
|
||||||
data = await request.text()
|
data = await request.text()
|
||||||
self.logger.bind(tag=TAG).debug(f"OTA请求方法: {request.method}")
|
self.logger.bind(tag=TAG).debug(f"OTA请求方法: {request.method}")
|
||||||
@@ -75,11 +46,9 @@ class SimpleOtaServer:
|
|||||||
data_json = json.loads(data)
|
data_json = json.loads(data)
|
||||||
|
|
||||||
server_config = self.config["server"]
|
server_config = self.config["server"]
|
||||||
host = server_config.get("ip", "0.0.0.0")
|
|
||||||
port = int(server_config.get("port", 8000))
|
port = int(server_config.get("port", 8000))
|
||||||
local_ip = get_local_ip()
|
local_ip = get_local_ip()
|
||||||
|
|
||||||
# OTA基础信息
|
|
||||||
return_json = {
|
return_json = {
|
||||||
"server_time": {
|
"server_time": {
|
||||||
"timestamp": int(round(time.time() * 1000)),
|
"timestamp": int(round(time.time() * 1000)),
|
||||||
@@ -104,16 +73,11 @@ class SimpleOtaServer:
|
|||||||
content_type="application/json",
|
content_type="application/json",
|
||||||
)
|
)
|
||||||
finally:
|
finally:
|
||||||
# 添加header,允许跨域访问
|
self._add_cors_headers(response)
|
||||||
response.headers["Access-Control-Allow-Headers"] = (
|
|
||||||
"client-id, content-type, device-id"
|
|
||||||
)
|
|
||||||
response.headers["Access-Control-Allow-Credentials"] = "true"
|
|
||||||
response.headers["Access-Control-Allow-Origin"] = "*"
|
|
||||||
return response
|
return response
|
||||||
|
|
||||||
async def _handle_ota_get_request(self, request):
|
async def handle_get(self, request):
|
||||||
"""处理 /xiaozhi/ota/ 的 GET 请求"""
|
"""处理 OTA GET 请求"""
|
||||||
try:
|
try:
|
||||||
server_config = self.config["server"]
|
server_config = self.config["server"]
|
||||||
local_ip = get_local_ip()
|
local_ip = get_local_ip()
|
||||||
@@ -125,10 +89,5 @@ class SimpleOtaServer:
|
|||||||
self.logger.bind(tag=TAG).error(f"OTA GET请求异常: {e}")
|
self.logger.bind(tag=TAG).error(f"OTA GET请求异常: {e}")
|
||||||
response = web.Response(text="OTA接口异常", content_type="text/plain")
|
response = web.Response(text="OTA接口异常", content_type="text/plain")
|
||||||
finally:
|
finally:
|
||||||
# 添加header,允许跨域访问
|
self._add_cors_headers(response)
|
||||||
response.headers["Access-Control-Allow-Headers"] = (
|
|
||||||
"client-id, content-type, device-id"
|
|
||||||
)
|
|
||||||
response.headers["Access-Control-Allow-Credentials"] = "true"
|
|
||||||
response.headers["Access-Control-Allow-Origin"] = "*"
|
|
||||||
return response
|
return response
|
||||||
@@ -0,0 +1,184 @@
|
|||||||
|
import json
|
||||||
|
import copy
|
||||||
|
from aiohttp import web
|
||||||
|
from config.logger import setup_logging
|
||||||
|
from core.utils.util import get_vision_url, is_valid_image_file
|
||||||
|
from core.utils.vllm import create_instance
|
||||||
|
from config.config_loader import get_private_config_from_api
|
||||||
|
from core.utils.auth import AuthToken
|
||||||
|
import base64
|
||||||
|
from typing import Tuple, Optional
|
||||||
|
|
||||||
|
TAG = __name__
|
||||||
|
|
||||||
|
# 设置最大文件大小为5MB
|
||||||
|
MAX_FILE_SIZE = 5 * 1024 * 1024
|
||||||
|
|
||||||
|
|
||||||
|
class VisionHandler:
|
||||||
|
def __init__(self, config: dict):
|
||||||
|
self.config = config
|
||||||
|
self.logger = setup_logging()
|
||||||
|
# 初始化认证工具
|
||||||
|
self.auth = AuthToken(config["server"]["auth_key"])
|
||||||
|
|
||||||
|
def _create_error_response(self, message: str) -> dict:
|
||||||
|
"""创建统一的错误响应格式"""
|
||||||
|
return {"success": False, "message": message}
|
||||||
|
|
||||||
|
def _verify_auth_token(self, request) -> Tuple[bool, Optional[str]]:
|
||||||
|
"""验证认证token"""
|
||||||
|
auth_header = request.headers.get("Authorization", "")
|
||||||
|
if not auth_header.startswith("Bearer "):
|
||||||
|
return False, None
|
||||||
|
|
||||||
|
token = auth_header[7:] # 移除"Bearer "前缀
|
||||||
|
return self.auth.verify_token(token)
|
||||||
|
|
||||||
|
async def handle_post(self, request):
|
||||||
|
"""处理 MCP Vision POST 请求"""
|
||||||
|
response = None # 初始化response变量
|
||||||
|
try:
|
||||||
|
# 验证token
|
||||||
|
is_valid, token_device_id = self._verify_auth_token(request)
|
||||||
|
if not is_valid:
|
||||||
|
response = web.Response(
|
||||||
|
text=json.dumps(
|
||||||
|
self._create_error_response("无效的认证token或token已过期")
|
||||||
|
),
|
||||||
|
content_type="application/json",
|
||||||
|
status=401,
|
||||||
|
)
|
||||||
|
return response
|
||||||
|
|
||||||
|
# 获取请求头信息
|
||||||
|
device_id = request.headers.get("Device-Id", "")
|
||||||
|
client_id = request.headers.get("Client-Id", "")
|
||||||
|
if device_id != token_device_id:
|
||||||
|
return web.Response(
|
||||||
|
text=json.dumps(self._create_error_response("设备ID与token不匹配")),
|
||||||
|
content_type="application/json",
|
||||||
|
status=401,
|
||||||
|
)
|
||||||
|
# 解析multipart/form-data请求
|
||||||
|
reader = await request.multipart()
|
||||||
|
|
||||||
|
# 读取question字段
|
||||||
|
question_field = await reader.next()
|
||||||
|
if question_field is None:
|
||||||
|
raise ValueError("缺少问题字段")
|
||||||
|
question = await question_field.text()
|
||||||
|
self.logger.bind(tag=TAG).debug(f"Question: {question}")
|
||||||
|
|
||||||
|
# 读取图片文件
|
||||||
|
image_field = await reader.next()
|
||||||
|
if image_field is None:
|
||||||
|
raise ValueError("缺少图片文件")
|
||||||
|
|
||||||
|
# 读取图片数据
|
||||||
|
image_data = await image_field.read()
|
||||||
|
if not image_data:
|
||||||
|
raise ValueError("图片数据为空")
|
||||||
|
|
||||||
|
# 检查文件大小
|
||||||
|
if len(image_data) > MAX_FILE_SIZE:
|
||||||
|
raise ValueError(
|
||||||
|
f"图片大小超过限制,最大允许{MAX_FILE_SIZE/1024/1024}MB"
|
||||||
|
)
|
||||||
|
|
||||||
|
# 检查文件格式
|
||||||
|
if not is_valid_image_file(image_data):
|
||||||
|
raise ValueError(
|
||||||
|
"不支持的文件格式,请上传有效的图片文件(支持JPEG、PNG、GIF、BMP、TIFF、WEBP格式)"
|
||||||
|
)
|
||||||
|
|
||||||
|
# 将图片转换为base64编码
|
||||||
|
image_base64 = base64.b64encode(image_data).decode("utf-8")
|
||||||
|
|
||||||
|
# 如果开启了智控台,则从智控台获取模型配置
|
||||||
|
current_config = copy.deepcopy(self.config)
|
||||||
|
read_config_from_api = current_config.get("read_config_from_api", False)
|
||||||
|
if read_config_from_api:
|
||||||
|
current_config = get_private_config_from_api(
|
||||||
|
current_config,
|
||||||
|
device_id,
|
||||||
|
client_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
select_vllm_module = current_config["selected_module"].get("VLLM")
|
||||||
|
if not select_vllm_module:
|
||||||
|
raise ValueError("您还未设置默认的视觉分析模块")
|
||||||
|
|
||||||
|
vllm_type = (
|
||||||
|
select_vllm_module
|
||||||
|
if "type" not in current_config["VLLM"][select_vllm_module]
|
||||||
|
else current_config["VLLM"][select_vllm_module]["type"]
|
||||||
|
)
|
||||||
|
|
||||||
|
if not vllm_type:
|
||||||
|
raise ValueError(f"无法找到VLLM模块对应的供应器{vllm_type}")
|
||||||
|
|
||||||
|
vllm = create_instance(
|
||||||
|
vllm_type, current_config["VLLM"][select_vllm_module]
|
||||||
|
)
|
||||||
|
|
||||||
|
result = vllm.response(question, image_base64)
|
||||||
|
|
||||||
|
return_json = {
|
||||||
|
"success": True,
|
||||||
|
"result": result,
|
||||||
|
}
|
||||||
|
|
||||||
|
response = web.Response(
|
||||||
|
text=json.dumps(return_json, separators=(",", ":")),
|
||||||
|
content_type="application/json",
|
||||||
|
)
|
||||||
|
except ValueError as e:
|
||||||
|
self.logger.bind(tag=TAG).error(f"MCP Vision POST请求异常: {e}")
|
||||||
|
return_json = self._create_error_response(str(e))
|
||||||
|
response = web.Response(
|
||||||
|
text=json.dumps(return_json, separators=(",", ":")),
|
||||||
|
content_type="application/json",
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
self.logger.bind(tag=TAG).error(f"MCP Vision POST请求异常: {e}")
|
||||||
|
return_json = self._create_error_response("处理请求时发生错误")
|
||||||
|
response = web.Response(
|
||||||
|
text=json.dumps(return_json, separators=(",", ":")),
|
||||||
|
content_type="application/json",
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
if response:
|
||||||
|
self._add_cors_headers(response)
|
||||||
|
return response
|
||||||
|
|
||||||
|
async def handle_get(self, request):
|
||||||
|
"""处理 MCP Vision GET 请求"""
|
||||||
|
try:
|
||||||
|
vision_explain = get_vision_url(self.config)
|
||||||
|
if vision_explain and len(vision_explain) > 0 and "null" != vision_explain:
|
||||||
|
message = (
|
||||||
|
f"MCP Vision 接口运行正常,视觉解释接口地址是:{vision_explain}"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
message = "MCP Vision 接口运行不正常,请打开data目录下的.config.yaml文件,找到【server.vision_explain】,设置好地址"
|
||||||
|
|
||||||
|
response = web.Response(text=message, content_type="text/plain")
|
||||||
|
except Exception as e:
|
||||||
|
self.logger.bind(tag=TAG).error(f"MCP Vision GET请求异常: {e}")
|
||||||
|
return_json = self._create_error_response("服务器内部错误")
|
||||||
|
response = web.Response(
|
||||||
|
text=json.dumps(return_json, separators=(",", ":")),
|
||||||
|
content_type="application/json",
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
self._add_cors_headers(response)
|
||||||
|
return response
|
||||||
|
|
||||||
|
def _add_cors_headers(self, response):
|
||||||
|
"""添加CORS头信息"""
|
||||||
|
response.headers["Access-Control-Allow-Headers"] = (
|
||||||
|
"client-id, content-type, device-id"
|
||||||
|
)
|
||||||
|
response.headers["Access-Control-Allow-Credentials"] = "true"
|
||||||
|
response.headers["Access-Control-Allow-Origin"] = "*"
|
||||||
@@ -35,7 +35,6 @@ from plugins_func.loadplugins import auto_import_modules
|
|||||||
from plugins_func.register import Action, ActionResponse
|
from plugins_func.register import Action, ActionResponse
|
||||||
from core.auth import AuthMiddleware, AuthenticationError
|
from core.auth import AuthMiddleware, AuthenticationError
|
||||||
from config.config_loader import get_private_config_from_api
|
from config.config_loader import get_private_config_from_api
|
||||||
from core.handle.receiveAudioHandle import handleAudioMessage
|
|
||||||
from core.providers.tts.dto.dto import ContentType, TTSMessageDTO, SentenceType
|
from core.providers.tts.dto.dto import ContentType, TTSMessageDTO, SentenceType
|
||||||
from config.logger import setup_logging, build_module_string, update_module_string
|
from config.logger import setup_logging, build_module_string, update_module_string
|
||||||
from config.manage_api_client import DeviceNotFoundException, DeviceBindException
|
from config.manage_api_client import DeviceNotFoundException, DeviceBindException
|
||||||
@@ -81,6 +80,7 @@ class ConnectionHandler:
|
|||||||
self.welcome_msg = None
|
self.welcome_msg = None
|
||||||
self.max_output_size = 0
|
self.max_output_size = 0
|
||||||
self.chat_history_conf = 0
|
self.chat_history_conf = 0
|
||||||
|
self.audio_format = "opus"
|
||||||
|
|
||||||
# 客户端状态相关
|
# 客户端状态相关
|
||||||
self.client_abort = False
|
self.client_abort = False
|
||||||
@@ -117,7 +117,10 @@ class ConnectionHandler:
|
|||||||
self.client_voice_stop = False
|
self.client_voice_stop = False
|
||||||
|
|
||||||
# asr相关变量
|
# asr相关变量
|
||||||
|
# 因为实际部署时可能会用到公共的本地ASR,不能把变量暴露给公共ASR
|
||||||
|
# 所以涉及到ASR的变量,需要在这里定义,属于connection的私有变量
|
||||||
self.asr_audio = []
|
self.asr_audio = []
|
||||||
|
self.asr_audio_queue = queue.Queue()
|
||||||
|
|
||||||
# llm相关变量
|
# llm相关变量
|
||||||
self.llm_finish_task = True
|
self.llm_finish_task = True
|
||||||
@@ -146,7 +149,6 @@ class ConnectionHandler:
|
|||||||
int(self.config.get("close_connection_no_voice_time", 120)) + 60
|
int(self.config.get("close_connection_no_voice_time", 120)) + 60
|
||||||
) # 在原来第一道关闭的基础上加60秒,进行二道关闭
|
) # 在原来第一道关闭的基础上加60秒,进行二道关闭
|
||||||
|
|
||||||
self.audio_format = "opus"
|
|
||||||
# {"mcp":true} 表示启用MCP功能
|
# {"mcp":true} 表示启用MCP功能
|
||||||
self.features = None
|
self.features = None
|
||||||
|
|
||||||
@@ -254,7 +256,11 @@ class ConnectionHandler:
|
|||||||
if isinstance(message, str):
|
if isinstance(message, str):
|
||||||
await handleTextMessage(self, message)
|
await handleTextMessage(self, message)
|
||||||
elif isinstance(message, bytes):
|
elif isinstance(message, bytes):
|
||||||
await handleAudioMessage(self, message)
|
if self.vad is None:
|
||||||
|
return
|
||||||
|
if self.asr is None:
|
||||||
|
return
|
||||||
|
self.asr_audio_queue.put(message)
|
||||||
|
|
||||||
async def handle_restart(self, message):
|
async def handle_restart(self, message):
|
||||||
"""处理服务器重启请求"""
|
"""处理服务器重启请求"""
|
||||||
@@ -478,6 +484,8 @@ class ConnectionHandler:
|
|||||||
self.memory = modules["memory"]
|
self.memory = modules["memory"]
|
||||||
|
|
||||||
def _initialize_memory(self):
|
def _initialize_memory(self):
|
||||||
|
if self.memory is None:
|
||||||
|
return
|
||||||
"""初始化记忆模块"""
|
"""初始化记忆模块"""
|
||||||
self.memory.init_memory(
|
self.memory.init_memory(
|
||||||
role_id=self.device_id,
|
role_id=self.device_id,
|
||||||
@@ -518,6 +526,8 @@ class ConnectionHandler:
|
|||||||
self.logger.bind(tag=TAG).info("使用主LLM作为意图识别模型")
|
self.logger.bind(tag=TAG).info("使用主LLM作为意图识别模型")
|
||||||
|
|
||||||
def _initialize_intent(self):
|
def _initialize_intent(self):
|
||||||
|
if self.intent is None:
|
||||||
|
return
|
||||||
self.intent_type = self.config["Intent"][
|
self.intent_type = self.config["Intent"][
|
||||||
self.config["selected_module"]["Intent"]
|
self.config["selected_module"]["Intent"]
|
||||||
]["type"]
|
]["type"]
|
||||||
@@ -641,7 +651,7 @@ class ConnectionHandler:
|
|||||||
# print("content_arguments", content_arguments)
|
# print("content_arguments", content_arguments)
|
||||||
tool_call_flag = True
|
tool_call_flag = True
|
||||||
|
|
||||||
if tools_call is not None:
|
if tools_call is not None and len(tools_call) > 0:
|
||||||
tool_call_flag = True
|
tool_call_flag = True
|
||||||
if tools_call[0].id is not None:
|
if tools_call[0].id is not None:
|
||||||
function_id = tools_call[0].id
|
function_id = tools_call[0].id
|
||||||
@@ -708,15 +718,24 @@ class ConnectionHandler:
|
|||||||
# 处理Server端MCP工具调用
|
# 处理Server端MCP工具调用
|
||||||
if self.mcp_manager.is_mcp_tool(function_name):
|
if self.mcp_manager.is_mcp_tool(function_name):
|
||||||
result = self._handle_mcp_tool_call(function_call_data)
|
result = self._handle_mcp_tool_call(function_call_data)
|
||||||
elif hasattr(self, "mcp_client") and self.mcp_client.has_tool(function_name):
|
elif hasattr(self, "mcp_client") and self.mcp_client.has_tool(
|
||||||
# 如果是小智端MCP工具调用
|
function_name
|
||||||
|
):
|
||||||
|
# 如果是小智端MCP工具调用
|
||||||
self.logger.bind(tag=TAG).debug(
|
self.logger.bind(tag=TAG).debug(
|
||||||
f"调用小智端MCP工具: {function_name}, 参数: {function_arguments}"
|
f"调用小智端MCP工具: {function_name}, 参数: {function_arguments}"
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
result = asyncio.run_coroutine_threadsafe(call_mcp_tool(self, self.mcp_client, function_name, function_arguments), self.loop).result()
|
result = asyncio.run_coroutine_threadsafe(
|
||||||
|
call_mcp_tool(
|
||||||
|
self, self.mcp_client, function_name, function_arguments
|
||||||
|
),
|
||||||
|
self.loop,
|
||||||
|
).result()
|
||||||
self.logger.bind(tag=TAG).debug(f"MCP工具调用结果: {result}")
|
self.logger.bind(tag=TAG).debug(f"MCP工具调用结果: {result}")
|
||||||
result = ActionResponse(action=Action.REQLLM, result=result, response="")
|
result = ActionResponse(
|
||||||
|
action=Action.REQLLM, result=result, response=""
|
||||||
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.logger.bind(tag=TAG).error(f"MCP工具调用失败: {e}")
|
self.logger.bind(tag=TAG).error(f"MCP工具调用失败: {e}")
|
||||||
result = ActionResponse(
|
result = ActionResponse(
|
||||||
|
|||||||
@@ -1,6 +1,4 @@
|
|||||||
import json
|
import json
|
||||||
import queue
|
|
||||||
from config.logger import setup_logging
|
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
|
|
||||||
|
|||||||
@@ -1,26 +1,31 @@
|
|||||||
import os
|
|
||||||
import time
|
import time
|
||||||
import json
|
import json
|
||||||
import random
|
import random
|
||||||
import shutil
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from core.handle.mcpHandle import MCPClient, send_mcp_initialize_message, send_mcp_tools_list_request
|
from core.utils.util import audio_to_data
|
||||||
from core.handle.sendAudioHandle import send_stt_message
|
from core.handle.sendAudioHandle import sendAudioMessage, send_stt_message
|
||||||
from core.utils.util import remove_punctuation_and_length
|
from core.utils.util import remove_punctuation_and_length, opus_datas_to_wav_bytes
|
||||||
from core.providers.tts.dto.dto import ContentType, InterfaceType
|
from core.providers.tts.dto.dto import ContentType, SentenceType
|
||||||
|
from core.handle.mcpHandle import (
|
||||||
|
MCPClient,
|
||||||
|
send_mcp_initialize_message,
|
||||||
|
send_mcp_tools_list_request,
|
||||||
|
)
|
||||||
|
from core.utils.wakeup_word import WakeupWordsConfig
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
|
|
||||||
WAKEUP_CONFIG = {
|
WAKEUP_CONFIG = {
|
||||||
"dir": "config/assets/",
|
"refresh_time": 5,
|
||||||
"file_name": "wakeup_words",
|
"words": ["你好", "你好啊", "嘿,你好", "嗨"],
|
||||||
"create_time": time.time(),
|
|
||||||
"refresh_time": 10,
|
|
||||||
"words": ["你好小智", "你好啊小智", "小智你好", "小智"],
|
|
||||||
"text": "",
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
# 创建全局的唤醒词配置管理器
|
||||||
|
wakeup_words_config = WakeupWordsConfig()
|
||||||
|
|
||||||
|
# 用于防止并发调用wakeupWordsResponse的锁
|
||||||
|
_wakeup_response_lock = asyncio.Lock()
|
||||||
|
|
||||||
|
|
||||||
async def handleHelloMessage(conn, msg_json):
|
async def handleHelloMessage(conn, msg_json):
|
||||||
"""处理hello消息"""
|
"""处理hello消息"""
|
||||||
@@ -29,8 +34,6 @@ async def handleHelloMessage(conn, msg_json):
|
|||||||
format = audio_params.get("format")
|
format = audio_params.get("format")
|
||||||
conn.logger.bind(tag=TAG).info(f"客户端音频格式: {format}")
|
conn.logger.bind(tag=TAG).info(f"客户端音频格式: {format}")
|
||||||
conn.audio_format = format
|
conn.audio_format = format
|
||||||
if conn.asr is not None:
|
|
||||||
conn.asr.set_audio_format(format)
|
|
||||||
conn.welcome_msg["audio_params"] = audio_params
|
conn.welcome_msg["audio_params"] = audio_params
|
||||||
features = msg_json.get("features")
|
features = msg_json.get("features")
|
||||||
if features:
|
if features:
|
||||||
@@ -51,81 +54,75 @@ async def checkWakeupWords(conn, text):
|
|||||||
enable_wakeup_words_response_cache = conn.config[
|
enable_wakeup_words_response_cache = conn.config[
|
||||||
"enable_wakeup_words_response_cache"
|
"enable_wakeup_words_response_cache"
|
||||||
]
|
]
|
||||||
"""是否用的是非流式tts"""
|
|
||||||
if conn.tts and conn.tts.interface_type != InterfaceType.NON_STREAM:
|
if not enable_wakeup_words_response_cache or not conn.tts:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
"""是否开启唤醒词加速"""
|
|
||||||
if not enable_wakeup_words_response_cache:
|
|
||||||
return False
|
|
||||||
"""检查是否是唤醒词"""
|
|
||||||
_, filtered_text = remove_punctuation_and_length(text)
|
_, filtered_text = remove_punctuation_and_length(text)
|
||||||
if filtered_text in conn.config.get("wakeup_words"):
|
if filtered_text not in conn.config.get("wakeup_words"):
|
||||||
await send_stt_message(conn, text)
|
return False
|
||||||
|
|
||||||
file = getWakeupWordFile(WAKEUP_CONFIG["file_name"])
|
conn.just_woken_up = True
|
||||||
if file is None:
|
await send_stt_message(conn, text)
|
||||||
|
|
||||||
|
# 获取当前音色
|
||||||
|
voice = getattr(conn.tts, "voice", "default")
|
||||||
|
|
||||||
|
# 获取唤醒词回复配置
|
||||||
|
response = wakeup_words_config.get_wakeup_response(voice)
|
||||||
|
|
||||||
|
# 播放唤醒词回复
|
||||||
|
conn.client_abort = False
|
||||||
|
opus_packets, _ = audio_to_data(response["file_path"])
|
||||||
|
|
||||||
|
conn.logger.bind(tag=TAG).info(f"播放唤醒词回复: {response['text']}")
|
||||||
|
await sendAudioMessage(conn, SentenceType.FIRST, opus_packets, response["text"])
|
||||||
|
await sendAudioMessage(conn, SentenceType.LAST, [], None)
|
||||||
|
|
||||||
|
# 检查是否需要更新唤醒词回复
|
||||||
|
if time.time() - response["time"] > WAKEUP_CONFIG["refresh_time"]:
|
||||||
|
if not _wakeup_response_lock.locked():
|
||||||
asyncio.create_task(wakeupWordsResponse(conn))
|
asyncio.create_task(wakeupWordsResponse(conn))
|
||||||
return False
|
return True
|
||||||
text_hello = WAKEUP_CONFIG["text"]
|
|
||||||
if not text_hello:
|
|
||||||
text_hello = text
|
|
||||||
conn.tts.tts_one_sentence(
|
|
||||||
conn, ContentType.FILE, content_file=file, content_detail=text_hello
|
|
||||||
)
|
|
||||||
if time.time() - WAKEUP_CONFIG["create_time"] > WAKEUP_CONFIG["refresh_time"]:
|
|
||||||
asyncio.create_task(wakeupWordsResponse(conn))
|
|
||||||
return True
|
|
||||||
return False
|
|
||||||
|
|
||||||
|
|
||||||
def getWakeupWordFile(file_name):
|
|
||||||
for file in os.listdir(WAKEUP_CONFIG["dir"]):
|
|
||||||
if file.startswith("my_" + file_name):
|
|
||||||
"""避免缓存文件是一个空文件"""
|
|
||||||
if os.stat(f"config/assets/{file}").st_size > (15 * 1024):
|
|
||||||
return f"config/assets/{file}"
|
|
||||||
|
|
||||||
"""查找config/assets/目录下名称为wakeup_words的文件"""
|
|
||||||
for file in os.listdir(WAKEUP_CONFIG["dir"]):
|
|
||||||
if file.startswith(file_name):
|
|
||||||
return f"config/assets/{file}"
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
async def wakeupWordsResponse(conn):
|
async def wakeupWordsResponse(conn):
|
||||||
wait_max_time = 5
|
if not conn.tts or not conn.llm or not conn.llm.response_no_stream:
|
||||||
while conn.llm is None or not conn.llm.response_no_stream:
|
return
|
||||||
await asyncio.sleep(1)
|
|
||||||
wait_max_time -= 1
|
try:
|
||||||
if wait_max_time <= 0:
|
# 尝试获取锁,如果获取不到就返回
|
||||||
conn.logger.bind(tag=TAG).error("连接对象没有llm")
|
if not await _wakeup_response_lock.acquire():
|
||||||
return
|
return
|
||||||
|
|
||||||
"""唤醒词响应"""
|
# 生成唤醒词回复
|
||||||
wakeup_word = random.choice(WAKEUP_CONFIG["words"])
|
wakeup_word = random.choice(WAKEUP_CONFIG["words"])
|
||||||
question = (
|
question = (
|
||||||
"此刻用户正在和你说```"
|
"此刻用户正在和你说```"
|
||||||
+ wakeup_word
|
+ wakeup_word
|
||||||
+ "```。\n请你根据以上用户的内容,进行简短回复,文字内容控制在15个字以内。\n"
|
+ "```。\n请你根据以上用户的内容进行简短回复。要像一个人正常人一样说话,不要像机器人一样说话。\n"
|
||||||
+ "请勿对这条内容本身进行任何解释和回应,仅返回对用户的内容的回复。"
|
+ "请勿对这条内容本身进行任何解释和回应,请勿返回表情符号,仅返回对用户的内容的回复。"
|
||||||
)
|
|
||||||
result = conn.llm.response_no_stream(conn.config["prompt"], question)
|
|
||||||
if result is None or result == "":
|
|
||||||
return
|
|
||||||
tts_file = await asyncio.to_thread(conn.tts.to_tts, result)
|
|
||||||
|
|
||||||
if tts_file is not None and os.path.exists(tts_file):
|
|
||||||
file_type = os.path.splitext(tts_file)[1]
|
|
||||||
if file_type:
|
|
||||||
file_type = file_type.lstrip(".")
|
|
||||||
old_file = getWakeupWordFile("my_" + WAKEUP_CONFIG["file_name"])
|
|
||||||
if old_file is not None:
|
|
||||||
os.remove(old_file)
|
|
||||||
"""将文件挪到"wakeup_words.mp3"""
|
|
||||||
shutil.move(
|
|
||||||
tts_file,
|
|
||||||
WAKEUP_CONFIG["dir"] + "my_" + WAKEUP_CONFIG["file_name"] + "." + file_type,
|
|
||||||
)
|
)
|
||||||
WAKEUP_CONFIG["create_time"] = time.time()
|
|
||||||
WAKEUP_CONFIG["text"] = result
|
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()
|
||||||
|
|||||||
@@ -96,6 +96,7 @@ async def process_intent_result(conn, intent_result, original_text):
|
|||||||
}
|
}
|
||||||
|
|
||||||
await send_stt_message(conn, original_text)
|
await send_stt_message(conn, original_text)
|
||||||
|
conn.client_abort = False
|
||||||
|
|
||||||
# 使用executor执行函数调用和结果处理
|
# 使用executor执行函数调用和结果处理
|
||||||
def process_function_call():
|
def process_function_call():
|
||||||
|
|||||||
@@ -1,6 +1,8 @@
|
|||||||
import json
|
import json
|
||||||
import asyncio
|
import asyncio
|
||||||
from concurrent.futures import Future
|
from concurrent.futures import Future
|
||||||
|
from core.utils.util import get_vision_url, sanitize_tool_name
|
||||||
|
from core.utils.auth import AuthToken
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
|
|
||||||
@@ -9,7 +11,8 @@ class MCPClient:
|
|||||||
"""MCPClient,用于管理MCP状态和工具"""
|
"""MCPClient,用于管理MCP状态和工具"""
|
||||||
|
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.tools = {} # Dictionary for O(1) lookup
|
self.tools = {} # sanitized_name -> tool_data
|
||||||
|
self.name_mapping = {}
|
||||||
self.ready = False
|
self.ready = False
|
||||||
self.call_results = {} # To store Futures for tool call responses
|
self.call_results = {} # To store Futures for tool call responses
|
||||||
self.next_id = 1
|
self.next_id = 1
|
||||||
@@ -28,7 +31,7 @@ class MCPClient:
|
|||||||
result = []
|
result = []
|
||||||
for tool_name, tool_data in self.tools.items():
|
for tool_name, tool_data in self.tools.items():
|
||||||
function_def = {
|
function_def = {
|
||||||
"name": tool_data["name"],
|
"name": tool_name,
|
||||||
"description": tool_data["description"],
|
"description": tool_data["description"],
|
||||||
"parameters": {
|
"parameters": {
|
||||||
"type": tool_data["inputSchema"].get("type", "object"),
|
"type": tool_data["inputSchema"].get("type", "object"),
|
||||||
@@ -51,7 +54,9 @@ class MCPClient:
|
|||||||
|
|
||||||
async def add_tool(self, tool_data: dict):
|
async def add_tool(self, tool_data: dict):
|
||||||
async with self.lock:
|
async with self.lock:
|
||||||
self.tools[tool_data["name"]] = tool_data
|
sanitized_name = sanitize_tool_name(tool_data["name"])
|
||||||
|
self.tools[sanitized_name] = tool_data
|
||||||
|
self.name_mapping[sanitized_name] = tool_data["name"]
|
||||||
self._cached_available_tools = (
|
self._cached_available_tools = (
|
||||||
None # Invalidate the cache when a tool is added
|
None # Invalidate the cache when a tool is added
|
||||||
)
|
)
|
||||||
@@ -131,9 +136,6 @@ async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict):
|
|||||||
conn.logger.bind(tag=TAG).info(
|
conn.logger.bind(tag=TAG).info(
|
||||||
f"客户端MCP服务器信息: name={name}, version={version}"
|
f"客户端MCP服务器信息: name={name}, version={version}"
|
||||||
)
|
)
|
||||||
await send_mcp_tools_list_request(
|
|
||||||
conn
|
|
||||||
) # After initialization, request tool list
|
|
||||||
return
|
return
|
||||||
|
|
||||||
elif msg_id == 2: # mcpToolsListID
|
elif msg_id == 2: # mcpToolsListID
|
||||||
@@ -172,6 +174,20 @@ async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict):
|
|||||||
await mcp_client.add_tool(new_tool)
|
await mcp_client.add_tool(new_tool)
|
||||||
conn.logger.bind(tag=TAG).debug(f"客户端工具 #{i+1}: {name}")
|
conn.logger.bind(tag=TAG).debug(f"客户端工具 #{i+1}: {name}")
|
||||||
|
|
||||||
|
# 替换所有工具描述中的工具名称
|
||||||
|
for tool_data in mcp_client.tools.values():
|
||||||
|
if "description" in tool_data:
|
||||||
|
description = tool_data["description"]
|
||||||
|
# 遍历所有工具名称进行替换
|
||||||
|
for (
|
||||||
|
sanitized_name,
|
||||||
|
original_name,
|
||||||
|
) in mcp_client.name_mapping.items():
|
||||||
|
description = description.replace(
|
||||||
|
original_name, sanitized_name
|
||||||
|
)
|
||||||
|
tool_data["description"] = description
|
||||||
|
|
||||||
next_cursor = result.get("nextCursor", "")
|
next_cursor = result.get("nextCursor", "")
|
||||||
if next_cursor:
|
if next_cursor:
|
||||||
conn.logger.bind(tag=TAG).info(
|
conn.logger.bind(tag=TAG).info(
|
||||||
@@ -205,6 +221,18 @@ async def handle_mcp_message(conn, mcp_client: MCPClient, payload: dict):
|
|||||||
|
|
||||||
async def send_mcp_initialize_message(conn):
|
async def send_mcp_initialize_message(conn):
|
||||||
"""发送MCP初始化消息"""
|
"""发送MCP初始化消息"""
|
||||||
|
|
||||||
|
vision_url = get_vision_url(conn.config)
|
||||||
|
|
||||||
|
# 密钥生成token
|
||||||
|
auth = AuthToken(conn.config["server"]["auth_key"])
|
||||||
|
token = auth.generate_token(conn.headers.get("device-id"))
|
||||||
|
|
||||||
|
vision = {
|
||||||
|
"url": vision_url,
|
||||||
|
"token": token,
|
||||||
|
}
|
||||||
|
|
||||||
payload = {
|
payload = {
|
||||||
"jsonrpc": "2.0",
|
"jsonrpc": "2.0",
|
||||||
"id": 1, # mcpInitializeID
|
"id": 1, # mcpInitializeID
|
||||||
@@ -214,6 +242,7 @@ async def send_mcp_initialize_message(conn):
|
|||||||
"capabilities": {
|
"capabilities": {
|
||||||
"roots": {"listChanged": True},
|
"roots": {"listChanged": True},
|
||||||
"sampling": {},
|
"sampling": {},
|
||||||
|
"vision": vision,
|
||||||
},
|
},
|
||||||
"clientInfo": {
|
"clientInfo": {
|
||||||
"name": "XiaozhiClient",
|
"name": "XiaozhiClient",
|
||||||
@@ -316,15 +345,16 @@ async def call_mcp_tool(
|
|||||||
raise ValueError(f"参数处理失败: {str(e)}")
|
raise ValueError(f"参数处理失败: {str(e)}")
|
||||||
raise e
|
raise e
|
||||||
|
|
||||||
|
actual_name = mcp_client.name_mapping.get(tool_name, tool_name)
|
||||||
payload = {
|
payload = {
|
||||||
"jsonrpc": "2.0",
|
"jsonrpc": "2.0",
|
||||||
"id": tool_call_id,
|
"id": tool_call_id,
|
||||||
"method": "tools/call",
|
"method": "tools/call",
|
||||||
"params": {"name": tool_name, "arguments": arguments},
|
"params": {"name": actual_name, "arguments": arguments},
|
||||||
}
|
}
|
||||||
|
|
||||||
conn.logger.bind(tag=TAG).info(
|
conn.logger.bind(tag=TAG).info(
|
||||||
f"发送客户端mcp工具调用请求: {tool_name},参数: {args}"
|
f"发送客户端mcp工具调用请求: {actual_name},参数: {args}"
|
||||||
)
|
)
|
||||||
await send_mcp_message(conn, payload)
|
await send_mcp_message(conn, payload)
|
||||||
|
|
||||||
@@ -332,7 +362,7 @@ async def call_mcp_tool(
|
|||||||
# Wait for response or timeout
|
# Wait for response or timeout
|
||||||
raw_result = await asyncio.wait_for(result_future, timeout=timeout)
|
raw_result = await asyncio.wait_for(result_future, timeout=timeout)
|
||||||
conn.logger.bind(tag=TAG).info(
|
conn.logger.bind(tag=TAG).info(
|
||||||
f"客户端mcp工具调用 {tool_name} 成功,原始结果: {raw_result}"
|
f"客户端mcp工具调用 {actual_name} 成功,原始结果: {raw_result}"
|
||||||
)
|
)
|
||||||
|
|
||||||
if isinstance(raw_result, dict):
|
if isinstance(raw_result, dict):
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ from core.handle.intentHandler import handle_user_intent
|
|||||||
from core.utils.output_counter import check_device_output_limit
|
from core.utils.output_counter import check_device_output_limit
|
||||||
from core.handle.abortHandle import handleAbortMessage
|
from core.handle.abortHandle import handleAbortMessage
|
||||||
import time
|
import time
|
||||||
|
import asyncio
|
||||||
from core.handle.sendAudioHandle import SentenceType
|
from core.handle.sendAudioHandle import SentenceType
|
||||||
from core.utils.util import audio_to_data
|
from core.utils.util import audio_to_data
|
||||||
|
|
||||||
@@ -10,21 +11,30 @@ TAG = __name__
|
|||||||
|
|
||||||
|
|
||||||
async def handleAudioMessage(conn, audio):
|
async def handleAudioMessage(conn, audio):
|
||||||
if conn.vad is None:
|
|
||||||
conn.logger.bind(tag=TAG).warning("VAD模块未初始化,继续等待")
|
|
||||||
return
|
|
||||||
if conn.asr is None or not hasattr(conn.asr, "conn") or conn.asr.conn is None:
|
|
||||||
conn.logger.bind(tag=TAG).warning("ASR模块未初始化或通道未就绪,继续等待")
|
|
||||||
return
|
|
||||||
# 当前片段是否有人说话
|
# 当前片段是否有人说话
|
||||||
have_voice = conn.vad.is_vad(conn, audio)
|
have_voice = conn.vad.is_vad(conn, audio)
|
||||||
|
# 如果设备刚刚被唤醒,短暂忽略VAD检测
|
||||||
|
if have_voice and hasattr(conn, "just_woken_up") and conn.just_woken_up:
|
||||||
|
have_voice = False
|
||||||
|
# 设置一个短暂延迟后恢复VAD检测
|
||||||
|
conn.asr_audio.clear()
|
||||||
|
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 have_voice:
|
||||||
if conn.client_is_speaking:
|
if conn.client_is_speaking:
|
||||||
await handleAbortMessage(conn)
|
await handleAbortMessage(conn)
|
||||||
# 设备长时间空闲检测,用于say goodbye
|
# 设备长时间空闲检测,用于say goodbye
|
||||||
await no_voice_close_connect(conn, have_voice)
|
await no_voice_close_connect(conn, have_voice)
|
||||||
# 接收音频
|
# 接收音频
|
||||||
await conn.asr.receive_audio(audio, have_voice)
|
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):
|
async def startToChat(conn, text):
|
||||||
|
|||||||
@@ -34,6 +34,7 @@ emoji_map = {
|
|||||||
|
|
||||||
async def sendAudioMessage(conn, sentenceType, audios, text):
|
async def sendAudioMessage(conn, sentenceType, audios, text):
|
||||||
# 发送句子开始消息
|
# 发送句子开始消息
|
||||||
|
conn.logger.bind(tag=TAG).info(f"发送音频消息: {sentenceType}, {text}")
|
||||||
if text is not None:
|
if text is not None:
|
||||||
emotion = analyze_emotion(text)
|
emotion = analyze_emotion(text)
|
||||||
emoji = emoji_map.get(emotion, "🙂") # 默认使用笑脸
|
emoji = emoji_map.get(emotion, "🙂") # 默认使用笑脸
|
||||||
@@ -89,8 +90,7 @@ async def sendAudio(conn, audios, pre_buffer=True):
|
|||||||
# 播放剩余音频帧
|
# 播放剩余音频帧
|
||||||
for opus_packet in remaining_audios:
|
for opus_packet in remaining_audios:
|
||||||
if conn.client_abort:
|
if conn.client_abort:
|
||||||
conn.client_abort = False
|
break
|
||||||
return
|
|
||||||
|
|
||||||
# 每分钟重置一次计时器
|
# 每分钟重置一次计时器
|
||||||
if time.perf_counter() - last_reset_time > 60:
|
if time.perf_counter() - last_reset_time > 60:
|
||||||
|
|||||||
@@ -61,6 +61,7 @@ async def handleTextMessage(conn, message):
|
|||||||
await send_tts_message(conn, "stop", None)
|
await send_tts_message(conn, "stop", None)
|
||||||
conn.client_is_speaking = False
|
conn.client_is_speaking = False
|
||||||
elif is_wakeup_words:
|
elif is_wakeup_words:
|
||||||
|
conn.just_woken_up = True
|
||||||
# 上报纯文字数据(复用ASR上报功能,但不提供音频数据)
|
# 上报纯文字数据(复用ASR上报功能,但不提供音频数据)
|
||||||
enqueue_asr_report(conn, "嘿,你好呀", [])
|
enqueue_asr_report(conn, "嘿,你好呀", [])
|
||||||
await startToChat(conn, "嘿,你好呀")
|
await startToChat(conn, "嘿,你好呀")
|
||||||
@@ -78,7 +79,9 @@ async def handleTextMessage(conn, message):
|
|||||||
elif msg_json["type"] == "mcp":
|
elif msg_json["type"] == "mcp":
|
||||||
conn.logger.bind(tag=TAG).info(f"收到mcp消息:{message}")
|
conn.logger.bind(tag=TAG).info(f"收到mcp消息:{message}")
|
||||||
if "payload" in msg_json:
|
if "payload" in msg_json:
|
||||||
asyncio.create_task(handle_mcp_message(conn, conn.mcp_client, msg_json["payload"]))
|
asyncio.create_task(
|
||||||
|
handle_mcp_message(conn, conn.mcp_client, msg_json["payload"])
|
||||||
|
)
|
||||||
elif msg_json["type"] == "server":
|
elif msg_json["type"] == "server":
|
||||||
# 记录日志时过滤敏感信息
|
# 记录日志时过滤敏感信息
|
||||||
conn.logger.bind(tag=TAG).info(
|
conn.logger.bind(tag=TAG).info(
|
||||||
|
|||||||
@@ -0,0 +1,71 @@
|
|||||||
|
import asyncio
|
||||||
|
from aiohttp import web
|
||||||
|
from config.logger import setup_logging
|
||||||
|
from core.api.ota_handler import OTAHandler
|
||||||
|
from core.api.vision_handler import VisionHandler
|
||||||
|
|
||||||
|
TAG = __name__
|
||||||
|
|
||||||
|
|
||||||
|
class SimpleHttpServer:
|
||||||
|
def __init__(self, config: dict):
|
||||||
|
self.config = config
|
||||||
|
self.logger = setup_logging()
|
||||||
|
self.ota_handler = OTAHandler(config)
|
||||||
|
self.vision_handler = VisionHandler(config)
|
||||||
|
|
||||||
|
def _get_websocket_url(self, local_ip: str, port: int) -> str:
|
||||||
|
"""获取websocket地址
|
||||||
|
|
||||||
|
Args:
|
||||||
|
local_ip: 本地IP地址
|
||||||
|
port: 端口号
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
str: websocket地址
|
||||||
|
"""
|
||||||
|
server_config = self.config["server"]
|
||||||
|
websocket_config = server_config.get("websocket")
|
||||||
|
|
||||||
|
if websocket_config and "你" not in websocket_config:
|
||||||
|
return websocket_config
|
||||||
|
else:
|
||||||
|
return f"ws://{local_ip}:{port}/xiaozhi/v1/"
|
||||||
|
|
||||||
|
async def start(self):
|
||||||
|
server_config = self.config["server"]
|
||||||
|
host = server_config.get("ip", "0.0.0.0")
|
||||||
|
port = int(server_config.get("http_port", 8003))
|
||||||
|
|
||||||
|
if port:
|
||||||
|
app = web.Application()
|
||||||
|
|
||||||
|
read_config_from_api = server_config.get("read_config_from_api", False)
|
||||||
|
|
||||||
|
if not read_config_from_api:
|
||||||
|
# 如果没有开启智控台,只是单模块运行,就需要再添加简单OTA接口,用于下发websocket接口
|
||||||
|
app.add_routes(
|
||||||
|
[
|
||||||
|
web.get("/xiaozhi/ota/", self.ota_handler.handle_get),
|
||||||
|
web.post("/xiaozhi/ota/", self.ota_handler.handle_post),
|
||||||
|
web.options("/xiaozhi/ota/", self.ota_handler.handle_post),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
# 添加路由
|
||||||
|
app.add_routes(
|
||||||
|
[
|
||||||
|
web.get("/mcp/vision/explain", self.vision_handler.handle_get),
|
||||||
|
web.post("/mcp/vision/explain", self.vision_handler.handle_post),
|
||||||
|
web.options("/mcp/vision/explain", self.vision_handler.handle_post),
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
# 运行服务
|
||||||
|
runner = web.AppRunner(app)
|
||||||
|
await runner.setup()
|
||||||
|
site = web.TCPSite(runner, host, port)
|
||||||
|
await site.start()
|
||||||
|
|
||||||
|
# 保持服务运行
|
||||||
|
while True:
|
||||||
|
await asyncio.sleep(3600) # 每隔 1 小时检查一次
|
||||||
@@ -9,6 +9,7 @@ from mcp import ClientSession, StdioServerParameters
|
|||||||
from mcp.client.stdio import stdio_client
|
from mcp.client.stdio import stdio_client
|
||||||
from mcp.client.sse import sse_client
|
from mcp.client.sse import sse_client
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
|
from core.utils.util import sanitize_tool_name
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
|
|
||||||
@@ -23,7 +24,9 @@ class MCPClient:
|
|||||||
self._shutdown_evt = asyncio.Event()
|
self._shutdown_evt = asyncio.Event()
|
||||||
|
|
||||||
self.session: Optional[ClientSession] = None
|
self.session: Optional[ClientSession] = None
|
||||||
self.tools: List = []
|
self.tools: List = [] # original tool objects
|
||||||
|
self.tools_dict: Dict[str, Any] = {}
|
||||||
|
self.name_mapping: Dict[str, str] = {}
|
||||||
|
|
||||||
async def initialize(self):
|
async def initialize(self):
|
||||||
if self._worker_task:
|
if self._worker_task:
|
||||||
@@ -32,7 +35,7 @@ class MCPClient:
|
|||||||
await self._ready_evt.wait()
|
await self._ready_evt.wait()
|
||||||
|
|
||||||
self.logger.bind(tag=TAG).info(
|
self.logger.bind(tag=TAG).info(
|
||||||
f"Connected, tools = {[t.name for t in self.tools]}"
|
f"Connected, tools = {[name for name in self.name_mapping.values()]}"
|
||||||
)
|
)
|
||||||
|
|
||||||
async def cleanup(self):
|
async def cleanup(self):
|
||||||
@@ -48,27 +51,28 @@ class MCPClient:
|
|||||||
self._worker_task = None
|
self._worker_task = None
|
||||||
|
|
||||||
def has_tool(self, name: str) -> bool:
|
def has_tool(self, name: str) -> bool:
|
||||||
return any(t.name == name for t in self.tools)
|
return name in self.tools_dict
|
||||||
|
|
||||||
def get_available_tools(self):
|
def get_available_tools(self):
|
||||||
return [
|
return [
|
||||||
{
|
{
|
||||||
"type": "function",
|
"type": "function",
|
||||||
"function": {
|
"function": {
|
||||||
"name": t.name,
|
"name": name,
|
||||||
"description": t.description,
|
"description": tool.description,
|
||||||
"parameters": t.inputSchema,
|
"parameters": tool.inputSchema,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
for t in self.tools
|
for name, tool in self.tools_dict.items()
|
||||||
]
|
]
|
||||||
|
|
||||||
async def call_tool(self, name: str, args: dict):
|
async def call_tool(self, name: str, args: dict):
|
||||||
if not self.session:
|
if not self.session:
|
||||||
raise RuntimeError("MCPClient not initialized")
|
raise RuntimeError("MCPClient not initialized")
|
||||||
|
|
||||||
|
real_name = self.name_mapping.get(name, name)
|
||||||
loop = self._worker_task.get_loop()
|
loop = self._worker_task.get_loop()
|
||||||
coro = self.session.call_tool(name, args)
|
coro = self.session.call_tool(real_name, args)
|
||||||
|
|
||||||
if loop is asyncio.get_running_loop():
|
if loop is asyncio.get_running_loop():
|
||||||
return await coro
|
return await coro
|
||||||
@@ -92,11 +96,21 @@ class MCPClient:
|
|||||||
args=self.config.get("args", []),
|
args=self.config.get("args", []),
|
||||||
env=env,
|
env=env,
|
||||||
)
|
)
|
||||||
stdio_r, stdio_w = await stack.enter_async_context(stdio_client(params))
|
stdio_r, stdio_w = await stack.enter_async_context(
|
||||||
|
stdio_client(params)
|
||||||
|
)
|
||||||
read_stream, write_stream = stdio_r, stdio_w
|
read_stream, write_stream = stdio_r, stdio_w
|
||||||
# 建立SSEClient
|
# 建立SSEClient
|
||||||
elif "url" in self.config:
|
elif "url" in self.config:
|
||||||
sse_r, sse_w = await stack.enter_async_context(sse_client(self.config["url"]))
|
if "API_ACCESS_TOKEN" in self.config:
|
||||||
|
headers = {
|
||||||
|
"Authorization": f"Bearer {self.config['API_ACCESS_TOKEN']}"
|
||||||
|
}
|
||||||
|
else:
|
||||||
|
headers = {}
|
||||||
|
sse_r, sse_w = await stack.enter_async_context(
|
||||||
|
sse_client(self.config["url"], headers=headers)
|
||||||
|
)
|
||||||
read_stream, write_stream = sse_r, sse_w
|
read_stream, write_stream = sse_r, sse_w
|
||||||
|
|
||||||
else:
|
else:
|
||||||
@@ -113,6 +127,10 @@ class MCPClient:
|
|||||||
|
|
||||||
# 获取工具
|
# 获取工具
|
||||||
self.tools = (await self.session.list_tools()).tools
|
self.tools = (await self.session.list_tools()).tools
|
||||||
|
for t in self.tools:
|
||||||
|
sanitized = sanitize_tool_name(t.name)
|
||||||
|
self.tools_dict[sanitized] = t
|
||||||
|
self.name_mapping[sanitized] = t.name
|
||||||
|
|
||||||
self._ready_evt.set()
|
self._ready_evt.set()
|
||||||
|
|
||||||
|
|||||||
@@ -213,7 +213,7 @@ class ASRProvider(ASRProviderBase):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
async def speech_to_text(
|
async def speech_to_text(
|
||||||
self, opus_data: List[bytes], session_id: str
|
self, opus_data: List[bytes], session_id: str, audio_format="opus"
|
||||||
) -> Tuple[Optional[str], Optional[str]]:
|
) -> Tuple[Optional[str], Optional[str]]:
|
||||||
"""将语音数据转换为文本"""
|
"""将语音数据转换为文本"""
|
||||||
if self._is_token_expired():
|
if self._is_token_expired():
|
||||||
@@ -223,7 +223,7 @@ class ASRProvider(ASRProviderBase):
|
|||||||
file_path = None
|
file_path = None
|
||||||
try:
|
try:
|
||||||
# 解码Opus为PCM
|
# 解码Opus为PCM
|
||||||
if self.audio_format == "pcm":
|
if audio_format == "pcm":
|
||||||
pcm_data = opus_data
|
pcm_data = opus_data
|
||||||
else:
|
else:
|
||||||
pcm_data = self.decode_opus(opus_data)
|
pcm_data = self.decode_opus(opus_data)
|
||||||
|
|||||||
@@ -30,7 +30,7 @@ class ASRProvider(ASRProviderBase):
|
|||||||
os.makedirs(self.output_dir, exist_ok=True)
|
os.makedirs(self.output_dir, exist_ok=True)
|
||||||
|
|
||||||
async def speech_to_text(
|
async def speech_to_text(
|
||||||
self, opus_data: List[bytes], session_id: str
|
self, opus_data: List[bytes], session_id: str, audio_format="opus"
|
||||||
) -> Tuple[Optional[str], Optional[str]]:
|
) -> Tuple[Optional[str], Optional[str]]:
|
||||||
"""将语音数据转换为文本"""
|
"""将语音数据转换为文本"""
|
||||||
if not opus_data:
|
if not opus_data:
|
||||||
@@ -45,7 +45,7 @@ class ASRProvider(ASRProviderBase):
|
|||||||
return None, file_path
|
return None, file_path
|
||||||
|
|
||||||
# 将Opus音频数据解码为PCM
|
# 将Opus音频数据解码为PCM
|
||||||
if self.audio_format == "pcm":
|
if audio_format == "pcm":
|
||||||
pcm_data = opus_data
|
pcm_data = opus_data
|
||||||
else:
|
else:
|
||||||
pcm_data = self.decode_opus(opus_data)
|
pcm_data = self.decode_opus(opus_data)
|
||||||
|
|||||||
@@ -1,15 +1,19 @@
|
|||||||
import os
|
import os
|
||||||
import time
|
import wave
|
||||||
import copy
|
import copy
|
||||||
import uuid
|
import uuid
|
||||||
import wave
|
import queue
|
||||||
|
import asyncio
|
||||||
|
import traceback
|
||||||
|
import threading
|
||||||
import opuslib_next
|
import opuslib_next
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from typing import Optional, Tuple, List
|
from typing import Optional, Tuple, List
|
||||||
from core.utils.util import remove_punctuation_and_length
|
|
||||||
from core.handle.reportHandle import enqueue_asr_report
|
|
||||||
from core.handle.receiveAudioHandle import startToChat
|
from core.handle.receiveAudioHandle import startToChat
|
||||||
|
from core.handle.reportHandle import enqueue_asr_report
|
||||||
|
from core.utils.util import remove_punctuation_and_length
|
||||||
|
from core.handle.receiveAudioHandle import handleAudioMessage
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
@@ -17,53 +21,75 @@ logger = setup_logging()
|
|||||||
|
|
||||||
class ASRProviderBase(ABC):
|
class ASRProviderBase(ABC):
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.audio_format = "opus"
|
pass
|
||||||
self.conn = None
|
|
||||||
|
|
||||||
# 打开音频通道
|
# 打开音频通道
|
||||||
# 这里默认是非流式的处理方式
|
# 这里默认是非流式的处理方式
|
||||||
# 流式处理方式请在子类中重写
|
# 流式处理方式请在子类中重写
|
||||||
async def open_audio_channels(self, conn):
|
async def open_audio_channels(self, conn):
|
||||||
self.conn = conn
|
# tts 消化线程
|
||||||
|
conn.asr_priority_thread = threading.Thread(
|
||||||
|
target=self.asr_text_priority_thread, args=(conn,), daemon=True
|
||||||
|
)
|
||||||
|
conn.asr_priority_thread.start()
|
||||||
|
|
||||||
|
# 有序处理ASR音频
|
||||||
|
def asr_text_priority_thread(self, conn):
|
||||||
|
while not conn.stop_event.is_set():
|
||||||
|
try:
|
||||||
|
message = conn.asr_audio_queue.get(timeout=1)
|
||||||
|
future = asyncio.run_coroutine_threadsafe(
|
||||||
|
handleAudioMessage(conn, message),
|
||||||
|
conn.loop,
|
||||||
|
)
|
||||||
|
future.result()
|
||||||
|
except queue.Empty:
|
||||||
|
continue
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(
|
||||||
|
f"处理ASR文本失败: {str(e)}, 类型: {type(e).__name__}, 堆栈: {traceback.format_exc()}"
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
# 接收音频
|
# 接收音频
|
||||||
# 这里默认是非流式的处理方式
|
# 这里默认是非流式的处理方式
|
||||||
# 流式处理方式请在子类中重写
|
# 流式处理方式请在子类中重写
|
||||||
async def receive_audio(self, audio, audio_have_voice):
|
async def receive_audio(self, conn, audio, audio_have_voice):
|
||||||
if (
|
if conn.client_listen_mode == "auto" or conn.client_listen_mode == "realtime":
|
||||||
self.conn.client_listen_mode == "auto"
|
|
||||||
or self.conn.client_listen_mode == "realtime"
|
|
||||||
):
|
|
||||||
have_voice = audio_have_voice
|
have_voice = audio_have_voice
|
||||||
else:
|
else:
|
||||||
have_voice = self.conn.client_have_voice
|
have_voice = conn.client_have_voice
|
||||||
# 如果本次没有声音,本段也没声音,就把声音丢弃了
|
# 如果本次没有声音,本段也没声音,就把声音丢弃了
|
||||||
self.conn.asr_audio.append(audio)
|
conn.asr_audio.append(audio)
|
||||||
if have_voice == False and self.conn.client_have_voice == False:
|
if have_voice == False and conn.client_have_voice == False:
|
||||||
self.conn.asr_audio = self.conn.asr_audio[-10:]
|
conn.asr_audio = conn.asr_audio[-10:]
|
||||||
return
|
return
|
||||||
|
|
||||||
# 如果本段有声音,且已经停止了
|
# 如果本段有声音,且已经停止了
|
||||||
if self.conn.client_voice_stop:
|
if conn.client_voice_stop:
|
||||||
asr_audio_task = copy.deepcopy(self.conn.asr_audio)
|
asr_audio_task = copy.deepcopy(conn.asr_audio)
|
||||||
self.conn.asr_audio.clear()
|
conn.asr_audio.clear()
|
||||||
|
|
||||||
# 音频太短了,无法识别
|
# 音频太短了,无法识别
|
||||||
self.conn.reset_vad_states()
|
conn.reset_vad_states()
|
||||||
if len(asr_audio_task) > 15:
|
if len(asr_audio_task) > 15:
|
||||||
await self.handle_voice_stop(asr_audio_task)
|
await self.handle_voice_stop(conn, asr_audio_task)
|
||||||
|
|
||||||
# 处理语音停止
|
# 处理语音停止
|
||||||
async def handle_voice_stop(self, asr_audio_task):
|
async def handle_voice_stop(self, conn, asr_audio_task):
|
||||||
raw_text, _ = await self.speech_to_text(
|
raw_text, _ = await self.speech_to_text(
|
||||||
asr_audio_task, self.conn.session_id
|
asr_audio_task, conn.session_id, conn.audio_format
|
||||||
) # 确保ASR模块返回原始文本
|
) # 确保ASR模块返回原始文本
|
||||||
self.conn.logger.bind(tag=TAG).info(f"识别文本: {raw_text}")
|
conn.logger.bind(tag=TAG).info(f"识别文本: {raw_text}")
|
||||||
text_len, _ = remove_punctuation_and_length(raw_text)
|
text_len, _ = remove_punctuation_and_length(raw_text)
|
||||||
|
self.stop_ws_connection()
|
||||||
if text_len > 0:
|
if text_len > 0:
|
||||||
# 使用自定义模块进行上报
|
# 使用自定义模块进行上报
|
||||||
await startToChat(self.conn, raw_text)
|
await startToChat(conn, raw_text)
|
||||||
enqueue_asr_report(self.conn, raw_text, asr_audio_task)
|
enqueue_asr_report(conn, raw_text, asr_audio_task)
|
||||||
|
|
||||||
|
def stop_ws_connection(self):
|
||||||
|
pass
|
||||||
|
|
||||||
def save_audio_to_file(self, pcm_data: List[bytes], session_id: str) -> str:
|
def save_audio_to_file(self, pcm_data: List[bytes], session_id: str) -> str:
|
||||||
"""PCM数据保存为WAV文件"""
|
"""PCM数据保存为WAV文件"""
|
||||||
@@ -81,15 +107,11 @@ class ASRProviderBase(ABC):
|
|||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
async def speech_to_text(
|
async def speech_to_text(
|
||||||
self, opus_data: List[bytes], session_id: str
|
self, opus_data: List[bytes], session_id: str, audio_format="opus"
|
||||||
) -> Tuple[Optional[str], Optional[str]]:
|
) -> Tuple[Optional[str], Optional[str]]:
|
||||||
"""将语音数据转换为文本"""
|
"""将语音数据转换为文本"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def set_audio_format(self, format: str) -> None:
|
|
||||||
"""设置音频格式"""
|
|
||||||
self.audio_format = format
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def decode_opus(opus_data: List[bytes]) -> bytes:
|
def decode_opus(opus_data: List[bytes]) -> bytes:
|
||||||
"""将Opus音频数据解码为PCM数据"""
|
"""将Opus音频数据解码为PCM数据"""
|
||||||
|
|||||||
@@ -1,534 +1,271 @@
|
|||||||
|
import time
|
||||||
|
import os
|
||||||
|
import uuid
|
||||||
import json
|
import json
|
||||||
import gzip
|
import gzip
|
||||||
import uuid
|
|
||||||
import asyncio
|
|
||||||
import websockets
|
import websockets
|
||||||
import opuslib_next
|
|
||||||
from core.providers.asr.base import ASRProviderBase
|
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
|
from typing import Optional, Tuple, List
|
||||||
|
from core.providers.asr.base import ASRProviderBase
|
||||||
from core.providers.asr.dto.dto import InterfaceType
|
from core.providers.asr.dto.dto import InterfaceType
|
||||||
import threading
|
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
logger = setup_logging()
|
logger = setup_logging()
|
||||||
|
|
||||||
CLIENT_FULL_REQUEST = 0b0001
|
CLIENT_FULL_REQUEST = 0b0001
|
||||||
CLIENT_AUDIO_ONLY_REQUEST = 0b0010
|
CLIENT_AUDIO_ONLY_REQUEST = 0b0010
|
||||||
|
|
||||||
|
NO_SEQUENCE = 0b0000
|
||||||
|
NEG_SEQUENCE = 0b0010
|
||||||
|
|
||||||
SERVER_FULL_RESPONSE = 0b1001
|
SERVER_FULL_RESPONSE = 0b1001
|
||||||
SERVER_ACK = 0b1011
|
SERVER_ACK = 0b1011
|
||||||
SERVER_ERROR_RESPONSE = 0b1111
|
SERVER_ERROR_RESPONSE = 0b1111
|
||||||
NO_SEQUENCE = 0b0000
|
|
||||||
NEG_SEQUENCE = 0b0010
|
NO_SERIALIZATION = 0b0000
|
||||||
JSON_SERIALIZATION = 0b0001
|
JSON = 0b0001
|
||||||
GZIP_COMPRESSION = 0b0001
|
THRIFT = 0b0011
|
||||||
PROTOCOL_VERSION = 0b0001
|
CUSTOM_TYPE = 0b1111
|
||||||
|
NO_COMPRESSION = 0b0000
|
||||||
|
GZIP = 0b0001
|
||||||
|
CUSTOM_COMPRESSION = 0b1111
|
||||||
|
|
||||||
|
|
||||||
|
def parse_response(res):
|
||||||
|
"""
|
||||||
|
protocol_version(4 bits), header_size(4 bits),
|
||||||
|
message_type(4 bits), message_type_specific_flags(4 bits)
|
||||||
|
serialization_method(4 bits) message_compression(4 bits)
|
||||||
|
reserved (8bits) 保留字段
|
||||||
|
header_extensions 扩展头(大小等于 8 * 4 * (header_size - 1) )
|
||||||
|
payload 类似与http 请求体
|
||||||
|
"""
|
||||||
|
protocol_version = res[0] >> 4
|
||||||
|
header_size = res[0] & 0x0F
|
||||||
|
message_type = res[1] >> 4
|
||||||
|
message_type_specific_flags = res[1] & 0x0F
|
||||||
|
serialization_method = res[2] >> 4
|
||||||
|
message_compression = res[2] & 0x0F
|
||||||
|
reserved = res[3]
|
||||||
|
header_extensions = res[4 : header_size * 4]
|
||||||
|
payload = res[header_size * 4 :]
|
||||||
|
result = {}
|
||||||
|
payload_msg = None
|
||||||
|
payload_size = 0
|
||||||
|
if message_type == SERVER_FULL_RESPONSE:
|
||||||
|
payload_size = int.from_bytes(payload[:4], "big", signed=True)
|
||||||
|
payload_msg = payload[4:]
|
||||||
|
elif message_type == SERVER_ACK:
|
||||||
|
seq = int.from_bytes(payload[:4], "big", signed=True)
|
||||||
|
result["seq"] = seq
|
||||||
|
if len(payload) >= 8:
|
||||||
|
payload_size = int.from_bytes(payload[4:8], "big", signed=False)
|
||||||
|
payload_msg = payload[8:]
|
||||||
|
elif message_type == SERVER_ERROR_RESPONSE:
|
||||||
|
code = int.from_bytes(payload[:4], "big", signed=False)
|
||||||
|
result["code"] = code
|
||||||
|
payload_size = int.from_bytes(payload[4:8], "big", signed=False)
|
||||||
|
payload_msg = payload[8:]
|
||||||
|
if payload_msg is None:
|
||||||
|
return result
|
||||||
|
if message_compression == GZIP:
|
||||||
|
payload_msg = gzip.decompress(payload_msg)
|
||||||
|
if serialization_method == JSON:
|
||||||
|
payload_msg = json.loads(str(payload_msg, "utf-8"))
|
||||||
|
elif serialization_method != NO_SERIALIZATION:
|
||||||
|
payload_msg = str(payload_msg, "utf-8")
|
||||||
|
result["payload_msg"] = payload_msg
|
||||||
|
result["payload_size"] = payload_size
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
class ASRProvider(ASRProviderBase):
|
class ASRProvider(ASRProviderBase):
|
||||||
def __init__(self, config, delete_audio_file):
|
def __init__(self, config: dict, delete_audio_file: bool):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.interface_type = InterfaceType.STREAM
|
self.interface_type = InterfaceType.NON_STREAM
|
||||||
self.config = config
|
self.appid = config.get("appid")
|
||||||
self.text = ""
|
|
||||||
self.max_retries = 3
|
|
||||||
self.retry_delay = 2 # 重试延迟秒数
|
|
||||||
self.recv_lock = asyncio.Lock() # 添加接收锁
|
|
||||||
self.reconnect_lock = asyncio.Lock() # 添加重连锁
|
|
||||||
self.last_reconnect_time = 0 # 上次重连时间
|
|
||||||
self.reconnect_cooldown = 1 # 增加重连冷却时间到10秒
|
|
||||||
self.reconnect_count = 0 # 当前重连次数
|
|
||||||
self.max_reconnect_count = 3 # 减少最大重连次数到3次
|
|
||||||
self.asr_thread = None # ASR监听线程
|
|
||||||
self.thread_lock = threading.Lock() # 线程管理锁
|
|
||||||
self.is_reconnecting = False # 添加重连状态标志
|
|
||||||
|
|
||||||
# 添加会话管理相关属性
|
|
||||||
self._session_lock = asyncio.Lock() # 会话操作的并发锁
|
|
||||||
self._current_session_id = None # 当前会话ID
|
|
||||||
self._session_started = False # 会话是否已开始
|
|
||||||
self._session_finished = False # 会话是否已结束
|
|
||||||
self._session_close_event = asyncio.Event() # 添加会话关闭事件
|
|
||||||
|
|
||||||
self.appid = str(config.get("appid"))
|
|
||||||
self.cluster = config.get("cluster")
|
self.cluster = config.get("cluster")
|
||||||
self.access_token = config.get("access_token")
|
self.access_token = config.get("access_token")
|
||||||
self.boosting_table_name = config.get("boosting_table_name", "")
|
self.boosting_table_name = config.get("boosting_table_name", "")
|
||||||
self.correct_table_name = config.get("correct_table_name", "")
|
self.correct_table_name = config.get("correct_table_name", "")
|
||||||
self.output_dir = config.get("output_dir", "temp/")
|
self.output_dir = config.get("output_dir")
|
||||||
self.delete_audio_file = delete_audio_file
|
self.delete_audio_file = delete_audio_file
|
||||||
|
|
||||||
self.ws_url = "wss://openspeech.bytedance.com/api/v2/asr"
|
self.host = "openspeech.bytedance.com"
|
||||||
self.uid = config.get("uid", "streaming_asr_service")
|
self.ws_url = f"wss://{self.host}/api/v2/asr"
|
||||||
self.workflow = config.get(
|
self.success_code = 1000
|
||||||
"workflow", "audio_in,resample,partition,vad,fe,decode,itn,nlu_punctuate"
|
self.seg_duration = 15000
|
||||||
)
|
|
||||||
self.result_type = config.get("result_type", "single")
|
|
||||||
self.format = config.get("format", "raw")
|
|
||||||
self.codec = config.get("codec", "pcm")
|
|
||||||
self.rate = config.get("sample_rate", 16000)
|
|
||||||
self.language = config.get("language", "zh-CN")
|
|
||||||
self.bits = config.get("bits", 16)
|
|
||||||
self.channel = config.get("channel", 1)
|
|
||||||
self.auth_method = config.get("auth_method", "token")
|
|
||||||
self.secret = config.get("secret", "access_secret")
|
|
||||||
self.decoder = opuslib_next.Decoder(16000, 1)
|
|
||||||
self.asr_ws = None
|
|
||||||
self.forward_task = None
|
|
||||||
self.conn = None
|
|
||||||
|
|
||||||
###################################################################################
|
# 确保输出目录存在
|
||||||
# 豆包流式ASR重写父类的方法--开始
|
os.makedirs(self.output_dir, exist_ok=True)
|
||||||
###################################################################################
|
|
||||||
async def open_audio_channels(self, conn):
|
|
||||||
await super().open_audio_channels(conn)
|
|
||||||
|
|
||||||
async with self._session_lock:
|
@staticmethod
|
||||||
# 如果正在重连,等待重连完成
|
def _generate_header(
|
||||||
if self.is_reconnecting:
|
message_type=CLIENT_FULL_REQUEST, message_type_specific_flags=NO_SEQUENCE
|
||||||
logger.bind(tag=TAG).info("等待当前重连完成...")
|
) -> bytearray:
|
||||||
await self._session_close_event.wait()
|
"""Generate protocol header."""
|
||||||
self._session_close_event.clear()
|
header = bytearray()
|
||||||
|
header_size = 1
|
||||||
|
header.append((0b0001 << 4) | header_size) # Protocol version
|
||||||
|
header.append((message_type << 4) | message_type_specific_flags)
|
||||||
|
header.append((0b0001 << 4) | 0b0001) # JSON serialization & GZIP compression
|
||||||
|
header.append(0x00) # reserved
|
||||||
|
return header
|
||||||
|
|
||||||
# 如果已有会话未结束,先关闭它
|
def _construct_request(self, reqid) -> dict:
|
||||||
if self._session_started and not self._session_finished:
|
"""Construct the request payload."""
|
||||||
logger.bind(tag=TAG).warning(
|
return {
|
||||||
f"发现未关闭的会话 {self._current_session_id},正在关闭..."
|
|
||||||
)
|
|
||||||
if self.asr_ws is not None:
|
|
||||||
try:
|
|
||||||
await self.asr_ws.close()
|
|
||||||
except Exception as e:
|
|
||||||
logger.bind(tag=TAG).warning(f"关闭旧连接时发生错误: {e}")
|
|
||||||
finally:
|
|
||||||
self.asr_ws = None
|
|
||||||
self._session_finished = True
|
|
||||||
self._session_close_event.set()
|
|
||||||
|
|
||||||
# 重置会话状态
|
|
||||||
self._current_session_id = str(uuid.uuid4())
|
|
||||||
self._session_started = True
|
|
||||||
self._session_finished = False
|
|
||||||
self.is_reconnecting = True
|
|
||||||
|
|
||||||
try:
|
|
||||||
retry_count = 0
|
|
||||||
while retry_count < self.max_retries:
|
|
||||||
try:
|
|
||||||
headers = (
|
|
||||||
self.token_auth() if self.auth_method == "token" else None
|
|
||||||
)
|
|
||||||
self.asr_ws = await websockets.connect(
|
|
||||||
self.ws_url,
|
|
||||||
additional_headers=headers,
|
|
||||||
max_size=1000000000,
|
|
||||||
ping_interval=None,
|
|
||||||
ping_timeout=None,
|
|
||||||
close_timeout=10,
|
|
||||||
)
|
|
||||||
|
|
||||||
# 发送初始化请求
|
|
||||||
request_params = self.construct_request(
|
|
||||||
self._current_session_id
|
|
||||||
)
|
|
||||||
try:
|
|
||||||
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 self.asr_ws.send(full_client_request)
|
|
||||||
except Exception as e:
|
|
||||||
logger.bind(tag=TAG).error(f"发送初始化请求失败: {e}")
|
|
||||||
raise e
|
|
||||||
|
|
||||||
# 等待初始化响应
|
|
||||||
try:
|
|
||||||
init_res = await self.asr_ws.recv()
|
|
||||||
self.parse_response(init_res)
|
|
||||||
except Exception as e:
|
|
||||||
logger.bind(tag=TAG).error(f"ASR服务初始化失败: {e}")
|
|
||||||
raise e
|
|
||||||
|
|
||||||
# 启动接收ASR结果的异步任务
|
|
||||||
with self.thread_lock:
|
|
||||||
if (
|
|
||||||
self.asr_thread is None
|
|
||||||
or not self.asr_thread.is_alive()
|
|
||||||
):
|
|
||||||
logger.bind(tag=TAG).info("创建新的ASR监听线程...")
|
|
||||||
self.asr_thread = threading.Thread(
|
|
||||||
target=self._start_monitor_asr_response_thread,
|
|
||||||
daemon=True,
|
|
||||||
)
|
|
||||||
self.asr_thread.start()
|
|
||||||
# 等待一小段时间确保线程启动
|
|
||||||
await asyncio.sleep(0.1)
|
|
||||||
if not self.asr_thread.is_alive():
|
|
||||||
logger.bind(tag=TAG).error("ASR监听线程启动失败")
|
|
||||||
raise Exception("ASR监听线程启动失败")
|
|
||||||
logger.bind(tag=TAG).info("ASR监听线程已启动")
|
|
||||||
return
|
|
||||||
|
|
||||||
except websockets.exceptions.WebSocketException as e:
|
|
||||||
retry_count += 1
|
|
||||||
if retry_count < self.max_retries:
|
|
||||||
logger.bind(tag=TAG).warning(
|
|
||||||
f"WebSocket连接失败,正在进行第{retry_count}次重试: {e}"
|
|
||||||
)
|
|
||||||
await asyncio.sleep(self.retry_delay)
|
|
||||||
else:
|
|
||||||
logger.bind(tag=TAG).warning(
|
|
||||||
f"WebSocket连接失败,已达到最大重试次数: {e}"
|
|
||||||
)
|
|
||||||
raise
|
|
||||||
except Exception as e:
|
|
||||||
logger.bind(tag=TAG).error(f"WebSocket连接发生未知错误: {e}")
|
|
||||||
raise
|
|
||||||
finally:
|
|
||||||
self.is_reconnecting = False
|
|
||||||
self._session_close_event.set()
|
|
||||||
|
|
||||||
async def receive_audio(self, audio, _):
|
|
||||||
if not isinstance(audio, bytes):
|
|
||||||
return
|
|
||||||
|
|
||||||
try:
|
|
||||||
# 解码opus得到PCM数据
|
|
||||||
pcm_frame = self.decoder.decode(audio, 960)
|
|
||||||
payload = gzip.compress(pcm_frame)
|
|
||||||
audio_request = bytearray(self.generate_audio_default_header())
|
|
||||||
audio_request.extend(len(payload).to_bytes(4, "big"))
|
|
||||||
audio_request.extend(payload)
|
|
||||||
if self.asr_ws:
|
|
||||||
await self.asr_ws.send(audio_request)
|
|
||||||
except Exception as e:
|
|
||||||
logger.bind(tag=TAG).debug(f"发送音频数据时发生错误: {e}")
|
|
||||||
|
|
||||||
###################################################################################
|
|
||||||
# 豆包流式ASR重写父类的方法--结束
|
|
||||||
###################################################################################
|
|
||||||
|
|
||||||
def construct_request(self, reqid):
|
|
||||||
req = {
|
|
||||||
"app": {
|
"app": {
|
||||||
"appid": self.appid,
|
"appid": f"{self.appid}",
|
||||||
"cluster": self.cluster,
|
"cluster": self.cluster,
|
||||||
"token": self.access_token,
|
"token": self.access_token,
|
||||||
},
|
},
|
||||||
"user": {"uid": self.uid},
|
"user": {
|
||||||
|
"uid": str(uuid.uuid4()),
|
||||||
|
},
|
||||||
"request": {
|
"request": {
|
||||||
"reqid": reqid,
|
"reqid": reqid,
|
||||||
"workflow": self.workflow,
|
"show_utterances": False,
|
||||||
"show_utterances": True,
|
|
||||||
"result_type": self.result_type,
|
|
||||||
"sequence": 1,
|
"sequence": 1,
|
||||||
"boosting_table_name": self.boosting_table_name,
|
"boosting_table_name": self.boosting_table_name,
|
||||||
"correct_table_name": self.correct_table_name,
|
"correct_table_name": self.correct_table_name,
|
||||||
},
|
},
|
||||||
"audio": {
|
"audio": {
|
||||||
"format": self.format,
|
"format": "raw",
|
||||||
"codec": self.codec,
|
"rate": 16000,
|
||||||
"rate": self.rate,
|
"language": "zh-CN",
|
||||||
"language": self.language,
|
"bits": 16,
|
||||||
"bits": self.bits,
|
"channel": 1,
|
||||||
"channel": self.channel,
|
"codec": "raw",
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
return req
|
|
||||||
|
|
||||||
def token_auth(self):
|
async def _send_request(
|
||||||
return {"Authorization": f"Bearer; {self.access_token}"}
|
self, audio_data: List[bytes], segment_size: int
|
||||||
|
) -> Optional[str]:
|
||||||
def generate_header(
|
"""Send request to Volcano ASR service."""
|
||||||
self,
|
|
||||||
version=PROTOCOL_VERSION,
|
|
||||||
message_type=CLIENT_FULL_REQUEST,
|
|
||||||
message_type_specific_flags=NO_SEQUENCE,
|
|
||||||
serial_method=JSON_SERIALIZATION,
|
|
||||||
compression_type=GZIP_COMPRESSION,
|
|
||||||
reserved_data=0x00,
|
|
||||||
extension_header: bytes = b"",
|
|
||||||
):
|
|
||||||
"""
|
|
||||||
生成协议头:
|
|
||||||
- 第1字节:高4位:协议版本,低4位:头部大小(单位 4 字节)
|
|
||||||
- 第2字节:高4位:消息类型,低4位:消息类型特定标志
|
|
||||||
- 第3字节:高4位:序列化方式,低4位:压缩方式
|
|
||||||
- 第4字节:保留字段
|
|
||||||
- 后续:扩展头(如果有)
|
|
||||||
"""
|
|
||||||
header = bytearray()
|
|
||||||
header_size = int(len(extension_header) / 4) + 1
|
|
||||||
header.append((version << 4) | header_size)
|
|
||||||
header.append((message_type << 4) | message_type_specific_flags)
|
|
||||||
header.append((serial_method << 4) | compression_type)
|
|
||||||
header.append(reserved_data)
|
|
||||||
header.extend(extension_header)
|
|
||||||
return header
|
|
||||||
|
|
||||||
def generate_full_default_header(self):
|
|
||||||
# full client request 默认头
|
|
||||||
return self.generate_header(
|
|
||||||
version=PROTOCOL_VERSION,
|
|
||||||
message_type=CLIENT_FULL_REQUEST,
|
|
||||||
message_type_specific_flags=NO_SEQUENCE,
|
|
||||||
serial_method=JSON_SERIALIZATION,
|
|
||||||
compression_type=GZIP_COMPRESSION,
|
|
||||||
)
|
|
||||||
|
|
||||||
def generate_audio_default_header(self):
|
|
||||||
# 普通音频片段请求
|
|
||||||
return self.generate_header(
|
|
||||||
version=PROTOCOL_VERSION,
|
|
||||||
message_type=CLIENT_AUDIO_ONLY_REQUEST,
|
|
||||||
message_type_specific_flags=NO_SEQUENCE,
|
|
||||||
serial_method=JSON_SERIALIZATION,
|
|
||||||
compression_type=GZIP_COMPRESSION,
|
|
||||||
)
|
|
||||||
|
|
||||||
def generate_last_audio_default_header(self):
|
|
||||||
# 最后一个音频片段标志
|
|
||||||
return self.generate_header(
|
|
||||||
version=PROTOCOL_VERSION,
|
|
||||||
message_type=CLIENT_AUDIO_ONLY_REQUEST,
|
|
||||||
message_type_specific_flags=NEG_SEQUENCE, # 用 NEG_SEQUENCE 表示结束
|
|
||||||
serial_method=JSON_SERIALIZATION,
|
|
||||||
compression_type=GZIP_COMPRESSION,
|
|
||||||
)
|
|
||||||
|
|
||||||
def _start_monitor_asr_response_thread(self):
|
|
||||||
# 初始化链接
|
|
||||||
try:
|
try:
|
||||||
with self.thread_lock:
|
auth_header = {"Authorization": "Bearer; {}".format(self.access_token)}
|
||||||
if self.conn is None or self.conn.loop is None:
|
async with websockets.connect(
|
||||||
logger.bind(tag=TAG).error(
|
self.ws_url, additional_headers=auth_header
|
||||||
"无法启动ASR监听线程:conn或loop未初始化"
|
) as websocket:
|
||||||
)
|
# Prepare request data
|
||||||
return
|
request_params = self._construct_request(str(uuid.uuid4()))
|
||||||
|
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")
|
||||||
|
) # payload size(4 bytes)
|
||||||
|
full_client_request.extend(payload_bytes) # payload
|
||||||
|
|
||||||
try:
|
# Send header and metadata
|
||||||
logger.bind(tag=TAG).info("开始启动ASR监听...")
|
# full_client_request
|
||||||
asyncio.run_coroutine_threadsafe(
|
await websocket.send(full_client_request)
|
||||||
self._forward_asr_results(), loop=self.conn.loop
|
res = await websocket.recv()
|
||||||
)
|
result = parse_response(res)
|
||||||
logger.bind(tag=TAG).info("ASR监听已启动")
|
if (
|
||||||
except Exception as e:
|
"payload_msg" in result
|
||||||
logger.bind(tag=TAG).error(f"启动ASR监听线程失败: {e}")
|
and result["payload_msg"]["code"] != self.success_code
|
||||||
except Exception as e:
|
and result["payload_msg"]["code"] != 1013 # 忽略无有效语音的错误
|
||||||
logger.bind(tag=TAG).error(f"ASR监听线程发生未预期的错误: {e}")
|
):
|
||||||
|
logger.bind(tag=TAG).error(f"ASR error: {result}")
|
||||||
|
return None
|
||||||
|
|
||||||
async def _forward_asr_results(self):
|
for seq, (chunk, last) in enumerate(
|
||||||
try:
|
self.slice_data(audio_data, segment_size), 1
|
||||||
while not self.conn.stop_event.is_set():
|
):
|
||||||
try:
|
if last:
|
||||||
if self.asr_ws is None:
|
audio_only_request = self._generate_header(
|
||||||
# 检查是否需要重连
|
message_type=CLIENT_AUDIO_ONLY_REQUEST,
|
||||||
async with self.reconnect_lock:
|
message_type_specific_flags=NEG_SEQUENCE,
|
||||||
current_time = asyncio.get_event_loop().time()
|
|
||||||
if (
|
|
||||||
current_time - self.last_reconnect_time
|
|
||||||
< self.reconnect_cooldown
|
|
||||||
):
|
|
||||||
await asyncio.sleep(1)
|
|
||||||
continue
|
|
||||||
|
|
||||||
if self.reconnect_count >= self.max_reconnect_count:
|
|
||||||
logger.bind(tag=TAG).error(
|
|
||||||
"达到最大重连次数限制,停止重连"
|
|
||||||
)
|
|
||||||
await asyncio.sleep(self.reconnect_cooldown)
|
|
||||||
self.reconnect_count = 0
|
|
||||||
continue
|
|
||||||
|
|
||||||
self.last_reconnect_time = current_time
|
|
||||||
self.reconnect_count += 1
|
|
||||||
logger.bind(tag=TAG).info(
|
|
||||||
f"尝试重新连接ASR服务... (第{self.reconnect_count}次)"
|
|
||||||
)
|
|
||||||
await self.open_audio_channels(self.conn)
|
|
||||||
continue
|
|
||||||
|
|
||||||
# 使用锁来确保同一时间只有一个协程在接收数据
|
|
||||||
async with self.recv_lock:
|
|
||||||
response = await self.asr_ws.recv()
|
|
||||||
result = self.parse_response(response)
|
|
||||||
|
|
||||||
# 检查是否需要重连
|
|
||||||
if result.get("need_reconnect", False):
|
|
||||||
logger.bind(tag=TAG).info(
|
|
||||||
"检测到需要重连的错误,准备重新连接..."
|
|
||||||
)
|
)
|
||||||
if self.asr_ws is not None:
|
else:
|
||||||
try:
|
audio_only_request = self._generate_header(
|
||||||
await self.asr_ws.close()
|
message_type=CLIENT_AUDIO_ONLY_REQUEST
|
||||||
except Exception as e:
|
)
|
||||||
logger.bind(tag=TAG).warning(
|
payload_bytes = gzip.compress(chunk)
|
||||||
f"关闭旧连接时发生错误: {e}"
|
audio_only_request.extend(
|
||||||
)
|
(len(payload_bytes)).to_bytes(4, "big")
|
||||||
finally:
|
) # payload size(4 bytes)
|
||||||
self.asr_ws = None
|
audio_only_request.extend(payload_bytes) # payload
|
||||||
continue
|
# Send audio data
|
||||||
|
await websocket.send(audio_only_request)
|
||||||
|
|
||||||
if "payload_msg" in result:
|
# Receive response
|
||||||
if "result" in result["payload_msg"]:
|
response = await websocket.recv()
|
||||||
# 检查是否有utterances并且definite为True
|
result = parse_response(response)
|
||||||
utterances = result["payload_msg"]["result"][0].get(
|
|
||||||
"utterances", []
|
|
||||||
)
|
|
||||||
for utterance in utterances:
|
|
||||||
if utterance.get("definite", False):
|
|
||||||
self.text = utterance["text"]
|
|
||||||
await self.handle_voice_stop(None)
|
|
||||||
break
|
|
||||||
|
|
||||||
except websockets.ConnectionClosed:
|
if (
|
||||||
logger.bind(tag=TAG).debug("ASR服务连接已关闭,准备重连...")
|
"payload_msg" in result
|
||||||
# 确保关闭旧连接
|
and result["payload_msg"]["code"] == self.success_code
|
||||||
if self.asr_ws is not None:
|
):
|
||||||
try:
|
if len(result["payload_msg"]["result"]) > 0:
|
||||||
await self.asr_ws.close()
|
return result["payload_msg"]["result"][0]["text"]
|
||||||
except Exception as e:
|
return None
|
||||||
logger.bind(tag=TAG).warning(f"关闭旧连接时发生错误: {e}")
|
elif "payload_msg" in result and result["payload_msg"]["code"] == 1013:
|
||||||
finally:
|
# 无有效语音,返回空字符串
|
||||||
self.asr_ws = None
|
return ""
|
||||||
|
else:
|
||||||
# 等待冷却时间
|
logger.bind(tag=TAG).error(f"ASR error: {result}")
|
||||||
await asyncio.sleep(self.reconnect_cooldown)
|
return None
|
||||||
continue
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
if not self.conn.stop_event.is_set():
|
|
||||||
logger.bind(tag=TAG).error(f"ASR监听发生错误: {e}")
|
|
||||||
await asyncio.sleep(self.retry_delay)
|
|
||||||
continue
|
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"ASR监听线程发生错误: {e}")
|
logger.bind(tag=TAG).error(f"ASR request failed: {e}", exc_info=True)
|
||||||
# 确保在发生严重错误时也能继续尝试重连
|
return None
|
||||||
if not self.conn.stop_event.is_set():
|
|
||||||
await asyncio.sleep(self.retry_delay)
|
|
||||||
await self._forward_asr_results() # 递归重试
|
|
||||||
|
|
||||||
async def speech_to_text(self, opus_data, session_id):
|
@staticmethod
|
||||||
result = self.text
|
def slice_data(data: bytes, chunk_size: int) -> (list, bool):
|
||||||
self.text = "" # 清空text
|
|
||||||
return result, None
|
|
||||||
|
|
||||||
def parse_response(self, res: bytes) -> dict:
|
|
||||||
"""
|
"""
|
||||||
解析 ASR 服务返回的二进制响应。
|
slice data
|
||||||
根据协议格式解析头部和 payload,若采用 GZIP 压缩则先解压,再根据 JSON 反序列化。
|
:param data: wav data
|
||||||
|
:param chunk_size: the segment size in one request
|
||||||
|
:return: segment data, last flag
|
||||||
"""
|
"""
|
||||||
protocol_version = res[0] >> 4
|
data_len = len(data)
|
||||||
header_size = res[0] & 0x0F
|
offset = 0
|
||||||
message_type = res[1] >> 4
|
while offset + chunk_size < data_len:
|
||||||
serialization_method = res[2] >> 4
|
yield data[offset : offset + chunk_size], False
|
||||||
message_compression = res[2] & 0x0F
|
offset += chunk_size
|
||||||
payload = res[header_size * 4 :]
|
|
||||||
result = {}
|
|
||||||
payload_msg = None
|
|
||||||
payload_size = 0
|
|
||||||
|
|
||||||
if message_type == SERVER_FULL_RESPONSE:
|
|
||||||
payload_size = int.from_bytes(payload[:4], "big", signed=True)
|
|
||||||
payload_msg = payload[4:]
|
|
||||||
elif message_type == SERVER_ACK:
|
|
||||||
seq = int.from_bytes(payload[:4], "big", signed=True)
|
|
||||||
result["seq"] = seq
|
|
||||||
if len(payload) >= 8:
|
|
||||||
payload_size = int.from_bytes(payload[4:8], "big", signed=False)
|
|
||||||
payload_msg = payload[8:]
|
|
||||||
elif message_type == SERVER_ERROR_RESPONSE:
|
|
||||||
code = int.from_bytes(payload[:4], "big", signed=False)
|
|
||||||
result["code"] = code
|
|
||||||
payload_size = int.from_bytes(payload[4:8], "big", signed=False)
|
|
||||||
payload_msg = payload[8:]
|
|
||||||
|
|
||||||
if payload_msg is None:
|
|
||||||
return result
|
|
||||||
if message_compression == GZIP_COMPRESSION:
|
|
||||||
payload_msg = gzip.decompress(payload_msg)
|
|
||||||
if serialization_method == JSON_SERIALIZATION:
|
|
||||||
payload_msg = json.loads(payload_msg.decode("utf-8"))
|
|
||||||
else:
|
else:
|
||||||
payload_msg = payload_msg.decode("utf-8")
|
yield data[offset:data_len], True
|
||||||
result["payload_msg"] = payload_msg
|
|
||||||
result["payload_size"] = payload_size
|
|
||||||
|
|
||||||
# 错误码处理
|
async def speech_to_text(
|
||||||
if "code" in result:
|
self, opus_data: List[bytes], session_id: str, audio_format="opus"
|
||||||
error_code = result["code"]
|
) -> Tuple[Optional[str], Optional[str]]:
|
||||||
error_message = ""
|
"""将语音数据转换为文本"""
|
||||||
|
|
||||||
if error_code == 1000:
|
file_path = None
|
||||||
error_message = "成功"
|
try:
|
||||||
elif error_code == 1001:
|
# 合并所有opus数据包
|
||||||
error_message = "请求参数无效:请求参数缺失必需字段/字段值无效/重复请求"
|
if audio_format == "pcm":
|
||||||
elif error_code == 1002:
|
pcm_data = opus_data
|
||||||
error_message = "无访问权限:token无效/过期/无权访问指定服务"
|
|
||||||
elif error_code == 1003:
|
|
||||||
error_message = "访问超频:当前appid访问QPS超出设定阈值"
|
|
||||||
elif error_code == 1004:
|
|
||||||
error_message = "访问超额:当前appid访问次数超出限制"
|
|
||||||
elif error_code == 1005:
|
|
||||||
error_message = "服务器繁忙:服务过载,无法处理当前请求"
|
|
||||||
elif error_code == 1010:
|
|
||||||
error_message = "音频过长:音频数据时长超出阈值"
|
|
||||||
elif error_code == 1011:
|
|
||||||
error_message = "音频过大:音频数据大小超出阈值"
|
|
||||||
elif error_code == 1012:
|
|
||||||
error_message = "音频格式无效:音频header有误/无法进行音频解码"
|
|
||||||
elif error_code == 1013:
|
|
||||||
error_message = "音频静音:音频未识别出任何文本结果"
|
|
||||||
elif error_code >= 1020 and error_code <= 1022:
|
|
||||||
error_message = "识别相关错误:需要重连"
|
|
||||||
if error_code == 1020:
|
|
||||||
error_message = "识别等待超时:等待下一包就绪超时"
|
|
||||||
elif error_code == 1021:
|
|
||||||
error_message = "识别处理超时:识别处理过程超时"
|
|
||||||
elif error_code == 1022:
|
|
||||||
error_message = "识别错误:识别过程中发生错误"
|
|
||||||
else:
|
else:
|
||||||
error_message = "未知错误:未归类错误"
|
pcm_data = self.decode_opus(opus_data)
|
||||||
|
combined_pcm_data = b"".join(pcm_data)
|
||||||
|
|
||||||
logger.bind(tag=TAG).debug(
|
# 判断是否保存为WAV文件
|
||||||
f"ASR错误: {error_message} (错误码: {error_code})"
|
if self.delete_audio_file:
|
||||||
)
|
pass
|
||||||
|
else:
|
||||||
|
file_path = self.save_audio_to_file(pcm_data, session_id)
|
||||||
|
|
||||||
# 如果是识别相关错误,标记需要重连
|
# 直接使用PCM数据
|
||||||
if error_code >= 1020 or error_code == 1001:
|
# 计算分段大小 (单声道, 16bit, 16kHz采样率)
|
||||||
result["need_reconnect"] = True
|
size_per_sec = 1 * 2 * 16000 # nchannels * sampwidth * framerate
|
||||||
|
segment_size = int(size_per_sec * self.seg_duration / 1000)
|
||||||
|
|
||||||
return result
|
# 语音识别
|
||||||
|
start_time = time.time()
|
||||||
async def close_session(self):
|
text = await self._send_request(combined_pcm_data, segment_size)
|
||||||
"""关闭当前会话"""
|
if text:
|
||||||
async with self._session_lock:
|
logger.bind(tag=TAG).debug(
|
||||||
if not self._session_started:
|
f"语音识别耗时: {time.time() - start_time:.3f}s | 结果: {text}"
|
||||||
logger.bind(tag=TAG).warning("尝试关闭未开始的会话")
|
|
||||||
return
|
|
||||||
|
|
||||||
if self._session_finished:
|
|
||||||
logger.bind(tag=TAG).warning(
|
|
||||||
f"会话 {self._current_session_id} 已经关闭"
|
|
||||||
)
|
)
|
||||||
return
|
return text, file_path
|
||||||
|
return "", file_path
|
||||||
|
|
||||||
try:
|
except Exception as e:
|
||||||
if self.asr_ws is not None:
|
logger.bind(tag=TAG).error(f"语音识别失败: {e}", exc_info=True)
|
||||||
await self.asr_ws.close()
|
return "", file_path
|
||||||
except Exception as e:
|
|
||||||
logger.bind(tag=TAG).warning(f"关闭WebSocket连接时发生错误: {e}")
|
|
||||||
finally:
|
|
||||||
self.asr_ws = None
|
|
||||||
self._session_finished = True
|
|
||||||
self._session_started = False
|
|
||||||
self._current_session_id = None
|
|
||||||
# 重置重连计数
|
|
||||||
self.reconnect_count = 0
|
|
||||||
|
|
||||||
async def close(self):
|
|
||||||
"""资源清理方法"""
|
|
||||||
await self.close_session()
|
|
||||||
|
|||||||
@@ -0,0 +1,349 @@
|
|||||||
|
import json
|
||||||
|
import gzip
|
||||||
|
import uuid
|
||||||
|
import asyncio
|
||||||
|
import websockets
|
||||||
|
import opuslib_next
|
||||||
|
from core.providers.asr.base import ASRProviderBase
|
||||||
|
from config.logger import setup_logging
|
||||||
|
from core.providers.asr.dto.dto import InterfaceType
|
||||||
|
|
||||||
|
TAG = __name__
|
||||||
|
logger = setup_logging()
|
||||||
|
|
||||||
|
|
||||||
|
class ASRProvider(ASRProviderBase):
|
||||||
|
def __init__(self, config, delete_audio_file):
|
||||||
|
super().__init__()
|
||||||
|
self.interface_type = InterfaceType.STREAM
|
||||||
|
self.config = config
|
||||||
|
self.text = ""
|
||||||
|
self.max_retries = 3
|
||||||
|
self.retry_delay = 2
|
||||||
|
self.decoder = opuslib_next.Decoder(16000, 1)
|
||||||
|
self.asr_ws = None
|
||||||
|
self.forward_task = None
|
||||||
|
self.is_processing = False # 添加处理状态标志
|
||||||
|
|
||||||
|
# 配置参数
|
||||||
|
self.appid = str(config.get("appid"))
|
||||||
|
self.cluster = config.get("cluster")
|
||||||
|
self.access_token = config.get("access_token")
|
||||||
|
self.boosting_table_name = config.get("boosting_table_name", "")
|
||||||
|
self.correct_table_name = config.get("correct_table_name", "")
|
||||||
|
self.output_dir = config.get("output_dir", "tmp/")
|
||||||
|
self.delete_audio_file = delete_audio_file
|
||||||
|
|
||||||
|
# 火山引擎ASR配置
|
||||||
|
self.ws_url = "wss://openspeech.bytedance.com/api/v3/sauc/bigmodel"
|
||||||
|
self.uid = config.get("uid", "streaming_asr_service")
|
||||||
|
self.workflow = config.get(
|
||||||
|
"workflow", "audio_in,resample,partition,vad,fe,decode,itn,nlu_punctuate"
|
||||||
|
)
|
||||||
|
self.result_type = config.get("result_type", "single")
|
||||||
|
self.format = config.get("format", "pcm")
|
||||||
|
self.codec = config.get("codec", "pcm")
|
||||||
|
self.rate = config.get("sample_rate", 16000)
|
||||||
|
self.language = config.get("language", "zh-CN")
|
||||||
|
self.bits = config.get("bits", 16)
|
||||||
|
self.channel = config.get("channel", 1)
|
||||||
|
self.auth_method = config.get("auth_method", "token")
|
||||||
|
self.secret = config.get("secret", "access_secret")
|
||||||
|
|
||||||
|
async def open_audio_channels(self, conn):
|
||||||
|
await super().open_audio_channels(conn)
|
||||||
|
|
||||||
|
async def receive_audio(self, conn, audio, audio_have_voice):
|
||||||
|
conn.asr_audio.append(audio)
|
||||||
|
conn.asr_audio = conn.asr_audio[-10:]
|
||||||
|
|
||||||
|
# 如果本次有声音,且之前没有建立连接
|
||||||
|
if audio_have_voice and self.asr_ws is None and not self.is_processing:
|
||||||
|
try:
|
||||||
|
self.is_processing = True
|
||||||
|
# 建立新的WebSocket连接
|
||||||
|
headers = self.token_auth() if self.auth_method == "token" else None
|
||||||
|
logger.bind(tag=TAG).info(f"正在连接ASR服务,headers: {headers}")
|
||||||
|
|
||||||
|
self.asr_ws = await websockets.connect(
|
||||||
|
self.ws_url,
|
||||||
|
additional_headers=headers,
|
||||||
|
max_size=1000000000,
|
||||||
|
ping_interval=None,
|
||||||
|
ping_timeout=None,
|
||||||
|
close_timeout=10,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 发送初始化请求
|
||||||
|
request_params = self.construct_request(str(uuid.uuid4()))
|
||||||
|
try:
|
||||||
|
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)
|
||||||
|
|
||||||
|
logger.bind(tag=TAG).info(f"发送初始化请求: {request_params}")
|
||||||
|
await self.asr_ws.send(full_client_request)
|
||||||
|
|
||||||
|
# 等待初始化响应
|
||||||
|
init_res = await self.asr_ws.recv()
|
||||||
|
result = self.parse_response(init_res)
|
||||||
|
logger.bind(tag=TAG).info(f"收到初始化响应: {result}")
|
||||||
|
|
||||||
|
# 检查初始化响应
|
||||||
|
if "code" in result and result["code"] != 1000:
|
||||||
|
error_msg = f"ASR服务初始化失败: {result.get('payload_msg', {}).get('message', '未知错误')}"
|
||||||
|
if "payload_msg" in result:
|
||||||
|
error_msg += f"\n详细错误信息: {json.dumps(result['payload_msg'], ensure_ascii=False)}"
|
||||||
|
logger.bind(tag=TAG).error(error_msg)
|
||||||
|
raise Exception(error_msg)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"发送初始化请求失败: {str(e)}")
|
||||||
|
if hasattr(e, "__cause__") and e.__cause__:
|
||||||
|
logger.bind(tag=TAG).error(f"错误原因: {str(e.__cause__)}")
|
||||||
|
raise e
|
||||||
|
|
||||||
|
# 启动接收ASR结果的异步任务
|
||||||
|
self.forward_task = asyncio.create_task(self._forward_asr_results(conn))
|
||||||
|
|
||||||
|
# 发送缓存的音频数据
|
||||||
|
if conn.asr_audio and len(conn.asr_audio) > 0:
|
||||||
|
for cached_audio in conn.asr_audio[-10:]:
|
||||||
|
try:
|
||||||
|
pcm_frame = self.decoder.decode(cached_audio, 960)
|
||||||
|
payload = gzip.compress(pcm_frame)
|
||||||
|
audio_request = bytearray(
|
||||||
|
self.generate_audio_default_header()
|
||||||
|
)
|
||||||
|
audio_request.extend(len(payload).to_bytes(4, "big"))
|
||||||
|
audio_request.extend(payload)
|
||||||
|
await self.asr_ws.send(audio_request)
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).info(
|
||||||
|
f"发送缓存音频数据时发生错误: {e}"
|
||||||
|
)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"建立ASR连接失败: {str(e)}")
|
||||||
|
if hasattr(e, "__cause__") and e.__cause__:
|
||||||
|
logger.bind(tag=TAG).error(f"错误原因: {str(e.__cause__)}")
|
||||||
|
if self.asr_ws:
|
||||||
|
await self.asr_ws.close()
|
||||||
|
self.asr_ws = None
|
||||||
|
self.is_processing = False
|
||||||
|
return
|
||||||
|
|
||||||
|
# 发送当前音频数据
|
||||||
|
if self.asr_ws and self.is_processing:
|
||||||
|
try:
|
||||||
|
pcm_frame = self.decoder.decode(audio, 960)
|
||||||
|
payload = gzip.compress(pcm_frame)
|
||||||
|
audio_request = bytearray(self.generate_audio_default_header())
|
||||||
|
audio_request.extend(len(payload).to_bytes(4, "big"))
|
||||||
|
audio_request.extend(payload)
|
||||||
|
await self.asr_ws.send(audio_request)
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).info(f"发送音频数据时发生错误: {e}")
|
||||||
|
|
||||||
|
async def _forward_asr_results(self, conn):
|
||||||
|
try:
|
||||||
|
while self.asr_ws and not conn.stop_event.is_set():
|
||||||
|
try:
|
||||||
|
response = await self.asr_ws.recv()
|
||||||
|
result = self.parse_response(response)
|
||||||
|
logger.bind(tag=TAG).debug(f"收到ASR结果: {result}")
|
||||||
|
|
||||||
|
if "payload_msg" in result:
|
||||||
|
payload = result["payload_msg"]
|
||||||
|
# 检查是否是错误码1013(无有效语音)
|
||||||
|
if "code" in payload and payload["code"] == 1013:
|
||||||
|
# 静默处理,不记录错误日志
|
||||||
|
continue
|
||||||
|
|
||||||
|
if "result" in payload:
|
||||||
|
utterances = payload["result"].get("utterances", [])
|
||||||
|
# 检查duration和空文本的情况
|
||||||
|
if (
|
||||||
|
payload.get("audio_info", {}).get("duration", 0) > 2000
|
||||||
|
and not utterances
|
||||||
|
and not payload["result"].get("text")
|
||||||
|
):
|
||||||
|
logger.bind(tag=TAG).error(f"识别文本:空")
|
||||||
|
self.text = ""
|
||||||
|
conn.reset_vad_states()
|
||||||
|
await self.handle_voice_stop(conn, None)
|
||||||
|
break
|
||||||
|
|
||||||
|
for utterance in utterances:
|
||||||
|
if utterance.get("definite", False):
|
||||||
|
self.text = utterance["text"]
|
||||||
|
logger.bind(tag=TAG).info(
|
||||||
|
f"识别到文本: {self.text}"
|
||||||
|
)
|
||||||
|
conn.reset_vad_states()
|
||||||
|
await self.handle_voice_stop(conn, None)
|
||||||
|
break
|
||||||
|
elif "error" in payload:
|
||||||
|
error_msg = payload.get("error", "未知错误")
|
||||||
|
logger.bind(tag=TAG).error(f"ASR服务返回错误: {error_msg}")
|
||||||
|
break
|
||||||
|
|
||||||
|
except websockets.ConnectionClosed:
|
||||||
|
logger.bind(tag=TAG).info("ASR服务连接已关闭")
|
||||||
|
self.is_processing = False
|
||||||
|
break
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"处理ASR结果时发生错误: {str(e)}")
|
||||||
|
if hasattr(e, "__cause__") and e.__cause__:
|
||||||
|
logger.bind(tag=TAG).error(f"错误原因: {str(e.__cause__)}")
|
||||||
|
self.is_processing = False
|
||||||
|
break
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"ASR结果转发任务发生错误: {str(e)}")
|
||||||
|
if hasattr(e, "__cause__") and e.__cause__:
|
||||||
|
logger.bind(tag=TAG).error(f"错误原因: {str(e.__cause__)}")
|
||||||
|
finally:
|
||||||
|
if self.asr_ws:
|
||||||
|
await self.asr_ws.close()
|
||||||
|
self.asr_ws = None
|
||||||
|
self.is_processing = False
|
||||||
|
|
||||||
|
def stop_ws_connection(self):
|
||||||
|
if self.asr_ws:
|
||||||
|
asyncio.create_task(self.asr_ws.close())
|
||||||
|
self.asr_ws = None
|
||||||
|
self.is_processing = False
|
||||||
|
|
||||||
|
def construct_request(self, reqid):
|
||||||
|
req = {
|
||||||
|
"app": {
|
||||||
|
"appid": self.appid,
|
||||||
|
"cluster": self.cluster,
|
||||||
|
"token": self.access_token,
|
||||||
|
},
|
||||||
|
"user": {"uid": self.uid},
|
||||||
|
"request": {
|
||||||
|
"reqid": reqid,
|
||||||
|
"workflow": self.workflow,
|
||||||
|
"show_utterances": True,
|
||||||
|
"result_type": self.result_type,
|
||||||
|
"sequence": 1,
|
||||||
|
"boosting_table_name": self.boosting_table_name,
|
||||||
|
"correct_table_name": self.correct_table_name,
|
||||||
|
"end_window_size": 200,
|
||||||
|
},
|
||||||
|
"audio": {
|
||||||
|
"format": self.format,
|
||||||
|
"codec": self.codec,
|
||||||
|
"rate": self.rate,
|
||||||
|
"language": self.language,
|
||||||
|
"bits": self.bits,
|
||||||
|
"channel": self.channel,
|
||||||
|
"sample_rate": self.rate,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
logger.bind(tag=TAG).debug(
|
||||||
|
f"构造请求参数: {json.dumps(req, ensure_ascii=False)}"
|
||||||
|
)
|
||||||
|
return req
|
||||||
|
|
||||||
|
def token_auth(self):
|
||||||
|
return {
|
||||||
|
"X-Api-App-Key": self.appid,
|
||||||
|
"X-Api-Access-Key": self.access_token,
|
||||||
|
"X-Api-Resource-Id": "volc.bigasr.sauc.duration",
|
||||||
|
"X-Api-Connect-Id": str(uuid.uuid4()),
|
||||||
|
"Host": "openspeech.bytedance.com",
|
||||||
|
}
|
||||||
|
|
||||||
|
def generate_header(
|
||||||
|
self,
|
||||||
|
version=0x01,
|
||||||
|
message_type=0x01,
|
||||||
|
message_type_specific_flags=0x00,
|
||||||
|
serial_method=0x01,
|
||||||
|
compression_type=0x01,
|
||||||
|
reserved_data=0x00,
|
||||||
|
extension_header: bytes = b"",
|
||||||
|
):
|
||||||
|
header = bytearray()
|
||||||
|
header_size = int(len(extension_header) / 4) + 1
|
||||||
|
header.append((version << 4) | header_size)
|
||||||
|
header.append((message_type << 4) | message_type_specific_flags)
|
||||||
|
header.append((serial_method << 4) | compression_type)
|
||||||
|
header.append(reserved_data)
|
||||||
|
header.extend(extension_header)
|
||||||
|
return header
|
||||||
|
|
||||||
|
def generate_audio_default_header(self):
|
||||||
|
return self.generate_header(
|
||||||
|
version=0x01,
|
||||||
|
message_type=0x02,
|
||||||
|
message_type_specific_flags=0x00,
|
||||||
|
serial_method=0x01,
|
||||||
|
compression_type=0x01,
|
||||||
|
)
|
||||||
|
|
||||||
|
def generate_last_audio_default_header(self):
|
||||||
|
return self.generate_header(
|
||||||
|
version=0x01,
|
||||||
|
message_type=0x02,
|
||||||
|
message_type_specific_flags=0x02,
|
||||||
|
serial_method=0x01,
|
||||||
|
compression_type=0x01,
|
||||||
|
)
|
||||||
|
|
||||||
|
def parse_response(self, res: bytes) -> dict:
|
||||||
|
try:
|
||||||
|
# 检查响应长度
|
||||||
|
if len(res) < 4:
|
||||||
|
logger.bind(tag=TAG).error(f"响应数据长度不足: {len(res)}")
|
||||||
|
return {"error": "响应数据长度不足"}
|
||||||
|
|
||||||
|
# 获取消息头
|
||||||
|
header = res[:4]
|
||||||
|
message_type = header[1] >> 4
|
||||||
|
|
||||||
|
# 如果是错误响应
|
||||||
|
if message_type == 0x0F: # SERVER_ERROR_RESPONSE
|
||||||
|
code = int.from_bytes(header[4:8], "big", signed=False)
|
||||||
|
error_msg = res[8:].decode("utf-8")
|
||||||
|
return {"code": code, "error": error_msg}
|
||||||
|
|
||||||
|
# 获取JSON数据(跳过12字节头部)
|
||||||
|
try:
|
||||||
|
json_data = res[12:].decode("utf-8")
|
||||||
|
result = json.loads(json_data)
|
||||||
|
logger.bind(tag=TAG).debug(f"成功解析JSON响应: {result}")
|
||||||
|
return {"payload_msg": result}
|
||||||
|
except (UnicodeDecodeError, json.JSONDecodeError) as e:
|
||||||
|
logger.bind(tag=TAG).error(f"JSON解析失败: {str(e)}")
|
||||||
|
logger.bind(tag=TAG).error(f"原始数据: {res}")
|
||||||
|
raise
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"解析响应失败: {str(e)}")
|
||||||
|
logger.bind(tag=TAG).error(f"原始响应数据: {res.hex()}")
|
||||||
|
raise
|
||||||
|
|
||||||
|
async def speech_to_text(self, opus_data, session_id, audio_format):
|
||||||
|
result = self.text
|
||||||
|
self.text = "" # 清空text
|
||||||
|
return result, None
|
||||||
|
|
||||||
|
async def close(self):
|
||||||
|
"""资源清理方法"""
|
||||||
|
if self.asr_ws:
|
||||||
|
await self.asr_ws.close()
|
||||||
|
self.asr_ws = None
|
||||||
|
if self.forward_task:
|
||||||
|
self.forward_task.cancel()
|
||||||
|
try:
|
||||||
|
await self.forward_task
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
pass
|
||||||
|
self.forward_task = None
|
||||||
|
self.is_processing = False
|
||||||
@@ -2,6 +2,7 @@ import time
|
|||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
import io
|
import io
|
||||||
|
import psutil
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from typing import Optional, Tuple, List
|
from typing import Optional, Tuple, List
|
||||||
from core.providers.asr.base import ASRProviderBase
|
from core.providers.asr.base import ASRProviderBase
|
||||||
@@ -37,6 +38,13 @@ class CaptureOutput:
|
|||||||
class ASRProvider(ASRProviderBase):
|
class ASRProvider(ASRProviderBase):
|
||||||
def __init__(self, config: dict, delete_audio_file: bool):
|
def __init__(self, config: dict, delete_audio_file: bool):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
|
# 内存检测,要求大于2G
|
||||||
|
min_mem_bytes = 2 * 1024 * 1024 * 1024
|
||||||
|
total_mem = psutil.virtual_memory().total
|
||||||
|
if total_mem < min_mem_bytes:
|
||||||
|
logger.bind(tag=TAG).error(f"可用内存不足2G,当前仅有 {total_mem / (1024*1024):.2f} MB,可能无法启动FunASR")
|
||||||
|
|
||||||
self.interface_type = InterfaceType.LOCAL
|
self.interface_type = InterfaceType.LOCAL
|
||||||
self.model_dir = config.get("model_dir")
|
self.model_dir = config.get("model_dir")
|
||||||
self.output_dir = config.get("output_dir") # 修正配置键名
|
self.output_dir = config.get("output_dir") # 修正配置键名
|
||||||
@@ -54,7 +62,7 @@ class ASRProvider(ASRProviderBase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
async def speech_to_text(
|
async def speech_to_text(
|
||||||
self, opus_data: List[bytes], session_id: str
|
self, opus_data: List[bytes], session_id: str, audio_format="opus"
|
||||||
) -> Tuple[Optional[str], Optional[str]]:
|
) -> Tuple[Optional[str], Optional[str]]:
|
||||||
"""语音转文本主处理逻辑"""
|
"""语音转文本主处理逻辑"""
|
||||||
file_path = None
|
file_path = None
|
||||||
@@ -63,7 +71,7 @@ class ASRProvider(ASRProviderBase):
|
|||||||
while retry_count < MAX_RETRIES:
|
while retry_count < MAX_RETRIES:
|
||||||
try:
|
try:
|
||||||
# 合并所有opus数据包
|
# 合并所有opus数据包
|
||||||
if self.audio_format == "pcm":
|
if audio_format == "pcm":
|
||||||
pcm_data = opus_data
|
pcm_data = opus_data
|
||||||
else:
|
else:
|
||||||
pcm_data = self.decode_opus(opus_data)
|
pcm_data = self.decode_opus(opus_data)
|
||||||
|
|||||||
@@ -100,7 +100,7 @@ class ASRProvider(ASRProviderBase):
|
|||||||
logger.bind(tag=TAG).debug(f"Sent end message: {end_message}")
|
logger.bind(tag=TAG).debug(f"Sent end message: {end_message}")
|
||||||
|
|
||||||
async def speech_to_text(
|
async def speech_to_text(
|
||||||
self, opus_data: List[bytes], session_id: str
|
self, opus_data: List[bytes], session_id: str, audio_format="opus"
|
||||||
) -> Tuple[Optional[str], Optional[str]]:
|
) -> Tuple[Optional[str], Optional[str]]:
|
||||||
"""
|
"""
|
||||||
Convert speech data to text using FunASR.
|
Convert speech data to text using FunASR.
|
||||||
@@ -109,7 +109,7 @@ class ASRProvider(ASRProviderBase):
|
|||||||
:return: Tuple containing recognized text and optional timestamp.
|
:return: Tuple containing recognized text and optional timestamp.
|
||||||
"""
|
"""
|
||||||
file_path = None
|
file_path = None
|
||||||
if self.audio_format == "pcm":
|
if audio_format == "pcm":
|
||||||
pcm_data = opus_data
|
pcm_data = opus_data
|
||||||
else:
|
else:
|
||||||
pcm_data = self.decode_opus(opus_data)
|
pcm_data = self.decode_opus(opus_data)
|
||||||
|
|||||||
@@ -109,14 +109,14 @@ class ASRProvider(ASRProviderBase):
|
|||||||
return samples_float32, f.getframerate()
|
return samples_float32, f.getframerate()
|
||||||
|
|
||||||
async def speech_to_text(
|
async def speech_to_text(
|
||||||
self, opus_data: List[bytes], session_id: str
|
self, opus_data: List[bytes], session_id: str, audio_format="opus"
|
||||||
) -> Tuple[Optional[str], Optional[str]]:
|
) -> Tuple[Optional[str], Optional[str]]:
|
||||||
"""语音转文本主处理逻辑"""
|
"""语音转文本主处理逻辑"""
|
||||||
file_path = None
|
file_path = None
|
||||||
try:
|
try:
|
||||||
# 保存音频文件
|
# 保存音频文件
|
||||||
start_time = time.time()
|
start_time = time.time()
|
||||||
if self.audio_format == "pcm":
|
if audio_format == "pcm":
|
||||||
pcm_data = opus_data
|
pcm_data = opus_data
|
||||||
else:
|
else:
|
||||||
pcm_data = self.decode_opus(opus_data)
|
pcm_data = self.decode_opus(opus_data)
|
||||||
|
|||||||
@@ -32,7 +32,7 @@ class ASRProvider(ASRProviderBase):
|
|||||||
os.makedirs(self.output_dir, exist_ok=True)
|
os.makedirs(self.output_dir, exist_ok=True)
|
||||||
|
|
||||||
async def speech_to_text(
|
async def speech_to_text(
|
||||||
self, opus_data: List[bytes], session_id: str
|
self, opus_data: List[bytes], session_id: str, audio_format="opus"
|
||||||
) -> Tuple[Optional[str], Optional[str]]:
|
) -> Tuple[Optional[str], Optional[str]]:
|
||||||
"""将语音数据转换为文本"""
|
"""将语音数据转换为文本"""
|
||||||
if not opus_data:
|
if not opus_data:
|
||||||
@@ -47,7 +47,7 @@ class ASRProvider(ASRProviderBase):
|
|||||||
return None, file_path
|
return None, file_path
|
||||||
|
|
||||||
# 将Opus音频数据解码为PCM
|
# 将Opus音频数据解码为PCM
|
||||||
if self.audio_format == "pcm":
|
if audio_format == "pcm":
|
||||||
pcm_data = opus_data
|
pcm_data = opus_data
|
||||||
else:
|
else:
|
||||||
pcm_data = self.decode_opus(opus_data)
|
pcm_data = self.decode_opus(opus_data)
|
||||||
|
|||||||
@@ -122,6 +122,8 @@ class IntentProvider(IntentProviderBase):
|
|||||||
async def detect_intent(self, conn, dialogue_history: List[Dict], text: str) -> str:
|
async def detect_intent(self, conn, dialogue_history: List[Dict], text: str) -> str:
|
||||||
if not self.llm:
|
if not self.llm:
|
||||||
raise ValueError("LLM provider not set")
|
raise ValueError("LLM provider not set")
|
||||||
|
if conn.func_handler is None:
|
||||||
|
return '{"function_call": {"name": "continue_chat"}}'
|
||||||
|
|
||||||
# 记录整体开始时间
|
# 记录整体开始时间
|
||||||
total_start_time = time.time()
|
total_start_time = time.time()
|
||||||
@@ -148,9 +150,8 @@ class IntentProvider(IntentProviderBase):
|
|||||||
self.clean_cache()
|
self.clean_cache()
|
||||||
|
|
||||||
if self.promot == "":
|
if self.promot == "":
|
||||||
if hasattr(conn, "func_handler"):
|
functions = conn.func_handler.get_functions()
|
||||||
functions = conn.func_handler.get_functions()
|
self.promot = self.get_intent_system_prompt(functions)
|
||||||
self.promot = self.get_intent_system_prompt(functions)
|
|
||||||
|
|
||||||
music_config = initialize_music_handler(conn)
|
music_config = initialize_music_handler(conn)
|
||||||
music_file_names = music_config["music_file_names"]
|
music_file_names = music_config["music_file_names"]
|
||||||
|
|||||||
@@ -91,7 +91,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
|
|
||||||
self.appkey = config.get("appkey")
|
self.appkey = config.get("appkey")
|
||||||
self.format = config.get("format", "wav")
|
self.format = config.get("format", "wav")
|
||||||
|
self.audio_file_type = config.get("format", "wav")
|
||||||
sample_rate = config.get("sample_rate", "16000")
|
sample_rate = config.get("sample_rate", "16000")
|
||||||
self.sample_rate = int(sample_rate) if sample_rate else 16000
|
self.sample_rate = int(sample_rate) if sample_rate else 16000
|
||||||
|
|
||||||
@@ -188,9 +188,12 @@ class TTSProvider(TTSProviderBase):
|
|||||||
)
|
)
|
||||||
# 检查返回请求数据的mime类型是否是audio/***,是则保存到指定路径下;返回的是binary格式的
|
# 检查返回请求数据的mime类型是否是audio/***,是则保存到指定路径下;返回的是binary格式的
|
||||||
if resp.headers["Content-Type"].startswith("audio/"):
|
if resp.headers["Content-Type"].startswith("audio/"):
|
||||||
with open(output_file, "wb") as f:
|
if output_file:
|
||||||
f.write(resp.content)
|
with open(output_file, "wb") as f:
|
||||||
return output_file
|
f.write(resp.content)
|
||||||
|
return output_file
|
||||||
|
else:
|
||||||
|
return resp.content
|
||||||
else:
|
else:
|
||||||
raise Exception(
|
raise Exception(
|
||||||
f"{__name__} status_code: {resp.status_code} response: {resp.content}"
|
f"{__name__} status_code: {resp.status_code} response: {resp.content}"
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ from datetime import datetime
|
|||||||
from core.utils import textUtils
|
from core.utils import textUtils
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from core.utils.util import audio_to_data
|
from core.utils.util import audio_to_data, audio_bytes_to_data
|
||||||
from core.utils.tts import MarkdownCleaner
|
from core.utils.tts import MarkdownCleaner
|
||||||
from core.utils.output_counter import add_device_output
|
from core.utils.output_counter import add_device_output
|
||||||
from core.handle.reportHandle import enqueue_tts_report
|
from core.handle.reportHandle import enqueue_tts_report
|
||||||
@@ -20,7 +20,6 @@ from core.providers.tts.dto.dto import (
|
|||||||
InterfaceType,
|
InterfaceType,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
import traceback
|
import traceback
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
@@ -33,10 +32,12 @@ class TTSProviderBase(ABC):
|
|||||||
self.conn = None
|
self.conn = None
|
||||||
self.tts_timeout = 10
|
self.tts_timeout = 10
|
||||||
self.delete_audio_file = delete_audio_file
|
self.delete_audio_file = delete_audio_file
|
||||||
|
self.audio_file_type = "wav"
|
||||||
self.output_file = config.get("output_dir", "tmp/")
|
self.output_file = config.get("output_dir", "tmp/")
|
||||||
self.tts_text_queue = queue.Queue()
|
self.tts_text_queue = queue.Queue()
|
||||||
self.tts_audio_queue = queue.Queue()
|
self.tts_audio_queue = queue.Queue()
|
||||||
self.tts_audio_first_sentence = True
|
self.tts_audio_first_sentence = True
|
||||||
|
self.before_stop_play_files = []
|
||||||
|
|
||||||
self.tts_text_buff = []
|
self.tts_text_buff = []
|
||||||
self.punctuations = (
|
self.punctuations = (
|
||||||
@@ -52,6 +53,7 @@ class TTSProviderBase(ABC):
|
|||||||
)
|
)
|
||||||
self.first_sentence_punctuations = (
|
self.first_sentence_punctuations = (
|
||||||
",",
|
",",
|
||||||
|
"~",
|
||||||
"~",
|
"~",
|
||||||
"、",
|
"、",
|
||||||
",",
|
",",
|
||||||
@@ -76,35 +78,62 @@ class TTSProviderBase(ABC):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def to_tts(self, text):
|
def to_tts(self, text):
|
||||||
tmp_file = self.generate_filename()
|
text = MarkdownCleaner.clean_markdown(text)
|
||||||
try:
|
max_repeat_time = 5
|
||||||
max_repeat_time = 5
|
if self.delete_audio_file:
|
||||||
text = MarkdownCleaner.clean_markdown(text)
|
# 需要删除文件的直接转为音频数据
|
||||||
while not os.path.exists(tmp_file) and max_repeat_time > 0:
|
while max_repeat_time > 0:
|
||||||
try:
|
try:
|
||||||
asyncio.run(self.text_to_speak(text, tmp_file))
|
audio_bytes = asyncio.run(self.text_to_speak(text, None))
|
||||||
|
if audio_bytes:
|
||||||
|
audio_datas, _ = audio_bytes_to_data(
|
||||||
|
audio_bytes, file_type=self.audio_file_type, is_opus=True
|
||||||
|
)
|
||||||
|
return audio_datas
|
||||||
|
else:
|
||||||
|
max_repeat_time -= 1
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).warning(
|
logger.bind(tag=TAG).warning(
|
||||||
f"语音生成失败{5 - max_repeat_time + 1}次: {text},错误: {e}"
|
f"语音生成失败{5 - max_repeat_time + 1}次: {text},错误: {e}"
|
||||||
)
|
)
|
||||||
# 未执行成功,删除文件
|
|
||||||
if os.path.exists(tmp_file):
|
|
||||||
os.remove(tmp_file)
|
|
||||||
max_repeat_time -= 1
|
max_repeat_time -= 1
|
||||||
|
|
||||||
if max_repeat_time > 0:
|
if max_repeat_time > 0:
|
||||||
logger.bind(tag=TAG).info(
|
logger.bind(tag=TAG).info(
|
||||||
f"语音生成成功: {text}:{tmp_file},重试{5 - max_repeat_time}次"
|
f"语音生成成功: {text},重试{5 - max_repeat_time}次"
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
logger.bind(tag=TAG).error(
|
logger.bind(tag=TAG).error(
|
||||||
f"语音生成失败: {text},请检查网络或服务是否正常"
|
f"语音生成失败: {text},请检查网络或服务是否正常"
|
||||||
)
|
)
|
||||||
|
|
||||||
return tmp_file
|
|
||||||
except Exception as e:
|
|
||||||
logger.bind(tag=TAG).error(f"Failed to generate TTS file: {e}")
|
|
||||||
return None
|
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
|
@abstractmethod
|
||||||
async def text_to_speak(self, text, output_file):
|
async def text_to_speak(self, text, output_file):
|
||||||
@@ -178,6 +207,9 @@ class TTSProviderBase(ABC):
|
|||||||
while not self.conn.stop_event.is_set():
|
while not self.conn.stop_event.is_set():
|
||||||
try:
|
try:
|
||||||
message = self.tts_text_queue.get(timeout=1)
|
message = self.tts_text_queue.get(timeout=1)
|
||||||
|
if self.conn.client_abort:
|
||||||
|
logger.bind(tag=TAG).info("收到打断信息,终止TTS文本处理线程")
|
||||||
|
continue
|
||||||
if message.sentence_type == SentenceType.FIRST:
|
if message.sentence_type == SentenceType.FIRST:
|
||||||
# 初始化参数
|
# 初始化参数
|
||||||
self.tts_stop_request = False
|
self.tts_stop_request = False
|
||||||
@@ -189,12 +221,19 @@ class TTSProviderBase(ABC):
|
|||||||
self.tts_text_buff.append(message.content_detail)
|
self.tts_text_buff.append(message.content_detail)
|
||||||
segment_text = self._get_segment_text()
|
segment_text = self._get_segment_text()
|
||||||
if segment_text:
|
if segment_text:
|
||||||
tts_file = self.to_tts(segment_text)
|
if self.delete_audio_file:
|
||||||
if tts_file:
|
audio_datas = self.to_tts(segment_text)
|
||||||
audio_datas = self._process_audio_file(tts_file)
|
if audio_datas:
|
||||||
self.tts_audio_queue.put(
|
self.tts_audio_queue.put(
|
||||||
(message.sentence_type, audio_datas, segment_text)
|
(message.sentence_type, audio_datas, segment_text)
|
||||||
)
|
)
|
||||||
|
else:
|
||||||
|
tts_file = self.to_tts(segment_text)
|
||||||
|
if tts_file:
|
||||||
|
audio_datas = self._process_audio_file(tts_file)
|
||||||
|
self.tts_audio_queue.put(
|
||||||
|
(message.sentence_type, audio_datas, segment_text)
|
||||||
|
)
|
||||||
elif ContentType.FILE == message.content_type:
|
elif ContentType.FILE == message.content_type:
|
||||||
self._process_remaining_text()
|
self._process_remaining_text()
|
||||||
tts_file = message.content_file
|
tts_file = message.content_file
|
||||||
@@ -320,6 +359,14 @@ class TTSProviderBase(ABC):
|
|||||||
os.remove(tts_file)
|
os.remove(tts_file)
|
||||||
return audio_datas
|
return audio_datas
|
||||||
|
|
||||||
|
def _process_before_stop_play_files(self):
|
||||||
|
for tts_file, text in self.before_stop_play_files:
|
||||||
|
if tts_file and os.path.exists(tts_file):
|
||||||
|
audio_datas = self._process_audio_file(tts_file)
|
||||||
|
self.tts_audio_queue.put((SentenceType.MIDDLE, audio_datas, text))
|
||||||
|
self.before_stop_play_files.clear()
|
||||||
|
self.tts_audio_queue.put((SentenceType.LAST, [], None))
|
||||||
|
|
||||||
def _process_remaining_text(self):
|
def _process_remaining_text(self):
|
||||||
"""处理剩余的文本并生成语音
|
"""处理剩余的文本并生成语音
|
||||||
|
|
||||||
@@ -331,11 +378,18 @@ class TTSProviderBase(ABC):
|
|||||||
if remaining_text:
|
if remaining_text:
|
||||||
segment_text = textUtils.get_string_no_punctuation_or_emoji(remaining_text)
|
segment_text = textUtils.get_string_no_punctuation_or_emoji(remaining_text)
|
||||||
if segment_text:
|
if segment_text:
|
||||||
tts_file = self.to_tts(segment_text)
|
if self.delete_audio_file:
|
||||||
audio_datas = self._process_audio_file(tts_file)
|
audio_datas = self.to_tts(segment_text)
|
||||||
self.tts_audio_queue.put(
|
if audio_datas:
|
||||||
(SentenceType.MIDDLE, audio_datas, segment_text)
|
self.tts_audio_queue.put(
|
||||||
)
|
(SentenceType.MIDDLE, audio_datas, segment_text)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
tts_file = self.to_tts(segment_text)
|
||||||
|
audio_datas = self._process_audio_file(tts_file)
|
||||||
|
self.tts_audio_queue.put(
|
||||||
|
(SentenceType.MIDDLE, audio_datas, segment_text)
|
||||||
|
)
|
||||||
self.processed_chars += len(full_text)
|
self.processed_chars += len(full_text)
|
||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|||||||
@@ -11,8 +11,8 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.voice = config.get("private_voice")
|
self.voice = config.get("private_voice")
|
||||||
else:
|
else:
|
||||||
self.voice = config.get("voice")
|
self.voice = config.get("voice")
|
||||||
self.response_format = config.get("response_format")
|
self.response_format = config.get("response_format", "mp3")
|
||||||
|
self.audio_file_type = config.get("response_format", "mp3")
|
||||||
self.host = "api.coze.cn"
|
self.host = "api.coze.cn"
|
||||||
self.api_url = f"https://{self.host}/v1/audio/speech"
|
self.api_url = f"https://{self.host}/v1/audio/speech"
|
||||||
|
|
||||||
@@ -33,7 +33,10 @@ class TTSProvider(TTSProviderBase):
|
|||||||
"POST", self.api_url, json=request_json, headers=headers
|
"POST", self.api_url, json=request_json, headers=headers
|
||||||
)
|
)
|
||||||
data = response.content
|
data = response.content
|
||||||
file_to_save = open(output_file, "wb")
|
if output_file:
|
||||||
file_to_save.write(data)
|
with open(output_file, "wb") as file_to_save:
|
||||||
|
file_to_save.write(data)
|
||||||
|
else:
|
||||||
|
return data
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
raise Exception(f"{__name__} error: {e}")
|
raise Exception(f"{__name__} error: {e}")
|
||||||
|
|||||||
@@ -16,8 +16,8 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.method = config.get("method", "GET")
|
self.method = config.get("method", "GET")
|
||||||
self.headers = config.get("headers", {})
|
self.headers = config.get("headers", {})
|
||||||
self.format = config.get("format", "wav")
|
self.format = config.get("format", "wav")
|
||||||
|
self.audio_file_type = config.get("format", "wav")
|
||||||
self.output_file = config.get("output_dir", "tmp/")
|
self.output_file = config.get("output_dir", "tmp/")
|
||||||
|
|
||||||
self.params = config.get("params")
|
self.params = config.get("params")
|
||||||
|
|
||||||
if isinstance(self.params, str):
|
if isinstance(self.params, str):
|
||||||
@@ -43,8 +43,11 @@ class TTSProvider(TTSProviderBase):
|
|||||||
else:
|
else:
|
||||||
resp = requests.get(self.url, params=request_params, headers=self.headers)
|
resp = requests.get(self.url, params=request_params, headers=self.headers)
|
||||||
if resp.status_code == 200:
|
if resp.status_code == 200:
|
||||||
with open(output_file, "wb") as file:
|
if output_file:
|
||||||
file.write(resp.content)
|
with open(output_file, "wb") as file:
|
||||||
|
file.write(resp.content)
|
||||||
|
else:
|
||||||
|
return resp.content
|
||||||
else:
|
else:
|
||||||
error_msg = f"Custom TTS请求失败: {resp.status_code} - {resp.text}"
|
error_msg = f"Custom TTS请求失败: {resp.status_code} - {resp.text}"
|
||||||
logger.bind(tag=TAG).error(error_msg)
|
logger.bind(tag=TAG).error(error_msg)
|
||||||
|
|||||||
@@ -29,7 +29,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
speed_ratio = config.get("speed_ratio", "1.0")
|
speed_ratio = config.get("speed_ratio", "1.0")
|
||||||
volume_ratio = config.get("volume_ratio", "1.0")
|
volume_ratio = config.get("volume_ratio", "1.0")
|
||||||
pitch_ratio = config.get("pitch_ratio", "1.0")
|
pitch_ratio = config.get("pitch_ratio", "1.0")
|
||||||
|
self.audio_file_type = config.get("format", "wav")
|
||||||
self.speed_ratio = float(speed_ratio) if speed_ratio else 1.0
|
self.speed_ratio = float(speed_ratio) if speed_ratio else 1.0
|
||||||
self.volume_ratio = float(volume_ratio) if volume_ratio else 1.0
|
self.volume_ratio = float(volume_ratio) if volume_ratio else 1.0
|
||||||
self.pitch_ratio = float(pitch_ratio) if pitch_ratio else 1.0
|
self.pitch_ratio = float(pitch_ratio) if pitch_ratio else 1.0
|
||||||
@@ -49,7 +49,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
"user": {"uid": "1"},
|
"user": {"uid": "1"},
|
||||||
"audio": {
|
"audio": {
|
||||||
"voice_type": self.voice,
|
"voice_type": self.voice,
|
||||||
"encoding": "wav",
|
"encoding": self.audio_file_type,
|
||||||
"speed_ratio": self.speed_ratio,
|
"speed_ratio": self.speed_ratio,
|
||||||
"volume_ratio": self.volume_ratio,
|
"volume_ratio": self.volume_ratio,
|
||||||
"pitch_ratio": self.pitch_ratio,
|
"pitch_ratio": self.pitch_ratio,
|
||||||
@@ -70,8 +70,12 @@ class TTSProvider(TTSProviderBase):
|
|||||||
)
|
)
|
||||||
if "data" in resp.json():
|
if "data" in resp.json():
|
||||||
data = resp.json()["data"]
|
data = resp.json()["data"]
|
||||||
file_to_save = open(output_file, "wb")
|
audio_bytes = base64.b64decode(data)
|
||||||
file_to_save.write(base64.b64decode(data))
|
if output_file:
|
||||||
|
with open(output_file, "wb") as file_to_save:
|
||||||
|
file_to_save.write(audio_bytes)
|
||||||
|
else:
|
||||||
|
return audio_bytes
|
||||||
else:
|
else:
|
||||||
raise Exception(
|
raise Exception(
|
||||||
f"{__name__} status_code: {resp.status_code} response: {resp.content}"
|
f"{__name__} status_code: {resp.status_code} response: {resp.content}"
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.voice = config.get("private_voice")
|
self.voice = config.get("private_voice")
|
||||||
else:
|
else:
|
||||||
self.voice = config.get("voice")
|
self.voice = config.get("voice")
|
||||||
|
self.audio_file_type = config.get("format", "mp3")
|
||||||
|
|
||||||
def generate_filename(self, extension=".mp3"):
|
def generate_filename(self, extension=".mp3"):
|
||||||
return os.path.join(
|
return os.path.join(
|
||||||
@@ -22,16 +23,24 @@ class TTSProvider(TTSProviderBase):
|
|||||||
async def text_to_speak(self, text, output_file):
|
async def text_to_speak(self, text, output_file):
|
||||||
try:
|
try:
|
||||||
communicate = edge_tts.Communicate(text, voice=self.voice)
|
communicate = edge_tts.Communicate(text, voice=self.voice)
|
||||||
# 确保目录存在并创建空文件
|
if output_file:
|
||||||
os.makedirs(os.path.dirname(output_file), exist_ok=True)
|
# 确保目录存在并创建空文件
|
||||||
with open(output_file, "wb") as f:
|
os.makedirs(os.path.dirname(output_file), exist_ok=True)
|
||||||
pass
|
with open(output_file, "wb") as f:
|
||||||
|
pass
|
||||||
|
|
||||||
# 流式写入音频数据
|
# 流式写入音频数据
|
||||||
with open(output_file, "ab") as f: # 改为追加模式避免覆盖
|
with open(output_file, "ab") as f: # 改为追加模式避免覆盖
|
||||||
|
async for chunk in communicate.stream():
|
||||||
|
if chunk["type"] == "audio": # 只处理音频数据块
|
||||||
|
f.write(chunk["data"])
|
||||||
|
else:
|
||||||
|
# 返回音频二进制数据
|
||||||
|
audio_bytes = b""
|
||||||
async for chunk in communicate.stream():
|
async for chunk in communicate.stream():
|
||||||
if chunk["type"] == "audio": # 只处理音频数据块
|
if chunk["type"] == "audio":
|
||||||
f.write(chunk["data"])
|
audio_bytes += chunk["data"]
|
||||||
|
return audio_bytes
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
error_msg = f"Edge TTS请求失败: {e}"
|
error_msg = f"Edge TTS请求失败: {e}"
|
||||||
raise Exception(error_msg) # 抛出异常,让调用方捕获
|
raise Exception(error_msg) # 抛出异常,让调用方捕获
|
||||||
@@ -88,7 +88,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.reference_audio = parse_string_to_list(config.get("reference_audio"))
|
self.reference_audio = parse_string_to_list(config.get("reference_audio"))
|
||||||
self.reference_text = parse_string_to_list(config.get("reference_text"))
|
self.reference_text = parse_string_to_list(config.get("reference_text"))
|
||||||
self.format = config.get("response_format", "wav")
|
self.format = config.get("response_format", "wav")
|
||||||
|
self.audio_file_type = config.get("response_format", "wav")
|
||||||
self.api_key = config.get("api_key", "YOUR_API_KEY")
|
self.api_key = config.get("api_key", "YOUR_API_KEY")
|
||||||
have_key = check_model_key("FishSpeech TTS", self.api_key)
|
have_key = check_model_key("FishSpeech TTS", self.api_key)
|
||||||
if not have_key:
|
if not have_key:
|
||||||
@@ -170,8 +170,11 @@ class TTSProvider(TTSProviderBase):
|
|||||||
if response.status_code == 200:
|
if response.status_code == 200:
|
||||||
audio_content = response.content
|
audio_content = response.content
|
||||||
|
|
||||||
with open(output_file, "wb") as audio_file:
|
if output_file:
|
||||||
audio_file.write(audio_content)
|
with open(output_file, "wb") as audio_file:
|
||||||
|
audio_file.write(audio_content)
|
||||||
|
else:
|
||||||
|
return audio_content
|
||||||
|
|
||||||
else:
|
else:
|
||||||
error_msg = f"Request failed with status code {response.status_code}"
|
error_msg = f"Request failed with status code {response.status_code}"
|
||||||
|
|||||||
@@ -65,6 +65,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.aux_ref_audio_paths = parse_string_to_list(
|
self.aux_ref_audio_paths = parse_string_to_list(
|
||||||
config.get("aux_ref_audio_paths")
|
config.get("aux_ref_audio_paths")
|
||||||
)
|
)
|
||||||
|
self.audio_file_type = config.get("format", "wav")
|
||||||
|
|
||||||
async def text_to_speak(self, text, output_file):
|
async def text_to_speak(self, text, output_file):
|
||||||
request_json = {
|
request_json = {
|
||||||
@@ -91,8 +92,11 @@ class TTSProvider(TTSProviderBase):
|
|||||||
|
|
||||||
resp = requests.post(self.url, json=request_json)
|
resp = requests.post(self.url, json=request_json)
|
||||||
if resp.status_code == 200:
|
if resp.status_code == 200:
|
||||||
with open(output_file, "wb") as file:
|
if output_file:
|
||||||
file.write(resp.content)
|
with open(output_file, "wb") as file:
|
||||||
|
file.write(resp.content)
|
||||||
|
else:
|
||||||
|
return resp.content
|
||||||
else:
|
else:
|
||||||
error_msg = f"GPT_SoVITS_V2 TTS请求失败: {resp.status_code} - {resp.text}"
|
error_msg = f"GPT_SoVITS_V2 TTS请求失败: {resp.status_code} - {resp.text}"
|
||||||
logger.bind(tag=TAG).error(error_msg)
|
logger.bind(tag=TAG).error(error_msg)
|
||||||
|
|||||||
@@ -32,6 +32,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.cut_punc = config.get("cut_punc", "")
|
self.cut_punc = config.get("cut_punc", "")
|
||||||
self.inp_refs = parse_string_to_list(config.get("inp_refs"))
|
self.inp_refs = parse_string_to_list(config.get("inp_refs"))
|
||||||
self.if_sr = str(config.get("if_sr", False)).lower() in ("true", "1", "yes")
|
self.if_sr = str(config.get("if_sr", False)).lower() in ("true", "1", "yes")
|
||||||
|
self.audio_file_type = config.get("format", "wav")
|
||||||
|
|
||||||
async def text_to_speak(self, text, output_file):
|
async def text_to_speak(self, text, output_file):
|
||||||
request_params = {
|
request_params = {
|
||||||
@@ -52,8 +53,11 @@ class TTSProvider(TTSProviderBase):
|
|||||||
|
|
||||||
resp = requests.get(self.url, params=request_params)
|
resp = requests.get(self.url, params=request_params)
|
||||||
if resp.status_code == 200:
|
if resp.status_code == 200:
|
||||||
with open(output_file, "wb") as file:
|
if output_file:
|
||||||
file.write(resp.content)
|
with open(output_file, "wb") as file:
|
||||||
|
file.write(resp.content)
|
||||||
|
else:
|
||||||
|
return resp.content
|
||||||
else:
|
else:
|
||||||
error_msg = f"GPT_SoVITS_V3 TTS请求失败: {resp.status_code} - {resp.text}"
|
error_msg = f"GPT_SoVITS_V3 TTS请求失败: {resp.status_code} - {resp.text}"
|
||||||
logger.bind(tag=TAG).error(error_msg)
|
logger.bind(tag=TAG).error(error_msg)
|
||||||
|
|||||||
@@ -3,14 +3,13 @@ import uuid
|
|||||||
import json
|
import json
|
||||||
import queue
|
import queue
|
||||||
import asyncio
|
import asyncio
|
||||||
import threading
|
|
||||||
import traceback
|
import traceback
|
||||||
import websockets
|
import websockets
|
||||||
import time
|
|
||||||
from config.logger import setup_logging
|
from config.logger import setup_logging
|
||||||
from core.utils import opus_encoder_utils
|
from core.utils import opus_encoder_utils
|
||||||
from core.utils.util import check_model_key
|
from core.utils.util import check_model_key
|
||||||
from core.providers.tts.base import TTSProviderBase
|
from core.providers.tts.base import TTSProviderBase
|
||||||
|
from core.handle.abortHandle import handleAbortMessage
|
||||||
from core.providers.tts.dto.dto import SentenceType, ContentType, InterfaceType
|
from core.providers.tts.dto.dto import SentenceType, ContentType, InterfaceType
|
||||||
|
|
||||||
TAG = __name__
|
TAG = __name__
|
||||||
@@ -139,106 +138,66 @@ class Response:
|
|||||||
class TTSProvider(TTSProviderBase):
|
class TTSProvider(TTSProviderBase):
|
||||||
def __init__(self, config, delete_audio_file):
|
def __init__(self, config, delete_audio_file):
|
||||||
super().__init__(config, delete_audio_file)
|
super().__init__(config, delete_audio_file)
|
||||||
self.ws = None # 初始化ws属性
|
self.ws = None
|
||||||
self.interface_type = InterfaceType.DUAL_STREAM
|
self.interface_type = InterfaceType.DUAL_STREAM
|
||||||
self.appId = config.get("appid")
|
self.appId = config.get("appid")
|
||||||
self.access_token = config.get("access_token")
|
self.access_token = config.get("access_token")
|
||||||
self.cluster = config.get("cluster")
|
self.cluster = config.get("cluster")
|
||||||
self.resource_id = config.get("resource_id")
|
self.resource_id = config.get("resource_id")
|
||||||
if config.get("private_voice"):
|
if config.get("private_voice"):
|
||||||
self.speaker = config.get("private_voice")
|
self.voice = config.get("private_voice")
|
||||||
else:
|
else:
|
||||||
self.speaker = config.get("speaker")
|
self.voice = config.get("speaker")
|
||||||
self.voice = config.get("voice")
|
|
||||||
self.ws_url = config.get("ws_url")
|
self.ws_url = config.get("ws_url")
|
||||||
self.authorization = config.get("authorization")
|
self.authorization = config.get("authorization")
|
||||||
self.header = {"Authorization": f"{self.authorization}{self.access_token}"}
|
self.header = {"Authorization": f"{self.authorization}{self.access_token}"}
|
||||||
self.enable_two_way = True
|
self.enable_two_way = True
|
||||||
self.start_connection_flag = False
|
|
||||||
self.tts_text = ""
|
self.tts_text = ""
|
||||||
# 合成文字语音后,播放的音频文件列表
|
|
||||||
self.before_stop_play_files = []
|
|
||||||
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
|
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
|
||||||
sample_rate=16000, channels=1, frame_size_ms=60
|
sample_rate=16000, channels=1, frame_size_ms=60
|
||||||
)
|
)
|
||||||
check_model_key("TTS", self.access_token)
|
check_model_key("TTS", self.access_token)
|
||||||
|
|
||||||
# 添加会话状态控制
|
|
||||||
self._session_lock = asyncio.Lock() # 会话操作的并发锁
|
|
||||||
self._current_session_id = None # 当前会话ID
|
|
||||||
self._session_started = False # 会话是否已开始
|
|
||||||
self._session_finished = False # 会话是否已结束
|
|
||||||
self._connection_ready = False # 连接是否就绪
|
|
||||||
self._reconnect_attempts = 0 # 重连尝试次数
|
|
||||||
self._max_reconnect_attempts = 3 # 最大重连次数
|
|
||||||
|
|
||||||
###################################################################################
|
|
||||||
# 火山双流式TTS重写父类的方法--开始
|
|
||||||
###################################################################################
|
|
||||||
|
|
||||||
async def open_audio_channels(self, conn):
|
async def open_audio_channels(self, conn):
|
||||||
try:
|
try:
|
||||||
await super().open_audio_channels(conn)
|
await super().open_audio_channels(conn)
|
||||||
await self._ensure_connection()
|
|
||||||
tts_priority = threading.Thread(
|
|
||||||
target=self._start_monitor_tts_response_thread, daemon=True
|
|
||||||
)
|
|
||||||
tts_priority.start()
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"Failed to open audio channels: {str(e)}")
|
logger.bind(tag=TAG).error(f"Failed to open audio channels: {str(e)}")
|
||||||
self.ws = None
|
self.ws = None
|
||||||
raise
|
raise
|
||||||
|
|
||||||
async def _ensure_connection(self):
|
async def _ensure_connection(self):
|
||||||
"""确保WebSocket连接可用"""
|
"""建立新的WebSocket连接"""
|
||||||
try:
|
try:
|
||||||
if self.ws is None:
|
logger.bind(tag=TAG).info("开始建立新连接...")
|
||||||
logger.bind(tag=TAG).info("WebSocket连接不存在,开始建立新连接...")
|
ws_header = {
|
||||||
ws_header = {
|
"X-Api-App-Key": self.appId,
|
||||||
"X-Api-App-Key": self.appId,
|
"X-Api-Access-Key": self.access_token,
|
||||||
"X-Api-Access-Key": self.access_token,
|
"X-Api-Resource-Id": self.resource_id,
|
||||||
"X-Api-Resource-Id": self.resource_id,
|
"X-Api-Connect-Id": uuid.uuid4(),
|
||||||
"X-Api-Connect-Id": uuid.uuid4(),
|
}
|
||||||
}
|
self.ws = await websockets.connect(
|
||||||
self.ws = await websockets.connect(
|
self.ws_url, additional_headers=ws_header, max_size=1000000000
|
||||||
self.ws_url, additional_headers=ws_header, max_size=1000000000
|
)
|
||||||
)
|
logger.bind(tag=TAG).info("WebSocket连接建立成功")
|
||||||
self._connection_ready = True
|
return self.ws
|
||||||
self._reconnect_attempts = 0
|
|
||||||
logger.bind(tag=TAG).info("WebSocket连接建立成功")
|
|
||||||
else:
|
|
||||||
# 尝试发送ping来检查连接是否还活着
|
|
||||||
try:
|
|
||||||
logger.bind(tag=TAG).debug("检查WebSocket连接状态...")
|
|
||||||
pong_waiter = await self.ws.ping()
|
|
||||||
await asyncio.wait_for(pong_waiter, timeout=1.0)
|
|
||||||
logger.bind(tag=TAG).debug("WebSocket连接状态正常")
|
|
||||||
except (asyncio.TimeoutError, websockets.ConnectionClosed):
|
|
||||||
# 如果ping失败,重新建立连接
|
|
||||||
logger.bind(tag=TAG).warning("WebSocket连接已断开,准备重新连接...")
|
|
||||||
try:
|
|
||||||
await self.ws.close()
|
|
||||||
except:
|
|
||||||
pass
|
|
||||||
self.ws = None
|
|
||||||
self._connection_ready = False
|
|
||||||
# 重新建立连接
|
|
||||||
await self._ensure_connection()
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"确保连接失败: {str(e)}")
|
logger.bind(tag=TAG).error(f"建立连接失败: {str(e)}")
|
||||||
self._connection_ready = False
|
|
||||||
self.ws = None
|
self.ws = None
|
||||||
raise
|
raise
|
||||||
|
|
||||||
def tts_text_priority_thread(self):
|
def tts_text_priority_thread(self):
|
||||||
logger.bind(tag=TAG).info("TTS文本处理线程启动")
|
"""火山引擎双流式TTS的文本处理线程"""
|
||||||
while not self.conn.stop_event.is_set():
|
while not self.conn.stop_event.is_set():
|
||||||
try:
|
try:
|
||||||
logger.bind(tag=TAG).debug("等待TTS文本队列消息...")
|
|
||||||
message = self.tts_text_queue.get(timeout=1)
|
message = self.tts_text_queue.get(timeout=1)
|
||||||
logger.bind(tag=TAG).info(
|
logger.bind(tag=TAG).debug(
|
||||||
f"收到TTS任务|{message.sentence_type.name} | {message.content_type.name} | 会话ID: {self.conn.sentence_id}"
|
f"收到TTS任务|{message.sentence_type.name} | {message.content_type.name} | 会话ID: {self.conn.sentence_id}"
|
||||||
)
|
)
|
||||||
|
if self.conn.client_abort:
|
||||||
|
logger.bind(tag=TAG).info("收到打断信息,终止TTS文本处理线程")
|
||||||
|
continue
|
||||||
|
|
||||||
if message.sentence_type == SentenceType.FIRST:
|
if message.sentence_type == SentenceType.FIRST:
|
||||||
# 初始化参数
|
# 初始化参数
|
||||||
try:
|
try:
|
||||||
@@ -253,13 +212,12 @@ class TTSProvider(TTSProviderBase):
|
|||||||
logger.bind(tag=TAG).info("TTS会话启动成功")
|
logger.bind(tag=TAG).info("TTS会话启动成功")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"启动TTS会话失败: {str(e)}")
|
logger.bind(tag=TAG).error(f"启动TTS会话失败: {str(e)}")
|
||||||
# 直接跳过当前消息,不重新入队
|
|
||||||
time.sleep(1)
|
|
||||||
continue
|
continue
|
||||||
|
|
||||||
elif ContentType.TEXT == message.content_type:
|
elif ContentType.TEXT == message.content_type:
|
||||||
if message.content_detail:
|
if message.content_detail:
|
||||||
try:
|
try:
|
||||||
logger.bind(tag=TAG).info(
|
logger.bind(tag=TAG).debug(
|
||||||
f"开始发送TTS文本: {message.content_detail}"
|
f"开始发送TTS文本: {message.content_detail}"
|
||||||
)
|
)
|
||||||
future = asyncio.run_coroutine_threadsafe(
|
future = asyncio.run_coroutine_threadsafe(
|
||||||
@@ -267,12 +225,11 @@ class TTSProvider(TTSProviderBase):
|
|||||||
loop=self.conn.loop,
|
loop=self.conn.loop,
|
||||||
)
|
)
|
||||||
future.result()
|
future.result()
|
||||||
logger.bind(tag=TAG).info("TTS文本发送成功")
|
logger.bind(tag=TAG).debug("TTS文本发送成功")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"发送TTS文本失败: {str(e)}")
|
logger.bind(tag=TAG).error(f"发送TTS文本失败: {str(e)}")
|
||||||
# 直接跳过当前消息,不重新入队
|
|
||||||
time.sleep(1)
|
|
||||||
continue
|
continue
|
||||||
|
|
||||||
elif ContentType.FILE == message.content_type:
|
elif ContentType.FILE == message.content_type:
|
||||||
logger.bind(tag=TAG).info(
|
logger.bind(tag=TAG).info(
|
||||||
f"添加音频文件到待播放列表: {message.content_file}"
|
f"添加音频文件到待播放列表: {message.content_file}"
|
||||||
@@ -289,11 +246,8 @@ class TTSProvider(TTSProviderBase):
|
|||||||
loop=self.conn.loop,
|
loop=self.conn.loop,
|
||||||
)
|
)
|
||||||
future.result()
|
future.result()
|
||||||
logger.bind(tag=TAG).info("TTS会话结束成功")
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"结束TTS会话失败: {str(e)}")
|
logger.bind(tag=TAG).error(f"结束TTS会话失败: {str(e)}")
|
||||||
# 直接跳过当前消息,不重新入队
|
|
||||||
time.sleep(1)
|
|
||||||
continue
|
continue
|
||||||
|
|
||||||
except queue.Empty:
|
except queue.Empty:
|
||||||
@@ -302,106 +256,209 @@ class TTSProvider(TTSProviderBase):
|
|||||||
logger.bind(tag=TAG).error(
|
logger.bind(tag=TAG).error(
|
||||||
f"处理TTS文本失败: {str(e)}, 类型: {type(e).__name__}, 堆栈: {traceback.format_exc()}"
|
f"处理TTS文本失败: {str(e)}, 类型: {type(e).__name__}, 堆栈: {traceback.format_exc()}"
|
||||||
)
|
)
|
||||||
# 如果是WebSocket连接关闭错误,等待一段时间后继续
|
|
||||||
if "non-exist session" in str(e):
|
|
||||||
time.sleep(1)
|
|
||||||
continue
|
continue
|
||||||
|
|
||||||
async def text_to_speak(self, text, _):
|
async def text_to_speak(self, text, _):
|
||||||
"""发送文本到TTS服务"""
|
"""发送文本到TTS服务"""
|
||||||
try:
|
try:
|
||||||
# 确保WebSocket连接可用
|
# 建立新连接
|
||||||
if not self._connection_ready or self.ws is None:
|
if self.ws is None:
|
||||||
logger.bind(tag=TAG).warning("WebSocket连接不可用,尝试重新连接...")
|
await handleAbortMessage(self.conn)
|
||||||
await self._ensure_connection()
|
logger.bind(tag=TAG).error(f"WebSocket连接不存在,终止发送文本")
|
||||||
|
return
|
||||||
# 发送文本
|
# 发送文本
|
||||||
await self.send_text(self.speaker, text, self.conn.sentence_id)
|
await self.send_text(self.voice, text, self.conn.sentence_id)
|
||||||
return
|
return
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"发送TTS文本失败: {str(e)}")
|
logger.bind(tag=TAG).error(f"发送TTS文本失败: {str(e)}")
|
||||||
# 如果是连接问题,尝试重新连接
|
if self.ws:
|
||||||
if isinstance(e, websockets.ConnectionClosed):
|
try:
|
||||||
self._connection_ready = False
|
await self.ws.close()
|
||||||
|
except:
|
||||||
|
pass
|
||||||
self.ws = None
|
self.ws = None
|
||||||
await self._handle_connection_error()
|
|
||||||
raise
|
raise
|
||||||
|
|
||||||
###################################################################################
|
async def start_session(self, session_id):
|
||||||
# 火山双流式TTS重写父类的方法--结束
|
logger.bind(tag=TAG).info(f"开始会话~~{session_id}")
|
||||||
###################################################################################
|
try:
|
||||||
def _start_monitor_tts_response_thread(self):
|
# 建立新连接
|
||||||
# 初始化链接
|
await self._ensure_connection()
|
||||||
asyncio.run_coroutine_threadsafe(
|
|
||||||
self._start_monitor_tts_response(), loop=self.conn.loop
|
# 启动监听任务
|
||||||
)
|
self._monitor_task = asyncio.create_task(self._start_monitor_tts_response())
|
||||||
|
|
||||||
|
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(self.ws, header, optional, payload)
|
||||||
|
logger.bind(tag=TAG).info("会话启动请求已发送")
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"启动会话失败: {str(e)}")
|
||||||
|
# 确保清理资源
|
||||||
|
if hasattr(self, "_monitor_task"):
|
||||||
|
try:
|
||||||
|
self._monitor_task.cancel()
|
||||||
|
await self._monitor_task
|
||||||
|
except:
|
||||||
|
pass
|
||||||
|
self._monitor_task = None
|
||||||
|
if self.ws:
|
||||||
|
try:
|
||||||
|
await self.ws.close()
|
||||||
|
except:
|
||||||
|
pass
|
||||||
|
self.ws = None
|
||||||
|
raise
|
||||||
|
|
||||||
|
async def finish_session(self, session_id):
|
||||||
|
logger.bind(tag=TAG).info(f"关闭会话~~{session_id}")
|
||||||
|
try:
|
||||||
|
if self.ws:
|
||||||
|
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(self.ws, header, optional, payload)
|
||||||
|
logger.bind(tag=TAG).info("会话结束请求已发送")
|
||||||
|
|
||||||
|
# 等待监听任务完成
|
||||||
|
if hasattr(self, "_monitor_task"):
|
||||||
|
try:
|
||||||
|
await self._monitor_task
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(
|
||||||
|
f"等待监听任务完成时发生错误: {str(e)}"
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
self._monitor_task = None
|
||||||
|
|
||||||
|
# 关闭连接
|
||||||
|
await self.close()
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"关闭会话失败: {str(e)}")
|
||||||
|
# 确保清理资源
|
||||||
|
if hasattr(self, "_monitor_task"):
|
||||||
|
try:
|
||||||
|
self._monitor_task.cancel()
|
||||||
|
await self._monitor_task
|
||||||
|
except:
|
||||||
|
pass
|
||||||
|
self._monitor_task = None
|
||||||
|
if self.ws:
|
||||||
|
try:
|
||||||
|
await self.ws.close()
|
||||||
|
except:
|
||||||
|
pass
|
||||||
|
self.ws = None
|
||||||
|
raise
|
||||||
|
|
||||||
|
async def close(self):
|
||||||
|
"""资源清理方法"""
|
||||||
|
if self.ws:
|
||||||
|
try:
|
||||||
|
await self.ws.close()
|
||||||
|
except:
|
||||||
|
pass
|
||||||
|
self.ws = None
|
||||||
|
|
||||||
async def _start_monitor_tts_response(self):
|
async def _start_monitor_tts_response(self):
|
||||||
|
"""监听TTS响应"""
|
||||||
opus_datas_cache = []
|
opus_datas_cache = []
|
||||||
# 添加标志来区分是否是第一句话
|
|
||||||
is_first_sentence = True
|
is_first_sentence = True
|
||||||
while not self.conn.stop_event.is_set():
|
first_sentence_segment_count = 0 # 添加计数器
|
||||||
try:
|
try:
|
||||||
# 确保 `recv()` 运行在同一个 event loop
|
while not self.conn.stop_event.is_set():
|
||||||
msg = await self.ws.recv()
|
try:
|
||||||
res = self.parser_response(msg)
|
# 确保 `recv()` 运行在同一个 event loop
|
||||||
self.print_response(res, "send_text res:")
|
msg = await self.ws.recv()
|
||||||
|
res = self.parser_response(msg)
|
||||||
|
self.print_response(res, "send_text res:")
|
||||||
|
|
||||||
if res.optional.event == EVENT_TTSSentenceStart:
|
# 检查客户端是否中止
|
||||||
json_data = json.loads(res.payload.decode("utf-8"))
|
if self.conn.client_abort:
|
||||||
self.tts_text = json_data.get("text", "")
|
logger.bind(tag=TAG).info("收到打断信息,终止监听TTS响应")
|
||||||
logger.bind(tag=TAG).debug(f"句子语音生成开始: {self.tts_text}")
|
break
|
||||||
self.tts_audio_queue.put((SentenceType.FIRST, [], self.tts_text))
|
|
||||||
opus_datas_cache = []
|
if res.optional.event == EVENT_TTSSentenceStart:
|
||||||
elif (
|
json_data = json.loads(res.payload.decode("utf-8"))
|
||||||
res.optional.event == EVENT_TTSResponse
|
self.tts_text = json_data.get("text", "")
|
||||||
and res.header.message_type == AUDIO_ONLY_RESPONSE
|
logger.bind(tag=TAG).debug(f"句子语音生成开始: {self.tts_text}")
|
||||||
):
|
|
||||||
logger.bind(tag=TAG).debug(f"推送数据到队列里面~~")
|
|
||||||
opus_datas = self.wav_to_opus_data_audio_raw(res.payload)
|
|
||||||
logger.bind(tag=TAG).debug(
|
|
||||||
f"推送数据到队列里面帧数~~{len(opus_datas)}"
|
|
||||||
)
|
|
||||||
if is_first_sentence:
|
|
||||||
# 第一句话直接发送
|
|
||||||
self.tts_audio_queue.put(
|
self.tts_audio_queue.put(
|
||||||
(SentenceType.MIDDLE, opus_datas, self.tts_text)
|
(SentenceType.FIRST, [], self.tts_text)
|
||||||
)
|
)
|
||||||
else:
|
opus_datas_cache = []
|
||||||
# 后续句子缓存
|
first_sentence_segment_count = 0 # 重置计数器
|
||||||
opus_datas_cache = opus_datas_cache + opus_datas
|
elif (
|
||||||
elif res.optional.event == EVENT_TTSSentenceEnd:
|
res.optional.event == EVENT_TTSResponse
|
||||||
logger.bind(tag=TAG).info(f"句子语音生成成功:{self.tts_text}")
|
and res.header.message_type == AUDIO_ONLY_RESPONSE
|
||||||
if not is_first_sentence:
|
):
|
||||||
# 只有非第一句话才发送缓存的数据
|
logger.bind(tag=TAG).debug(f"推送数据到队列里面~~")
|
||||||
self.tts_audio_queue.put(
|
opus_datas = self.wav_to_opus_data_audio_raw(res.payload)
|
||||||
(SentenceType.MIDDLE, opus_datas_cache, self.tts_text)
|
logger.bind(tag=TAG).debug(
|
||||||
|
f"推送数据到队列里面帧数~~{len(opus_datas)}"
|
||||||
)
|
)
|
||||||
# 第一句话结束后,将标志设置为False
|
if is_first_sentence:
|
||||||
is_first_sentence = False
|
first_sentence_segment_count += 1
|
||||||
elif res.optional.event == EVENT_SessionFinished:
|
if first_sentence_segment_count <= 6:
|
||||||
logger.bind(tag=TAG).debug(f"会话结束~~")
|
self.tts_audio_queue.put(
|
||||||
for tts_file, text in self.before_stop_play_files:
|
(SentenceType.MIDDLE, opus_datas, None)
|
||||||
if tts_file and os.path.exists(tts_file):
|
)
|
||||||
audio_datas = self._process_audio_file(tts_file)
|
else:
|
||||||
|
opus_datas_cache = opus_datas_cache + opus_datas
|
||||||
|
else:
|
||||||
|
# 后续句子缓存
|
||||||
|
opus_datas_cache = opus_datas_cache + opus_datas
|
||||||
|
elif res.optional.event == EVENT_TTSSentenceEnd:
|
||||||
|
logger.bind(tag=TAG).info(f"句子语音生成成功:{self.tts_text}")
|
||||||
|
if not is_first_sentence or first_sentence_segment_count > 10:
|
||||||
|
# 发送缓存的数据
|
||||||
self.tts_audio_queue.put(
|
self.tts_audio_queue.put(
|
||||||
(SentenceType.MIDDLE, audio_datas, text)
|
(SentenceType.MIDDLE, opus_datas_cache, None)
|
||||||
)
|
)
|
||||||
self.before_stop_play_files.clear()
|
# 第一句话结束后,将标志设置为False
|
||||||
self.tts_audio_queue.put((SentenceType.LAST, [], None))
|
is_first_sentence = False
|
||||||
|
elif res.optional.event == EVENT_SessionFinished:
|
||||||
opus_datas_cache = []
|
logger.bind(tag=TAG).debug(f"会话结束~~")
|
||||||
is_first_sentence = True
|
self._process_before_stop_play_files()
|
||||||
continue
|
break
|
||||||
except websockets.ConnectionClosed:
|
except websockets.ConnectionClosed:
|
||||||
break # 连接关闭时退出监听
|
logger.bind(tag=TAG).warning("WebSocket连接已关闭")
|
||||||
except Exception as e:
|
break
|
||||||
logger.bind(tag=TAG).error(f"Error in _start_monitor_tts_response: {e}")
|
except Exception as e:
|
||||||
traceback.print_exc()
|
logger.bind(tag=TAG).error(
|
||||||
continue
|
f"Error in _start_monitor_tts_response: {e}"
|
||||||
|
)
|
||||||
|
traceback.print_exc()
|
||||||
|
break
|
||||||
|
finally:
|
||||||
|
# 确保清理资源
|
||||||
|
if self.ws:
|
||||||
|
try:
|
||||||
|
await self.ws.close()
|
||||||
|
except:
|
||||||
|
pass
|
||||||
|
self.ws = None
|
||||||
|
|
||||||
async def send_event(
|
async def send_event(
|
||||||
self, header: bytes, optional: bytes | None = None, payload: bytes = None
|
self,
|
||||||
|
ws: websockets.WebSocketClientProtocol,
|
||||||
|
header: bytes,
|
||||||
|
optional: bytes | None = None,
|
||||||
|
payload: bytes = None,
|
||||||
):
|
):
|
||||||
try:
|
try:
|
||||||
full_client_request = bytearray(header)
|
full_client_request = bytearray(header)
|
||||||
@@ -411,13 +468,10 @@ class TTSProvider(TTSProviderBase):
|
|||||||
payload_size = len(payload).to_bytes(4, "big", signed=True)
|
payload_size = len(payload).to_bytes(4, "big", signed=True)
|
||||||
full_client_request.extend(payload_size)
|
full_client_request.extend(payload_size)
|
||||||
full_client_request.extend(payload)
|
full_client_request.extend(payload)
|
||||||
await self.ws.send(full_client_request)
|
await ws.send(full_client_request)
|
||||||
except websockets.ConnectionClosed:
|
except websockets.ConnectionClosed:
|
||||||
if await self._handle_connection_error():
|
logger.bind(tag=TAG).error(f"ConnectionClosed")
|
||||||
# 重连成功后重试发送
|
raise
|
||||||
await self.ws.send(full_client_request)
|
|
||||||
else:
|
|
||||||
raise
|
|
||||||
|
|
||||||
async def send_text(self, speaker: str, text: str, session_id):
|
async def send_text(self, speaker: str, text: str, session_id):
|
||||||
header = Header(
|
header = Header(
|
||||||
@@ -429,7 +483,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
payload = self.get_payload_bytes(
|
payload = self.get_payload_bytes(
|
||||||
event=EVENT_TaskRequest, text=text, speaker=speaker
|
event=EVENT_TaskRequest, text=text, speaker=speaker
|
||||||
)
|
)
|
||||||
return await self.send_event(header, optional, payload)
|
return await self.send_event(self.ws, header, optional, payload)
|
||||||
|
|
||||||
# 读取 res 数组某段 字符串内容
|
# 读取 res 数组某段 字符串内容
|
||||||
def read_res_content(self, res: bytes, offset: int):
|
def read_res_content(self, res: bytes, offset: int):
|
||||||
@@ -507,7 +561,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
).as_bytes()
|
).as_bytes()
|
||||||
optional = Optional(event=EVENT_Start_Connection).as_bytes()
|
optional = Optional(event=EVENT_Start_Connection).as_bytes()
|
||||||
payload = str.encode("{}")
|
payload = str.encode("{}")
|
||||||
return await self.send_event(header, optional, payload)
|
return await self.send_event(self.ws, header, optional, payload)
|
||||||
|
|
||||||
def print_response(self, res, tag_msg: str):
|
def print_response(self, res, tag_msg: str):
|
||||||
logger.bind(tag=TAG).debug(f"===>{tag_msg} header:{res.header.__dict__}")
|
logger.bind(tag=TAG).debug(f"===>{tag_msg} header:{res.header.__dict__}")
|
||||||
@@ -540,50 +594,44 @@ class TTSProvider(TTSProviderBase):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
async def finish_connection(self):
|
def wav_to_opus_data_audio_raw(self, raw_data_var, is_end=False):
|
||||||
header = Header(
|
opus_datas = self.opus_encoder.encode_pcm_to_opus(raw_data_var, is_end)
|
||||||
message_type=FULL_CLIENT_REQUEST,
|
return opus_datas
|
||||||
message_type_specific_flags=MsgTypeFlagWithEvent,
|
|
||||||
serial_method=JSON,
|
|
||||||
).as_bytes()
|
|
||||||
optional = Optional(event=EVENT_FinishConnection).as_bytes()
|
|
||||||
payload = str.encode("{}")
|
|
||||||
await self.send_event(header, optional, payload)
|
|
||||||
return
|
|
||||||
|
|
||||||
async def start_session(self, session_id):
|
def to_tts(self, text: str) -> list:
|
||||||
logger.bind(tag=TAG).info(f"开始会话~~{session_id}")
|
"""非流式生成音频数据,用于生成音频及测试场景
|
||||||
|
|
||||||
|
Args:
|
||||||
|
text: 要转换的文本
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
list: 音频数据列表
|
||||||
|
"""
|
||||||
try:
|
try:
|
||||||
async with self._session_lock:
|
# 创建事件循环
|
||||||
|
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:
|
try:
|
||||||
# 确保连接可用
|
# 启动会话
|
||||||
logger.bind(tag=TAG).info("检查WebSocket连接状态...")
|
|
||||||
await asyncio.wait_for(self._ensure_connection(), timeout=5)
|
|
||||||
|
|
||||||
# 如果已有会话未结束,先关闭它
|
|
||||||
if self._session_started and not self._session_finished:
|
|
||||||
logger.bind(tag=TAG).warning(
|
|
||||||
f"发现未关闭的会话 {self._current_session_id},正在关闭..."
|
|
||||||
)
|
|
||||||
try:
|
|
||||||
await asyncio.wait_for(
|
|
||||||
self.finish_session(self._current_session_id), timeout=5
|
|
||||||
)
|
|
||||||
except Exception as e:
|
|
||||||
logger.bind(tag=TAG).error(f"关闭旧会话失败: {str(e)}")
|
|
||||||
# 强制重置会话状态
|
|
||||||
self._session_started = False
|
|
||||||
self._session_finished = True
|
|
||||||
self._current_session_id = None
|
|
||||||
|
|
||||||
# 重置会话状态
|
|
||||||
self._current_session_id = session_id
|
|
||||||
self._session_started = True
|
|
||||||
self._session_finished = False
|
|
||||||
logger.bind(tag=TAG).info(
|
|
||||||
f"会话状态已更新 - 开始: {self._session_started}, 结束: {self._session_finished}"
|
|
||||||
)
|
|
||||||
|
|
||||||
header = Header(
|
header = Header(
|
||||||
message_type=FULL_CLIENT_REQUEST,
|
message_type=FULL_CLIENT_REQUEST,
|
||||||
message_type_specific_flags=MsgTypeFlagWithEvent,
|
message_type_specific_flags=MsgTypeFlagWithEvent,
|
||||||
@@ -593,69 +641,25 @@ class TTSProvider(TTSProviderBase):
|
|||||||
event=EVENT_StartSession, sessionId=session_id
|
event=EVENT_StartSession, sessionId=session_id
|
||||||
).as_bytes()
|
).as_bytes()
|
||||||
payload = self.get_payload_bytes(
|
payload = self.get_payload_bytes(
|
||||||
event=EVENT_StartSession, speaker=self.speaker
|
event=EVENT_StartSession, speaker=self.voice
|
||||||
)
|
)
|
||||||
await asyncio.wait_for(
|
await self.send_event(ws, header, optional, payload)
|
||||||
self.send_event(header, optional, payload), timeout=5
|
|
||||||
|
# 发送文本
|
||||||
|
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
|
||||||
)
|
)
|
||||||
logger.bind(tag=TAG).info("会话启动请求已发送")
|
await self.send_event(ws, header, optional, payload)
|
||||||
except Exception as e:
|
|
||||||
logger.bind(tag=TAG).error(f"启动会话失败: {str(e)}")
|
|
||||||
self._session_started = False
|
|
||||||
self._session_finished = True
|
|
||||||
self._current_session_id = None
|
|
||||||
raise
|
|
||||||
except asyncio.TimeoutError:
|
|
||||||
logger.bind(tag=TAG).error(f"启动会话超时: {session_id}")
|
|
||||||
# 超时后强制重置会话状态
|
|
||||||
self._session_started = False
|
|
||||||
self._session_finished = True
|
|
||||||
self._current_session_id = None
|
|
||||||
# 尝试关闭WebSocket连接
|
|
||||||
if self.ws:
|
|
||||||
try:
|
|
||||||
await self.ws.close()
|
|
||||||
except:
|
|
||||||
pass
|
|
||||||
self.ws = None
|
|
||||||
except Exception as e:
|
|
||||||
logger.bind(tag=TAG).error(f"启动会话时发生未知错误: {str(e)}")
|
|
||||||
# 发生未知错误时也重置会话状态
|
|
||||||
self._session_started = False
|
|
||||||
self._session_finished = True
|
|
||||||
self._current_session_id = None
|
|
||||||
|
|
||||||
async def finish_session(self, session_id):
|
|
||||||
logger.bind(tag=TAG).info(f"关闭会话~~{session_id}")
|
|
||||||
try:
|
|
||||||
async with self._session_lock:
|
|
||||||
try:
|
|
||||||
# 检查会话状态
|
|
||||||
if not self._session_started:
|
|
||||||
logger.bind(tag=TAG).warning(
|
|
||||||
f"尝试关闭未开始的会话 {session_id}"
|
|
||||||
)
|
|
||||||
return
|
|
||||||
|
|
||||||
if self._session_finished:
|
|
||||||
logger.bind(tag=TAG).warning(f"会话 {session_id} 已经关闭")
|
|
||||||
return
|
|
||||||
|
|
||||||
if self._current_session_id != session_id:
|
|
||||||
logger.bind(tag=TAG).warning(
|
|
||||||
f"尝试关闭错误的会话 {session_id},当前会话为 {self._current_session_id}"
|
|
||||||
)
|
|
||||||
# 即使会话ID不匹配,也尝试关闭当前会话
|
|
||||||
if self._current_session_id:
|
|
||||||
session_id = self._current_session_id
|
|
||||||
|
|
||||||
# 确保WebSocket连接可用
|
|
||||||
if self.ws is None:
|
|
||||||
logger.bind(tag=TAG).warning(
|
|
||||||
"WebSocket连接不存在,尝试重新连接..."
|
|
||||||
)
|
|
||||||
await asyncio.wait_for(self._ensure_connection(), timeout=5)
|
|
||||||
|
|
||||||
|
# 发送结束会话请求
|
||||||
header = Header(
|
header = Header(
|
||||||
message_type=FULL_CLIENT_REQUEST,
|
message_type=FULL_CLIENT_REQUEST,
|
||||||
message_type_specific_flags=MsgTypeFlagWithEvent,
|
message_type_specific_flags=MsgTypeFlagWithEvent,
|
||||||
@@ -665,75 +669,35 @@ class TTSProvider(TTSProviderBase):
|
|||||||
event=EVENT_FinishSession, sessionId=session_id
|
event=EVENT_FinishSession, sessionId=session_id
|
||||||
).as_bytes()
|
).as_bytes()
|
||||||
payload = str.encode("{}")
|
payload = str.encode("{}")
|
||||||
await asyncio.wait_for(
|
await self.send_event(ws, header, optional, payload)
|
||||||
self.send_event(header, optional, payload), timeout=5
|
|
||||||
)
|
# 接收音频数据
|
||||||
logger.bind(tag=TAG).info("会话结束请求已发送")
|
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
|
||||||
|
):
|
||||||
|
opus_datas = self.wav_to_opus_data_audio_raw(res.payload)
|
||||||
|
audio_data.extend(opus_datas)
|
||||||
|
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
|
||||||
|
|
||||||
# 更新会话状态
|
|
||||||
self._session_finished = True
|
|
||||||
self._session_started = False
|
|
||||||
self._current_session_id = None
|
|
||||||
logger.bind(tag=TAG).info(
|
|
||||||
"会话状态已更新 - 开始: False, 结束: True"
|
|
||||||
)
|
|
||||||
except Exception as e:
|
|
||||||
logger.bind(tag=TAG).error(f"关闭会话失败: {str(e)}")
|
|
||||||
# 即使发生错误,也要重置会话状态
|
|
||||||
self._session_finished = True
|
|
||||||
self._session_started = False
|
|
||||||
self._current_session_id = None
|
|
||||||
raise
|
|
||||||
except asyncio.TimeoutError:
|
|
||||||
logger.bind(tag=TAG).error(f"关闭会话超时: {session_id}")
|
|
||||||
# 超时后强制重置会话状态
|
|
||||||
self._session_finished = True
|
|
||||||
self._session_started = False
|
|
||||||
self._current_session_id = None
|
|
||||||
# 尝试关闭WebSocket连接
|
|
||||||
if self.ws:
|
|
||||||
try:
|
|
||||||
await self.ws.close()
|
|
||||||
except:
|
|
||||||
pass
|
|
||||||
self.ws = None
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.bind(tag=TAG).error(f"关闭会话时发生未知错误: {str(e)}")
|
logger.bind(tag=TAG).error(f"生成音频数据失败: {str(e)}")
|
||||||
# 发生未知错误时也重置会话状态
|
return []
|
||||||
self._session_finished = True
|
|
||||||
self._session_started = False
|
|
||||||
self._current_session_id = None
|
|
||||||
|
|
||||||
async def reset(self):
|
|
||||||
# 关闭之前的对话
|
|
||||||
if self.start_connection_flag:
|
|
||||||
await self.finish_connection()
|
|
||||||
self.start_connection_flag = False
|
|
||||||
await self.start_connection()
|
|
||||||
self.start_connection_flag = True
|
|
||||||
|
|
||||||
async def close(self):
|
|
||||||
"""资源清理方法"""
|
|
||||||
await self.finish_connection()
|
|
||||||
await self.ws.close()
|
|
||||||
|
|
||||||
def wav_to_opus_data_audio_raw(self, raw_data_var, is_end=False):
|
|
||||||
opus_datas = self.opus_encoder.encode_pcm_to_opus(raw_data_var, is_end)
|
|
||||||
return opus_datas
|
|
||||||
|
|
||||||
async def _handle_connection_error(self):
|
|
||||||
"""处理连接错误"""
|
|
||||||
if self._reconnect_attempts < self._max_reconnect_attempts:
|
|
||||||
self._reconnect_attempts += 1
|
|
||||||
logger.bind(tag=TAG).warning(
|
|
||||||
f"尝试重新连接 (第{self._reconnect_attempts}次)"
|
|
||||||
)
|
|
||||||
try:
|
|
||||||
await self._ensure_connection()
|
|
||||||
return True
|
|
||||||
except Exception as e:
|
|
||||||
logger.bind(tag=TAG).error(f"重新连接失败: {str(e)}")
|
|
||||||
return False
|
|
||||||
else:
|
|
||||||
logger.bind(tag=TAG).error("达到最大重连次数,放弃重连")
|
|
||||||
return False
|
|
||||||
|
|||||||
@@ -0,0 +1,303 @@
|
|||||||
|
import queue
|
||||||
|
import asyncio
|
||||||
|
import traceback
|
||||||
|
import aiohttp
|
||||||
|
import requests
|
||||||
|
import time
|
||||||
|
from config.logger import setup_logging
|
||||||
|
from core.utils.tts import MarkdownCleaner
|
||||||
|
from core.providers.tts.base import TTSProviderBase
|
||||||
|
from core.utils import opus_encoder_utils, textUtils
|
||||||
|
from core.providers.tts.dto.dto import SentenceType, ContentType, InterfaceType
|
||||||
|
|
||||||
|
TAG = __name__
|
||||||
|
logger = setup_logging()
|
||||||
|
|
||||||
|
|
||||||
|
class TTSProvider(TTSProviderBase):
|
||||||
|
def __init__(self, config, delete_audio_file):
|
||||||
|
super().__init__(config, delete_audio_file)
|
||||||
|
self.interface_type = InterfaceType.SINGLE_STREAM
|
||||||
|
self.access_token = config.get("access_token")
|
||||||
|
self.voice = config.get("voice")
|
||||||
|
self.api_url = config.get("api_url")
|
||||||
|
self.audio_format = "pcm"
|
||||||
|
self.before_stop_play_files = []
|
||||||
|
self.segment_count = 0 # 添加片段计数器
|
||||||
|
|
||||||
|
# 创建Opus编码器
|
||||||
|
self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(
|
||||||
|
sample_rate=16000, channels=1, frame_size_ms=60
|
||||||
|
)
|
||||||
|
|
||||||
|
# 添加文本缓冲区
|
||||||
|
self.text_buffer = ""
|
||||||
|
|
||||||
|
# PCM缓冲区
|
||||||
|
self.pcm_buffer = bytearray()
|
||||||
|
|
||||||
|
###################################################################################
|
||||||
|
# linkerai单流式TTS重写父类的方法--开始
|
||||||
|
###################################################################################
|
||||||
|
|
||||||
|
def tts_text_priority_thread(self):
|
||||||
|
"""流式文本处理线程"""
|
||||||
|
while not self.conn.stop_event.is_set():
|
||||||
|
try:
|
||||||
|
message = self.tts_text_queue.get(timeout=1)
|
||||||
|
if message.sentence_type == SentenceType.FIRST:
|
||||||
|
# 初始化参数
|
||||||
|
self.tts_stop_request = False
|
||||||
|
self.processed_chars = 0
|
||||||
|
self.tts_text_buff = []
|
||||||
|
self.segment_count = 0
|
||||||
|
self.tts_audio_first_sentence = True
|
||||||
|
self.before_stop_play_files.clear()
|
||||||
|
elif ContentType.TEXT == message.content_type:
|
||||||
|
self.tts_text_buff.append(message.content_detail)
|
||||||
|
segment_text = self._get_segment_text()
|
||||||
|
if segment_text:
|
||||||
|
self.to_tts_single_stream(segment_text)
|
||||||
|
|
||||||
|
elif ContentType.FILE == message.content_type:
|
||||||
|
logger.bind(tag=TAG).info(
|
||||||
|
f"添加音频文件到待播放列表: {message.content_file}"
|
||||||
|
)
|
||||||
|
self.before_stop_play_files.append(
|
||||||
|
(message.content_file, message.content_detail)
|
||||||
|
)
|
||||||
|
|
||||||
|
if message.sentence_type == SentenceType.LAST:
|
||||||
|
# 处理剩余的文本
|
||||||
|
self._process_remaining_text(True)
|
||||||
|
|
||||||
|
except queue.Empty:
|
||||||
|
continue
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(
|
||||||
|
f"处理TTS文本失败: {str(e)}, 类型: {type(e).__name__}, 堆栈: {traceback.format_exc()}"
|
||||||
|
)
|
||||||
|
|
||||||
|
def _process_remaining_text(self, is_last=False):
|
||||||
|
"""处理剩余的文本并生成语音
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
bool: 是否成功处理了文本
|
||||||
|
"""
|
||||||
|
full_text = "".join(self.tts_text_buff)
|
||||||
|
remaining_text = full_text[self.processed_chars :]
|
||||||
|
if remaining_text:
|
||||||
|
segment_text = textUtils.get_string_no_punctuation_or_emoji(remaining_text)
|
||||||
|
if segment_text:
|
||||||
|
self.to_tts_single_stream(segment_text, is_last)
|
||||||
|
self.processed_chars += len(full_text)
|
||||||
|
else:
|
||||||
|
self._process_before_stop_play_files()
|
||||||
|
else:
|
||||||
|
self._process_before_stop_play_files()
|
||||||
|
|
||||||
|
def to_tts_single_stream(self, text, is_last=False):
|
||||||
|
try:
|
||||||
|
max_repeat_time = 5
|
||||||
|
text = MarkdownCleaner.clean_markdown(text)
|
||||||
|
try:
|
||||||
|
asyncio.run(self.text_to_speak(text, is_last))
|
||||||
|
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},请检查网络或服务是否正常"
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"Failed to generate TTS file: {e}")
|
||||||
|
finally:
|
||||||
|
return None
|
||||||
|
|
||||||
|
###################################################################################
|
||||||
|
# linkerai单流式TTS重写父类的方法--结束
|
||||||
|
###################################################################################
|
||||||
|
|
||||||
|
async def text_to_speak(self, text, is_last):
|
||||||
|
"""流式处理TTS音频,每句只推送一次音频列表"""
|
||||||
|
await self._tts_request(text, is_last)
|
||||||
|
|
||||||
|
async def close(self):
|
||||||
|
"""资源清理"""
|
||||||
|
await super().close()
|
||||||
|
if hasattr(self, "opus_encoder"):
|
||||||
|
self.opus_encoder.close()
|
||||||
|
|
||||||
|
async def _tts_request(self, text: str, is_last: bool) -> None:
|
||||||
|
params = {
|
||||||
|
"tts_text": text,
|
||||||
|
"spk_id": self.voice,
|
||||||
|
"frame_durition": 60,
|
||||||
|
"stream": "true",
|
||||||
|
"target_sr": 16000,
|
||||||
|
"audio_format": "pcm",
|
||||||
|
"instruct_text": "请生成一段自然流畅的语音",
|
||||||
|
}
|
||||||
|
headers = {
|
||||||
|
"Authorization": f"Bearer {self.access_token}",
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
}
|
||||||
|
|
||||||
|
# 一帧 PCM 所需字节数:60 ms × 16 kHz × 1 ch × 2 B = 1 920
|
||||||
|
frame_bytes = int(
|
||||||
|
self.opus_encoder.sample_rate
|
||||||
|
* self.opus_encoder.channels # 1
|
||||||
|
* self.opus_encoder.frame_size_ms
|
||||||
|
/ 1000
|
||||||
|
* 2
|
||||||
|
) # 16-bit = 2 bytes
|
||||||
|
|
||||||
|
try:
|
||||||
|
async with aiohttp.ClientSession() as session:
|
||||||
|
async with session.get(
|
||||||
|
self.api_url, params=params, headers=headers, timeout=10
|
||||||
|
) as resp:
|
||||||
|
|
||||||
|
if resp.status != 200:
|
||||||
|
logger.error(f"TTS请求失败: {resp.status}, {await resp.text()}")
|
||||||
|
self.tts_audio_queue.put((SentenceType.LAST, [], None))
|
||||||
|
return
|
||||||
|
|
||||||
|
self.pcm_buffer.clear()
|
||||||
|
opus_datas_cache = []
|
||||||
|
|
||||||
|
self.tts_audio_queue.put((SentenceType.FIRST, [], text))
|
||||||
|
|
||||||
|
# 兼容 iter_chunked / iter_chunks / iter_any
|
||||||
|
async for chunk in resp.content.iter_any():
|
||||||
|
data = chunk[0] if isinstance(chunk, (list, tuple)) else chunk
|
||||||
|
if not data:
|
||||||
|
continue
|
||||||
|
|
||||||
|
# 拼到 buffer
|
||||||
|
self.pcm_buffer.extend(data)
|
||||||
|
|
||||||
|
# 够一帧就编码
|
||||||
|
while len(self.pcm_buffer) >= frame_bytes:
|
||||||
|
frame = bytes(self.pcm_buffer[:frame_bytes])
|
||||||
|
del self.pcm_buffer[:frame_bytes]
|
||||||
|
|
||||||
|
opus = self.opus_encoder.encode_pcm_to_opus(
|
||||||
|
frame, end_of_stream=False
|
||||||
|
)
|
||||||
|
if opus:
|
||||||
|
if self.segment_count < 10: # 前10个片段直接发送
|
||||||
|
self.tts_audio_queue.put(
|
||||||
|
(SentenceType.MIDDLE, opus, None)
|
||||||
|
)
|
||||||
|
self.segment_count += 1
|
||||||
|
else:
|
||||||
|
opus_datas_cache.extend(opus)
|
||||||
|
|
||||||
|
# flush 剩余不足一帧的数据
|
||||||
|
if self.pcm_buffer:
|
||||||
|
opus = self.opus_encoder.encode_pcm_to_opus(
|
||||||
|
bytes(self.pcm_buffer), end_of_stream=True
|
||||||
|
)
|
||||||
|
if opus:
|
||||||
|
if self.segment_count < 10: # 前10个片段直接发送
|
||||||
|
# 直接发送
|
||||||
|
self.tts_audio_queue.put(
|
||||||
|
(SentenceType.MIDDLE, opus, None)
|
||||||
|
)
|
||||||
|
self.segment_count += 1
|
||||||
|
else:
|
||||||
|
# 后续片段缓存
|
||||||
|
opus_datas_cache.extend(opus)
|
||||||
|
self.pcm_buffer.clear()
|
||||||
|
|
||||||
|
# 如果不是前10个片段,发送缓存的数据
|
||||||
|
if self.segment_count >= 10 and opus_datas_cache:
|
||||||
|
self.tts_audio_queue.put(
|
||||||
|
(SentenceType.MIDDLE, opus_datas_cache, None)
|
||||||
|
)
|
||||||
|
|
||||||
|
# 如果是最后一段,输出音频获取完毕
|
||||||
|
if is_last:
|
||||||
|
self._process_before_stop_play_files()
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.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.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))
|
||||||
|
|
||||||
|
opus = self.opus_encoder.encode_pcm_to_opus(
|
||||||
|
frame, end_of_stream=(i + frame_bytes >= len(pcm_data))
|
||||||
|
)
|
||||||
|
if opus:
|
||||||
|
opus_datas.extend(opus)
|
||||||
|
|
||||||
|
return opus_datas
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"TTS请求异常: {e}")
|
||||||
|
return []
|
||||||
@@ -14,9 +14,9 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.api_key = config.get("api_key")
|
self.api_key = config.get("api_key")
|
||||||
self.model = config.get("model")
|
self.model = config.get("model")
|
||||||
if config.get("private_voice"):
|
if config.get("private_voice"):
|
||||||
self.voice_id = config.get("private_voice")
|
self.voice = config.get("private_voice")
|
||||||
else:
|
else:
|
||||||
self.voice_id = config.get("voice_id")
|
self.voice = config.get("voice_id")
|
||||||
|
|
||||||
default_voice_setting = {
|
default_voice_setting = {
|
||||||
"voice_id": "female-shaonv",
|
"voice_id": "female-shaonv",
|
||||||
@@ -43,8 +43,8 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.audio_setting = {**defult_audio_setting, **config.get("audio_setting", {})}
|
self.audio_setting = {**defult_audio_setting, **config.get("audio_setting", {})}
|
||||||
self.timber_weights = parse_string_to_list(config.get("timber_weights"))
|
self.timber_weights = parse_string_to_list(config.get("timber_weights"))
|
||||||
|
|
||||||
if self.voice_id:
|
if self.voice:
|
||||||
self.voice_setting["voice_id"] = self.voice_id
|
self.voice_setting["voice_id"] = self.voice
|
||||||
|
|
||||||
self.host = "api.minimax.chat"
|
self.host = "api.minimax.chat"
|
||||||
self.api_url = f"https://{self.host}/v1/t2a_v2?GroupId={self.group_id}"
|
self.api_url = f"https://{self.host}/v1/t2a_v2?GroupId={self.group_id}"
|
||||||
@@ -52,6 +52,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
"Content-Type": "application/json",
|
"Content-Type": "application/json",
|
||||||
"Authorization": f"Bearer {self.api_key}",
|
"Authorization": f"Bearer {self.api_key}",
|
||||||
}
|
}
|
||||||
|
self.audio_file_type = defult_audio_setting.get("format", "mp3")
|
||||||
|
|
||||||
def generate_filename(self, extension=".mp3"):
|
def generate_filename(self, extension=".mp3"):
|
||||||
return os.path.join(
|
return os.path.join(
|
||||||
@@ -80,8 +81,12 @@ class TTSProvider(TTSProviderBase):
|
|||||||
# 检查返回请求数据的status_code是否为0
|
# 检查返回请求数据的status_code是否为0
|
||||||
if resp.json()["base_resp"]["status_code"] == 0:
|
if resp.json()["base_resp"]["status_code"] == 0:
|
||||||
data = resp.json()["data"]["audio"]
|
data = resp.json()["data"]["audio"]
|
||||||
file_to_save = open(output_file, "wb")
|
audio_bytes = bytes.fromhex(data)
|
||||||
file_to_save.write(bytes.fromhex(data))
|
if output_file:
|
||||||
|
with open(output_file, "wb") as file_to_save:
|
||||||
|
file_to_save.write(audio_bytes)
|
||||||
|
else:
|
||||||
|
return audio_bytes
|
||||||
else:
|
else:
|
||||||
raise Exception(
|
raise Exception(
|
||||||
f"{__name__} status_code: {resp.status_code} response: {resp.content}"
|
f"{__name__} status_code: {resp.status_code} response: {resp.content}"
|
||||||
|
|||||||
@@ -17,7 +17,8 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.voice = config.get("private_voice")
|
self.voice = config.get("private_voice")
|
||||||
else:
|
else:
|
||||||
self.voice = config.get("voice", "alloy")
|
self.voice = config.get("voice", "alloy")
|
||||||
self.response_format = "wav"
|
self.response_format = config.get("format", "wav")
|
||||||
|
self.audio_file_type = config.get("format", "wav")
|
||||||
|
|
||||||
# 处理空字符串的情况
|
# 处理空字符串的情况
|
||||||
speed = config.get("speed", "1.0")
|
speed = config.get("speed", "1.0")
|
||||||
@@ -40,8 +41,11 @@ class TTSProvider(TTSProviderBase):
|
|||||||
}
|
}
|
||||||
response = requests.post(self.api_url, json=data, headers=headers)
|
response = requests.post(self.api_url, json=data, headers=headers)
|
||||||
if response.status_code == 200:
|
if response.status_code == 200:
|
||||||
with open(output_file, "wb") as audio_file:
|
if output_file:
|
||||||
audio_file.write(response.content)
|
with open(output_file, "wb") as audio_file:
|
||||||
|
audio_file.write(response.content)
|
||||||
|
else:
|
||||||
|
return response.content
|
||||||
else:
|
else:
|
||||||
raise Exception(
|
raise Exception(
|
||||||
f"OpenAI TTS请求失败: {response.status_code} - {response.text}"
|
f"OpenAI TTS请求失败: {response.status_code} - {response.text}"
|
||||||
|
|||||||
@@ -11,7 +11,8 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.voice = config.get("private_voice")
|
self.voice = config.get("private_voice")
|
||||||
else:
|
else:
|
||||||
self.voice = config.get("voice")
|
self.voice = config.get("voice")
|
||||||
self.response_format = config.get("response_format")
|
self.response_format = config.get("response_format", "mp3")
|
||||||
|
self.audio_file_type = config.get("response_format", "mp3")
|
||||||
self.sample_rate = config.get("sample_rate")
|
self.sample_rate = config.get("sample_rate")
|
||||||
self.speed = float(config.get("speed", 1.0))
|
self.speed = float(config.get("speed", 1.0))
|
||||||
self.gain = config.get("gain")
|
self.gain = config.get("gain")
|
||||||
@@ -35,7 +36,10 @@ class TTSProvider(TTSProviderBase):
|
|||||||
"POST", self.api_url, json=request_json, headers=headers
|
"POST", self.api_url, json=request_json, headers=headers
|
||||||
)
|
)
|
||||||
data = response.content
|
data = response.content
|
||||||
file_to_save = open(output_file, "wb")
|
if output_file:
|
||||||
file_to_save.write(data)
|
with open(output_file, "wb") as file_to_save:
|
||||||
|
file_to_save.write(data)
|
||||||
|
else:
|
||||||
|
return data
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
raise Exception(f"{__name__} error: {e}")
|
raise Exception(f"{__name__} error: {e}")
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.api_url = "https://tts.tencentcloudapi.com" # 正确的API端点
|
self.api_url = "https://tts.tencentcloudapi.com" # 正确的API端点
|
||||||
self.region = config.get("region")
|
self.region = config.get("region")
|
||||||
self.output_file = config.get("output_dir")
|
self.output_file = config.get("output_dir")
|
||||||
|
self.audio_file_type = config.get("format", "wav")
|
||||||
|
|
||||||
def _get_auth_headers(self, request_body):
|
def _get_auth_headers(self, request_body):
|
||||||
"""生成鉴权请求头"""
|
"""生成鉴权请求头"""
|
||||||
@@ -148,12 +149,14 @@ class TTSProvider(TTSProviderBase):
|
|||||||
f"API返回错误: {error_info['Code']}: {error_info['Message']}"
|
f"API返回错误: {error_info['Code']}: {error_info['Message']}"
|
||||||
)
|
)
|
||||||
|
|
||||||
# 提取音频数据
|
# 解码Base64音频数据
|
||||||
audio_data = response_data["Response"].get("Audio")
|
audio_bytes = base64.b64decode(response_data["Response"].get("Audio"))
|
||||||
if audio_data:
|
if audio_bytes:
|
||||||
# 解码Base64音频数据并保存
|
if output_file:
|
||||||
with open(output_file, "wb") as f:
|
with open(output_file, "wb") as f:
|
||||||
f.write(base64.b64decode(audio_data))
|
f.write(audio_bytes)
|
||||||
|
else:
|
||||||
|
return audio_bytes
|
||||||
else:
|
else:
|
||||||
raise Exception(f"{__name__}: 没有返回音频数据: {response_data}")
|
raise Exception(f"{__name__}: 没有返回音频数据: {response_data}")
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -19,9 +19,9 @@ class TTSProvider(TTSProviderBase):
|
|||||||
"https://u95167-bd74-2aef8085.westx.seetacloud.com:8443/flashsummary/tts?token=",
|
"https://u95167-bd74-2aef8085.westx.seetacloud.com:8443/flashsummary/tts?token=",
|
||||||
)
|
)
|
||||||
if config.get("private_voice"):
|
if config.get("private_voice"):
|
||||||
self.voice_id = int(config.get("private_voice"))
|
self.voice = int(config.get("private_voice"))
|
||||||
else:
|
else:
|
||||||
self.voice_id = int(config.get("voice_id", 1695))
|
self.voice = int(config.get("voice_id", 1695))
|
||||||
self.token = config.get("token")
|
self.token = config.get("token")
|
||||||
self.to_lang = config.get("to_lang")
|
self.to_lang = config.get("to_lang")
|
||||||
self.volume_change_dB = int(config.get("volume_change_dB", 0))
|
self.volume_change_dB = int(config.get("volume_change_dB", 0))
|
||||||
@@ -30,6 +30,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
self.output_file = config.get("output_dir")
|
self.output_file = config.get("output_dir")
|
||||||
self.pitch_factor = int(config.get("pitch_factor", 0))
|
self.pitch_factor = int(config.get("pitch_factor", 0))
|
||||||
self.format = config.get("format", "mp3")
|
self.format = config.get("format", "mp3")
|
||||||
|
self.audio_file_type = config.get("format", "mp3")
|
||||||
self.emotion = int(config.get("emotion", 1))
|
self.emotion = int(config.get("emotion", 1))
|
||||||
self.header = {"Content-Type": "application/json"}
|
self.header = {"Content-Type": "application/json"}
|
||||||
|
|
||||||
@@ -49,7 +50,7 @@ class TTSProvider(TTSProviderBase):
|
|||||||
"emotion": self.emotion,
|
"emotion": self.emotion,
|
||||||
"format": self.format,
|
"format": self.format,
|
||||||
"volume_change_dB": self.volume_change_dB,
|
"volume_change_dB": self.volume_change_dB,
|
||||||
"voice_id": self.voice_id,
|
"voice_id": self.voice,
|
||||||
"pitch_factor": self.pitch_factor,
|
"pitch_factor": self.pitch_factor,
|
||||||
"speed_factor": self.speed_factor,
|
"speed_factor": self.speed_factor,
|
||||||
"token": self.token,
|
"token": self.token,
|
||||||
@@ -73,9 +74,11 @@ class TTSProvider(TTSProviderBase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
audio_content = requests.get(result)
|
audio_content = requests.get(result)
|
||||||
with open(output_file, "wb") as f:
|
if output_file:
|
||||||
f.write(audio_content.content)
|
with open(output_file, "wb") as f:
|
||||||
return True
|
f.write(audio_content.content)
|
||||||
|
else:
|
||||||
|
return audio_content.content
|
||||||
voice_path = resp_json.get("voice_path")
|
voice_path = resp_json.get("voice_path")
|
||||||
des_path = output_file
|
des_path = output_file
|
||||||
shutil.move(voice_path, des_path)
|
shutil.move(voice_path, des_path)
|
||||||
|
|||||||
@@ -0,0 +1,12 @@
|
|||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from config.logger import setup_logging
|
||||||
|
|
||||||
|
TAG = __name__
|
||||||
|
logger = setup_logging()
|
||||||
|
|
||||||
|
|
||||||
|
class VLLMProviderBase(ABC):
|
||||||
|
@abstractmethod
|
||||||
|
def response(self, question, base64_image):
|
||||||
|
"""VLLM response generator"""
|
||||||
|
pass
|
||||||
@@ -0,0 +1,65 @@
|
|||||||
|
import openai
|
||||||
|
import json
|
||||||
|
from config.logger import setup_logging
|
||||||
|
from core.utils.util import check_model_key
|
||||||
|
from core.providers.vllm.base import VLLMProviderBase
|
||||||
|
|
||||||
|
TAG = __name__
|
||||||
|
logger = setup_logging()
|
||||||
|
|
||||||
|
|
||||||
|
class VLLMProvider(VLLMProviderBase):
|
||||||
|
def __init__(self, config):
|
||||||
|
self.model_name = config.get("model_name")
|
||||||
|
self.api_key = config.get("api_key")
|
||||||
|
if "base_url" in config:
|
||||||
|
self.base_url = config.get("base_url")
|
||||||
|
else:
|
||||||
|
self.base_url = config.get("url")
|
||||||
|
|
||||||
|
param_defaults = {
|
||||||
|
"max_tokens": (500, int),
|
||||||
|
"temperature": (0.7, lambda x: round(float(x), 1)),
|
||||||
|
"top_p": (1.0, lambda x: round(float(x), 1)),
|
||||||
|
}
|
||||||
|
|
||||||
|
for param, (default, converter) in param_defaults.items():
|
||||||
|
value = config.get(param)
|
||||||
|
try:
|
||||||
|
setattr(
|
||||||
|
self,
|
||||||
|
param,
|
||||||
|
converter(value) if value not in (None, "") else default,
|
||||||
|
)
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
setattr(self, param, default)
|
||||||
|
|
||||||
|
check_model_key("VLLM", self.api_key)
|
||||||
|
self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url)
|
||||||
|
|
||||||
|
def response(self, question, base64_image):
|
||||||
|
try:
|
||||||
|
messages = [
|
||||||
|
{
|
||||||
|
"role": "user",
|
||||||
|
"content": [
|
||||||
|
{"type": "text", "text": question},
|
||||||
|
{
|
||||||
|
"type": "image_url",
|
||||||
|
"image_url": {
|
||||||
|
"url": f"data:image/jpeg;base64,{base64_image}"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
]
|
||||||
|
|
||||||
|
response = self.client.chat.completions.create(
|
||||||
|
model=self.model_name, messages=messages, stream=False
|
||||||
|
)
|
||||||
|
|
||||||
|
return response.choices[0].message.content
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.bind(tag=TAG).error(f"Error in response generation: {e}")
|
||||||
|
raise
|
||||||
@@ -0,0 +1,43 @@
|
|||||||
|
import jwt
|
||||||
|
import time
|
||||||
|
from datetime import datetime, timedelta, timezone
|
||||||
|
from typing import Optional, Tuple
|
||||||
|
|
||||||
|
|
||||||
|
class AuthToken:
|
||||||
|
def __init__(self, secret_key: str):
|
||||||
|
self.secret_key = secret_key
|
||||||
|
|
||||||
|
def generate_token(self, device_id: str) -> str:
|
||||||
|
"""
|
||||||
|
生成JWT token
|
||||||
|
:param device_id: 设备ID
|
||||||
|
:return: JWT token字符串
|
||||||
|
"""
|
||||||
|
# 设置过期时间为1小时后
|
||||||
|
expire_time = datetime.now(timezone.utc) + timedelta(hours=1)
|
||||||
|
|
||||||
|
# 创建payload
|
||||||
|
payload = {"device_id": device_id, "exp": expire_time.timestamp()}
|
||||||
|
|
||||||
|
# 使用JWT进行编码
|
||||||
|
token = jwt.encode(payload, self.secret_key, algorithm="HS256")
|
||||||
|
return token
|
||||||
|
|
||||||
|
def verify_token(self, token: str) -> Tuple[bool, Optional[str]]:
|
||||||
|
"""
|
||||||
|
验证token
|
||||||
|
:param token: JWT token字符串
|
||||||
|
:return: (是否有效, 设备ID)
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
# 解码token
|
||||||
|
payload = jwt.decode(token, self.secret_key, algorithms=["HS256"])
|
||||||
|
|
||||||
|
# 检查是否过期
|
||||||
|
if payload["exp"] < time.time():
|
||||||
|
return False, None
|
||||||
|
|
||||||
|
return True, payload["device_id"]
|
||||||
|
except jwt.InvalidTokenError:
|
||||||
|
return False, None
|
||||||
@@ -29,5 +29,31 @@ def decode_opus_from_file(input_file):
|
|||||||
total_frames += 1
|
total_frames += 1
|
||||||
|
|
||||||
# 计算总时长
|
# 计算总时长
|
||||||
|
total_duration = (total_frames * frame_duration_ms) / 1000.0
|
||||||
|
return opus_datas, total_duration
|
||||||
|
|
||||||
|
def decode_opus_from_bytes(input_bytes):
|
||||||
|
"""
|
||||||
|
从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:
|
||||||
|
header = f.read(4)
|
||||||
|
if not header:
|
||||||
|
break
|
||||||
|
_, _, data_len = struct.unpack('>BBH', header)
|
||||||
|
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.")
|
||||||
|
opus_datas.append(opus_data)
|
||||||
|
total_frames += 1
|
||||||
|
|
||||||
total_duration = (total_frames * frame_duration_ms) / 1000.0
|
total_duration = (total_frames * frame_duration_ms) / 1000.0
|
||||||
return opus_datas, total_duration
|
return opus_datas, total_duration
|
||||||
@@ -3,6 +3,9 @@ import socket
|
|||||||
import subprocess
|
import subprocess
|
||||||
import re
|
import re
|
||||||
import os
|
import os
|
||||||
|
import wave
|
||||||
|
from io import BytesIO
|
||||||
|
from core.utils import p3
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import requests
|
import requests
|
||||||
import opuslib_next
|
import opuslib_next
|
||||||
@@ -773,6 +776,22 @@ def audio_to_data(audio_file_path, is_opus=True):
|
|||||||
return pcm_to_data(raw_data, is_opus), duration
|
return pcm_to_data(raw_data, is_opus), duration
|
||||||
|
|
||||||
|
|
||||||
|
def audio_bytes_to_data(audio_bytes, file_type, is_opus=True):
|
||||||
|
"""
|
||||||
|
直接用音频二进制数据转为opus/pcm数据,支持wav、mp3、p3
|
||||||
|
"""
|
||||||
|
if file_type == "p3":
|
||||||
|
# 直接用p3解码
|
||||||
|
return p3.decode_opus_from_bytes(audio_bytes)
|
||||||
|
else:
|
||||||
|
# 其他格式用pydub
|
||||||
|
audio = AudioSegment.from_file(BytesIO(audio_bytes), format=file_type, parameters=["-nostdin"])
|
||||||
|
audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2)
|
||||||
|
duration = len(audio) / 1000.0
|
||||||
|
raw_data = audio.raw_data
|
||||||
|
return pcm_to_data(raw_data, is_opus), duration
|
||||||
|
|
||||||
|
|
||||||
def pcm_to_data(raw_data, is_opus=True):
|
def pcm_to_data(raw_data, is_opus=True):
|
||||||
# 初始化Opus编码器
|
# 初始化Opus编码器
|
||||||
encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO)
|
encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO)
|
||||||
@@ -804,6 +823,33 @@ def pcm_to_data(raw_data, is_opus=True):
|
|||||||
return datas
|
return datas
|
||||||
|
|
||||||
|
|
||||||
|
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):
|
def check_vad_update(before_config, new_config):
|
||||||
if (
|
if (
|
||||||
new_config.get("selected_module") is None
|
new_config.get("selected_module") is None
|
||||||
@@ -882,3 +928,56 @@ def filter_sensitive_info(config: dict) -> dict:
|
|||||||
return filtered
|
return filtered
|
||||||
|
|
||||||
return _filter_dict(copy.deepcopy(config))
|
return _filter_dict(copy.deepcopy(config))
|
||||||
|
|
||||||
|
|
||||||
|
def get_vision_url(config: dict) -> str:
|
||||||
|
"""获取 vision URL
|
||||||
|
|
||||||
|
Args:
|
||||||
|
config: 配置字典
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
str: vision URL
|
||||||
|
"""
|
||||||
|
server_config = config["server"]
|
||||||
|
vision_explain = server_config.get("vision_explain", "")
|
||||||
|
if "你的" in vision_explain:
|
||||||
|
local_ip = get_local_ip()
|
||||||
|
port = int(server_config.get("http_port", 8003))
|
||||||
|
vision_explain = f"http://{local_ip}:{port}/mcp/vision/explain"
|
||||||
|
return vision_explain
|
||||||
|
|
||||||
|
|
||||||
|
def is_valid_image_file(file_data: bytes) -> bool:
|
||||||
|
"""
|
||||||
|
检查文件数据是否为有效的图片格式
|
||||||
|
|
||||||
|
Args:
|
||||||
|
file_data: 文件的二进制数据
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
bool: 如果是有效的图片格式返回True,否则返回False
|
||||||
|
"""
|
||||||
|
# 常见图片格式的魔数(文件头)
|
||||||
|
image_signatures = {
|
||||||
|
b"\xff\xd8\xff": "JPEG",
|
||||||
|
b"\x89PNG\r\n\x1a\n": "PNG",
|
||||||
|
b"GIF87a": "GIF",
|
||||||
|
b"GIF89a": "GIF",
|
||||||
|
b"BM": "BMP",
|
||||||
|
b"II*\x00": "TIFF",
|
||||||
|
b"MM\x00*": "TIFF",
|
||||||
|
b"RIFF": "WEBP",
|
||||||
|
}
|
||||||
|
|
||||||
|
# 检查文件头是否匹配任何已知的图片格式
|
||||||
|
for signature in image_signatures:
|
||||||
|
if file_data.startswith(signature):
|
||||||
|
return True
|
||||||
|
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def sanitize_tool_name(name: str) -> str:
|
||||||
|
"""Sanitize tool names for OpenAI compatibility."""
|
||||||
|
return re.sub(r"[^a-zA-Z0-9_-]", "_", name)
|
||||||
|
|||||||
@@ -0,0 +1,23 @@
|
|||||||
|
import os
|
||||||
|
import sys
|
||||||
|
|
||||||
|
# 添加项目根目录到Python路径
|
||||||
|
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
project_root = os.path.abspath(os.path.join(current_dir, "..", ".."))
|
||||||
|
sys.path.insert(0, project_root)
|
||||||
|
|
||||||
|
from config.logger import setup_logging
|
||||||
|
import importlib
|
||||||
|
|
||||||
|
logger = setup_logging()
|
||||||
|
|
||||||
|
|
||||||
|
def create_instance(class_name, *args, **kwargs):
|
||||||
|
# 创建LLM实例
|
||||||
|
if os.path.exists(os.path.join("core", "providers", "vllm", f"{class_name}.py")):
|
||||||
|
lib_name = f"core.providers.vllm.{class_name}"
|
||||||
|
if lib_name not in sys.modules:
|
||||||
|
sys.modules[lib_name] = importlib.import_module(f"{lib_name}")
|
||||||
|
return sys.modules[lib_name].VLLMProvider(*args, **kwargs)
|
||||||
|
|
||||||
|
raise ValueError(f"不支持的VLLM类型: {class_name},请检查该配置的type是否设置正确")
|
||||||
@@ -0,0 +1,143 @@
|
|||||||
|
import os
|
||||||
|
import yaml
|
||||||
|
import time
|
||||||
|
import hashlib
|
||||||
|
import fcntl
|
||||||
|
from typing import Dict
|
||||||
|
from contextlib import contextmanager
|
||||||
|
|
||||||
|
|
||||||
|
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:
|
||||||
|
fcntl.flock(self.file, fcntl.LOCK_EX | fcntl.LOCK_NB)
|
||||||
|
return self.file
|
||||||
|
except IOError:
|
||||||
|
if time.time() - self.start_time > self.timeout:
|
||||||
|
raise TimeoutError("获取文件锁超时")
|
||||||
|
time.sleep(0.1)
|
||||||
|
|
||||||
|
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||||
|
fcntl.flock(self.file, fcntl.LOCK_UN)
|
||||||
|
|
||||||
|
|
||||||
|
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()
|
||||||
|
default_response = {
|
||||||
|
"voice": "default",
|
||||||
|
"file_path": "config/assets/wakeup_words.wav",
|
||||||
|
"time": 0,
|
||||||
|
"text": "哈啰啊,我是小智啦,声音好听的台湾女孩一枚,超开心认识你耶,最近在忙啥,别忘了给我来点有趣的料哦,我超爱听八卦的啦",
|
||||||
|
}
|
||||||
|
|
||||||
|
if not config or voice not in config:
|
||||||
|
return default_response
|
||||||
|
|
||||||
|
# 检查文件大小
|
||||||
|
file_path = config[voice]["file_path"]
|
||||||
|
if not os.path.exists(file_path) or os.stat(file_path).st_size < (15 * 1024):
|
||||||
|
return default_response
|
||||||
|
|
||||||
|
return config[voice]
|
||||||
|
|
||||||
|
def update_wakeup_response(self, voice: str, file_path: str, text: str):
|
||||||
|
"""更新唤醒词回复配置"""
|
||||||
|
try:
|
||||||
|
config = self._load_config()
|
||||||
|
voice_hash = hashlib.md5(voice.encode()).hexdigest()
|
||||||
|
config[voice_hash] = {
|
||||||
|
"voice": voice,
|
||||||
|
"file_path": file_path,
|
||||||
|
"time": time.time(),
|
||||||
|
"text": 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
|
||||||
@@ -13,8 +13,8 @@ services:
|
|||||||
ports:
|
ports:
|
||||||
# ws服务端
|
# ws服务端
|
||||||
- "8000:8000"
|
- "8000:8000"
|
||||||
# ota服务端
|
# http服务的端口,用于简单OTA接口(单服务部署),以及视觉分析接口
|
||||||
- "8002:8002"
|
- "8003:8003"
|
||||||
volumes:
|
volumes:
|
||||||
# 配置文件目录
|
# 配置文件目录
|
||||||
- ./data:/opt/xiaozhi-esp32-server/data
|
- ./data:/opt/xiaozhi-esp32-server/data
|
||||||
|
|||||||
@@ -15,6 +15,8 @@ services:
|
|||||||
ports:
|
ports:
|
||||||
# ws服务端
|
# ws服务端
|
||||||
- "8000:8000"
|
- "8000:8000"
|
||||||
|
# http服务的端口,用于视觉分析接口
|
||||||
|
- "8003:8003"
|
||||||
security_opt:
|
security_opt:
|
||||||
- seccomp:unconfined
|
- seccomp:unconfined
|
||||||
environment:
|
environment:
|
||||||
|
|||||||
@@ -140,7 +140,9 @@ class AsyncPerformanceTester:
|
|||||||
|
|
||||||
print(f"🎵 测试 STT: {stt_name}")
|
print(f"🎵 测试 STT: {stt_name}")
|
||||||
|
|
||||||
text, _ = await stt.speech_to_text([self.test_wav_list[0]], "1")
|
text, _ = await stt.speech_to_text(
|
||||||
|
[self.test_wav_list[0]], "1", stt.audio_format
|
||||||
|
)
|
||||||
|
|
||||||
if text is None:
|
if text is None:
|
||||||
print(f"❌ {stt_name} 连接失败")
|
print(f"❌ {stt_name} 连接失败")
|
||||||
@@ -151,7 +153,7 @@ class AsyncPerformanceTester:
|
|||||||
|
|
||||||
for i, sentence in enumerate(self.test_wav_list, 1):
|
for i, sentence in enumerate(self.test_wav_list, 1):
|
||||||
start = time.time()
|
start = time.time()
|
||||||
text, _ = await stt.speech_to_text([sentence], "1")
|
text, _ = await stt.speech_to_text([sentence], "1", stt.audio_format)
|
||||||
duration = time.time() - start
|
duration = time.time() - start
|
||||||
total_time += duration
|
total_time += duration
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,189 @@
|
|||||||
|
import time
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
import statistics
|
||||||
|
import base64
|
||||||
|
from typing import Dict
|
||||||
|
from tabulate import tabulate
|
||||||
|
from config.settings import load_config
|
||||||
|
from core.utils.vllm import create_instance
|
||||||
|
|
||||||
|
# 设置全局日志级别为WARNING,抑制INFO级别日志
|
||||||
|
logging.basicConfig(level=logging.WARNING)
|
||||||
|
|
||||||
|
|
||||||
|
class AsyncVisionPerformanceTester:
|
||||||
|
def __init__(self):
|
||||||
|
self.config = load_config()
|
||||||
|
self.test_images = [
|
||||||
|
"../../docs/images/demo1.png",
|
||||||
|
"../../docs/images/demo2.png",
|
||||||
|
]
|
||||||
|
self.test_questions = [
|
||||||
|
"这张图片里有什么?",
|
||||||
|
"请详细描述这张图片的内容",
|
||||||
|
]
|
||||||
|
|
||||||
|
# 加载测试图片
|
||||||
|
self.results = {"vllm": {}}
|
||||||
|
|
||||||
|
async def _test_vllm(self, vllm_name: str, config: Dict) -> Dict:
|
||||||
|
"""异步测试单个视觉大模型性能"""
|
||||||
|
try:
|
||||||
|
# 检查API密钥配置
|
||||||
|
if "api_key" in config and any(
|
||||||
|
x in config["api_key"] for x in ["你的", "placeholder", "sk-xxx"]
|
||||||
|
):
|
||||||
|
print(f"⏭️ VLLM {vllm_name} 未配置api_key,已跳过")
|
||||||
|
return {"name": vllm_name, "type": "vllm", "errors": 1}
|
||||||
|
|
||||||
|
# 获取实际类型(兼容旧配置)
|
||||||
|
module_type = config.get("type", vllm_name)
|
||||||
|
vllm = create_instance(module_type, config)
|
||||||
|
|
||||||
|
print(f"🖼️ 测试 VLLM: {vllm_name}")
|
||||||
|
|
||||||
|
# 创建所有测试任务
|
||||||
|
test_tasks = []
|
||||||
|
for question in self.test_questions:
|
||||||
|
for image in self.test_images:
|
||||||
|
test_tasks.append(
|
||||||
|
self._test_single_vision(vllm_name, vllm, question, image)
|
||||||
|
)
|
||||||
|
|
||||||
|
# 并发执行所有测试
|
||||||
|
test_results = await asyncio.gather(*test_tasks)
|
||||||
|
|
||||||
|
# 处理结果
|
||||||
|
valid_results = [r for r in test_results if r is not None]
|
||||||
|
if not valid_results:
|
||||||
|
print(f"⚠️ {vllm_name} 无有效数据,可能配置错误")
|
||||||
|
return {"name": vllm_name, "type": "vllm", "errors": 1}
|
||||||
|
|
||||||
|
response_times = [r["response_time"] for r in valid_results]
|
||||||
|
|
||||||
|
# 过滤异常数据
|
||||||
|
mean = statistics.mean(response_times)
|
||||||
|
stdev = statistics.stdev(response_times) if len(response_times) > 1 else 0
|
||||||
|
filtered_times = [t for t in response_times if t <= mean + 3 * stdev]
|
||||||
|
|
||||||
|
if len(filtered_times) < len(test_tasks) * 0.5:
|
||||||
|
print(f"⚠️ {vllm_name} 有效数据不足,可能网络不稳定")
|
||||||
|
return {"name": vllm_name, "type": "vllm", "errors": 1}
|
||||||
|
|
||||||
|
return {
|
||||||
|
"name": vllm_name,
|
||||||
|
"type": "vllm",
|
||||||
|
"avg_response": sum(response_times) / len(response_times),
|
||||||
|
"std_response": (
|
||||||
|
statistics.stdev(response_times) if len(response_times) > 1 else 0
|
||||||
|
),
|
||||||
|
"errors": 0,
|
||||||
|
}
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
print(f"⚠️ VLLM {vllm_name} 测试失败: {str(e)}")
|
||||||
|
return {"name": vllm_name, "type": "vllm", "errors": 1}
|
||||||
|
|
||||||
|
async def _test_single_vision(
|
||||||
|
self, vllm_name: str, vllm, question: str, image: str
|
||||||
|
) -> Dict:
|
||||||
|
"""测试单个视觉问题的性能"""
|
||||||
|
try:
|
||||||
|
print(f"📝 {vllm_name} 开始测试: {question[:20]}...")
|
||||||
|
start_time = time.time()
|
||||||
|
|
||||||
|
# 读取图片并转换为base64
|
||||||
|
with open(image, "rb") as image_file:
|
||||||
|
image_data = image_file.read()
|
||||||
|
image_base64 = base64.b64encode(image_data).decode("utf-8")
|
||||||
|
|
||||||
|
# 直接获取响应
|
||||||
|
response = vllm.response(question, image_base64)
|
||||||
|
response_time = time.time() - start_time
|
||||||
|
print(f"✓ {vllm_name} 完成响应: {response_time:.3f}s")
|
||||||
|
|
||||||
|
return {
|
||||||
|
"name": vllm_name,
|
||||||
|
"type": "vllm",
|
||||||
|
"response_time": response_time,
|
||||||
|
}
|
||||||
|
except Exception as e:
|
||||||
|
print(f"⚠️ {vllm_name} 测试失败: {str(e)}")
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _print_results(self):
|
||||||
|
"""打印测试结果"""
|
||||||
|
vllm_table = []
|
||||||
|
for name, data in self.results["vllm"].items():
|
||||||
|
if data["errors"] == 0:
|
||||||
|
stability = data["std_response"] / data["avg_response"]
|
||||||
|
vllm_table.append(
|
||||||
|
[
|
||||||
|
name,
|
||||||
|
f"{data['avg_response']:.3f}秒",
|
||||||
|
f"{stability:.3f}",
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
if vllm_table:
|
||||||
|
print("\n视觉大模型性能排行:\n")
|
||||||
|
print(
|
||||||
|
tabulate(
|
||||||
|
vllm_table,
|
||||||
|
headers=["模型名称", "响应耗时", "稳定性"],
|
||||||
|
tablefmt="github",
|
||||||
|
colalign=("left", "right", "right"),
|
||||||
|
disable_numparse=True,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
print("\n⚠️ 没有可用的视觉大模型进行测试。")
|
||||||
|
|
||||||
|
async def run(self):
|
||||||
|
"""执行全量异步测试"""
|
||||||
|
print("🔍 开始筛选可用视觉大模型...")
|
||||||
|
|
||||||
|
if not self.test_images:
|
||||||
|
print(f"\n⚠️ {self.image_root} 路径下没有图片文件,无法进行测试")
|
||||||
|
return
|
||||||
|
|
||||||
|
# 创建所有测试任务
|
||||||
|
all_tasks = []
|
||||||
|
|
||||||
|
# VLLM测试任务
|
||||||
|
if self.config.get("VLLM") is not None:
|
||||||
|
for vllm_name, config in self.config.get("VLLM", {}).items():
|
||||||
|
if "api_key" in config and any(
|
||||||
|
x in config["api_key"] for x in ["你的", "placeholder", "sk-xxx"]
|
||||||
|
):
|
||||||
|
print(f"⏭️ VLLM {vllm_name} 未配置api_key,已跳过")
|
||||||
|
continue
|
||||||
|
print(f"🖼️ 添加VLLM测试任务: {vllm_name}")
|
||||||
|
all_tasks.append(self._test_vllm(vllm_name, config))
|
||||||
|
|
||||||
|
print(f"\n✅ 找到 {len(all_tasks)} 个可用视觉大模型")
|
||||||
|
print(f"✅ 使用 {len(self.test_images)} 张测试图片")
|
||||||
|
print(f"✅ 使用 {len(self.test_questions)} 个测试问题")
|
||||||
|
print("\n⏳ 开始并发测试所有模型...\n")
|
||||||
|
|
||||||
|
# 并发执行所有测试任务
|
||||||
|
all_results = await asyncio.gather(*all_tasks, return_exceptions=True)
|
||||||
|
|
||||||
|
# 处理结果
|
||||||
|
for result in all_results:
|
||||||
|
if isinstance(result, dict) and result["errors"] == 0:
|
||||||
|
self.results["vllm"][result["name"]] = result
|
||||||
|
|
||||||
|
# 打印结果
|
||||||
|
print("\n📊 生成测试报告...")
|
||||||
|
self._print_results()
|
||||||
|
|
||||||
|
|
||||||
|
async def main():
|
||||||
|
tester = AsyncVisionPerformanceTester()
|
||||||
|
await tester.run()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
asyncio.run(main())
|
||||||
@@ -21,7 +21,7 @@ cozepy==0.12.0
|
|||||||
mem0ai==0.1.62
|
mem0ai==0.1.62
|
||||||
bs4==0.0.2
|
bs4==0.0.2
|
||||||
modelscope==1.23.2
|
modelscope==1.23.2
|
||||||
sherpa_onnx==1.11.0
|
sherpa_onnx==1.12.0
|
||||||
mcp==1.8.1
|
mcp==1.8.1
|
||||||
cnlunar==0.2.0
|
cnlunar==0.2.0
|
||||||
PySocks==1.7.1
|
PySocks==1.7.1
|
||||||
@@ -31,3 +31,5 @@ chardet==5.2.0
|
|||||||
aioconsole==0.8.1
|
aioconsole==0.8.1
|
||||||
markitdown==0.1.1
|
markitdown==0.1.1
|
||||||
mcp-proxy==0.6.0
|
mcp-proxy==0.6.0
|
||||||
|
PyJWT==2.8.0
|
||||||
|
psutil==7.0.0
|
||||||
Reference in New Issue
Block a user